update gui
This commit is contained in:
@@ -0,0 +1,209 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user