Files
2026-07-28 21:07:30 -05:00

305 lines
8.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 packageStatusRequest(
t *testing.T,
server *Server,
sessionToken string,
slug string,
enabled bool,
) *httptest.ResponseRecorder {
t.Helper()
values := url.Values{
"csrf_token": {"csrf-token"},
"slug": {slug},
}
if enabled {
values.Set("enabled", "1")
}
request := httptest.NewRequest(
http.MethodPost,
"/portal/packages/status",
strings.NewReader(values.Encode()),
)
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
request.AddCookie(&http.Cookie{Name: sessionCookieName, Value: sessionToken})
response := httptest.NewRecorder()
server.requireTechnician(server.handleSetPackageStatus)(response, request)
return response
}
func TestSetPackageStatusChangesOnlyDistributionAvailability(t *testing.T) {
server, sessionToken := newAuthorizationTestServer(t)
if _, err := server.db.Exec(
`INSERT INTO packages
(slug, display_name, package_name, package_version, file_name, sha256, enabled)
VALUES ('sentinelone-linux', 'SentinelOne', 'sentinelone-linux',
'26.1.1.31', 'agent.deb', ?, TRUE)`,
strings.Repeat("a", 64),
); err != nil {
t.Fatal(err)
}
response := packageStatusRequest(t, server, sessionToken, "sentinelone-linux", false)
if response.Code != http.StatusSeeOther {
t.Fatalf("disable status = %d, want %d: %s", response.Code, http.StatusSeeOther, response.Body.String())
}
if location := response.Header().Get("Location"); !strings.Contains(location, "disabled") {
t.Fatalf("disable redirect = %q", location)
}
var version, fileName, sha256 string
var enabled bool
if err := server.db.QueryRow(
`SELECT package_version, file_name, sha256, enabled
FROM packages WHERE slug = 'sentinelone-linux'`,
).Scan(&version, &fileName, &sha256, &enabled); err != nil {
t.Fatal(err)
}
if enabled || version != "26.1.1.31" || fileName != "agent.deb" ||
sha256 != strings.Repeat("a", 64) {
t.Fatalf("disable changed package metadata: enabled=%t version=%q file=%q sha256=%q", enabled, version, fileName, sha256)
}
response = packageStatusRequest(t, server, sessionToken, "sentinelone-linux", true)
if response.Code != http.StatusSeeOther {
t.Fatalf("enable status = %d, want %d: %s", response.Code, http.StatusSeeOther, response.Body.String())
}
if err := server.db.QueryRow(
`SELECT enabled FROM packages WHERE slug = 'sentinelone-linux'`,
).Scan(&enabled); err != nil {
t.Fatal(err)
}
if !enabled {
t.Fatal("package was not re-enabled")
}
rows, err := server.db.Query(
`SELECT event_type, actor, package_slug, details
FROM audit_events ORDER BY id`,
)
if err != nil {
t.Fatal(err)
}
defer rows.Close()
for index, expected := range []string{"package_disabled", "package_enabled"} {
if !rows.Next() {
t.Fatalf("missing audit event %d", index)
}
var eventType, actor, slug, details string
if err := rows.Scan(&eventType, &actor, &slug, &details); err != nil {
t.Fatal(err)
}
if eventType != expected || actor != "taiadmin" || slug != "sentinelone-linux" {
t.Fatalf("audit event = %q/%q/%q", eventType, actor, slug)
}
}
}
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
}