diff --git a/apps/api/internal/store/acceptance.go b/apps/api/internal/store/acceptance.go index d5834a4..ab1bfab 100644 --- a/apps/api/internal/store/acceptance.go +++ b/apps/api/internal/store/acceptance.go @@ -153,6 +153,12 @@ func (s *Store) CreateAcceptanceRun(ctx context.Context, input CreateAcceptanceR } config, _ := json.Marshal(sanitizeAcceptanceMetadata(sanitizeJSONForStorage(input.Config))) tokenHash := acceptanceTokenSHA256(input.Token) + if err := s.beginTransaction(ctx, func(tx pgx.Tx) error { + _, err := cancelSafeAcceptanceTasksTx(ctx, tx, "") + return err + }); err != nil { + return AcceptanceRun{}, err + } return scanAcceptanceRun(s.pool.QueryRow(ctx, ` INSERT INTO gateway_acceptance_runs ( release_sha, api_image_digest, worker_image_digest, api_key_id, user_id, @@ -370,6 +376,9 @@ FOR UPDATE`, SystemSettingGatewayTrafficMode).Scan(&value); err != nil { current.WorkerImageDigest != strings.TrimSpace(input.WorkerImageDigest) { return GatewayTrafficMode{}, ErrAcceptanceStateConflict } + if _, err := cancelSafeAcceptanceTasksTx(ctx, tx, current.RunID); err != nil { + return GatewayTrafficMode{}, err + } next := GatewayTrafficMode{Mode: "live", Revision: current.Revision + 1} nextValue, _ := json.Marshal(next) if _, err := tx.Exec(ctx, ` @@ -391,6 +400,169 @@ WHERE id = $1::uuid AND status IN ('running', 'failed', 'passed')`, current.RunI return next, nil } +func cancelSafeAcceptanceTasksTx(ctx context.Context, tx pgx.Tx, runID string) (int64, error) { + runID = strings.TrimSpace(runID) + tag, err := tx.Exec(ctx, ` +UPDATE gateway_tasks task +SET status = 'cancelled', + error = NULL, + error_code = 'acceptance_run_aborted', + error_message = 'acceptance run ended before upstream submission', + billing_status = CASE + WHEN gateway_user_id IS NULL THEN 'not_required' + WHEN reservation_amount > 0 THEN 'pending' + ELSE 'released' + END, + billing_updated_at = now(), + remote_task_payload = '{}'::jsonb, + locked_by = NULL, + locked_at = NULL, + heartbeat_at = NULL, + execution_token = NULL, + execution_lease_expires_at = NULL, + finished_at = now(), + updated_at = now() +WHERE task.run_mode = 'acceptance' + AND task.status IN ('queued', 'running') + AND COALESCE(task.remote_task_id, '') = '' + AND NOT EXISTS ( + SELECT 1 + FROM gateway_task_attempts attempt + WHERE attempt.task_id = task.id + AND attempt.upstream_submission_status <> 'not_submitted' + ) + AND ( + task.acceptance_run_id = NULLIF($1, '')::uuid + OR ( + NULLIF($1, '') IS NULL + AND EXISTS ( + SELECT 1 + FROM gateway_acceptance_runs run + WHERE run.id = task.acceptance_run_id + AND run.status IN ('failed', 'aborted') + ) + ) + )`, runID) + if err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, ` +DELETE FROM gateway_task_param_preprocessing_logs log +USING gateway_task_attempts attempt, gateway_tasks task +WHERE log.attempt_id = attempt.id + AND attempt.task_id = task.id + AND attempt.upstream_submission_status = 'not_submitted' + AND task.error_code = 'acceptance_run_aborted' + AND ( + task.acceptance_run_id = NULLIF($1, '')::uuid + OR ( + NULLIF($1, '') IS NULL + AND EXISTS ( + SELECT 1 + FROM gateway_acceptance_runs run + WHERE run.id = task.acceptance_run_id + AND run.status IN ('failed', 'aborted') + ) + ) + )`, runID); err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, ` +DELETE FROM gateway_task_attempts attempt +USING gateway_tasks task +WHERE attempt.task_id = task.id + AND attempt.upstream_submission_status = 'not_submitted' + AND task.error_code = 'acceptance_run_aborted' + AND ( + task.acceptance_run_id = NULLIF($1, '')::uuid + OR ( + NULLIF($1, '') IS NULL + AND EXISTS ( + SELECT 1 + FROM gateway_acceptance_runs run + WHERE run.id = task.acceptance_run_id + AND run.status IN ('failed', 'aborted') + ) + ) + )`, runID); err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, ` +INSERT INTO settlement_outbox ( + task_id, event_type, action, amount, currency, pricing_snapshot, payload, + status, next_attempt_at +) +SELECT task.id, 'task.billing.release', 'release', task.reservation_amount, + task.billing_currency, task.pricing_snapshot, + jsonb_build_object('taskId', task.id, 'reason', 'acceptance_run_aborted'), + 'pending', now() +FROM gateway_tasks task +WHERE task.error_code = 'acceptance_run_aborted' + AND task.billing_status = 'pending' + AND task.reservation_amount > 0 + AND ( + task.acceptance_run_id = NULLIF($1, '')::uuid + OR ( + NULLIF($1, '') IS NULL + AND EXISTS ( + SELECT 1 + FROM gateway_acceptance_runs run + WHERE run.id = task.acceptance_run_id + AND run.status IN ('failed', 'aborted') + ) + ) + ) +ON CONFLICT (task_id, event_type) DO NOTHING`, runID); err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, ` +UPDATE gateway_concurrency_leases lease +SET released_at = now() +FROM gateway_tasks task +WHERE lease.task_id = task.id + AND lease.released_at IS NULL + AND task.error_code = 'acceptance_run_aborted' + AND ( + task.acceptance_run_id = NULLIF($1, '')::uuid + OR ( + NULLIF($1, '') IS NULL + AND EXISTS ( + SELECT 1 + FROM gateway_acceptance_runs run + WHERE run.id = task.acceptance_run_id + AND run.status IN ('failed', 'aborted') + ) + ) + )`, runID); err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, ` +DELETE FROM gateway_task_admissions admission +USING gateway_tasks task +WHERE admission.task_id = task.id + AND task.error_code = 'acceptance_run_aborted' + AND ( + task.acceptance_run_id = NULLIF($1, '')::uuid + OR ( + NULLIF($1, '') IS NULL + AND EXISTS ( + SELECT 1 + FROM gateway_acceptance_runs run + WHERE run.id = task.acceptance_run_id + AND run.status IN ('failed', 'aborted') + ) + ) + )`, runID); err != nil { + return 0, err + } + if tag.RowsAffected() > 0 { + if err := notifyTaskAdmissionTx(ctx, tx, "*"); err != nil { + return 0, err + } + } + return tag.RowsAffected(), nil +} + func (s *Store) AuthorizeAcceptanceTask(ctx context.Context, runID string, token string, user *auth.User) (string, error) { mode, err := s.GetGatewayTrafficMode(ctx) if err != nil { diff --git a/apps/api/internal/store/acceptance_integration_test.go b/apps/api/internal/store/acceptance_integration_test.go index 405dce6..d7486a9 100644 --- a/apps/api/internal/store/acceptance_integration_test.go +++ b/apps/api/internal/store/acceptance_integration_test.go @@ -114,6 +114,68 @@ WHERE setting_key = $1`, SystemSettingGatewayTrafficMode) if err != nil { t.Fatalf("activate retry acceptance run: %v", err) } + createAcceptanceTask := func(label string) GatewayTask { + t.Helper() + task, createErr := db.CreateTask(ctx, CreateTaskInput{ + Kind: "images.edits", Model: "acceptance-cleanup-" + label, + RunMode: "acceptance", AcceptanceRunID: failedRun.ID, + Request: map[string]any{"prompt": "acceptance cleanup integration test"}, + }, &auth.User{ID: "acceptance-user", Source: "gateway"}) + if createErr != nil { + t.Fatalf("create acceptance cleanup task: %v", createErr) + } + return task + } + queuedTask := createAcceptanceTask("queued") + runningTask := createAcceptanceTask("running") + notSubmittedTask := createAcceptanceTask("not-submitted") + submittingTask := createAcceptanceTask("submitting") + var billingUserID string + if err := db.pool.QueryRow(ctx, ` +INSERT INTO gateway_users (user_key, username) +VALUES ($1, $1) +RETURNING id::text`, "acceptance-cleanup-"+time.Now().Format("150405.000000000")).Scan(&billingUserID); err != nil { + t.Fatalf("create acceptance cleanup billing user: %v", err) + } + if _, err := db.pool.Exec(ctx, ` +UPDATE gateway_tasks +SET gateway_user_id = $2::uuid, + reservation_amount = 1, + billing_status = 'not_started' +WHERE id = $1::uuid`, queuedTask.ID, billingUserID); err != nil { + t.Fatalf("reserve acceptance cleanup task billing: %v", err) + } + if _, err := db.ClaimTaskExecution(ctx, runningTask.ID, "10000000-0000-4000-8000-000000000001", time.Minute); err != nil { + t.Fatalf("claim safe running acceptance task: %v", err) + } + notSubmittedClaim, err := db.ClaimTaskExecution(ctx, notSubmittedTask.ID, "10000000-0000-4000-8000-000000000002", time.Minute) + if err != nil { + t.Fatalf("claim not-submitted acceptance task: %v", err) + } + notSubmittedAttemptID, err := db.CreateTaskAttempt(ctx, CreateTaskAttemptInput{ + TaskID: notSubmittedTask.ID, ExecutionToken: notSubmittedClaim.ExecutionToken, + AttemptNo: 1, Status: "running", Simulated: true, + }) + if err != nil { + t.Fatalf("create not-submitted acceptance task attempt: %v", err) + } + submittingClaim, err := db.ClaimTaskExecution(ctx, submittingTask.ID, "10000000-0000-4000-8000-000000000003", time.Minute) + if err != nil { + t.Fatalf("claim submitting acceptance task: %v", err) + } + submittingAttemptID, err := db.CreateTaskAttempt(ctx, CreateTaskAttemptInput{ + TaskID: submittingTask.ID, ExecutionToken: submittingClaim.ExecutionToken, + AttemptNo: 1, Status: "running", Simulated: true, + }) + if err != nil { + t.Fatalf("create submitting acceptance task attempt: %v", err) + } + if _, err := db.pool.Exec(ctx, ` +UPDATE gateway_task_attempts +SET upstream_submission_status = 'submitting', upstream_submission_updated_at = now() +WHERE id = $1::uuid`, submittingAttemptID); err != nil { + t.Fatalf("mark acceptance task attempt submitting: %v", err) + } if _, err := db.FinishAcceptanceRun(ctx, FinishAcceptanceRunInput{ RunID: failedRun.ID, Passed: false, FailureReason: "load failed", }); err != nil { @@ -133,4 +195,46 @@ WHERE setting_key = $1`, SystemSettingGatewayTrafficMode) if aborted.Mode != "live" || aborted.Revision != failedMode.Revision+1 { t.Fatalf("unexpected aborted mode: %+v", aborted) } + var safeCancelled, submittingRunning, notSubmittedAttempts, releaseEvents int + if err := db.pool.QueryRow(ctx, ` +SELECT + count(*) FILTER ( + WHERE id = ANY($1::uuid[]) + AND status = 'cancelled' + AND error_code = 'acceptance_run_aborted' + ), + count(*) FILTER ( + WHERE id = $2::uuid + AND status = 'running' + ), + (SELECT count(*) FROM gateway_task_attempts WHERE id = $3::uuid), + (SELECT count(*) + FROM settlement_outbox + WHERE task_id = $4::uuid + AND action = 'release' + AND status = 'pending') +FROM gateway_tasks`, + []string{queuedTask.ID, runningTask.ID, notSubmittedTask.ID}, + submittingTask.ID, + notSubmittedAttemptID, + queuedTask.ID, + ).Scan(&safeCancelled, &submittingRunning, ¬SubmittedAttempts, &releaseEvents); err != nil { + t.Fatalf("read aborted acceptance task states: %v", err) + } + if safeCancelled != 3 || submittingRunning != 1 || notSubmittedAttempts != 0 || releaseEvents != 1 { + t.Fatalf( + "abort cleanup safe_cancelled=%d submitting_running=%d not_submitted_attempts=%d release_events=%d", + safeCancelled, + submittingRunning, + notSubmittedAttempts, + releaseEvents, + ) + } + if _, err := db.pool.Exec(ctx, `DELETE FROM gateway_tasks WHERE id = ANY($1::uuid[])`, + []string{queuedTask.ID, runningTask.ID, notSubmittedTask.ID, submittingTask.ID}); err != nil { + t.Fatalf("delete acceptance cleanup tasks: %v", err) + } + if _, err := db.pool.Exec(ctx, `DELETE FROM gateway_users WHERE id = $1::uuid`, billingUserID); err != nil { + t.Fatalf("delete acceptance cleanup billing user: %v", err) + } } diff --git a/scripts/cluster/run-production-acceptance.sh b/scripts/cluster/run-production-acceptance.sh index c441cf7..4380c60 100755 --- a/scripts/cluster/run-production-acceptance.sh +++ b/scripts/cluster/run-production-acceptance.sh @@ -218,6 +218,26 @@ wait_for_existing_tasks_to_drain() { return 1 } +wait_for_terminal_acceptance_tasks_to_drain() { + local deadline=$((SECONDS + 300)) + local active + while (( SECONDS < deadline )); do + active=$(database_query " +SELECT count(*) +FROM gateway_tasks task +JOIN gateway_acceptance_runs run ON run.id=task.acceptance_run_id +WHERE run.status IN ('failed', 'aborted') + AND task.status IN ('queued', 'running');") + if [[ $active == 0 ]]; then + return 0 + fi + echo "waiting_for_terminal_acceptance_tasks=$active" + sleep 2 + done + echo "terminal acceptance tasks did not drain within 5 minutes: active=$active" >&2 + return 1 +} + ensure_acceptance_user_group() { local groups_response=$temporary_root/user-groups.json local group_response=$temporary_root/acceptance-user-group.json @@ -621,6 +641,7 @@ create_and_activate_run() { echo 'previous failed acceptance Run did not return traffic to live' >&2 return 1 } + wait_for_terminal_acceptance_tasks_to_drain elif [[ $traffic_mode != live ]]; then echo "unsupported production traffic mode before acceptance: $traffic_mode" >&2 return 1