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:
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user