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) }