210 lines
5.8 KiB
Go
210 lines
5.8 KiB
Go
package app
|
|
|
|
import (
|
|
"database/sql"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"os"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
_ "modernc.org/sqlite"
|
|
)
|
|
|
|
func newAuthorizationTestServer(t *testing.T) (*Server, string) {
|
|
t.Helper()
|
|
db, err := sql.Open("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
db.SetMaxOpenConns(1)
|
|
t.Cleanup(func() { _ = db.Close() })
|
|
migration, err := os.ReadFile("../../migrations/001_initial.sql")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := db.Exec(string(migration)); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
sessionToken := "authorization-test-session"
|
|
sessionHash := hashValue(sessionToken)
|
|
if _, err := db.Exec(
|
|
`INSERT INTO technician_sessions
|
|
(token_hash, csrf_token, gitea_login, display_name, expires_at)
|
|
VALUES (?, 'csrf-token', 'taiadmin', 'TAI Admin', ?)`,
|
|
sessionHash[:], time.Now().UTC().Add(time.Hour),
|
|
); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return &Server{
|
|
db: db,
|
|
cfg: Config{
|
|
MaxHostLimit: 25,
|
|
CookieSecret: []byte("0123456789abcdef0123456789abcdef"),
|
|
},
|
|
}, sessionToken
|
|
}
|
|
|
|
func seedAuthorizationForUpdate(t *testing.T, server *Server) time.Time {
|
|
t.Helper()
|
|
for _, values := range [][]any{
|
|
{1, "sentinelone-linux", "SentinelOne", "sentinelone-linux", "1.0", "s1.deb", strings.Repeat("a", 64), true},
|
|
{2, "support-tool", "Support Tool", "support-tool", "2.0", "support.deb", strings.Repeat("b", 64), true},
|
|
} {
|
|
if _, err := server.db.Exec(
|
|
`INSERT INTO packages
|
|
(id, slug, display_name, package_name, package_version,
|
|
file_name, sha256, enabled)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?)`,
|
|
values...,
|
|
); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
|
|
expiresAt := time.Now().UTC().Add(48 * time.Hour).Truncate(time.Second)
|
|
if _, err := server.db.Exec(
|
|
`INSERT INTO authorizations
|
|
(id, code_hash, code_hint, created_by, customer_label, host_limit, expires_at)
|
|
VALUES (1, ?, 'ABCDE', 'taiadmin', 'Test migration', 3, ?)`,
|
|
make([]byte, 32), expiresAt,
|
|
); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := server.db.Exec(
|
|
`INSERT INTO authorization_packages
|
|
(authorization_id, package_id, package_slug, display_name,
|
|
package_name, package_version, file_name, sha256)
|
|
SELECT 1, id, slug, display_name, package_name, package_version,
|
|
file_name, sha256
|
|
FROM packages WHERE id = 1`,
|
|
); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := server.db.Exec(
|
|
`INSERT INTO authorization_hosts
|
|
(id, authorization_id, host_fingerprint, hostname)
|
|
VALUES
|
|
(1, 1, ?, 'pve01'),
|
|
(2, 1, ?, 'pve02')`,
|
|
make([]byte, 32), bytesOf(1, 32),
|
|
); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := server.db.Exec(
|
|
`INSERT INTO download_sessions
|
|
(token_hash, authorization_id, authorization_host_id, expires_at)
|
|
VALUES (?, 1, 1, ?)`,
|
|
make([]byte, 32), expiresAt,
|
|
); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return expiresAt
|
|
}
|
|
|
|
func authorizationUpdateRequest(
|
|
t *testing.T,
|
|
server *Server,
|
|
sessionToken string,
|
|
values url.Values,
|
|
) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
request := httptest.NewRequest(
|
|
http.MethodPost,
|
|
"/portal/authorizations/1/update",
|
|
strings.NewReader(values.Encode()),
|
|
)
|
|
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
request.AddCookie(&http.Cookie{Name: sessionCookieName, Value: sessionToken})
|
|
request.SetPathValue("id", "1")
|
|
response := httptest.NewRecorder()
|
|
server.handleUpdateAuthorization(response, request)
|
|
return response
|
|
}
|
|
|
|
func TestUpdateAuthorizationExtendsAndReplacesAccess(t *testing.T) {
|
|
server, sessionToken := newAuthorizationTestServer(t)
|
|
oldExpiresAt := seedAuthorizationForUpdate(t, server)
|
|
|
|
response := authorizationUpdateRequest(t, server, sessionToken, url.Values{
|
|
"csrf_token": {"csrf-token"},
|
|
"host_limit": {"6"},
|
|
"extend_days": {"7"},
|
|
"package_id": {"2"},
|
|
"action_slug": {"install-rmm"},
|
|
})
|
|
if response.Code != http.StatusSeeOther {
|
|
t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
|
|
}
|
|
|
|
var hostLimit int
|
|
var expiresAt time.Time
|
|
if err := server.db.QueryRow(
|
|
`SELECT host_limit, expires_at FROM authorizations WHERE id = 1`,
|
|
).Scan(&hostLimit, &expiresAt); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if hostLimit != 6 {
|
|
t.Fatalf("host limit = %d, want 6", hostLimit)
|
|
}
|
|
wantExpiresAt := oldExpiresAt.Add(7 * 24 * time.Hour)
|
|
if !expiresAt.Equal(wantExpiresAt) {
|
|
t.Fatalf("expiration = %s, want %s", expiresAt, wantExpiresAt)
|
|
}
|
|
|
|
var selectedPackage uint64
|
|
if err := server.db.QueryRow(
|
|
`SELECT package_id FROM authorization_packages WHERE authorization_id = 1`,
|
|
).Scan(&selectedPackage); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if selectedPackage != 2 {
|
|
t.Fatalf("selected package = %d, want 2", selectedPackage)
|
|
}
|
|
var selectedAction string
|
|
if err := server.db.QueryRow(
|
|
`SELECT action_slug FROM authorization_actions WHERE authorization_id = 1`,
|
|
).Scan(&selectedAction); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if selectedAction != "install-rmm" {
|
|
t.Fatalf("selected action = %q, want install-rmm", selectedAction)
|
|
}
|
|
|
|
var sessionExpiresAt time.Time
|
|
if err := server.db.QueryRow(
|
|
`SELECT expires_at FROM download_sessions WHERE authorization_id = 1`,
|
|
).Scan(&sessionExpiresAt); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if !sessionExpiresAt.Equal(wantExpiresAt) {
|
|
t.Fatalf("session expiration = %s, want %s", sessionExpiresAt, wantExpiresAt)
|
|
}
|
|
}
|
|
|
|
func TestUpdateAuthorizationRejectsLimitBelowEnrolledHosts(t *testing.T) {
|
|
server, sessionToken := newAuthorizationTestServer(t)
|
|
seedAuthorizationForUpdate(t, server)
|
|
|
|
response := authorizationUpdateRequest(t, server, sessionToken, url.Values{
|
|
"csrf_token": {"csrf-token"},
|
|
"host_limit": {"1"},
|
|
"extend_days": {"0"},
|
|
"package_id": {"1"},
|
|
})
|
|
if response.Code != http.StatusBadRequest {
|
|
t.Fatalf("status = %d, want %d", response.Code, http.StatusBadRequest)
|
|
}
|
|
}
|
|
|
|
func bytesOf(value byte, size int) []byte {
|
|
result := make([]byte, size)
|
|
for index := range result {
|
|
result[index] = value
|
|
}
|
|
return result
|
|
}
|