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 }