diff --git a/internal/server/store/sessions.go b/internal/server/store/sessions.go new file mode 100644 index 0000000..f669f7a --- /dev/null +++ b/internal/server/store/sessions.go @@ -0,0 +1,303 @@ +package store + +import ( + "context" + "crypto/sha256" + "crypto/subtle" + "database/sql" + "encoding/binary" + "errors" + "fmt" + "time" +) + +var ( + ErrSessionUnknown = errors.New("unknown or fenced session") + ErrTakeoverRequired = errors.New("a different client instance is already live") + ErrTakeoverMismatch = errors.New("takeover does not match the pending client instance") + ErrTakeoverAlreadyGranted = errors.New("a takeover grant is already active") +) + +type ClientRegistration struct { + ClientID string + Platform uint32 + Architecture string + DaemonVersion string + DaemonCWD string + SupportedShells []byte + Capabilities []byte + ClientInstanceID [16]byte + SessionID [16]byte + ConnectedAt time.Time +} + +type SessionRegistration struct { + SessionID [16]byte + Generation uint64 +} + +type TakeoverAuthorization struct { + ClientID string + ClientInstanceID [16]byte + RequestID [16]byte + ExpiresAt time.Time +} + +func (store *Store) RegisterClientSession(ctx context.Context, registration ClientRegistration) (SessionRegistration, error) { + var result SessionRegistration + if err := validateClientRegistration(registration); err != nil { + return result, err + } + store.writeMu.Lock() + defer store.writeMu.Unlock() + database, err := store.openDatabase() + if err != nil { + return result, err + } + now := registration.ConnectedAt.UTC().UnixNano() + if registration.Capabilities == nil { + registration.Capabilities = []byte{} + } + tx, err := database.BeginTx(ctx, nil) + if err != nil { + return result, err + } + var generation uint64 + err = tx.QueryRowContext(ctx, `SELECT generation FROM clients WHERE client_id = ?`, registration.ClientID).Scan(&generation) + newClient := errors.Is(err, sql.ErrNoRows) + if err != nil && !newClient { + _ = tx.Rollback() + return result, err + } + if newClient { + _, err = tx.ExecContext(ctx, `INSERT INTO clients ( +client_id, platform, architecture, daemon_version, daemon_cwd, supported_shells, capabilities, client_instance_id, +generation, connected_at, last_seen_at, charged_bytes +) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, ?, ?, 0)`, registration.ClientID, registration.Platform, + registration.Architecture, registration.DaemonVersion, registration.DaemonCWD, registration.SupportedShells, + registration.Capabilities, registration.ClientInstanceID[:], now, now) + generation = 0 + } + if err != nil { + _ = tx.Rollback() + return result, err + } + + var liveSessionID []byte + var liveInstance []byte + var liveGeneration uint64 + liveErr := tx.QueryRowContext(ctx, `SELECT session_id, client_instance_id, generation FROM sessions +WHERE client_id = ? AND closed_at IS NULL AND fenced_at IS NULL`, registration.ClientID).Scan(&liveSessionID, &liveInstance, &liveGeneration) + live := liveErr == nil + if liveErr != nil && !errors.Is(liveErr, sql.ErrNoRows) { + _ = tx.Rollback() + return result, liveErr + } + differentLiveInstance := live && (len(liveInstance) != 16 || subtle.ConstantTimeCompare(liveInstance, registration.ClientInstanceID[:]) != 1) + if differentLiveInstance { + var grantInstance []byte + var expiresAt int64 + var consumedAt sql.NullInt64 + grantErr := tx.QueryRowContext(ctx, `SELECT pending_instance_id, expires_at, consumed_at FROM takeover_authorizations WHERE client_id = ?`, registration.ClientID).Scan(&grantInstance, &expiresAt, &consumedAt) + grantMatches := grantErr == nil && !consumedAt.Valid && expiresAt > now && len(grantInstance) == 16 && subtle.ConstantTimeCompare(grantInstance, registration.ClientInstanceID[:]) == 1 + if !grantMatches { + // An unconsumed live grant pins its exact pending instance; a third + // claimant cannot replace that identity before it expires or is used. + if grantErr == nil && !consumedAt.Valid && expiresAt > now { + err = tx.Rollback() + if err != nil { + return result, err + } + return result, ErrTakeoverRequired + } + _, err = tx.ExecContext(ctx, `UPDATE clients SET pending_instance_id = ?, pending_instance_seen_at = ? WHERE client_id = ?`, registration.ClientInstanceID[:], now, registration.ClientID) + if err == nil { + err = tx.Commit() + } else { + _ = tx.Rollback() + } + if err != nil { + return result, err + } + return result, ErrTakeoverRequired + } + _, err = tx.ExecContext(ctx, `UPDATE takeover_authorizations SET consumed_at = ? WHERE client_id = ? AND consumed_at IS NULL AND expires_at > ?`, now, registration.ClientID, now) + if err != nil { + _ = tx.Rollback() + return result, err + } + } + if live { + _, err = tx.ExecContext(ctx, `UPDATE sessions SET fenced_at = ? WHERE session_id = ? AND generation = ? AND fenced_at IS NULL AND closed_at IS NULL`, now, liveSessionID, liveGeneration) + if err != nil { + _ = tx.Rollback() + return result, err + } + } + if !live { + // A sessionless registration never needs a grant. Clear any stale grant + // rather than allowing it to pin a future, unrelated pending claimant. + _, err = tx.ExecContext(ctx, `DELETE FROM takeover_authorizations WHERE client_id = ?`, registration.ClientID) + if err != nil { + _ = tx.Rollback() + return result, err + } + } + if generation == ^uint64(0) { + _ = tx.Rollback() + return result, errors.New("client generation exhausted") + } + generation++ + _, err = tx.ExecContext(ctx, `UPDATE clients SET platform = ?, architecture = ?, daemon_version = ?, daemon_cwd = ?, +supported_shells = ?, capabilities = ?, client_instance_id = ?, generation = ?, connected_at = ?, last_seen_at = ?, +pending_instance_id = NULL, pending_instance_seen_at = NULL WHERE client_id = ?`, registration.Platform, registration.Architecture, + registration.DaemonVersion, registration.DaemonCWD, registration.SupportedShells, registration.Capabilities, + registration.ClientInstanceID[:], generation, now, now, registration.ClientID) + if err == nil { + _, err = tx.ExecContext(ctx, `INSERT INTO sessions ( +session_id, client_id, client_instance_id, generation, opened_at +) VALUES (?, ?, ?, ?, ?)`, registration.SessionID[:], registration.ClientID, registration.ClientInstanceID[:], generation, now) + } + if err != nil { + _ = tx.Rollback() + return result, err + } + if err := tx.Commit(); err != nil { + return result, err + } + return SessionRegistration{SessionID: registration.SessionID, Generation: generation}, nil +} + +func (store *Store) AuthorizeClientTakeover(ctx context.Context, authorization TakeoverAuthorization) (time.Time, bool, error) { + if authorization.ClientID == "" || isZeroUUID(authorization.ClientInstanceID) || isZeroUUID(authorization.RequestID) || authorization.ExpiresAt.IsZero() { + return time.Time{}, false, ErrTakeoverMismatch + } + store.writeMu.Lock() + defer store.writeMu.Unlock() + database, err := store.openDatabase() + if err != nil { + return time.Time{}, false, err + } + now := time.Now().UTC().UnixNano() + expiresAt := authorization.ExpiresAt.UTC().UnixNano() + if expiresAt <= now { + return time.Time{}, false, ErrTakeoverMismatch + } + hash := takeoverHash(authorization.ClientID, authorization.ClientInstanceID) + target := authorization.ClientID + var existingMethod, existingTarget string + var existingHash, existingResult []byte + err = database.QueryRowContext(ctx, `SELECT method, target, immutable_sha256, result FROM control_mutations WHERE request_uuid = ?`, authorization.RequestID[:]).Scan(&existingMethod, &existingTarget, &existingHash, &existingResult) + if err == nil { + if existingMethod != "authorize_client_takeover" || existingTarget != target || len(existingHash) != 32 || subtle.ConstantTimeCompare(existingHash, hash[:]) != 1 || len(existingResult) != 8 { + return time.Time{}, false, ErrMutationConflict + } + return time.Unix(0, int64(binary.BigEndian.Uint64(existingResult))), true, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return time.Time{}, false, err + } + tx, err := database.BeginTx(ctx, nil) + if err != nil { + return time.Time{}, false, err + } + var pending []byte + err = tx.QueryRowContext(ctx, `SELECT pending_instance_id FROM clients WHERE client_id = ?`, authorization.ClientID).Scan(&pending) + if errors.Is(err, sql.ErrNoRows) { + err = ErrTakeoverMismatch + } + if err == nil && (len(pending) != 16 || subtle.ConstantTimeCompare(pending, authorization.ClientInstanceID[:]) != 1) { + err = ErrTakeoverMismatch + } + if err == nil { + var existingExpiry int64 + var consumed sql.NullInt64 + grantErr := tx.QueryRowContext(ctx, `SELECT expires_at, consumed_at FROM takeover_authorizations WHERE client_id = ?`, authorization.ClientID).Scan(&existingExpiry, &consumed) + if grantErr == nil && !consumed.Valid && existingExpiry > now { + err = ErrTakeoverAlreadyGranted + } else if grantErr != nil && !errors.Is(grantErr, sql.ErrNoRows) { + err = grantErr + } + } + if err == nil { + _, err = tx.ExecContext(ctx, `INSERT INTO takeover_authorizations ( +client_id, pending_instance_id, request_uuid, created_at, expires_at, consumed_at +) VALUES (?, ?, ?, ?, ?, NULL) +ON CONFLICT(client_id) DO UPDATE SET pending_instance_id = excluded.pending_instance_id, +request_uuid = excluded.request_uuid, created_at = excluded.created_at, expires_at = excluded.expires_at, consumed_at = NULL`, + authorization.ClientID, authorization.ClientInstanceID[:], authorization.RequestID[:], now, expiresAt) + } + if err == nil { + result := make([]byte, 8) + binary.BigEndian.PutUint64(result, uint64(expiresAt)) + _, err = tx.ExecContext(ctx, `INSERT INTO control_mutations ( +request_uuid, method, owner_kind, owner_id, target, immutable_sha256, result, created_at +) VALUES (?, 'authorize_client_takeover', 'client', ?, ?, ?, ?, ?)`, authorization.RequestID[:], target, target, hash[:], result, now) + } + if err != nil { + _ = tx.Rollback() + return time.Time{}, false, err + } + if err := tx.Commit(); err != nil { + return time.Time{}, false, err + } + return authorization.ExpiresAt.UTC(), false, nil +} + +func (store *Store) ValidateLiveSession(ctx context.Context, sessionID [16]byte, generation uint64) (string, error) { + if isZeroUUID(sessionID) || generation == 0 { + return "", ErrSessionUnknown + } + store.writeMu.Lock() + defer store.writeMu.Unlock() + database, err := store.openDatabase() + if err != nil { + return "", err + } + var clientID string + err = database.QueryRowContext(ctx, `SELECT client_id FROM sessions WHERE session_id = ? AND generation = ? AND fenced_at IS NULL AND closed_at IS NULL`, sessionID[:], generation).Scan(&clientID) + if errors.Is(err, sql.ErrNoRows) { + return "", ErrSessionUnknown + } + return clientID, err +} + +func (store *Store) CloseLiveSession(ctx context.Context, sessionID [16]byte, generation uint64, reason string, closedAt time.Time) error { + if isZeroUUID(sessionID) || generation == 0 || closedAt.IsZero() || len(reason) > 1024 { + return ErrSessionUnknown + } + store.writeMu.Lock() + defer store.writeMu.Unlock() + database, err := store.openDatabase() + if err != nil { + return err + } + result, err := database.ExecContext(ctx, `UPDATE sessions SET closed_at = ?, close_reason = ? +WHERE session_id = ? AND generation = ? AND fenced_at IS NULL AND closed_at IS NULL`, closedAt.UTC().UnixNano(), reason, sessionID[:], generation) + if err != nil { + return err + } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected != 1 { + return ErrSessionUnknown + } + return nil +} + +func validateClientRegistration(registration ClientRegistration) error { + if registration.ClientID == "" || len(registration.ClientID) > 128 || registration.Platform == 0 || registration.Architecture == "" || registration.DaemonVersion == "" || registration.DaemonCWD == "" || len(registration.SupportedShells) == 0 || isZeroUUID(registration.ClientInstanceID) || isZeroUUID(registration.SessionID) || registration.ConnectedAt.IsZero() { + return fmt.Errorf("invalid client registration") + } + return nil +} + +func takeoverHash(clientID string, instance [16]byte) [sha256.Size]byte { + value := make([]byte, 0, len(clientID)+17) + value = append(value, instance[:]...) + value = append(value, 0) + value = append(value, clientID...) + return sha256.Sum256(value) +} diff --git a/test/coverage.toml b/test/coverage.toml index 03b5f71..ce9a129 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -345,3 +345,15 @@ id = "BH-STORE-09" layer = "integration" status = "implemented" tests = ["test/integration/store/store_integration_test.go:TestIrreparableIncidentRequiresExplicitAcknowledgement_BH_STORE_09"] + +[[requirements]] +id = "HP-SES-02" +layer = "integration" +status = "implemented" +tests = ["test/integration/store/store_integration_test.go:TestSessionRegistrationFencingAndExactTakeover_HP_SES_02"] + +[[requirements]] +id = "BH-SES-02" +layer = "integration" +status = "implemented" +tests = ["test/integration/store/store_integration_test.go:TestTakeoverBoundariesAndClosedSession_BH_SES_02"] diff --git a/test/integration/store/store_integration_test.go b/test/integration/store/store_integration_test.go index c2427cc..2a6640c 100644 --- a/test/integration/store/store_integration_test.go +++ b/test/integration/store/store_integration_test.go @@ -883,6 +883,127 @@ func TestIrreparableIncidentRequiresExplicitAcknowledgement_BH_STORE_09(t *testi } } +func TestSessionRegistrationFencingAndExactTakeover_HP_SES_02(t *testing.T) { + t.Parallel() + + opened := openStore(t, filepath.Join(t.TempDir(), "state")) + now := time.Unix(7_000_000, 0).UTC() + instanceA := uuidBytes(1) + first := registerClient(t, opened, "client-a", instanceA, uuidBytes(20), now) + if first.Generation != 1 { + t.Fatalf("first generation = %d", first.Generation) + } + if clientID, err := opened.ValidateLiveSession(context.Background(), first.SessionID, first.Generation); err != nil || clientID != "client-a" { + t.Fatalf("first session validation = (%q, %v)", clientID, err) + } + replacement := registerClient(t, opened, "client-a", instanceA, uuidBytes(40), now.Add(time.Second)) + if replacement.Generation != 2 { + t.Fatalf("replacement generation = %d", replacement.Generation) + } + if _, err := opened.ValidateLiveSession(context.Background(), first.SessionID, first.Generation); !errors.Is(err, store.ErrSessionUnknown) { + t.Fatalf("fenced first session error = %v", err) + } + + instanceB := uuidBytes(60) + _, err := opened.RegisterClientSession(context.Background(), store.ClientRegistration{ + ClientID: "client-a", Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\work`, + SupportedShells: []byte{1}, ClientInstanceID: instanceB, SessionID: uuidBytes(80), ConnectedAt: now.Add(2 * time.Second), + }) + if !errors.Is(err, store.ErrTakeoverRequired) { + t.Fatalf("different live instance error = %v", err) + } + var pending []byte + if err := opened.DB().QueryRow(`SELECT pending_instance_id FROM clients WHERE client_id = 'client-a'`).Scan(&pending); err != nil || !bytes.Equal(pending, instanceB[:]) { + t.Fatalf("pending instance = %x, err = %v", pending, err) + } + grantRequest := uuidBytes(100) + expires := time.Now().UTC().Add(5 * time.Minute) + granted, replay, err := opened.AuthorizeClientTakeover(context.Background(), store.TakeoverAuthorization{ + ClientID: "client-a", ClientInstanceID: instanceB, RequestID: grantRequest, ExpiresAt: expires, + }) + if err != nil || replay || !granted.Equal(expires) { + t.Fatalf("takeover grant = (%s, replay=%v, err=%v)", granted, replay, err) + } + granted, replay, err = opened.AuthorizeClientTakeover(context.Background(), store.TakeoverAuthorization{ + ClientID: "client-a", ClientInstanceID: instanceB, RequestID: grantRequest, ExpiresAt: expires, + }) + if err != nil || !replay || !granted.Equal(expires) { + t.Fatalf("takeover replay = (%s, replay=%v, err=%v)", granted, replay, err) + } + accepted := registerClient(t, opened, "client-a", instanceB, uuidBytes(120), now.Add(3*time.Second)) + if accepted.Generation != 3 { + t.Fatalf("takeover generation = %d", accepted.Generation) + } + if _, err := opened.ValidateLiveSession(context.Background(), replacement.SessionID, replacement.Generation); !errors.Is(err, store.ErrSessionUnknown) { + t.Fatalf("takeover did not fence prior session: %v", err) + } + if clientID, err := opened.ValidateLiveSession(context.Background(), accepted.SessionID, accepted.Generation); err != nil || clientID != "client-a" { + t.Fatalf("accepted session validation = (%q, %v)", clientID, err) + } + if err := opened.Close(); err != nil { + t.Fatal(err) + } +} + +func TestTakeoverBoundariesAndClosedSession_BH_SES_02(t *testing.T) { + t.Parallel() + + opened := openStore(t, filepath.Join(t.TempDir(), "state")) + now := time.Unix(8_000_000, 0).UTC() + instanceA := uuidBytes(140) + active := registerClient(t, opened, "client-b", instanceA, uuidBytes(160), now) + instanceB := uuidBytes(180) + _, err := opened.RegisterClientSession(context.Background(), store.ClientRegistration{ + ClientID: "client-b", Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\work`, + SupportedShells: []byte{1}, ClientInstanceID: instanceB, SessionID: uuidBytes(200), ConnectedAt: now.Add(time.Second), + }) + if !errors.Is(err, store.ErrTakeoverRequired) { + t.Fatalf("takeover-needed error = %v", err) + } + if _, _, err := opened.AuthorizeClientTakeover(context.Background(), store.TakeoverAuthorization{ + ClientID: "client-b", ClientInstanceID: uuidBytes(220), RequestID: uuidBytes(10), ExpiresAt: time.Now().UTC().Add(time.Minute), + }); !errors.Is(err, store.ErrTakeoverMismatch) { + t.Fatalf("mismatched grant error = %v", err) + } + request := uuidBytes(30) + if _, _, err := opened.AuthorizeClientTakeover(context.Background(), store.TakeoverAuthorization{ + ClientID: "client-b", ClientInstanceID: instanceB, RequestID: request, ExpiresAt: time.Now().UTC().Add(time.Minute), + }); err != nil { + t.Fatal(err) + } + if _, _, err := opened.AuthorizeClientTakeover(context.Background(), store.TakeoverAuthorization{ + ClientID: "client-b", ClientInstanceID: instanceB, RequestID: uuidBytes(50), ExpiresAt: time.Now().UTC().Add(time.Minute), + }); !errors.Is(err, store.ErrTakeoverAlreadyGranted) { + t.Fatalf("second active grant error = %v", err) + } + if err := opened.CloseLiveSession(context.Background(), active.SessionID, active.Generation, "test close", now.Add(2*time.Second)); err != nil { + t.Fatal(err) + } + if err := opened.CloseLiveSession(context.Background(), active.SessionID, active.Generation, "again", now.Add(3*time.Second)); !errors.Is(err, store.ErrSessionUnknown) { + t.Fatalf("duplicate close error = %v", err) + } + // With no live session, a new instance is accepted without consuming the prior grant. + accepted := registerClient(t, opened, "client-b", instanceB, uuidBytes(70), now.Add(4*time.Second)) + if accepted.Generation != 2 { + t.Fatalf("closed-session registration generation = %d", accepted.Generation) + } + instanceC := uuidBytes(90) + if _, err := opened.RegisterClientSession(context.Background(), store.ClientRegistration{ + ClientID: "client-b", Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\work`, + SupportedShells: []byte{1}, ClientInstanceID: instanceC, SessionID: uuidBytes(110), ConnectedAt: now.Add(5 * time.Second), + }); !errors.Is(err, store.ErrTakeoverRequired) { + t.Fatalf("post-cleanup takeover-needed error = %v", err) + } + if _, _, err := opened.AuthorizeClientTakeover(context.Background(), store.TakeoverAuthorization{ + ClientID: "client-b", ClientInstanceID: instanceC, RequestID: uuidBytes(130), ExpiresAt: time.Now().UTC().Add(time.Minute), + }); err != nil { + t.Fatalf("stale grant blocked fresh pending instance: %v", err) + } + if err := opened.Close(); err != nil { + t.Fatal(err) + } +} + func openStore(t *testing.T, dataDir string) *store.Store { t.Helper() return openStoreWithOptions(t, store.Options{DataDir: dataDir, BusyTimeout: busyTimeout}) @@ -999,3 +1120,15 @@ func markTerminal(t *testing.T, database *sql.DB, issue [16]byte, issueTime, ter t.Fatal(err) } } + +func registerClient(t *testing.T, opened *store.Store, clientID string, instance, session [16]byte, at time.Time) store.SessionRegistration { + t.Helper() + registration, err := opened.RegisterClientSession(context.Background(), store.ClientRegistration{ + ClientID: clientID, Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\work`, + SupportedShells: []byte{1}, ClientInstanceID: instance, SessionID: session, ConnectedAt: at, + }) + if err != nil { + t.Fatalf("RegisterClientSession: %v", err) + } + return registration +}