From 6bab0f0749d0b57e2b6f8a5fa6b912fb4659ad51 Mon Sep 17 00:00:00 2001 From: wangbo Date: Thu, 30 Jul 2026 17:27:11 +0800 Subject: [PATCH] =?UTF-8?q?fix(worker):=20=E9=98=B2=E6=AD=A2=E4=BA=8B?= =?UTF-8?q?=E5=8A=A1=E6=B3=84=E6=BC=8F=E5=B9=B6=E6=81=A2=E5=A4=8D=E6=BB=9E?= =?UTF-8?q?=E7=95=99=E9=98=9F=E5=88=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 为 PostgreSQL 连接增加可配置的事务空闲与锁等待超时,并在请求取消或 Worker 退出后使用独立有界上下文回滚事务。\n\n恢复过期任务时按批次使用 SKIP LOCKED,周期重建缺失或终态 River Job;锁等待超时只在确认尚未提交上游时安全重排,避免重复执行和重复结算。\n\n验证:完整 Go 测试、go vet、生产 Kustomize 渲染、gofmt 与 git diff --check 均通过。 --- apps/api/cmd/gateway/main.go | 6 +- apps/api/internal/config/config.go | 55 +++++-- apps/api/internal/config/config_test.go | 22 +++ apps/api/internal/runner/queue_worker.go | 38 +++++ apps/api/internal/store/access_rules.go | 2 +- apps/api/internal/store/admission_queue.go | 10 +- .../store/admission_queue_integration_test.go | 154 ++++++++++++++++++ apps/api/internal/store/base_models.go | 4 +- .../api/internal/store/billing_settlements.go | 4 +- .../internal/store/binary_result_backfill.go | 4 +- apps/api/internal/store/conversations.go | 2 +- apps/api/internal/store/identity_admin.go | 2 +- .../internal/store/identity_configurations.go | 8 +- apps/api/internal/store/identity_pairing.go | 10 +- .../internal/store/oidc_multi_tenant_users.go | 2 +- .../store/oidc_tenant_context_sync.go | 4 +- apps/api/internal/store/oidc_users.go | 2 +- apps/api/internal/store/platform_models.go | 4 +- apps/api/internal/store/postgres.go | 50 +++++- .../internal/store/postgres_config_test.go | 32 ++++ apps/api/internal/store/pricing_rules.go | 4 +- apps/api/internal/store/rate_limit_status.go | 2 +- apps/api/internal/store/rate_limits.go | 31 ++-- apps/api/internal/store/runtime_policies.go | 4 +- .../store/security_event_connections.go | 14 +- apps/api/internal/store/security_events.go | 8 +- apps/api/internal/store/task_history.go | 2 +- apps/api/internal/store/tasks_runtime.go | 123 ++++++++++++-- apps/api/internal/store/tx.go | 36 +++- apps/api/internal/store/wallet.go | 6 +- apps/api/internal/store/worker_registry.go | 4 +- .../production/application-config.yaml | 2 + 32 files changed, 545 insertions(+), 106 deletions(-) diff --git a/apps/api/cmd/gateway/main.go b/apps/api/cmd/gateway/main.go index 595b762..181f898 100644 --- a/apps/api/cmd/gateway/main.go +++ b/apps/api/cmd/gateway/main.go @@ -37,7 +37,11 @@ func main() { ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) - db, err := store.ConnectWithMaxConns(ctx, cfg.DatabaseURL, cfg.DatabaseMaxConns) + db, err := store.ConnectWithPoolOptions(ctx, cfg.DatabaseURL, store.PostgresPoolOptions{ + MaxConns: cfg.DatabaseMaxConns, + IdleInTransactionTimeout: time.Duration(cfg.DatabaseIdleInTransactionTimeoutSeconds) * time.Second, + LockTimeout: time.Duration(cfg.DatabaseLockTimeoutSeconds) * time.Second, + }) if err != nil { stop() logger.Error("connect postgres failed", "error", err) diff --git a/apps/api/internal/config/config.go b/apps/api/internal/config/config.go index 074c944..92ee0c2 100644 --- a/apps/api/internal/config/config.go +++ b/apps/api/internal/config/config.go @@ -59,6 +59,8 @@ type Config struct { BillingEngineMode string ProcessRole string DatabaseMaxConns int + DatabaseIdleInTransactionTimeoutSeconds int + DatabaseLockTimeoutSeconds int MediaRequestConcurrency int MediaMaterializationConcurrency int AsyncQueueWorkerEnabled bool @@ -106,20 +108,27 @@ func Load() Config { TaskProgressCallbackURL: env("TASK_PROGRESS_CALLBACK_URL", strings.TrimRight(env("SERVER_MAIN_BASE_URL", "http://localhost:3000"), "/")+"/internal/platform/task-progress-callbacks", ), - TaskProgressCallbackTimeoutMS: envIntValidated("TASK_PROGRESS_CALLBACK_TIMEOUT_MS", 5000), - TaskProgressCallbackMaxAttempts: envIntValidated("TASK_PROGRESS_CALLBACK_MAX_ATTEMPTS", 10), - TaskCleanupEnabled: env("AI_GATEWAY_TASK_CLEANUP_ENABLED", "false") == "true", - TaskRetentionDays: envIntValidated("AI_GATEWAY_TASK_RETENTION_DAYS", 30), - TaskAnalysisRetentionDays: envIntValidated("AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS", 7), - TaskCleanupIntervalSeconds: envIntValidated("AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS", 300), - TaskCleanupBatchSize: envIntValidated("AI_GATEWAY_TASK_CLEANUP_BATCH_SIZE", 1000), - CORSAllowedOrigin: env("CORS_ALLOWED_ORIGIN", "http://localhost:5178,http://127.0.0.1:5178"), - GlobalHTTPProxy: globalProxy.HTTPProxy, - GlobalHTTPProxySource: globalProxy.Source, - LogLevel: logLevel(env("LOG_LEVEL", "info")), - BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")), - ProcessRole: strings.ToLower(strings.TrimSpace(env("AI_GATEWAY_PROCESS_ROLE", "all"))), - DatabaseMaxConns: envInt("AI_GATEWAY_DATABASE_MAX_CONNS", 0), + TaskProgressCallbackTimeoutMS: envIntValidated("TASK_PROGRESS_CALLBACK_TIMEOUT_MS", 5000), + TaskProgressCallbackMaxAttempts: envIntValidated("TASK_PROGRESS_CALLBACK_MAX_ATTEMPTS", 10), + TaskCleanupEnabled: env("AI_GATEWAY_TASK_CLEANUP_ENABLED", "false") == "true", + TaskRetentionDays: envIntValidated("AI_GATEWAY_TASK_RETENTION_DAYS", 30), + TaskAnalysisRetentionDays: envIntValidated("AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS", 7), + TaskCleanupIntervalSeconds: envIntValidated("AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS", 300), + TaskCleanupBatchSize: envIntValidated("AI_GATEWAY_TASK_CLEANUP_BATCH_SIZE", 1000), + CORSAllowedOrigin: env("CORS_ALLOWED_ORIGIN", "http://localhost:5178,http://127.0.0.1:5178"), + GlobalHTTPProxy: globalProxy.HTTPProxy, + GlobalHTTPProxySource: globalProxy.Source, + LogLevel: logLevel(env("LOG_LEVEL", "info")), + BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")), + ProcessRole: strings.ToLower(strings.TrimSpace(env("AI_GATEWAY_PROCESS_ROLE", "all"))), + DatabaseMaxConns: envInt("AI_GATEWAY_DATABASE_MAX_CONNS", 0), + + DatabaseIdleInTransactionTimeoutSeconds: envOptionalIntValidated( + "AI_GATEWAY_DATABASE_IDLE_IN_TRANSACTION_TIMEOUT_SECONDS", + 60, + ), + DatabaseLockTimeoutSeconds: envOptionalIntValidated("AI_GATEWAY_DATABASE_LOCK_TIMEOUT_SECONDS", 30), + MediaRequestConcurrency: envIntValidated("AI_GATEWAY_MEDIA_REQUEST_CONCURRENCY", 16), MediaMaterializationConcurrency: envIntValidated("AI_GATEWAY_MEDIA_MATERIALIZATION_CONCURRENCY", 8), AsyncQueueWorkerEnabled: env("AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED", "true") == "true", @@ -138,6 +147,12 @@ func (c Config) Validate() error { if c.DatabaseMaxConns < 0 || c.DatabaseMaxConns > 1000 { return errors.New("AI_GATEWAY_DATABASE_MAX_CONNS must be between 1 and 1000 when configured") } + if c.DatabaseIdleInTransactionTimeoutSeconds < 0 || c.DatabaseIdleInTransactionTimeoutSeconds > 3600 { + return errors.New("AI_GATEWAY_DATABASE_IDLE_IN_TRANSACTION_TIMEOUT_SECONDS must be between 0 and 3600") + } + if c.DatabaseLockTimeoutSeconds < 0 || c.DatabaseLockTimeoutSeconds > 3600 { + return errors.New("AI_GATEWAY_DATABASE_LOCK_TIMEOUT_SECONDS must be between 0 and 3600") + } if c.MediaMaterializationConcurrency != 0 && (c.MediaMaterializationConcurrency < 1 || c.MediaMaterializationConcurrency > 256) { return errors.New("AI_GATEWAY_MEDIA_MATERIALIZATION_CONCURRENCY must be between 1 and 256") } @@ -336,6 +351,18 @@ func envIntValidated(key string, fallback int) int { return parsed } +func envOptionalIntValidated(key string, fallback int) int { + value := envValue(key) + if value == "" { + return fallback + } + parsed, err := strconv.Atoi(value) + if err != nil { + return -1 + } + return parsed +} + func envInt64Validated(key string, fallback int64) int64 { value := envValue(key) if value == "" { diff --git a/apps/api/internal/config/config_test.go b/apps/api/internal/config/config_test.go index c131c4d..72bc10b 100644 --- a/apps/api/internal/config/config_test.go +++ b/apps/api/internal/config/config_test.go @@ -144,6 +144,13 @@ func TestProcessRolePrecedenceAndCompatibility(t *testing.T) { func TestValidateProcessRoleAndDatabasePool(t *testing.T) { cfg := Load() + if cfg.DatabaseIdleInTransactionTimeoutSeconds != 60 || cfg.DatabaseLockTimeoutSeconds != 30 { + t.Fatalf( + "database transaction timeouts = %d/%d, want 60/30", + cfg.DatabaseIdleInTransactionTimeoutSeconds, + cfg.DatabaseLockTimeoutSeconds, + ) + } cfg.ProcessRole = "invalid" if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "PROCESS_ROLE") { t.Fatalf("Validate() error = %v, want invalid process role", err) @@ -154,6 +161,21 @@ func TestValidateProcessRoleAndDatabasePool(t *testing.T) { t.Fatalf("Validate() error = %v, want invalid database max conns", err) } cfg.DatabaseMaxConns = 16 + cfg.DatabaseIdleInTransactionTimeoutSeconds = 3601 + if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "IDLE_IN_TRANSACTION") { + t.Fatalf("Validate() error = %v, want invalid idle transaction timeout", err) + } + cfg.DatabaseIdleInTransactionTimeoutSeconds = 60 + cfg.DatabaseLockTimeoutSeconds = 3601 + if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "LOCK_TIMEOUT") { + t.Fatalf("Validate() error = %v, want invalid lock timeout", err) + } + cfg.DatabaseLockTimeoutSeconds = 30 + t.Setenv("AI_GATEWAY_DATABASE_LOCK_TIMEOUT_SECONDS", "not-an-integer") + if err := Load().Validate(); err == nil || !strings.Contains(err.Error(), "LOCK_TIMEOUT") { + t.Fatalf("Validate() error = %v, want invalid non-integer lock timeout", err) + } + t.Setenv("AI_GATEWAY_DATABASE_LOCK_TIMEOUT_SECONDS", "30") cfg.MediaMaterializationConcurrency = 257 if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "MEDIA_MATERIALIZATION_CONCURRENCY") { t.Fatalf("Validate() error = %v, want invalid media materialization concurrency", err) diff --git a/apps/api/internal/runner/queue_worker.go b/apps/api/internal/runner/queue_worker.go index 5bca6fc..aeea9a5 100644 --- a/apps/api/internal/runner/queue_worker.go +++ b/apps/api/internal/runner/queue_worker.go @@ -80,6 +80,38 @@ func (w *asyncTaskWorker) Work(ctx context.Context, job *river.Job[asyncTaskArgs w.service.logger.Debug("river async task interrupted and requeued", "taskID", task.ID, "status", queued.Status, "riverJobID", job.ID) return river.JobSnooze(0) } + if store.IsPostgresLockTimeout(runErr) { + queued, changed, queueErr := w.service.store.RequeueTaskBeforeUpstreamSubmission( + context.WithoutCancel(ctx), + task.ID, + executionToken, + time.Second, + ) + if queueErr != nil { + return queueErr + } + if changed { + if eventErr := w.service.emit( + context.WithoutCancel(ctx), + task.ID, + "task.queued", + "queued", + "database_lock_timeout", + 0.2, + "async task queued after database lock timeout", + map[string]any{"code": "database_lock_timeout"}, + task.RunMode == "simulation", + ); eventErr != nil { + w.service.logger.Warn("record database lock timeout requeue event failed", "taskID", task.ID, "error", eventErr) + } + w.service.logger.Warn("river async task requeued after database lock timeout", + "taskID", task.ID, + "status", queued.Status, + "riverJobID", job.ID, + ) + return river.JobSnooze(time.Second) + } + } w.service.logger.Warn("river async task completed with failure", "taskID", task.ID, "error", runErr, "riverJobID", job.ID) return nil } @@ -570,6 +602,12 @@ func (s *Service) recoverOrphanedAsyncRiverJobs(ctx context.Context) { if recovered > 0 { s.logger.Warn("orphaned river jobs recovered", "count", recovered) } + if err := s.recoverAsyncRiverJobs(ctx); err != nil { + if ctx.Err() == nil { + s.logger.Warn("recover queued tasks without active river jobs failed", "error", err) + } + return + } } recoverJobs() ticker := time.NewTicker(orphanedRiverJobScanInterval) diff --git a/apps/api/internal/store/access_rules.go b/apps/api/internal/store/access_rules.go index 3eaf1ba..73f1ef2 100644 --- a/apps/api/internal/store/access_rules.go +++ b/apps/api/internal/store/access_rules.go @@ -153,7 +153,7 @@ func (s *Store) BatchAccessRules(ctx context.Context, input AccessRuleBatchInput if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) for _, resource := range dedupeAccessRuleResources(input.DeleteResources) { resource = normalizeAccessRuleResource(resource, input.Effect) diff --git a/apps/api/internal/store/admission_queue.go b/apps/api/internal/store/admission_queue.go index f0f6bc0..307c8df 100644 --- a/apps/api/internal/store/admission_queue.go +++ b/apps/api/internal/store/admission_queue.go @@ -134,7 +134,7 @@ func (s *Store) QueueTaskAdmissionWithHook( if err != nil { return TaskAdmission{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil { return TaskAdmission{}, err @@ -232,7 +232,7 @@ func (s *Store) tryTaskAdmission( if err != nil { return TaskAdmissionResult{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil { return TaskAdmissionResult{}, err @@ -456,7 +456,7 @@ func (s *Store) RebindWaitingTaskAdmission(ctx context.Context, input TaskAdmiss if err != nil { return TaskAdmission{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil { return TaskAdmission{}, err } @@ -633,7 +633,7 @@ func (s *Store) DeleteTaskAdmission(ctx context.Context, taskID string) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil { return err } @@ -794,7 +794,7 @@ func (s *Store) ReapExpiredTaskAdmissions(ctx context.Context, limit int) (TaskA if err != nil { return TaskAdmissionReapResult{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) rows, err := tx.Query(ctx, ` SELECT task_id::text FROM gateway_task_admissions diff --git a/apps/api/internal/store/admission_queue_integration_test.go b/apps/api/internal/store/admission_queue_integration_test.go index 6a70358..1275255 100644 --- a/apps/api/internal/store/admission_queue_integration_test.go +++ b/apps/api/internal/store/admission_queue_integration_test.go @@ -543,6 +543,160 @@ SELECT t.Fatalf("recovered task left admissions=%d leases=%d", recoveredAdmissions, recoveredLeases) } + lockedRecoveryTask := createTask(true) + unlockedRecoveryTask := createTask(true) + if _, err := first.pool.Exec(ctx, ` +UPDATE gateway_tasks +SET status = 'running', + execution_token = gen_random_uuid(), + execution_lease_expires_at = now() - interval '1 second' +WHERE id = ANY($1::uuid[])`, []string{lockedRecoveryTask.ID, unlockedRecoveryTask.ID}); err != nil { + t.Fatalf("prepare locked recovery tasks: %v", err) + } + lockTx, err := second.pool.Begin(ctx) + if err != nil { + t.Fatalf("begin task row lock: %v", err) + } + var lockedTaskID string + if err := lockTx.QueryRow(ctx, ` +SELECT id::text +FROM gateway_tasks +WHERE id = $1::uuid +FOR UPDATE`, lockedRecoveryTask.ID).Scan(&lockedTaskID); err != nil { + rollbackTransaction(lockTx) + t.Fatalf("lock interrupted task row: %v", err) + } + recoveryCtx, recoveryCancel := context.WithTimeout(ctx, 3*time.Second) + recovery, err = first.RecoverInterruptedRuntimeState(recoveryCtx) + recoveryCancel() + if err != nil { + rollbackTransaction(lockTx) + t.Fatalf("recover unlocked task while another row is locked: %v", err) + } + if recovery.RequeuedAsyncTasks < 1 { + rollbackTransaction(lockTx) + t.Fatalf("SKIP LOCKED recovery = %+v, want at least one requeued task", recovery) + } + var lockedStatus, unlockedStatus string + if err := first.pool.QueryRow(ctx, ` +SELECT + (SELECT status FROM gateway_tasks WHERE id = $1::uuid), + (SELECT status FROM gateway_tasks WHERE id = $2::uuid)`, + lockedRecoveryTask.ID, + unlockedRecoveryTask.ID, + ).Scan(&lockedStatus, &unlockedStatus); err != nil { + rollbackTransaction(lockTx) + t.Fatalf("read SKIP LOCKED recovery states: %v", err) + } + if lockedStatus != "running" || unlockedStatus != "queued" { + rollbackTransaction(lockTx) + t.Fatalf("SKIP LOCKED recovery statuses = %s/%s, want running/queued", lockedStatus, unlockedStatus) + } + rollbackTransaction(lockTx) + recovery, err = first.RecoverInterruptedRuntimeState(ctx) + if err != nil { + t.Fatalf("recover previously locked task: %v", err) + } + if recovery.RequeuedAsyncTasks < 1 { + t.Fatalf("unlocked recovery = %+v, want requeued task", recovery) + } + if err := first.pool.QueryRow(ctx, ` +SELECT status +FROM gateway_tasks +WHERE id = $1::uuid`, lockedRecoveryTask.ID).Scan(&lockedStatus); err != nil { + t.Fatalf("read unlocked recovery status: %v", err) + } + if lockedStatus != "queued" { + t.Fatalf("unlocked recovery status = %s, want queued", lockedStatus) + } + + preSubmissionTask := createTask(true) + preSubmissionToken := uuid.NewString() + if _, err := first.pool.Exec(ctx, ` +UPDATE gateway_tasks +SET status = 'running', + execution_token = $2::uuid, + execution_lease_expires_at = now() + interval '5 minutes' +WHERE id = $1::uuid`, + preSubmissionTask.ID, + preSubmissionToken, + ); err != nil { + t.Fatalf("prepare pre-submission lock timeout task: %v", err) + } + if _, err := first.pool.Exec(ctx, ` +INSERT INTO gateway_task_attempts ( + task_id, attempt_no, queue_key, status, upstream_submission_status +) +VALUES ($1::uuid, 1, 'integration-test', 'running', 'not_submitted')`, + preSubmissionTask.ID, + ); err != nil { + t.Fatalf("create pre-submission attempt: %v", err) + } + requeuedTask, changed, err := first.RequeueTaskBeforeUpstreamSubmission( + ctx, + preSubmissionTask.ID, + preSubmissionToken, + time.Second, + ) + if err != nil || !changed || requeuedTask.Status != "queued" { + t.Fatalf("pre-submission lock timeout requeue task=%+v changed=%v err=%v", requeuedTask, changed, err) + } + var preSubmissionAttempts int + if err := first.pool.QueryRow(ctx, ` +SELECT count(*) +FROM gateway_task_attempts +WHERE task_id = $1::uuid`, preSubmissionTask.ID).Scan(&preSubmissionAttempts); err != nil { + t.Fatalf("count cleaned pre-submission attempts: %v", err) + } + if preSubmissionAttempts != 0 { + t.Fatalf("pre-submission lock timeout left %d attempts, want 0", preSubmissionAttempts) + } + + ambiguousTask := createTask(true) + ambiguousToken := uuid.NewString() + if _, err := first.pool.Exec(ctx, ` +UPDATE gateway_tasks +SET status = 'running', + execution_token = $2::uuid, + execution_lease_expires_at = now() + interval '5 minutes' +WHERE id = $1::uuid`, + ambiguousTask.ID, + ambiguousToken, + ); err != nil { + t.Fatalf("prepare ambiguous submission task: %v", err) + } + if _, err := first.pool.Exec(ctx, ` +INSERT INTO gateway_task_attempts ( + task_id, attempt_no, queue_key, status, upstream_submission_status +) +VALUES ($1::uuid, 1, 'integration-test', 'running', 'submitting')`, + ambiguousTask.ID, + ); err != nil { + t.Fatalf("create ambiguous submission attempt: %v", err) + } + _, changed, err = first.RequeueTaskBeforeUpstreamSubmission( + ctx, + ambiguousTask.ID, + ambiguousToken, + time.Second, + ) + if err != nil { + t.Fatalf("guard ambiguous submission task requeue: %v", err) + } + if changed { + t.Fatal("ambiguous submission task was requeued") + } + var ambiguousStatus string + if err := first.pool.QueryRow(ctx, ` +SELECT status +FROM gateway_tasks +WHERE id = $1::uuid`, ambiguousTask.ID).Scan(&ambiguousStatus); err != nil { + t.Fatalf("read ambiguous submission task: %v", err) + } + if ambiguousStatus != "running" { + t.Fatalf("ambiguous submission task status = %s, want running", ambiguousStatus) + } + terminalResidue := createTask(true) result, err = first.TryTaskAdmission(ctx, inputFor(terminalResidue, 100, "")) if err != nil || !result.Admitted { diff --git a/apps/api/internal/store/base_models.go b/apps/api/internal/store/base_models.go index 3b38452..620f097 100644 --- a/apps/api/internal/store/base_models.go +++ b/apps/api/internal/store/base_models.go @@ -101,7 +101,7 @@ func (s *Store) CreateBaseModel(ctx context.Context, input BaseModelInput) (Base if err != nil { return BaseModel{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) item, err := scanBaseModel(tx.QueryRow(ctx, ` INSERT INTO base_model_catalog ( @@ -162,7 +162,7 @@ func (s *Store) UpdateBaseModel(ctx context.Context, id string, input BaseModelI if err != nil { return BaseModel{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) item, err := scanBaseModel(tx.QueryRow(ctx, ` UPDATE base_model_catalog diff --git a/apps/api/internal/store/billing_settlements.go b/apps/api/internal/store/billing_settlements.go index 567df9b..c6c680e 100644 --- a/apps/api/internal/store/billing_settlements.go +++ b/apps/api/internal/store/billing_settlements.go @@ -157,7 +157,7 @@ RETURNING outbox.id::text, outbox.task_id::text, outbox.action, outbox.amount::t } func (s *Store) ProcessBillingSettlement(ctx context.Context, settlement BillingSettlement) error { - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.beginTransaction(ctx, func(tx pgx.Tx) error { var action string var amount string var currency string @@ -479,7 +479,7 @@ func (s *Store) MarkBillingSettlementFailed(ctx context.Context, settlement Bill taskStatus = "manual_review" manualReason = "maximum settlement attempts exceeded" } - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.beginTransaction(ctx, func(tx pgx.Tx) error { tag, err := tx.Exec(ctx, ` UPDATE settlement_outbox SET status = $3, diff --git a/apps/api/internal/store/binary_result_backfill.go b/apps/api/internal/store/binary_result_backfill.go index 5df9528..70872d7 100644 --- a/apps/api/internal/store/binary_result_backfill.go +++ b/apps/api/internal/store/binary_result_backfill.go @@ -19,7 +19,7 @@ func (s *Store) ListTaskBinaryResultBackfillBatch(ctx context.Context, afterID s batchSize = 100 } var items []TaskBinaryResultBackfillItem - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'; SET LOCAL statement_timeout = '5s'`); err != nil { return err } @@ -64,7 +64,7 @@ func (s *Store) UpdateTaskBinaryResultBackfill(ctx context.Context, taskID strin return false, err } updated := false - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.beginTransaction(ctx, func(tx pgx.Tx) error { if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'; SET LOCAL statement_timeout = '5s'`); err != nil { return err } diff --git a/apps/api/internal/store/conversations.go b/apps/api/internal/store/conversations.go index 1bf0292..b9dffe1 100644 --- a/apps/api/internal/store/conversations.go +++ b/apps/api/internal/store/conversations.go @@ -62,7 +62,7 @@ func (s *Store) UpsertConversationMessages(ctx context.Context, conversationID s if err != nil { return nil, 0, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) refs := make([]TaskMessageRefInput, 0, len(messages)) newCount := 0 diff --git a/apps/api/internal/store/identity_admin.go b/apps/api/internal/store/identity_admin.go index 7eb165b..97dcc81 100644 --- a/apps/api/internal/store/identity_admin.go +++ b/apps/api/internal/store/identity_admin.go @@ -272,7 +272,7 @@ func (s *Store) DeleteUserGroup(ctx context.Context, id string) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, ` DELETE FROM gateway_access_rules diff --git a/apps/api/internal/store/identity_configurations.go b/apps/api/internal/store/identity_configurations.go index f2d0407..90a81ab 100644 --- a/apps/api/internal/store/identity_configurations.go +++ b/apps/api/internal/store/identity_configurations.go @@ -59,7 +59,7 @@ func (s *Store) beginIdentityConfigurationLifecycleTx(ctx context.Context) (pgx. return nil, err } if err := lockIdentityConfigurationLifecycle(ctx, tx); err != nil { - _ = tx.Rollback(ctx) + rollbackTransaction(tx) return nil, err } return tx, nil @@ -171,7 +171,7 @@ func (s *Store) ApplyIdentityManifest(ctx context.Context, id string, expectedVe if err != nil { return identity.Revision{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) current, err := scanIdentityRevision(tx.QueryRow(ctx, `SELECT `+identityRevisionColumns+` FROM gateway_identity_configuration_revisions WHERE id=$1::uuid FOR UPDATE`, id)) if err != nil { @@ -282,7 +282,7 @@ func (s *Store) ActivateIdentityRevision(ctx context.Context, id string, expecte if err != nil { return identity.Revision{}, false, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var foreignPairingReservation bool if err := tx.QueryRow(ctx, `SELECT EXISTS( SELECT 1 FROM gateway_identity_pairing_start_reservation @@ -364,7 +364,7 @@ func (s *Store) DisableActiveIdentityRevision(ctx context.Context, expectedVersi if err != nil { return identity.Revision{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if ok, err := hasBreakGlassManager(ctx, tx); err != nil { return identity.Revision{}, err } else if !ok { diff --git a/apps/api/internal/store/identity_pairing.go b/apps/api/internal/store/identity_pairing.go index 8bace2c..4889204 100644 --- a/apps/api/internal/store/identity_pairing.go +++ b/apps/api/internal/store/identity_pairing.go @@ -20,7 +20,7 @@ func (s *Store) ReserveIdentityPairingStart(ctx context.Context, attemptID strin if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `DELETE FROM gateway_identity_pairing_start_reservation WHERE state='starting' AND expires_at <= now()`); err != nil { return err @@ -68,7 +68,7 @@ func (s *Store) CommitIdentityPairingStart(ctx context.Context, revision identit if err != nil { return identity.PairingExchange{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var reserved bool if err := tx.QueryRow(ctx, `SELECT true FROM gateway_identity_pairing_start_reservation WHERE singleton=true AND attempt_id=$1::uuid AND state='starting' FOR UPDATE`, exchange.ID).Scan(&reserved); errors.Is(err, pgx.ErrNoRows) { @@ -256,7 +256,7 @@ func (s *Store) UpdateIdentityPairingExchange(ctx context.Context, id string, ex if err != nil { return identity.PairingExchange{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) exchange, err := scanIdentityPairing(tx.QueryRow(ctx, ` UPDATE gateway_identity_onboarding_exchanges SET status=$3,remote_version=$4,last_error_category=NULLIF($5,''), auth_center_audit_id=COALESCE(NULLIF($6,''),auth_center_audit_id),version=version+1,updated_at=now() @@ -297,7 +297,7 @@ func (s *Store) CancelIdentityPairingExchange(ctx context.Context, id string, ex if err != nil { return identity.PairingExchange{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) exchange, err := scanIdentityPairing(tx.QueryRow(ctx, `SELECT `+identityPairingColumns+` FROM gateway_identity_onboarding_exchanges WHERE id=$1::uuid FOR UPDATE`, id)) if errors.Is(err, pgx.ErrNoRows) { @@ -352,7 +352,7 @@ func (s *Store) CompleteIdentityPairingCleanup(ctx context.Context, id string, e if err != nil { return identity.PairingExchange{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var revisionID string if err := tx.QueryRow(ctx, `SELECT revision_id::text FROM gateway_identity_onboarding_exchanges WHERE id=$1::uuid AND version=$2 AND status='cancelled' AND cleanup_status='pending' FOR UPDATE`, id, expectedVersion).Scan(&revisionID); errors.Is(err, pgx.ErrNoRows) { diff --git a/apps/api/internal/store/oidc_multi_tenant_users.go b/apps/api/internal/store/oidc_multi_tenant_users.go index a2b59f8..fc531be 100644 --- a/apps/api/internal/store/oidc_multi_tenant_users.go +++ b/apps/api/internal/store/oidc_multi_tenant_users.go @@ -24,7 +24,7 @@ func (s *Store) resolveOrProvisionOIDCMultiTenantUser(ctx context.Context, input if err != nil { return ResolveOrProvisionOIDCUserResult{}, err } - defer func() { _ = tx.Rollback(ctx) }() + defer rollbackTransaction(tx) bindingID, gatewayTenantID, tenantKey, groupID, groupKey, accessStatus, err := loadOIDCTenantBinding(ctx, tx, input) diff --git a/apps/api/internal/store/oidc_tenant_context_sync.go b/apps/api/internal/store/oidc_tenant_context_sync.go index 05b79e7..f6ae638 100644 --- a/apps/api/internal/store/oidc_tenant_context_sync.go +++ b/apps/api/internal/store/oidc_tenant_context_sync.go @@ -78,7 +78,7 @@ func (s *Store) ApplyOIDCTenantBindingSync(ctx context.Context, bindingID string if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer rollbackTransaction(tx) if unchanged { tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings SET last_sync_at=$2::timestamptz,next_sync_at=$2::timestamptz+interval '15 minutes',sync_failure_count=0, @@ -129,7 +129,7 @@ func (s *Store) RejectOIDCTenantBinding(ctx context.Context, bindingID, category if err != nil { return err } - defer func() { _ = tx.Rollback(ctx) }() + defer rollbackTransaction(tx) tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings SET access_status='disabled',metadata_status='rejected',last_sync_at=$2::timestamptz, next_sync_at=$2::timestamptz+interval '15 minutes', diff --git a/apps/api/internal/store/oidc_users.go b/apps/api/internal/store/oidc_users.go index daba193..b525076 100644 --- a/apps/api/internal/store/oidc_users.go +++ b/apps/api/internal/store/oidc_users.go @@ -69,7 +69,7 @@ func (s *Store) ResolveOrProvisionOIDCUser(ctx context.Context, input ResolveOrP if err != nil { return ResolveOrProvisionOIDCUserResult{}, err } - defer func() { _ = tx.Rollback(ctx) }() + defer rollbackTransaction(tx) projection, err := loadOIDCUserProjection(ctx, tx, input.Subject) if err == nil { diff --git a/apps/api/internal/store/platform_models.go b/apps/api/internal/store/platform_models.go index f8126ce..5475a4b 100644 --- a/apps/api/internal/store/platform_models.go +++ b/apps/api/internal/store/platform_models.go @@ -47,7 +47,7 @@ func (s *Store) ReplacePlatformModels(ctx context.Context, platformID string, in if err != nil { return nil, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) keptIDs := make([]string, 0, len(inputs)) for _, input := range inputs { @@ -372,7 +372,7 @@ func (s *Store) DeletePlatformModel(ctx context.Context, id string) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) result, err := tx.Exec(ctx, `DELETE FROM platform_models WHERE id = $1::uuid`, id) if err != nil { diff --git a/apps/api/internal/store/postgres.go b/apps/api/internal/store/postgres.go index 539aa17..07b9e96 100644 --- a/apps/api/internal/store/postgres.go +++ b/apps/api/internal/store/postgres.go @@ -8,6 +8,7 @@ import ( "encoding/json" "errors" "net" + "strconv" "strings" "time" "unicode" @@ -29,6 +30,12 @@ const ( postgresPoolWarmLimit = 64 ) +type PostgresPoolOptions struct { + MaxConns int + IdleInTransactionTimeout time.Duration + LockTimeout time.Duration +} + func defaultAPIKeyScopes() []string { return []string{"chat", "embedding", "rerank", "image", "image_vectorize", "video", "video_enhance", "music", "audio", "voice_clone"} } @@ -81,7 +88,11 @@ func Connect(ctx context.Context, databaseURL string) (*Store, error) { } func ConnectWithMaxConns(ctx context.Context, databaseURL string, maxConns int) (*Store, error) { - config, err := postgresPoolConfig(databaseURL, maxConns) + return ConnectWithPoolOptions(ctx, databaseURL, PostgresPoolOptions{MaxConns: maxConns}) +} + +func ConnectWithPoolOptions(ctx context.Context, databaseURL string, options PostgresPoolOptions) (*Store, error) { + config, err := postgresPoolConfigWithOptions(databaseURL, options) if err != nil { return nil, err } @@ -97,19 +108,37 @@ func ConnectWithMaxConns(ctx context.Context, databaseURL string, maxConns int) } func postgresPoolConfig(databaseURL string, maxConns ...int) (*pgxpool.Config, error) { + options := PostgresPoolOptions{} + if len(maxConns) > 0 { + options.MaxConns = maxConns[0] + } + return postgresPoolConfigWithOptions(databaseURL, options) +} + +func postgresPoolConfigWithOptions(databaseURL string, options PostgresPoolOptions) (*pgxpool.Config, error) { config, err := pgxpool.ParseConfig(databaseURL) if err != nil { return nil, err } config.ConnConfig.ConnectTimeout = postgresConnectTimeout config.ConnConfig.RuntimeParams["application_name"] = postgresApplicationName - if len(maxConns) > 0 && maxConns[0] > 0 { - config.MaxConns = int32(maxConns[0]) - config.MinIdleConns = int32(min(maxConns[0], postgresPoolWarmLimit)) + if options.IdleInTransactionTimeout > 0 { + config.ConnConfig.RuntimeParams["idle_in_transaction_session_timeout"] = postgresDuration(options.IdleInTransactionTimeout) + } + if options.LockTimeout > 0 { + config.ConnConfig.RuntimeParams["lock_timeout"] = postgresDuration(options.LockTimeout) + } + if options.MaxConns > 0 { + config.MaxConns = int32(options.MaxConns) + config.MinIdleConns = int32(min(options.MaxConns, postgresPoolWarmLimit)) } return config, nil } +func postgresDuration(value time.Duration) string { + return strconv.FormatInt(max(value.Milliseconds(), 1), 10) + "ms" +} + func warmPostgresPool(ctx context.Context, pool *pgxpool.Pool, count int) error { if count <= 0 { return pool.Ping(ctx) @@ -152,6 +181,11 @@ func IsPostgresUnavailable(err error) bool { return pgconn.SafeToRetry(err) } +func IsPostgresLockTimeout(err error) bool { + var postgresError *pgconn.PgError + return errors.As(err, &postgresError) && postgresError.Code == "55P03" +} + func (s *Store) Close() { s.pool.Close() } @@ -938,7 +972,7 @@ func (s *Store) DeletePlatform(ctx context.Context, id string) error { if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) rows, err := tx.Query(ctx, `SELECT id::text FROM platform_models WHERE platform_id = $1::uuid`, id) if err != nil { @@ -1638,7 +1672,7 @@ func (s *Store) DeleteAPIKey(ctx context.Context, apiKeyID string, user *auth.Us if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) result, err := tx.Exec(ctx, ` UPDATE gateway_api_keys @@ -1792,7 +1826,7 @@ func (s *Store) RegisterLocalUser(ctx context.Context, input LocalRegisterInput) if err != nil { return GatewayUser{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var tenantID string userGroupID := "" @@ -2065,7 +2099,7 @@ func (s *Store) CreateTaskIdempotent(ctx context.Context, input CreateTaskInput, if err != nil { return CreateTaskResult{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) task, err := scanGatewayTask(tx.QueryRow(ctx, ` INSERT INTO gateway_tasks ( diff --git a/apps/api/internal/store/postgres_config_test.go b/apps/api/internal/store/postgres_config_test.go index 904431c..e910f40 100644 --- a/apps/api/internal/store/postgres_config_test.go +++ b/apps/api/internal/store/postgres_config_test.go @@ -55,6 +55,29 @@ func TestPostgresPoolConfigWarmsBoundedIdleConnections(t *testing.T) { } } +func TestPostgresPoolConfigSetsBoundedTransactionTimeouts(t *testing.T) { + config, err := postgresPoolConfigWithOptions( + "postgresql://gateway:password@localhost:5432/gateway?sslmode=disable", + PostgresPoolOptions{ + MaxConns: 32, + IdleInTransactionTimeout: 60 * time.Second, + LockTimeout: 30 * time.Second, + }, + ) + if err != nil { + t.Fatalf("parse PostgreSQL pool config: %v", err) + } + if got := config.ConnConfig.RuntimeParams["idle_in_transaction_session_timeout"]; got != "60000ms" { + t.Fatalf("idle transaction timeout = %q, want 60000ms", got) + } + if got := config.ConnConfig.RuntimeParams["lock_timeout"]; got != "30000ms" { + t.Fatalf("lock timeout = %q, want 30000ms", got) + } + if config.MaxConns != 32 || config.MinIdleConns != 32 { + t.Fatalf("pool bounds max=%d minIdle=%d, want 32/32", config.MaxConns, config.MinIdleConns) + } +} + func TestIsPostgresUnavailableClassifiesConnectivityFailures(t *testing.T) { for _, testCase := range []struct { name string @@ -75,3 +98,12 @@ func TestIsPostgresUnavailableClassifiesConnectivityFailures(t *testing.T) { t.Fatal("SQL syntax error was incorrectly classified as PostgreSQL unavailable") } } + +func TestIsPostgresLockTimeout(t *testing.T) { + if !IsPostgresLockTimeout(&pgconn.PgError{Code: "55P03"}) { + t.Fatal("lock timeout was not classified") + } + if IsPostgresLockTimeout(&pgconn.PgError{Code: "57014"}) { + t.Fatal("query cancellation was classified as lock timeout") + } +} diff --git a/apps/api/internal/store/pricing_rules.go b/apps/api/internal/store/pricing_rules.go index a536230..0bd6085 100644 --- a/apps/api/internal/store/pricing_rules.go +++ b/apps/api/internal/store/pricing_rules.go @@ -113,7 +113,7 @@ func (s *Store) CreatePricingRuleSet(ctx context.Context, input PricingRuleSetIn if err != nil { return PricingRuleSet{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata)) item, err := scanPricingRuleSet(tx.QueryRow(ctx, ` @@ -141,7 +141,7 @@ func (s *Store) UpdatePricingRuleSet(ctx context.Context, id string, input Prici if err != nil { return PricingRuleSet{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata)) item, err := scanPricingRuleSet(tx.QueryRow(ctx, ` diff --git a/apps/api/internal/store/rate_limit_status.go b/apps/api/internal/store/rate_limit_status.go index 123b617..72caadd 100644 --- a/apps/api/internal/store/rate_limit_status.go +++ b/apps/api/internal/store/rate_limit_status.go @@ -98,7 +98,7 @@ func (s *Store) RestorePlatformModelRuntimeStatus(ctx context.Context, platformM if err != nil { return ModelRateLimitStatus{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var restoredModelID string if err := tx.QueryRow(ctx, ` diff --git a/apps/api/internal/store/rate_limits.go b/apps/api/internal/store/rate_limits.go index b830e4a..baa006d 100644 --- a/apps/api/internal/store/rate_limits.go +++ b/apps/api/internal/store/rate_limits.go @@ -19,6 +19,8 @@ type RuntimeRecoveryResult struct { CleanedTaskAdmissions int64 `json:"cleanedTaskAdmissions"` } +const runtimeRecoveryBatchSize = 100 + var ErrConcurrencyLeaseLost = errors.New("concurrency lease lost") // CheckRateLimits performs a non-consuming admission preflight for fixed-window @@ -104,7 +106,7 @@ func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID if err != nil { return RateLimitResult{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) lockKeys := make([]string, 0) lockKeySet := make(map[string]struct{}) @@ -414,7 +416,7 @@ func (s *Store) AttachRateLimitResultToAttempt(ctx context.Context, attemptID st if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) for _, reservation := range result.Reservations { if reservation.ReservationID == "" { @@ -458,7 +460,7 @@ func (s *Store) RecoverInterruptedRuntimeState(ctx context.Context) (RuntimeReco if err != nil { return RuntimeRecoveryResult{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-runtime-recovery', 0))`); err != nil { return RuntimeRecoveryResult{}, err } @@ -542,7 +544,18 @@ WHERE status = 'running' result.FailedAttempts = tag.RowsAffected() asyncTaskRows, err := tx.Query(ctx, ` -UPDATE gateway_tasks +WITH recoverable_async_tasks AS MATERIALIZED ( + SELECT id + FROM gateway_tasks + WHERE async_mode = true + AND status = 'running' + AND execution_lease_expires_at IS NOT NULL + AND execution_lease_expires_at <= now() + ORDER BY execution_lease_expires_at ASC, id ASC + LIMIT $1 + FOR UPDATE SKIP LOCKED +) +UPDATE gateway_tasks task SET status = 'queued', error = NULL, error_code = NULL, @@ -555,11 +568,9 @@ SET status = 'queued', next_run_at = now(), finished_at = NULL, updated_at = now() -WHERE async_mode = true - AND status = 'running' - AND execution_lease_expires_at IS NOT NULL - AND execution_lease_expires_at <= now() -RETURNING id::text`) +FROM recoverable_async_tasks recoverable +WHERE task.id = recoverable.id +RETURNING task.id::text`, runtimeRecoveryBatchSize) if err != nil { return RuntimeRecoveryResult{}, err } @@ -718,7 +729,7 @@ func (s *Store) finishRateLimitReservations(ctx context.Context, reservations [] if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) for _, reservation := range reservations { if reservation.ReservationID == "" { diff --git a/apps/api/internal/store/runtime_policies.go b/apps/api/internal/store/runtime_policies.go index 8fabc83..071ccbf 100644 --- a/apps/api/internal/store/runtime_policies.go +++ b/apps/api/internal/store/runtime_policies.go @@ -248,9 +248,7 @@ func (s *Store) ApplyCandidateFailureEffect(ctx context.Context, input Candidate if err != nil { return CandidateFailureEffectResult{}, err } - defer func() { - _ = tx.Rollback(ctx) - }() + defer rollbackTransaction(tx) lockKey := strings.TrimSpace(input.RequestedModel) + "\x1f" + strings.TrimSpace(input.ModelType) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1::text, 0))`, lockKey); err != nil { diff --git a/apps/api/internal/store/security_event_connections.go b/apps/api/internal/store/security_event_connections.go index 668cf10..606bb7c 100644 --- a/apps/api/internal/store/security_event_connections.go +++ b/apps/api/internal/store/security_event_connections.go @@ -185,7 +185,7 @@ func (s *Store) CreateSecurityEventConnection(ctx context.Context, input CreateS if err != nil { return SecurityEventConnection{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) _, err = tx.Exec(ctx, ` INSERT INTO gateway_security_event_connections( connection_id,transmitter_issuer,endpoint_url,credential_ref,management_client_id,management_credential_ref,lifecycle_status,idempotency_key @@ -196,7 +196,7 @@ INSERT INTO gateway_security_event_connections( ) if err != nil { if isUniqueViolation(err) { - _ = tx.Rollback(ctx) + rollbackTransaction(tx) existing, getErr := s.SecurityEventConnection(ctx) if getErr == nil && existing.IdempotencyKey == input.IdempotencyKey && existing.TransmitterIssuer == input.TransmitterIssuer { return existing, nil @@ -256,7 +256,7 @@ func (s *Store) SetSecurityEventManagementCredential(ctx context.Context, connec if err != nil { return SecurityEventConnection{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var oldReference sql.NullString var lifecycle string var version int64 @@ -315,7 +315,7 @@ func (s *Store) SetSecurityEventNextCredential(ctx context.Context, connectionID if err != nil { return SecurityEventConnection{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) tag, err := tx.Exec(ctx, `UPDATE gateway_security_event_connections SET next_credential_ref=$2,lifecycle_status='rotating',last_error_category=NULL,version=version+1,updated_at=now() WHERE connection_id=$1::uuid AND version=$3 AND next_credential_ref IS NULL @@ -340,7 +340,7 @@ func (s *Store) PromoteSecurityEventCredential(ctx context.Context, connectionID if err != nil { return SecurityEventConnection{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var oldReference string if err := tx.QueryRow(ctx, `SELECT credential_ref FROM gateway_security_event_connections WHERE connection_id=$1::uuid AND next_credential_ref=$2 AND lifecycle_status='rotating' FOR UPDATE`, connectionID, nextReference).Scan(&oldReference); errors.Is(err, pgx.ErrNoRows) { @@ -381,7 +381,7 @@ func (s *Store) DiscardPreparedSecurityEventConnection(ctx context.Context, conn if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var credentialReference string var nextReference, managementReference sql.NullString @@ -435,7 +435,7 @@ func (s *Store) FinalizeRetiringSecurityEventConnection(ctx context.Context, con if err != nil { return err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var credentialReference string var nextReference, managementReference sql.NullString diff --git a/apps/api/internal/store/security_events.go b/apps/api/internal/store/security_events.go index 209f3e0..1ab3007 100644 --- a/apps/api/internal/store/security_events.go +++ b/apps/api/internal/store/security_events.go @@ -13,7 +13,6 @@ import ( const ( securityEventTransactionLockTimeout = 5 * time.Second securityEventTransactionIdleTimeout = 15 * time.Second - transactionRollbackTimeout = 5 * time.Second ) func (s *Store) beginSecurityEventTransaction(ctx context.Context) (pgx.Tx, error) { @@ -38,12 +37,7 @@ func (s *Store) beginSecurityEventTransaction(ctx context.Context) (pgx.Tx, erro } func rollbackSecurityEventTransaction(tx pgx.Tx) { - if tx == nil { - return - } - ctx, cancel := context.WithTimeout(context.Background(), transactionRollbackTimeout) - defer cancel() - _ = tx.Rollback(ctx) + rollbackTransaction(tx) } type ApplySessionRevokedInput struct { diff --git a/apps/api/internal/store/task_history.go b/apps/api/internal/store/task_history.go index 022e154..d82095d 100644 --- a/apps/api/internal/store/task_history.go +++ b/apps/api/internal/store/task_history.go @@ -174,7 +174,7 @@ func (s *Store) CleanupTaskHistory(ctx context.Context, analysisCutoff time.Time batchSize = 1000 } var result TaskHistoryCleanupResult - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'; SET LOCAL statement_timeout = '5s'`); err != nil { return err } diff --git a/apps/api/internal/store/tasks_runtime.go b/apps/api/internal/store/tasks_runtime.go index 675a193..5894357 100644 --- a/apps/api/internal/store/tasks_runtime.go +++ b/apps/api/internal/store/tasks_runtime.go @@ -286,7 +286,7 @@ func (s *Store) ClaimTaskExecution(ctx context.Context, taskID string, execution } var task GatewayTask manualReview := false - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { var queuedReady bool var runningExpired bool var production bool @@ -573,6 +573,86 @@ WHERE id = $1::uuid RETURNING `+gatewayTaskColumns, taskID, executionToken, nextRunAt, strings.TrimSpace(queueKey))) } +func (s *Store) RequeueTaskBeforeUpstreamSubmission( + ctx context.Context, + taskID string, + executionToken string, + delay time.Duration, +) (GatewayTask, bool, error) { + if delay < time.Second { + delay = time.Second + } + if delay > 10*time.Minute { + delay = 10 * time.Minute + } + nextRunAt := time.Now().Add(delay) + var task GatewayTask + changed := false + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { + var err error + task, err = scanGatewayTask(tx.QueryRow(ctx, ` +UPDATE gateway_tasks task +SET status = 'queued', + locked_by = NULL, + locked_at = NULL, + heartbeat_at = NULL, + execution_token = NULL, + execution_lease_expires_at = NULL, + next_run_at = $3::timestamptz, + error = NULL, + error_code = NULL, + error_message = NULL, + updated_at = now() +WHERE task.id = $1::uuid + AND task.status IN ('queued', 'running') + AND task.execution_token = $2::uuid + AND COALESCE(task.remote_task_id, '') = '' + AND NOT EXISTS ( + SELECT 1 + FROM gateway_task_attempts attempt + WHERE attempt.task_id = task.id + AND COALESCE(attempt.upstream_submission_status, 'not_submitted') <> 'not_submitted' + ) +RETURNING `+gatewayTaskColumns, taskID, executionToken, nextRunAt)) + if errors.Is(err, pgx.ErrNoRows) { + return nil + } + if err != nil { + return err + } + changed = true + if _, err := tx.Exec(ctx, ` +DELETE FROM gateway_task_param_preprocessing_logs log +USING gateway_task_attempts attempt +WHERE log.attempt_id = attempt.id + AND attempt.task_id = $1::uuid + AND COALESCE(attempt.upstream_submission_status, 'not_submitted') = 'not_submitted'`, taskID); err != nil { + return err + } + if _, err := tx.Exec(ctx, ` +DELETE FROM gateway_task_attempts +WHERE task_id = $1::uuid + AND COALESCE(upstream_submission_status, 'not_submitted') = 'not_submitted'`, taskID); err != nil { + return err + } + if _, err := tx.Exec(ctx, ` +UPDATE gateway_concurrency_leases +SET released_at = now() +WHERE task_id = $1::uuid + AND released_at IS NULL`, taskID); err != nil { + return err + } + if _, err := tx.Exec(ctx, `DELETE FROM gateway_task_admissions WHERE task_id = $1::uuid`, taskID); err != nil { + return err + } + return notifyTaskAdmissionTx(ctx, tx, "*") + }) + if err != nil { + return GatewayTask{}, false, err + } + return task, changed, nil +} + func (s *Store) SetTaskRiverJobID(ctx context.Context, taskID string, riverJobID int64) error { if riverJobID <= 0 { return nil @@ -587,7 +667,7 @@ WHERE id = $1::uuid`, taskID, riverJobID) func (s *Store) SetTaskRemoteTask(ctx context.Context, taskID string, executionToken string, attemptID string, remoteTaskID string, payload map[string]any) error { payloadJSON, _ := json.Marshal(sanitizeJSONForStorage(emptyObjectIfNil(payload))) - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.beginTransaction(ctx, func(tx pgx.Tx) error { tag, err := tx.Exec(ctx, ` UPDATE gateway_tasks SET remote_task_id = NULLIF($3::text, ''), @@ -644,7 +724,7 @@ func (s *Store) CancelQueuedTask(ctx context.Context, taskID string, message str message = truncateUTF8Bytes(message, 2048) var task GatewayTask changed := false - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil { return err } @@ -705,7 +785,7 @@ func (s *Store) CancelTaskBeforeUpstreamSubmission( message = truncateUTF8Bytes(message, 2048) var task GatewayTask changed := false - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil { return err } @@ -789,7 +869,7 @@ func (s *Store) CancelSubmittedTask(ctx context.Context, taskID string, executio message = truncateUTF8Bytes(message, 2048) var task GatewayTask changed := false - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { var err error task, err = scanGatewayTask(tx.QueryRow(ctx, ` UPDATE gateway_tasks @@ -851,14 +931,23 @@ func (s *Store) ListRecoverableAsyncTasks(ctx context.Context, limit int) ([]Asy limit = 500 } rows, err := s.pool.Query(ctx, ` -SELECT id::text, priority, next_run_at -FROM gateway_tasks -WHERE async_mode = true +SELECT task.id::text, task.priority, task.next_run_at +FROM gateway_tasks task +LEFT JOIN river_job job ON job.id = task.river_job_id +WHERE task.async_mode = true AND ( - status = 'queued' - OR (status = 'running' AND (execution_lease_expires_at IS NULL OR execution_lease_expires_at <= now())) + task.status = 'queued' + OR ( + task.status = 'running' + AND (task.execution_lease_expires_at IS NULL OR task.execution_lease_expires_at <= now()) + ) ) -ORDER BY priority ASC, created_at ASC + AND ( + task.river_job_id IS NULL + OR job.id IS NULL + OR job.state NOT IN ('available', 'pending', 'retryable', 'running', 'scheduled') + ) +ORDER BY task.priority ASC, task.created_at ASC LIMIT $1`, limit) if err != nil { return nil, err @@ -883,7 +972,7 @@ func (s *Store) CreateTaskAttempt(ctx context.Context, input CreateTaskAttemptIn if err != nil { return "", err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) var attemptID string err = tx.QueryRow(ctx, ` @@ -1301,7 +1390,7 @@ func (s *Store) FinishTaskSuccess(ctx context.Context, input FinishTaskSuccessIn finalChargeAmount = strconv.FormatFloat(input.FinalChargeAmount, 'f', 9, 64) } currency := normalizeWalletCurrency(input.BillingCurrency) - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { if strings.TrimSpace(input.AttemptID) != "" { if _, err := tx.Exec(ctx, ` UPDATE gateway_task_attempts @@ -1430,7 +1519,7 @@ func (s *Store) FinishTaskManualReview(ctx context.Context, input FinishTaskManu } resultJSON, _ := json.Marshal(resultReport.Value) pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(emptyObjectIfNil(input.PricingSnapshot))) - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { if strings.TrimSpace(input.AttemptID) != "" { attemptStatus := "failed" if status == "succeeded" { @@ -1532,7 +1621,7 @@ func (s *Store) SettleTaskBilling(ctx context.Context, task GatewayTask) error { "billingSummary": task.BillingSummary, } metadata, _ := json.Marshal(sanitizeJSONForStorage(metadataMap)) - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.beginTransaction(ctx, func(tx pgx.Tx) error { if _, err := tx.Exec(ctx, ` INSERT INTO gateway_wallet_accounts ( gateway_tenant_id, gateway_user_id, tenant_id, tenant_key, user_id, currency @@ -1662,7 +1751,7 @@ func (s *Store) FinishTaskFailure(ctx context.Context, input FinishTaskFailureIn metricsJSON, _ := json.Marshal(sanitizeJSONForStorage(emptyObjectIfNil(input.Metrics))) resultJSON, _ := json.Marshal(minimalTaskResult(nil)) message := truncateUTF8Bytes(input.Message, 2048) - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { tag, err := tx.Exec(ctx, ` UPDATE gateway_tasks SET status = 'failed', @@ -1817,7 +1906,7 @@ func (s *Store) AddTaskEvent(ctx context.Context, taskID string, eventType strin if err != nil { return TaskEvent{}, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT 1 FROM gateway_tasks WHERE id = $1::uuid FOR UPDATE`, taskID); err != nil { return TaskEvent{}, err } diff --git a/apps/api/internal/store/tx.go b/apps/api/internal/store/tx.go index d4795c8..ceacdb5 100644 --- a/apps/api/internal/store/tx.go +++ b/apps/api/internal/store/tx.go @@ -2,18 +2,52 @@ package store import ( "context" + "errors" + "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgconn" ) +const transactionRollbackTimeout = 5 * time.Second + type Tx interface { Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) QueryRow(ctx context.Context, sql string, args ...any) pgx.Row } func (s *Store) InTx(ctx context.Context, fn func(Tx) error) error { - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.beginTransaction(ctx, func(tx pgx.Tx) error { return fn(tx) }) } + +func (s *Store) beginTransaction(ctx context.Context, fn func(pgx.Tx) error) (err error) { + tx, err := s.pool.Begin(ctx) + if err != nil { + return err + } + // A Worker shutdown cancels ctx before in-flight callbacks return. Use an + // independent bounded context so rollback still reaches PostgreSQL. + defer func() { + rollbackCtx, cancel := context.WithTimeout(context.Background(), transactionRollbackTimeout) + defer cancel() + rollbackErr := tx.Rollback(rollbackCtx) + if err == nil && rollbackErr != nil && !errors.Is(rollbackErr, pgx.ErrTxClosed) { + err = rollbackErr + } + }() + if err = fn(tx); err != nil { + return err + } + return tx.Commit(ctx) +} + +func rollbackTransaction(tx pgx.Tx) { + if tx == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), transactionRollbackTimeout) + defer cancel() + _ = tx.Rollback(ctx) +} diff --git a/apps/api/internal/store/wallet.go b/apps/api/internal/store/wallet.go index a3191ea..439ceff 100644 --- a/apps/api/internal/store/wallet.go +++ b/apps/api/internal/store/wallet.go @@ -163,7 +163,7 @@ func (s *Store) ReserveTaskBilling(ctx context.Context, task GatewayTask, user * reservations := make([]WalletBillingReservation, 0, len(amounts)) pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot)) requestFingerprint := walletString(pricingSnapshot["requestFingerprint"]) - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { for currency, rawAmount := range amounts { amount := roundMoney(rawAmount) if amount <= 0 { @@ -309,7 +309,7 @@ func (s *Store) reserveTaskBillingExact(ctx context.Context, task GatewayTask, g pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot)) requestFingerprint := walletString(pricingSnapshot["requestFingerprint"]) var reservations []WalletBillingReservation - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.beginTransaction(ctx, func(tx pgx.Tx) error { var positive bool if err := tx.QueryRow(ctx, `SELECT $1::numeric(38, 9) > 0`, amount).Scan(&positive); err != nil { return fmt.Errorf("invalid exact reservation amount: %w", err) @@ -408,7 +408,7 @@ func (s *Store) ReleaseTaskBillingReservations(ctx context.Context, reservations if reason == "" { reason = "task_not_settled" } - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.beginTransaction(ctx, func(tx pgx.Tx) error { taskIDs := map[string]struct{}{} for _, reservation := range reservations { if strings.TrimSpace(reservation.AccountID) == "" { diff --git a/apps/api/internal/store/worker_registry.go b/apps/api/internal/store/worker_registry.go index 5149094..96f067c 100644 --- a/apps/api/internal/store/worker_registry.go +++ b/apps/api/internal/store/worker_registry.go @@ -70,7 +70,7 @@ func (s *Store) RegisterWorkerInstance(ctx context.Context, input WorkerRegistra if err != nil { return WorkerAllocation{}, err } - defer tx.Rollback(ctx) + 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. @@ -245,7 +245,7 @@ func (s *Store) RecoverOrphanedAsyncRiverJobs( if err != nil { return 0, err } - defer tx.Rollback(ctx) + defer rollbackTransaction(tx) if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-river-orphan-recovery', 0))`); err != nil { return 0, err } diff --git a/deploy/kubernetes/production/application-config.yaml b/deploy/kubernetes/production/application-config.yaml index 5d6d490..750aca5 100644 --- a/deploy/kubernetes/production/application-config.yaml +++ b/deploy/kubernetes/production/application-config.yaml @@ -24,6 +24,8 @@ data: AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS: "7" AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS: "300" AI_GATEWAY_TASK_CLEANUP_BATCH_SIZE: "1000" + AI_GATEWAY_DATABASE_IDLE_IN_TRANSACTION_TIMEOUT_SECONDS: "60" + AI_GATEWAY_DATABASE_LOCK_TIMEOUT_SECONDS: "30" CORS_ALLOWED_ORIGIN: https://ai.51easyai.com AI_GATEWAY_PUBLIC_BASE_URL: https://ai.51easyai.com AI_GATEWAY_WEB_BASE_URL: https://ai.51easyai.com