fix(worker): 防止事务泄漏并恢复滞留队列
为 PostgreSQL 连接增加可配置的事务空闲与锁等待超时,并在请求取消或 Worker 退出后使用独立有界上下文回滚事务。\n\n恢复过期任务时按批次使用 SKIP LOCKED,周期重建缺失或终态 River Job;锁等待超时只在确认尚未提交上游时安全重排,避免重复执行和重复结算。\n\n验证:完整 Go 测试、go vet、生产 Kustomize 渲染、gofmt 与 git diff --check 均通过。
This commit is contained in:
@@ -37,7 +37,11 @@ func main() {
|
|||||||
|
|
||||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
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 {
|
if err != nil {
|
||||||
stop()
|
stop()
|
||||||
logger.Error("connect postgres failed", "error", err)
|
logger.Error("connect postgres failed", "error", err)
|
||||||
|
|||||||
@@ -59,6 +59,8 @@ type Config struct {
|
|||||||
BillingEngineMode string
|
BillingEngineMode string
|
||||||
ProcessRole string
|
ProcessRole string
|
||||||
DatabaseMaxConns int
|
DatabaseMaxConns int
|
||||||
|
DatabaseIdleInTransactionTimeoutSeconds int
|
||||||
|
DatabaseLockTimeoutSeconds int
|
||||||
MediaRequestConcurrency int
|
MediaRequestConcurrency int
|
||||||
MediaMaterializationConcurrency int
|
MediaMaterializationConcurrency int
|
||||||
AsyncQueueWorkerEnabled bool
|
AsyncQueueWorkerEnabled bool
|
||||||
@@ -106,20 +108,27 @@ func Load() Config {
|
|||||||
TaskProgressCallbackURL: env("TASK_PROGRESS_CALLBACK_URL",
|
TaskProgressCallbackURL: env("TASK_PROGRESS_CALLBACK_URL",
|
||||||
strings.TrimRight(env("SERVER_MAIN_BASE_URL", "http://localhost:3000"), "/")+"/internal/platform/task-progress-callbacks",
|
strings.TrimRight(env("SERVER_MAIN_BASE_URL", "http://localhost:3000"), "/")+"/internal/platform/task-progress-callbacks",
|
||||||
),
|
),
|
||||||
TaskProgressCallbackTimeoutMS: envIntValidated("TASK_PROGRESS_CALLBACK_TIMEOUT_MS", 5000),
|
TaskProgressCallbackTimeoutMS: envIntValidated("TASK_PROGRESS_CALLBACK_TIMEOUT_MS", 5000),
|
||||||
TaskProgressCallbackMaxAttempts: envIntValidated("TASK_PROGRESS_CALLBACK_MAX_ATTEMPTS", 10),
|
TaskProgressCallbackMaxAttempts: envIntValidated("TASK_PROGRESS_CALLBACK_MAX_ATTEMPTS", 10),
|
||||||
TaskCleanupEnabled: env("AI_GATEWAY_TASK_CLEANUP_ENABLED", "false") == "true",
|
TaskCleanupEnabled: env("AI_GATEWAY_TASK_CLEANUP_ENABLED", "false") == "true",
|
||||||
TaskRetentionDays: envIntValidated("AI_GATEWAY_TASK_RETENTION_DAYS", 30),
|
TaskRetentionDays: envIntValidated("AI_GATEWAY_TASK_RETENTION_DAYS", 30),
|
||||||
TaskAnalysisRetentionDays: envIntValidated("AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS", 7),
|
TaskAnalysisRetentionDays: envIntValidated("AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS", 7),
|
||||||
TaskCleanupIntervalSeconds: envIntValidated("AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS", 300),
|
TaskCleanupIntervalSeconds: envIntValidated("AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS", 300),
|
||||||
TaskCleanupBatchSize: envIntValidated("AI_GATEWAY_TASK_CLEANUP_BATCH_SIZE", 1000),
|
TaskCleanupBatchSize: envIntValidated("AI_GATEWAY_TASK_CLEANUP_BATCH_SIZE", 1000),
|
||||||
CORSAllowedOrigin: env("CORS_ALLOWED_ORIGIN", "http://localhost:5178,http://127.0.0.1:5178"),
|
CORSAllowedOrigin: env("CORS_ALLOWED_ORIGIN", "http://localhost:5178,http://127.0.0.1:5178"),
|
||||||
GlobalHTTPProxy: globalProxy.HTTPProxy,
|
GlobalHTTPProxy: globalProxy.HTTPProxy,
|
||||||
GlobalHTTPProxySource: globalProxy.Source,
|
GlobalHTTPProxySource: globalProxy.Source,
|
||||||
LogLevel: logLevel(env("LOG_LEVEL", "info")),
|
LogLevel: logLevel(env("LOG_LEVEL", "info")),
|
||||||
BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")),
|
BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")),
|
||||||
ProcessRole: strings.ToLower(strings.TrimSpace(env("AI_GATEWAY_PROCESS_ROLE", "all"))),
|
ProcessRole: strings.ToLower(strings.TrimSpace(env("AI_GATEWAY_PROCESS_ROLE", "all"))),
|
||||||
DatabaseMaxConns: envInt("AI_GATEWAY_DATABASE_MAX_CONNS", 0),
|
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),
|
MediaRequestConcurrency: envIntValidated("AI_GATEWAY_MEDIA_REQUEST_CONCURRENCY", 16),
|
||||||
MediaMaterializationConcurrency: envIntValidated("AI_GATEWAY_MEDIA_MATERIALIZATION_CONCURRENCY", 8),
|
MediaMaterializationConcurrency: envIntValidated("AI_GATEWAY_MEDIA_MATERIALIZATION_CONCURRENCY", 8),
|
||||||
AsyncQueueWorkerEnabled: env("AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED", "true") == "true",
|
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 {
|
if c.DatabaseMaxConns < 0 || c.DatabaseMaxConns > 1000 {
|
||||||
return errors.New("AI_GATEWAY_DATABASE_MAX_CONNS must be between 1 and 1000 when configured")
|
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) {
|
if c.MediaMaterializationConcurrency != 0 && (c.MediaMaterializationConcurrency < 1 || c.MediaMaterializationConcurrency > 256) {
|
||||||
return errors.New("AI_GATEWAY_MEDIA_MATERIALIZATION_CONCURRENCY must be between 1 and 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
|
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 {
|
func envInt64Validated(key string, fallback int64) int64 {
|
||||||
value := envValue(key)
|
value := envValue(key)
|
||||||
if value == "" {
|
if value == "" {
|
||||||
|
|||||||
@@ -144,6 +144,13 @@ func TestProcessRolePrecedenceAndCompatibility(t *testing.T) {
|
|||||||
|
|
||||||
func TestValidateProcessRoleAndDatabasePool(t *testing.T) {
|
func TestValidateProcessRoleAndDatabasePool(t *testing.T) {
|
||||||
cfg := Load()
|
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"
|
cfg.ProcessRole = "invalid"
|
||||||
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "PROCESS_ROLE") {
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "PROCESS_ROLE") {
|
||||||
t.Fatalf("Validate() error = %v, want invalid process role", err)
|
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)
|
t.Fatalf("Validate() error = %v, want invalid database max conns", err)
|
||||||
}
|
}
|
||||||
cfg.DatabaseMaxConns = 16
|
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
|
cfg.MediaMaterializationConcurrency = 257
|
||||||
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "MEDIA_MATERIALIZATION_CONCURRENCY") {
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "MEDIA_MATERIALIZATION_CONCURRENCY") {
|
||||||
t.Fatalf("Validate() error = %v, want invalid media materialization concurrency", err)
|
t.Fatalf("Validate() error = %v, want invalid media materialization concurrency", err)
|
||||||
|
|||||||
@@ -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)
|
w.service.logger.Debug("river async task interrupted and requeued", "taskID", task.ID, "status", queued.Status, "riverJobID", job.ID)
|
||||||
return river.JobSnooze(0)
|
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)
|
w.service.logger.Warn("river async task completed with failure", "taskID", task.ID, "error", runErr, "riverJobID", job.ID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -570,6 +602,12 @@ func (s *Service) recoverOrphanedAsyncRiverJobs(ctx context.Context) {
|
|||||||
if recovered > 0 {
|
if recovered > 0 {
|
||||||
s.logger.Warn("orphaned river jobs recovered", "count", recovered)
|
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()
|
recoverJobs()
|
||||||
ticker := time.NewTicker(orphanedRiverJobScanInterval)
|
ticker := time.NewTicker(orphanedRiverJobScanInterval)
|
||||||
|
|||||||
@@ -153,7 +153,7 @@ func (s *Store) BatchAccessRules(ctx context.Context, input AccessRuleBatchInput
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
for _, resource := range dedupeAccessRuleResources(input.DeleteResources) {
|
for _, resource := range dedupeAccessRuleResources(input.DeleteResources) {
|
||||||
resource = normalizeAccessRuleResource(resource, input.Effect)
|
resource = normalizeAccessRuleResource(resource, input.Effect)
|
||||||
|
|||||||
@@ -134,7 +134,7 @@ func (s *Store) QueueTaskAdmissionWithHook(
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return TaskAdmission{}, err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil {
|
||||||
return TaskAdmission{}, err
|
return TaskAdmission{}, err
|
||||||
@@ -232,7 +232,7 @@ func (s *Store) tryTaskAdmission(
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return TaskAdmissionResult{}, err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil {
|
||||||
return TaskAdmissionResult{}, err
|
return TaskAdmissionResult{}, err
|
||||||
@@ -456,7 +456,7 @@ func (s *Store) RebindWaitingTaskAdmission(ctx context.Context, input TaskAdmiss
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return TaskAdmission{}, err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+input.TaskID); err != nil {
|
||||||
return TaskAdmission{}, err
|
return TaskAdmission{}, err
|
||||||
}
|
}
|
||||||
@@ -633,7 +633,7 @@ func (s *Store) DeleteTaskAdmission(ctx context.Context, taskID string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -794,7 +794,7 @@ func (s *Store) ReapExpiredTaskAdmissions(ctx context.Context, limit int) (TaskA
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return TaskAdmissionReapResult{}, err
|
return TaskAdmissionReapResult{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
rows, err := tx.Query(ctx, `
|
rows, err := tx.Query(ctx, `
|
||||||
SELECT task_id::text
|
SELECT task_id::text
|
||||||
FROM gateway_task_admissions
|
FROM gateway_task_admissions
|
||||||
|
|||||||
@@ -543,6 +543,160 @@ SELECT
|
|||||||
t.Fatalf("recovered task left admissions=%d leases=%d", recoveredAdmissions, recoveredLeases)
|
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)
|
terminalResidue := createTask(true)
|
||||||
result, err = first.TryTaskAdmission(ctx, inputFor(terminalResidue, 100, ""))
|
result, err = first.TryTaskAdmission(ctx, inputFor(terminalResidue, 100, ""))
|
||||||
if err != nil || !result.Admitted {
|
if err != nil || !result.Admitted {
|
||||||
|
|||||||
@@ -101,7 +101,7 @@ func (s *Store) CreateBaseModel(ctx context.Context, input BaseModelInput) (Base
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return BaseModel{}, err
|
return BaseModel{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
item, err := scanBaseModel(tx.QueryRow(ctx, `
|
item, err := scanBaseModel(tx.QueryRow(ctx, `
|
||||||
INSERT INTO base_model_catalog (
|
INSERT INTO base_model_catalog (
|
||||||
@@ -162,7 +162,7 @@ func (s *Store) UpdateBaseModel(ctx context.Context, id string, input BaseModelI
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return BaseModel{}, err
|
return BaseModel{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
item, err := scanBaseModel(tx.QueryRow(ctx, `
|
item, err := scanBaseModel(tx.QueryRow(ctx, `
|
||||||
UPDATE base_model_catalog
|
UPDATE base_model_catalog
|
||||||
|
|||||||
@@ -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 {
|
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 action string
|
||||||
var amount string
|
var amount string
|
||||||
var currency string
|
var currency string
|
||||||
@@ -479,7 +479,7 @@ func (s *Store) MarkBillingSettlementFailed(ctx context.Context, settlement Bill
|
|||||||
taskStatus = "manual_review"
|
taskStatus = "manual_review"
|
||||||
manualReason = "maximum settlement attempts exceeded"
|
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, `
|
tag, err := tx.Exec(ctx, `
|
||||||
UPDATE settlement_outbox
|
UPDATE settlement_outbox
|
||||||
SET status = $3,
|
SET status = $3,
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ func (s *Store) ListTaskBinaryResultBackfillBatch(ctx context.Context, afterID s
|
|||||||
batchSize = 100
|
batchSize = 100
|
||||||
}
|
}
|
||||||
var items []TaskBinaryResultBackfillItem
|
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 {
|
if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'; SET LOCAL statement_timeout = '5s'`); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -64,7 +64,7 @@ func (s *Store) UpdateTaskBinaryResultBackfill(ctx context.Context, taskID strin
|
|||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
updated := false
|
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 {
|
if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'; SET LOCAL statement_timeout = '5s'`); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ func (s *Store) UpsertConversationMessages(ctx context.Context, conversationID s
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, err
|
return nil, 0, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
refs := make([]TaskMessageRefInput, 0, len(messages))
|
refs := make([]TaskMessageRefInput, 0, len(messages))
|
||||||
newCount := 0
|
newCount := 0
|
||||||
|
|||||||
@@ -272,7 +272,7 @@ func (s *Store) DeleteUserGroup(ctx context.Context, id string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
if _, err := tx.Exec(ctx, `
|
if _, err := tx.Exec(ctx, `
|
||||||
DELETE FROM gateway_access_rules
|
DELETE FROM gateway_access_rules
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ func (s *Store) beginIdentityConfigurationLifecycleTx(ctx context.Context) (pgx.
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := lockIdentityConfigurationLifecycle(ctx, tx); err != nil {
|
if err := lockIdentityConfigurationLifecycle(ctx, tx); err != nil {
|
||||||
_ = tx.Rollback(ctx)
|
rollbackTransaction(tx)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return tx, nil
|
return tx, nil
|
||||||
@@ -171,7 +171,7 @@ func (s *Store) ApplyIdentityManifest(ctx context.Context, id string, expectedVe
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return identity.Revision{}, err
|
return identity.Revision{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
current, err := scanIdentityRevision(tx.QueryRow(ctx, `SELECT `+identityRevisionColumns+`
|
current, err := scanIdentityRevision(tx.QueryRow(ctx, `SELECT `+identityRevisionColumns+`
|
||||||
FROM gateway_identity_configuration_revisions WHERE id=$1::uuid FOR UPDATE`, id))
|
FROM gateway_identity_configuration_revisions WHERE id=$1::uuid FOR UPDATE`, id))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -282,7 +282,7 @@ func (s *Store) ActivateIdentityRevision(ctx context.Context, id string, expecte
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return identity.Revision{}, false, err
|
return identity.Revision{}, false, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
var foreignPairingReservation bool
|
var foreignPairingReservation bool
|
||||||
if err := tx.QueryRow(ctx, `SELECT EXISTS(
|
if err := tx.QueryRow(ctx, `SELECT EXISTS(
|
||||||
SELECT 1 FROM gateway_identity_pairing_start_reservation
|
SELECT 1 FROM gateway_identity_pairing_start_reservation
|
||||||
@@ -364,7 +364,7 @@ func (s *Store) DisableActiveIdentityRevision(ctx context.Context, expectedVersi
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return identity.Revision{}, err
|
return identity.Revision{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
if ok, err := hasBreakGlassManager(ctx, tx); err != nil {
|
if ok, err := hasBreakGlassManager(ctx, tx); err != nil {
|
||||||
return identity.Revision{}, err
|
return identity.Revision{}, err
|
||||||
} else if !ok {
|
} else if !ok {
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ func (s *Store) ReserveIdentityPairingStart(ctx context.Context, attemptID strin
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
if _, err := tx.Exec(ctx, `DELETE FROM gateway_identity_pairing_start_reservation
|
if _, err := tx.Exec(ctx, `DELETE FROM gateway_identity_pairing_start_reservation
|
||||||
WHERE state='starting' AND expires_at <= now()`); err != nil {
|
WHERE state='starting' AND expires_at <= now()`); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -68,7 +68,7 @@ func (s *Store) CommitIdentityPairingStart(ctx context.Context, revision identit
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return identity.PairingExchange{}, err
|
return identity.PairingExchange{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
var reserved bool
|
var reserved bool
|
||||||
if err := tx.QueryRow(ctx, `SELECT true FROM gateway_identity_pairing_start_reservation
|
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) {
|
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 {
|
if err != nil {
|
||||||
return identity.PairingExchange{}, err
|
return identity.PairingExchange{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
exchange, err := scanIdentityPairing(tx.QueryRow(ctx, `
|
exchange, err := scanIdentityPairing(tx.QueryRow(ctx, `
|
||||||
UPDATE gateway_identity_onboarding_exchanges SET status=$3,remote_version=$4,last_error_category=NULLIF($5,''),
|
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()
|
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 {
|
if err != nil {
|
||||||
return identity.PairingExchange{}, err
|
return identity.PairingExchange{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
exchange, err := scanIdentityPairing(tx.QueryRow(ctx, `SELECT `+identityPairingColumns+`
|
exchange, err := scanIdentityPairing(tx.QueryRow(ctx, `SELECT `+identityPairingColumns+`
|
||||||
FROM gateway_identity_onboarding_exchanges WHERE id=$1::uuid FOR UPDATE`, id))
|
FROM gateway_identity_onboarding_exchanges WHERE id=$1::uuid FOR UPDATE`, id))
|
||||||
if errors.Is(err, pgx.ErrNoRows) {
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
@@ -352,7 +352,7 @@ func (s *Store) CompleteIdentityPairingCleanup(ctx context.Context, id string, e
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return identity.PairingExchange{}, err
|
return identity.PairingExchange{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
var revisionID string
|
var revisionID string
|
||||||
if err := tx.QueryRow(ctx, `SELECT revision_id::text FROM gateway_identity_onboarding_exchanges
|
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) {
|
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) {
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ func (s *Store) resolveOrProvisionOIDCMultiTenantUser(ctx context.Context, input
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return ResolveOrProvisionOIDCUserResult{}, err
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
}
|
}
|
||||||
defer func() { _ = tx.Rollback(ctx) }()
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
bindingID, gatewayTenantID, tenantKey, groupID, groupKey, accessStatus, err :=
|
bindingID, gatewayTenantID, tenantKey, groupID, groupKey, accessStatus, err :=
|
||||||
loadOIDCTenantBinding(ctx, tx, input)
|
loadOIDCTenantBinding(ctx, tx, input)
|
||||||
|
|||||||
@@ -78,7 +78,7 @@ func (s *Store) ApplyOIDCTenantBindingSync(ctx context.Context, bindingID string
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer func() { _ = tx.Rollback(ctx) }()
|
defer rollbackTransaction(tx)
|
||||||
if unchanged {
|
if unchanged {
|
||||||
tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
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,
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer func() { _ = tx.Rollback(ctx) }()
|
defer rollbackTransaction(tx)
|
||||||
tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
SET access_status='disabled',metadata_status='rejected',last_sync_at=$2::timestamptz,
|
SET access_status='disabled',metadata_status='rejected',last_sync_at=$2::timestamptz,
|
||||||
next_sync_at=$2::timestamptz+interval '15 minutes',
|
next_sync_at=$2::timestamptz+interval '15 minutes',
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ func (s *Store) ResolveOrProvisionOIDCUser(ctx context.Context, input ResolveOrP
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return ResolveOrProvisionOIDCUserResult{}, err
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
}
|
}
|
||||||
defer func() { _ = tx.Rollback(ctx) }()
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
projection, err := loadOIDCUserProjection(ctx, tx, input.Subject)
|
projection, err := loadOIDCUserProjection(ctx, tx, input.Subject)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ func (s *Store) ReplacePlatformModels(ctx context.Context, platformID string, in
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
keptIDs := make([]string, 0, len(inputs))
|
keptIDs := make([]string, 0, len(inputs))
|
||||||
for _, input := range inputs {
|
for _, input := range inputs {
|
||||||
@@ -372,7 +372,7 @@ func (s *Store) DeletePlatformModel(ctx context.Context, id string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
result, err := tx.Exec(ctx, `DELETE FROM platform_models WHERE id = $1::uuid`, id)
|
result, err := tx.Exec(ctx, `DELETE FROM platform_models WHERE id = $1::uuid`, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"net"
|
"net"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
"unicode"
|
"unicode"
|
||||||
@@ -29,6 +30,12 @@ const (
|
|||||||
postgresPoolWarmLimit = 64
|
postgresPoolWarmLimit = 64
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type PostgresPoolOptions struct {
|
||||||
|
MaxConns int
|
||||||
|
IdleInTransactionTimeout time.Duration
|
||||||
|
LockTimeout time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
func defaultAPIKeyScopes() []string {
|
func defaultAPIKeyScopes() []string {
|
||||||
return []string{"chat", "embedding", "rerank", "image", "image_vectorize", "video", "video_enhance", "music", "audio", "voice_clone"}
|
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) {
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
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) {
|
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)
|
config, err := pgxpool.ParseConfig(databaseURL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
config.ConnConfig.ConnectTimeout = postgresConnectTimeout
|
config.ConnConfig.ConnectTimeout = postgresConnectTimeout
|
||||||
config.ConnConfig.RuntimeParams["application_name"] = postgresApplicationName
|
config.ConnConfig.RuntimeParams["application_name"] = postgresApplicationName
|
||||||
if len(maxConns) > 0 && maxConns[0] > 0 {
|
if options.IdleInTransactionTimeout > 0 {
|
||||||
config.MaxConns = int32(maxConns[0])
|
config.ConnConfig.RuntimeParams["idle_in_transaction_session_timeout"] = postgresDuration(options.IdleInTransactionTimeout)
|
||||||
config.MinIdleConns = int32(min(maxConns[0], postgresPoolWarmLimit))
|
}
|
||||||
|
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
|
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 {
|
func warmPostgresPool(ctx context.Context, pool *pgxpool.Pool, count int) error {
|
||||||
if count <= 0 {
|
if count <= 0 {
|
||||||
return pool.Ping(ctx)
|
return pool.Ping(ctx)
|
||||||
@@ -152,6 +181,11 @@ func IsPostgresUnavailable(err error) bool {
|
|||||||
return pgconn.SafeToRetry(err)
|
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() {
|
func (s *Store) Close() {
|
||||||
s.pool.Close()
|
s.pool.Close()
|
||||||
}
|
}
|
||||||
@@ -938,7 +972,7 @@ func (s *Store) DeletePlatform(ctx context.Context, id string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
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)
|
rows, err := tx.Query(ctx, `SELECT id::text FROM platform_models WHERE platform_id = $1::uuid`, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -1638,7 +1672,7 @@ func (s *Store) DeleteAPIKey(ctx context.Context, apiKeyID string, user *auth.Us
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
result, err := tx.Exec(ctx, `
|
result, err := tx.Exec(ctx, `
|
||||||
UPDATE gateway_api_keys
|
UPDATE gateway_api_keys
|
||||||
@@ -1792,7 +1826,7 @@ func (s *Store) RegisterLocalUser(ctx context.Context, input LocalRegisterInput)
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return GatewayUser{}, err
|
return GatewayUser{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
var tenantID string
|
var tenantID string
|
||||||
userGroupID := ""
|
userGroupID := ""
|
||||||
@@ -2065,7 +2099,7 @@ func (s *Store) CreateTaskIdempotent(ctx context.Context, input CreateTaskInput,
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return CreateTaskResult{}, err
|
return CreateTaskResult{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
task, err := scanGatewayTask(tx.QueryRow(ctx, `
|
task, err := scanGatewayTask(tx.QueryRow(ctx, `
|
||||||
INSERT INTO gateway_tasks (
|
INSERT INTO gateway_tasks (
|
||||||
|
|||||||
@@ -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) {
|
func TestIsPostgresUnavailableClassifiesConnectivityFailures(t *testing.T) {
|
||||||
for _, testCase := range []struct {
|
for _, testCase := range []struct {
|
||||||
name string
|
name string
|
||||||
@@ -75,3 +98,12 @@ func TestIsPostgresUnavailableClassifiesConnectivityFailures(t *testing.T) {
|
|||||||
t.Fatal("SQL syntax error was incorrectly classified as PostgreSQL unavailable")
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -113,7 +113,7 @@ func (s *Store) CreatePricingRuleSet(ctx context.Context, input PricingRuleSetIn
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return PricingRuleSet{}, err
|
return PricingRuleSet{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata))
|
metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata))
|
||||||
item, err := scanPricingRuleSet(tx.QueryRow(ctx, `
|
item, err := scanPricingRuleSet(tx.QueryRow(ctx, `
|
||||||
@@ -141,7 +141,7 @@ func (s *Store) UpdatePricingRuleSet(ctx context.Context, id string, input Prici
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return PricingRuleSet{}, err
|
return PricingRuleSet{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata))
|
metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata))
|
||||||
item, err := scanPricingRuleSet(tx.QueryRow(ctx, `
|
item, err := scanPricingRuleSet(tx.QueryRow(ctx, `
|
||||||
|
|||||||
@@ -98,7 +98,7 @@ func (s *Store) RestorePlatformModelRuntimeStatus(ctx context.Context, platformM
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return ModelRateLimitStatus{}, err
|
return ModelRateLimitStatus{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
var restoredModelID string
|
var restoredModelID string
|
||||||
if err := tx.QueryRow(ctx, `
|
if err := tx.QueryRow(ctx, `
|
||||||
|
|||||||
@@ -19,6 +19,8 @@ type RuntimeRecoveryResult struct {
|
|||||||
CleanedTaskAdmissions int64 `json:"cleanedTaskAdmissions"`
|
CleanedTaskAdmissions int64 `json:"cleanedTaskAdmissions"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const runtimeRecoveryBatchSize = 100
|
||||||
|
|
||||||
var ErrConcurrencyLeaseLost = errors.New("concurrency lease lost")
|
var ErrConcurrencyLeaseLost = errors.New("concurrency lease lost")
|
||||||
|
|
||||||
// CheckRateLimits performs a non-consuming admission preflight for fixed-window
|
// 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 {
|
if err != nil {
|
||||||
return RateLimitResult{}, err
|
return RateLimitResult{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
lockKeys := make([]string, 0)
|
lockKeys := make([]string, 0)
|
||||||
lockKeySet := make(map[string]struct{})
|
lockKeySet := make(map[string]struct{})
|
||||||
@@ -414,7 +416,7 @@ func (s *Store) AttachRateLimitResultToAttempt(ctx context.Context, attemptID st
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
for _, reservation := range result.Reservations {
|
for _, reservation := range result.Reservations {
|
||||||
if reservation.ReservationID == "" {
|
if reservation.ReservationID == "" {
|
||||||
@@ -458,7 +460,7 @@ func (s *Store) RecoverInterruptedRuntimeState(ctx context.Context) (RuntimeReco
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return RuntimeRecoveryResult{}, err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-runtime-recovery', 0))`); err != nil {
|
||||||
return RuntimeRecoveryResult{}, err
|
return RuntimeRecoveryResult{}, err
|
||||||
}
|
}
|
||||||
@@ -542,7 +544,18 @@ WHERE status = 'running'
|
|||||||
result.FailedAttempts = tag.RowsAffected()
|
result.FailedAttempts = tag.RowsAffected()
|
||||||
|
|
||||||
asyncTaskRows, err := tx.Query(ctx, `
|
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',
|
SET status = 'queued',
|
||||||
error = NULL,
|
error = NULL,
|
||||||
error_code = NULL,
|
error_code = NULL,
|
||||||
@@ -555,11 +568,9 @@ SET status = 'queued',
|
|||||||
next_run_at = now(),
|
next_run_at = now(),
|
||||||
finished_at = NULL,
|
finished_at = NULL,
|
||||||
updated_at = now()
|
updated_at = now()
|
||||||
WHERE async_mode = true
|
FROM recoverable_async_tasks recoverable
|
||||||
AND status = 'running'
|
WHERE task.id = recoverable.id
|
||||||
AND execution_lease_expires_at IS NOT NULL
|
RETURNING task.id::text`, runtimeRecoveryBatchSize)
|
||||||
AND execution_lease_expires_at <= now()
|
|
||||||
RETURNING id::text`)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return RuntimeRecoveryResult{}, err
|
return RuntimeRecoveryResult{}, err
|
||||||
}
|
}
|
||||||
@@ -718,7 +729,7 @@ func (s *Store) finishRateLimitReservations(ctx context.Context, reservations []
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
for _, reservation := range reservations {
|
for _, reservation := range reservations {
|
||||||
if reservation.ReservationID == "" {
|
if reservation.ReservationID == "" {
|
||||||
|
|||||||
@@ -248,9 +248,7 @@ func (s *Store) ApplyCandidateFailureEffect(ctx context.Context, input Candidate
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return CandidateFailureEffectResult{}, err
|
return CandidateFailureEffectResult{}, err
|
||||||
}
|
}
|
||||||
defer func() {
|
defer rollbackTransaction(tx)
|
||||||
_ = tx.Rollback(ctx)
|
|
||||||
}()
|
|
||||||
|
|
||||||
lockKey := strings.TrimSpace(input.RequestedModel) + "\x1f" + strings.TrimSpace(input.ModelType)
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1::text, 0))`, lockKey); err != nil {
|
||||||
|
|||||||
@@ -185,7 +185,7 @@ func (s *Store) CreateSecurityEventConnection(ctx context.Context, input CreateS
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return SecurityEventConnection{}, err
|
return SecurityEventConnection{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
_, err = tx.Exec(ctx, `
|
_, err = tx.Exec(ctx, `
|
||||||
INSERT INTO gateway_security_event_connections(
|
INSERT INTO gateway_security_event_connections(
|
||||||
connection_id,transmitter_issuer,endpoint_url,credential_ref,management_client_id,management_credential_ref,lifecycle_status,idempotency_key
|
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 err != nil {
|
||||||
if isUniqueViolation(err) {
|
if isUniqueViolation(err) {
|
||||||
_ = tx.Rollback(ctx)
|
rollbackTransaction(tx)
|
||||||
existing, getErr := s.SecurityEventConnection(ctx)
|
existing, getErr := s.SecurityEventConnection(ctx)
|
||||||
if getErr == nil && existing.IdempotencyKey == input.IdempotencyKey && existing.TransmitterIssuer == input.TransmitterIssuer {
|
if getErr == nil && existing.IdempotencyKey == input.IdempotencyKey && existing.TransmitterIssuer == input.TransmitterIssuer {
|
||||||
return existing, nil
|
return existing, nil
|
||||||
@@ -256,7 +256,7 @@ func (s *Store) SetSecurityEventManagementCredential(ctx context.Context, connec
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return SecurityEventConnection{}, err
|
return SecurityEventConnection{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
var oldReference sql.NullString
|
var oldReference sql.NullString
|
||||||
var lifecycle string
|
var lifecycle string
|
||||||
var version int64
|
var version int64
|
||||||
@@ -315,7 +315,7 @@ func (s *Store) SetSecurityEventNextCredential(ctx context.Context, connectionID
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return SecurityEventConnection{}, err
|
return SecurityEventConnection{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
tag, err := tx.Exec(ctx, `UPDATE gateway_security_event_connections
|
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()
|
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
|
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 {
|
if err != nil {
|
||||||
return SecurityEventConnection{}, err
|
return SecurityEventConnection{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
var oldReference string
|
var oldReference string
|
||||||
if err := tx.QueryRow(ctx, `SELECT credential_ref FROM gateway_security_event_connections
|
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) {
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
var credentialReference string
|
var credentialReference string
|
||||||
var nextReference, managementReference sql.NullString
|
var nextReference, managementReference sql.NullString
|
||||||
@@ -435,7 +435,7 @@ func (s *Store) FinalizeRetiringSecurityEventConnection(ctx context.Context, con
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
var credentialReference string
|
var credentialReference string
|
||||||
var nextReference, managementReference sql.NullString
|
var nextReference, managementReference sql.NullString
|
||||||
|
|||||||
@@ -13,7 +13,6 @@ import (
|
|||||||
const (
|
const (
|
||||||
securityEventTransactionLockTimeout = 5 * time.Second
|
securityEventTransactionLockTimeout = 5 * time.Second
|
||||||
securityEventTransactionIdleTimeout = 15 * time.Second
|
securityEventTransactionIdleTimeout = 15 * time.Second
|
||||||
transactionRollbackTimeout = 5 * time.Second
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func (s *Store) beginSecurityEventTransaction(ctx context.Context) (pgx.Tx, error) {
|
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) {
|
func rollbackSecurityEventTransaction(tx pgx.Tx) {
|
||||||
if tx == nil {
|
rollbackTransaction(tx)
|
||||||
return
|
|
||||||
}
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), transactionRollbackTimeout)
|
|
||||||
defer cancel()
|
|
||||||
_ = tx.Rollback(ctx)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ApplySessionRevokedInput struct {
|
type ApplySessionRevokedInput struct {
|
||||||
|
|||||||
@@ -174,7 +174,7 @@ func (s *Store) CleanupTaskHistory(ctx context.Context, analysisCutoff time.Time
|
|||||||
batchSize = 1000
|
batchSize = 1000
|
||||||
}
|
}
|
||||||
var result TaskHistoryCleanupResult
|
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 {
|
if _, err := tx.Exec(ctx, `SET LOCAL lock_timeout = '5s'; SET LOCAL statement_timeout = '5s'`); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -286,7 +286,7 @@ func (s *Store) ClaimTaskExecution(ctx context.Context, taskID string, execution
|
|||||||
}
|
}
|
||||||
var task GatewayTask
|
var task GatewayTask
|
||||||
manualReview := false
|
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 queuedReady bool
|
||||||
var runningExpired bool
|
var runningExpired bool
|
||||||
var production bool
|
var production bool
|
||||||
@@ -573,6 +573,86 @@ WHERE id = $1::uuid
|
|||||||
RETURNING `+gatewayTaskColumns, taskID, executionToken, nextRunAt, strings.TrimSpace(queueKey)))
|
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 {
|
func (s *Store) SetTaskRiverJobID(ctx context.Context, taskID string, riverJobID int64) error {
|
||||||
if riverJobID <= 0 {
|
if riverJobID <= 0 {
|
||||||
return nil
|
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 {
|
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)))
|
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, `
|
tag, err := tx.Exec(ctx, `
|
||||||
UPDATE gateway_tasks
|
UPDATE gateway_tasks
|
||||||
SET remote_task_id = NULLIF($3::text, ''),
|
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)
|
message = truncateUTF8Bytes(message, 2048)
|
||||||
var task GatewayTask
|
var task GatewayTask
|
||||||
changed := false
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -705,7 +785,7 @@ func (s *Store) CancelTaskBeforeUpstreamSubmission(
|
|||||||
message = truncateUTF8Bytes(message, 2048)
|
message = truncateUTF8Bytes(message, 2048)
|
||||||
var task GatewayTask
|
var task GatewayTask
|
||||||
changed := false
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -789,7 +869,7 @@ func (s *Store) CancelSubmittedTask(ctx context.Context, taskID string, executio
|
|||||||
message = truncateUTF8Bytes(message, 2048)
|
message = truncateUTF8Bytes(message, 2048)
|
||||||
var task GatewayTask
|
var task GatewayTask
|
||||||
changed := false
|
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
|
var err error
|
||||||
task, err = scanGatewayTask(tx.QueryRow(ctx, `
|
task, err = scanGatewayTask(tx.QueryRow(ctx, `
|
||||||
UPDATE gateway_tasks
|
UPDATE gateway_tasks
|
||||||
@@ -851,14 +931,23 @@ func (s *Store) ListRecoverableAsyncTasks(ctx context.Context, limit int) ([]Asy
|
|||||||
limit = 500
|
limit = 500
|
||||||
}
|
}
|
||||||
rows, err := s.pool.Query(ctx, `
|
rows, err := s.pool.Query(ctx, `
|
||||||
SELECT id::text, priority, next_run_at
|
SELECT task.id::text, task.priority, task.next_run_at
|
||||||
FROM gateway_tasks
|
FROM gateway_tasks task
|
||||||
WHERE async_mode = true
|
LEFT JOIN river_job job ON job.id = task.river_job_id
|
||||||
|
WHERE task.async_mode = true
|
||||||
AND (
|
AND (
|
||||||
status = 'queued'
|
task.status = 'queued'
|
||||||
OR (status = 'running' AND (execution_lease_expires_at IS NULL OR execution_lease_expires_at <= now()))
|
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)
|
LIMIT $1`, limit)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -883,7 +972,7 @@ func (s *Store) CreateTaskAttempt(ctx context.Context, input CreateTaskAttemptIn
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
|
|
||||||
var attemptID string
|
var attemptID string
|
||||||
err = tx.QueryRow(ctx, `
|
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)
|
finalChargeAmount = strconv.FormatFloat(input.FinalChargeAmount, 'f', 9, 64)
|
||||||
}
|
}
|
||||||
currency := normalizeWalletCurrency(input.BillingCurrency)
|
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 strings.TrimSpace(input.AttemptID) != "" {
|
||||||
if _, err := tx.Exec(ctx, `
|
if _, err := tx.Exec(ctx, `
|
||||||
UPDATE gateway_task_attempts
|
UPDATE gateway_task_attempts
|
||||||
@@ -1430,7 +1519,7 @@ func (s *Store) FinishTaskManualReview(ctx context.Context, input FinishTaskManu
|
|||||||
}
|
}
|
||||||
resultJSON, _ := json.Marshal(resultReport.Value)
|
resultJSON, _ := json.Marshal(resultReport.Value)
|
||||||
pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(emptyObjectIfNil(input.PricingSnapshot)))
|
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) != "" {
|
if strings.TrimSpace(input.AttemptID) != "" {
|
||||||
attemptStatus := "failed"
|
attemptStatus := "failed"
|
||||||
if status == "succeeded" {
|
if status == "succeeded" {
|
||||||
@@ -1532,7 +1621,7 @@ func (s *Store) SettleTaskBilling(ctx context.Context, task GatewayTask) error {
|
|||||||
"billingSummary": task.BillingSummary,
|
"billingSummary": task.BillingSummary,
|
||||||
}
|
}
|
||||||
metadata, _ := json.Marshal(sanitizeJSONForStorage(metadataMap))
|
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, `
|
if _, err := tx.Exec(ctx, `
|
||||||
INSERT INTO gateway_wallet_accounts (
|
INSERT INTO gateway_wallet_accounts (
|
||||||
gateway_tenant_id, gateway_user_id, tenant_id, tenant_key, user_id, currency
|
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)))
|
metricsJSON, _ := json.Marshal(sanitizeJSONForStorage(emptyObjectIfNil(input.Metrics)))
|
||||||
resultJSON, _ := json.Marshal(minimalTaskResult(nil))
|
resultJSON, _ := json.Marshal(minimalTaskResult(nil))
|
||||||
message := truncateUTF8Bytes(input.Message, 2048)
|
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, `
|
tag, err := tx.Exec(ctx, `
|
||||||
UPDATE gateway_tasks
|
UPDATE gateway_tasks
|
||||||
SET status = 'failed',
|
SET status = 'failed',
|
||||||
@@ -1817,7 +1906,7 @@ func (s *Store) AddTaskEvent(ctx context.Context, taskID string, eventType strin
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return TaskEvent{}, err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT 1 FROM gateway_tasks WHERE id = $1::uuid FOR UPDATE`, taskID); err != nil {
|
||||||
return TaskEvent{}, err
|
return TaskEvent{}, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,18 +2,52 @@ package store
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
"github.com/jackc/pgx/v5/pgconn"
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const transactionRollbackTimeout = 5 * time.Second
|
||||||
|
|
||||||
type Tx interface {
|
type Tx interface {
|
||||||
Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error)
|
Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error)
|
||||||
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
|
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) InTx(ctx context.Context, fn func(Tx) error) error {
|
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)
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -163,7 +163,7 @@ func (s *Store) ReserveTaskBilling(ctx context.Context, task GatewayTask, user *
|
|||||||
reservations := make([]WalletBillingReservation, 0, len(amounts))
|
reservations := make([]WalletBillingReservation, 0, len(amounts))
|
||||||
pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot))
|
pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot))
|
||||||
requestFingerprint := walletString(pricingSnapshot["requestFingerprint"])
|
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 {
|
for currency, rawAmount := range amounts {
|
||||||
amount := roundMoney(rawAmount)
|
amount := roundMoney(rawAmount)
|
||||||
if amount <= 0 {
|
if amount <= 0 {
|
||||||
@@ -309,7 +309,7 @@ func (s *Store) reserveTaskBillingExact(ctx context.Context, task GatewayTask, g
|
|||||||
pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot))
|
pricingSnapshotJSON, _ := json.Marshal(sanitizeJSONForStorage(pricingSnapshot))
|
||||||
requestFingerprint := walletString(pricingSnapshot["requestFingerprint"])
|
requestFingerprint := walletString(pricingSnapshot["requestFingerprint"])
|
||||||
var reservations []WalletBillingReservation
|
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
|
var positive bool
|
||||||
if err := tx.QueryRow(ctx, `SELECT $1::numeric(38, 9) > 0`, amount).Scan(&positive); err != nil {
|
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)
|
return fmt.Errorf("invalid exact reservation amount: %w", err)
|
||||||
@@ -408,7 +408,7 @@ func (s *Store) ReleaseTaskBillingReservations(ctx context.Context, reservations
|
|||||||
if reason == "" {
|
if reason == "" {
|
||||||
reason = "task_not_settled"
|
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{}{}
|
taskIDs := map[string]struct{}{}
|
||||||
for _, reservation := range reservations {
|
for _, reservation := range reservations {
|
||||||
if strings.TrimSpace(reservation.AccountID) == "" {
|
if strings.TrimSpace(reservation.AccountID) == "" {
|
||||||
|
|||||||
@@ -70,7 +70,7 @@ func (s *Store) RegisterWorkerInstance(ctx context.Context, input WorkerRegistra
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return WorkerAllocation{}, err
|
return WorkerAllocation{}, err
|
||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer rollbackTransaction(tx)
|
||||||
// Worker heartbeats and allocations are ephemeral coordination state that is
|
// Worker heartbeats and allocations are ephemeral coordination state that is
|
||||||
// refreshed every few seconds. Do not hold the global allocation lock while
|
// refreshed every few seconds. Do not hold the global allocation lock while
|
||||||
// waiting for a synchronous replica to acknowledge the commit.
|
// waiting for a synchronous replica to acknowledge the commit.
|
||||||
@@ -245,7 +245,7 @@ func (s *Store) RecoverOrphanedAsyncRiverJobs(
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
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 {
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-river-orphan-recovery', 0))`); err != nil {
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,6 +24,8 @@ data:
|
|||||||
AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS: "7"
|
AI_GATEWAY_TASK_ANALYSIS_RETENTION_DAYS: "7"
|
||||||
AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS: "300"
|
AI_GATEWAY_TASK_CLEANUP_INTERVAL_SECONDS: "300"
|
||||||
AI_GATEWAY_TASK_CLEANUP_BATCH_SIZE: "1000"
|
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
|
CORS_ALLOWED_ORIGIN: https://ai.51easyai.com
|
||||||
AI_GATEWAY_PUBLIC_BASE_URL: https://ai.51easyai.com
|
AI_GATEWAY_PUBLIC_BASE_URL: https://ai.51easyai.com
|
||||||
AI_GATEWAY_WEB_BASE_URL: https://ai.51easyai.com
|
AI_GATEWAY_WEB_BASE_URL: https://ai.51easyai.com
|
||||||
|
|||||||
Reference in New Issue
Block a user