From b625edcd710973e89479dcd65a239895b2623c71 Mon Sep 17 00:00:00 2001 From: wangbo Date: Fri, 31 Jul 2026 07:37:37 +0800 Subject: [PATCH] =?UTF-8?q?fix(worker):=20=E6=94=B6=E6=95=9B=E5=BC=82?= =?UTF-8?q?=E6=AD=A5=E9=87=8D=E5=87=86=E5=85=A5=E5=88=B0=E8=B0=83=E5=BA=A6?= =?UTF-8?q?=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 准入租约失效时 Worker 仅释放执行准备并短暂让出 River 槽位,由异步调度器统一恢复 waiting 或无有效租约的任务,避免多 Worker 同时争抢全局容量锁并放大数据库同步提交等待。\n\n调度扫描覆盖到期且可运行的 scheduled/available/retryable River 任务,但跳过正在执行或仍持有有效准入租约的任务;新增 PostgreSQL/River 状态集成回归测试。\n\n验证:Go 全量测试、go vet、runner/store race、gofmt、迁移安全检查通过;专用数据库集成用例因本机未配置 AI_GATEWAY_TEST_DATABASE_URL 明确跳过。 --- apps/api/internal/runner/queue_client_test.go | 139 ++++++++++++++++++ apps/api/internal/runner/service.go | 3 - apps/api/internal/store/admission_queue.go | 19 ++- 3 files changed, 155 insertions(+), 6 deletions(-) diff --git a/apps/api/internal/runner/queue_client_test.go b/apps/api/internal/runner/queue_client_test.go index d167bdc..b427fae 100644 --- a/apps/api/internal/runner/queue_client_test.go +++ b/apps/api/internal/runner/queue_client_test.go @@ -73,6 +73,145 @@ WHERE id = $1`, queued.RiverJobID).Scan(&attemptedByCount); err != nil { } } +func TestWaitingAsyncAdmissionDispatcherOwnsScheduledRiverJobs(t *testing.T) { + databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) + if databaseURL == "" { + t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run the async admission dispatcher integration test") + } + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + db, err := store.Connect(ctx, databaseURL) + if err != nil { + t.Fatalf("connect store: %v", err) + } + t.Cleanup(db.Close) + + var databaseName string + if err := db.Pool().QueryRow(ctx, `SELECT current_database()`).Scan(&databaseName); err != nil { + t.Fatalf("read test database name: %v", err) + } + if !strings.Contains(strings.ToLower(databaseName), "test") { + t.Fatalf("refusing to use non-test database %q", databaseName) + } + + suffix := strconv.FormatInt(time.Now().UnixNano(), 10) + platform, err := db.CreatePlatform(ctx, store.CreatePlatformInput{ + Provider: "async-dispatch-test", + PlatformKey: "async-dispatch-test-" + suffix, + Name: "Async Dispatch Test " + suffix, + AuthType: "none", + Status: "enabled", + }) + if err != nil { + t.Fatalf("create platform: %v", err) + } + modelName := "async-dispatch-model-" + suffix + var platformModelID string + if err := db.Pool().QueryRow(ctx, ` +INSERT INTO platform_models ( + platform_id, model_name, provider_model_name, model_alias, model_type, + display_name, pricing_mode, enabled +) +VALUES ($1::uuid, $2, $2, $2, '["image_generate"]'::jsonb, $2, 'inherit_discount', true) +RETURNING id::text`, platform.ID, modelName).Scan(&platformModelID); err != nil { + t.Fatalf("create platform model: %v", err) + } + + task, err := db.CreateTask(ctx, store.CreateTaskInput{ + Kind: "images.edits", + Model: modelName, + Request: map[string]any{"prompt": "scheduled River job re-admission"}, + Async: true, + RunMode: "simulation", + }, &auth.User{ID: "async-dispatch-user-" + suffix, Source: "gateway"}) + if err != nil { + t.Fatalf("create task: %v", err) + } + t.Cleanup(func() { + cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cleanupCancel() + _ = db.DeleteTaskAdmission(cleanupCtx, task.ID) + _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM river_job WHERE id = (SELECT river_job_id FROM gateway_tasks WHERE id = $1::uuid)`, task.ID) + _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = $1::uuid`, task.ID) + _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM platform_models WHERE id = $1::uuid`, platformModelID) + _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM integration_platforms WHERE id = $1::uuid`, platform.ID) + }) + + service := New(config.Config{AppEnv: "test"}, db, slog.New(slog.NewTextHandler(io.Discard, nil))) + service.StartAsyncQueueClient(ctx) + if err := service.EnqueueAsyncTask(ctx, task); err != nil { + t.Fatalf("enqueue task: %v", err) + } + queued, err := db.GetTask(ctx, task.ID) + if err != nil { + t.Fatalf("load queued task: %v", err) + } + if queued.RiverJobID <= 0 { + t.Fatal("queued task did not persist a River job ID") + } + admissionInput := store.TaskAdmissionInput{ + TaskID: task.ID, + PlatformID: platform.ID, + PlatformModelID: platformModelID, + QueueKey: "async-dispatch-test:" + suffix, + Mode: "async", + Priority: 100, + Scopes: []store.AdmissionScope{{ + ScopeType: "worker_capacity", + ScopeKey: "global", + ScopeName: "Worker capacity", + ConcurrentLimit: 1, + Amount: 1, + LeaseTTLSeconds: 120, + }}, + } + if _, err := db.QueueTaskAdmissionWithHook(ctx, admissionInput, nil); err != nil { + t.Fatalf("queue task admission: %v", err) + } + + assertDispatchable := func(want bool) { + t.Helper() + taskIDs, listErr := db.ListWaitingAsyncAdmissionTaskIDs(ctx, 1000) + if listErr != nil { + t.Fatalf("list waiting async admissions: %v", listErr) + } + found := false + for _, taskID := range taskIDs { + if taskID == task.ID { + found = true + break + } + } + if found != want { + t.Fatalf("task dispatchable=%v, want %v", found, want) + } + } + + assertDispatchable(true) + if _, err := db.Pool().Exec(ctx, `UPDATE river_job SET state = 'running' WHERE id = $1`, queued.RiverJobID); err != nil { + t.Fatalf("mark River job running: %v", err) + } + assertDispatchable(false) + if _, err := db.Pool().Exec(ctx, `UPDATE river_job SET state = 'scheduled' WHERE id = $1`, queued.RiverJobID); err != nil { + t.Fatalf("snooze River job: %v", err) + } + assertDispatchable(true) + admitted, err := db.TryTaskAdmission(ctx, admissionInput) + if err != nil || !admitted.Admitted || len(admitted.Leases) != 1 { + t.Fatalf("admit scheduled task: result=%+v err=%v", admitted, err) + } + assertDispatchable(false) + if _, err := db.Pool().Exec(ctx, ` +UPDATE gateway_concurrency_leases +SET expires_at = now() - interval '1 second' +WHERE task_id = $1::uuid + AND released_at IS NULL`, task.ID); err != nil { + t.Fatalf("expire admission lease: %v", err) + } + assertDispatchable(true) +} + func TestRecoverOrphanedAsyncRiverJobAndYieldStaleAdmission(t *testing.T) { databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) if databaseURL == "" { diff --git a/apps/api/internal/runner/service.go b/apps/api/internal/runner/service.go index dc21835..103e47e 100644 --- a/apps/api/internal/runner/service.go +++ b/apps/api/internal/runner/service.go @@ -577,9 +577,6 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask, var alreadyAdmitted bool admissionResult, alreadyAdmitted, admissionErr = s.activeAsyncTaskAdmission(ctx, task, plan, asyncAdmission) if admissionErr == nil && !alreadyAdmitted { - admissionResult, admissionErr = s.tryTaskAdmission(ctx, task, plan, "") - } - if admissionErr == nil && !admissionResult.Admitted { _ = s.store.ReleaseTaskPreparation(context.WithoutCancel(ctx), task.ID, task.ExecutionToken) return Result{Task: task, Output: task.Result}, &TaskQueuedError{Delay: time.Second} } diff --git a/apps/api/internal/store/admission_queue.go b/apps/api/internal/store/admission_queue.go index 1f30866..0a80d7b 100644 --- a/apps/api/internal/store/admission_queue.go +++ b/apps/api/internal/store/admission_queue.go @@ -721,13 +721,26 @@ SELECT admission.task_id::text FROM gateway_task_admissions admission JOIN gateway_tasks task ON task.id = admission.task_id LEFT JOIN river_job job ON job.id = task.river_job_id -WHERE admission.status = 'waiting' - AND admission.mode = 'async' +WHERE admission.mode = 'async' AND task.status = 'queued' + AND task.next_run_at <= now() + AND ( + admission.status = 'waiting' + OR ( + admission.status = 'admitted' + AND NOT EXISTS ( + SELECT 1 + FROM gateway_concurrency_leases lease + WHERE lease.task_id = admission.task_id + AND lease.released_at IS NULL + AND lease.expires_at > now() + ) + ) + ) AND ( task.river_job_id IS NULL OR job.id IS NULL - OR job.state NOT IN ('available', 'pending', 'retryable', 'running', 'scheduled') + OR job.state <> 'running' ) ORDER BY admission.priority ASC, admission.enqueued_at ASC, admission.task_id ASC LIMIT $1`, limit)