fix(worker): 收敛异步重准入到调度器

准入租约失效时 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 明确跳过。
This commit is contained in:
2026-07-31 07:37:37 +08:00
parent d1b482c2f3
commit b625edcd71
3 changed files with 155 additions and 6 deletions
@@ -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 == "" {
-3
View File
@@ -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}
}