package store import ( "context" "errors" "strings" "time" "github.com/jackc/pgx/v5" ) const workerHeartbeatStaleAfter = 30 * time.Second type WorkerRegistrationInput struct { InstanceID string PodUID string PodName string Site string Revision string DesiredCapacity int CapacityLimit int HeartbeatStaleAfter time.Duration } type WorkerAllocation struct { InstanceID string DesiredCapacity int Allocated int GlobalAllocated int ActiveInstances int HeartbeatAt time.Time } type activeWorkerCapacity struct { InstanceID string CapacityLimit int } // ActiveWorkerCapacity returns the cluster-wide execution capacity currently // owned by live Worker instances. func (s *Store) ActiveWorkerCapacity(ctx context.Context) (int, error) { var capacity int err := s.pool.QueryRow(ctx, ` SELECT COALESCE(SUM(allocated_capacity), 0)::int FROM gateway_worker_instances WHERE status = 'active' AND heartbeat_at > now() - $1::interval`, workerHeartbeatStaleAfter.String()).Scan(&capacity) return capacity, err } func (s *Store) RegisterWorkerInstance(ctx context.Context, input WorkerRegistrationInput) (WorkerAllocation, error) { input.InstanceID = strings.TrimSpace(input.InstanceID) if input.InstanceID == "" { return WorkerAllocation{}, errors.New("worker instance ID is required") } if input.DesiredCapacity < 0 { return WorkerAllocation{}, errors.New("worker desired capacity cannot be negative") } if input.CapacityLimit < 0 { return WorkerAllocation{}, errors.New("worker capacity limit cannot be negative") } if input.CapacityLimit == 0 { input.CapacityLimit = input.DesiredCapacity } staleAfter := input.HeartbeatStaleAfter if staleAfter < workerHeartbeatStaleAfter { staleAfter = workerHeartbeatStaleAfter } tx, err := s.pool.Begin(ctx) if err != nil { return WorkerAllocation{}, err } defer rollbackTransaction(tx) // Worker heartbeats and allocations are ephemeral coordination state that is // refreshed every few seconds. Do not hold the global allocation lock while // waiting for a synchronous replica to acknowledge the commit. if _, err := tx.Exec(ctx, `SELECT set_config('synchronous_commit', 'off', true)`); err != nil { return WorkerAllocation{}, err } if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-worker-capacity', 0))`); err != nil { return WorkerAllocation{}, err } if _, err := tx.Exec(ctx, ` INSERT INTO gateway_worker_instances ( instance_id, pod_uid, pod_name, site, revision, status, desired_capacity, capacity_limit, allocated_capacity, started_at, heartbeat_at, updated_at ) VALUES ($1, $2, $3, $4, $5, 'active', $6, $7, 0, now(), now(), now()) ON CONFLICT (instance_id) DO UPDATE SET pod_uid = EXCLUDED.pod_uid, pod_name = EXCLUDED.pod_name, site = EXCLUDED.site, revision = EXCLUDED.revision, status = 'active', desired_capacity = EXCLUDED.desired_capacity, capacity_limit = EXCLUDED.capacity_limit, heartbeat_at = now(), updated_at = now()`, input.InstanceID, strings.TrimSpace(input.PodUID), strings.TrimSpace(input.PodName), strings.TrimSpace(input.Site), strings.TrimSpace(input.Revision), input.DesiredCapacity, input.CapacityLimit, ); err != nil { return WorkerAllocation{}, err } if _, err := tx.Exec(ctx, ` UPDATE gateway_worker_instances SET allocated_capacity = 0, updated_at = now() WHERE status <> 'active' OR heartbeat_at <= now() - $1::interval`, staleAfter.String()); err != nil { return WorkerAllocation{}, err } rows, err := tx.Query(ctx, ` SELECT instance_id, capacity_limit FROM gateway_worker_instances WHERE status = 'active' AND heartbeat_at > now() - $1::interval ORDER BY instance_id ASC FOR UPDATE`, staleAfter.String()) if err != nil { return WorkerAllocation{}, err } activeWorkers := make([]activeWorkerCapacity, 0) for rows.Next() { var worker activeWorkerCapacity if err := rows.Scan(&worker.InstanceID, &worker.CapacityLimit); err != nil { rows.Close() return WorkerAllocation{}, err } activeWorkers = append(activeWorkers, worker) } if err := rows.Err(); err != nil { rows.Close() return WorkerAllocation{}, err } rows.Close() if len(activeWorkers) == 0 { return WorkerAllocation{}, errors.New("worker registration was not active after heartbeat") } allocations, globalAllocated := allocateWorkerCapacities(activeWorkers, input.DesiredCapacity) allocated := 0 for _, worker := range activeWorkers { capacity := allocations[worker.InstanceID] if _, err := tx.Exec(ctx, ` UPDATE gateway_worker_instances SET desired_capacity = $2, allocated_capacity = $3, updated_at = now() WHERE instance_id = $1`, worker.InstanceID, input.DesiredCapacity, capacity); err != nil { return WorkerAllocation{}, err } if worker.InstanceID == input.InstanceID { allocated = capacity } } var heartbeatAt time.Time if err := tx.QueryRow(ctx, ` SELECT heartbeat_at FROM gateway_worker_instances WHERE instance_id = $1`, input.InstanceID).Scan(&heartbeatAt); err != nil { return WorkerAllocation{}, err } if err := tx.Commit(ctx); err != nil { return WorkerAllocation{}, err } return WorkerAllocation{ InstanceID: input.InstanceID, DesiredCapacity: input.DesiredCapacity, Allocated: allocated, GlobalAllocated: globalAllocated, ActiveInstances: len(activeWorkers), HeartbeatAt: heartbeatAt, }, nil } func allocateWorkerCapacities(workers []activeWorkerCapacity, desired int) (map[string]int, int) { allocations := make(map[string]int, len(workers)) if desired <= 0 || len(workers) == 0 { return allocations, 0 } remaining := desired for remaining > 0 { progressed := false for _, worker := range workers { limit := worker.CapacityLimit if limit <= 0 { limit = desired } if allocations[worker.InstanceID] >= limit { continue } allocations[worker.InstanceID]++ remaining-- progressed = true if remaining == 0 { break } } if !progressed { break } } return allocations, desired - remaining } func (s *Store) MarkWorkerDraining(ctx context.Context, instanceID string) error { result, err := s.pool.Exec(ctx, ` UPDATE gateway_worker_instances SET status = 'draining', allocated_capacity = 0, heartbeat_at = now(), updated_at = now() WHERE instance_id = $1`, strings.TrimSpace(instanceID)) if err != nil { return err } if result.RowsAffected() == 0 { return pgx.ErrNoRows } return nil } // RecoverOrphanedAsyncRiverJobs retries River jobs whose owning worker is no // longer alive after the gateway execution lease has already returned the task // to queued. Normal long-running jobs remain protected by their running task // state and the active worker heartbeat. func (s *Store) RecoverOrphanedAsyncRiverJobs( ctx context.Context, workerStaleAfter time.Duration, limit int, ) (int64, error) { if workerStaleAfter < workerHeartbeatStaleAfter { workerStaleAfter = workerHeartbeatStaleAfter } if limit <= 0 || limit > 1000 { limit = 100 } tx, err := s.pool.Begin(ctx) if err != nil { return 0, err } defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-river-orphan-recovery', 0))`); err != nil { return 0, err } var recovered int64 if err := tx.QueryRow(ctx, ` WITH orphaned AS MATERIALIZED ( SELECT job.id AS job_id, task.id AS task_id, job.attempt FROM river_job job JOIN gateway_tasks task ON task.river_job_id = job.id WHERE job.queue = 'gateway_tasks' AND job.kind = 'gateway_task_run' AND job.state = 'running' AND job.attempted_at <= now() - $1::interval AND task.async_mode = true AND task.status = 'queued' AND (task.execution_lease_expires_at IS NULL OR task.execution_lease_expires_at <= now()) AND NOT EXISTS ( SELECT 1 FROM gateway_worker_instances worker WHERE worker.status = 'active' AND worker.heartbeat_at > now() - $1::interval AND COALESCE(job.attempted_by[array_length(job.attempted_by, 1)], '') LIKE worker.instance_id || '-exec-%' ) ORDER BY job.attempted_at ASC, job.id ASC LIMIT $2 FOR UPDATE OF job SKIP LOCKED ), released_leases AS ( UPDATE gateway_concurrency_leases lease SET released_at = now() FROM orphaned WHERE lease.task_id = orphaned.task_id AND lease.released_at IS NULL ), cleared_admissions AS ( DELETE FROM gateway_task_admissions admission USING orphaned WHERE admission.task_id = orphaned.task_id ), reset_tasks AS ( UPDATE gateway_tasks task SET locked_by = NULL, locked_at = NULL, heartbeat_at = NULL, execution_token = NULL, execution_lease_expires_at = NULL, next_run_at = now(), updated_at = now() FROM orphaned WHERE task.id = orphaned.task_id ), recovered_jobs AS ( UPDATE river_job job SET errors = array_append( job.errors, jsonb_build_object( 'at', now(), 'attempt', orphaned.attempt, 'error', 'Orphaned gateway job recovered after worker heartbeat expired', 'trace', '' ) ), finalized_at = NULL, scheduled_at = now(), state = 'retryable' FROM orphaned WHERE job.id = orphaned.job_id RETURNING job.id ) SELECT count(*)::bigint FROM recovered_jobs`, workerStaleAfter.String(), limit).Scan(&recovered); err != nil { return 0, err } if err := tx.Commit(ctx); err != nil { return 0, err } return recovered, nil }