feat: persist client session fencing
This commit is contained in:
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user