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