305 lines
8.8 KiB
Go
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
|
|
}
|