From e0f841e8fb45663e04d51d0a26811ae545227c5a Mon Sep 17 00:00:00 2001 From: wangbo Date: Fri, 31 Jul 2026 02:18:22 +0800 Subject: [PATCH] =?UTF-8?q?fix(admission):=20=E6=8C=89=E5=80=99=E9=80=89?= =?UTF-8?q?=E5=94=A4=E9=86=92=E5=90=8C=E6=AD=A5=E9=98=9F=E5=88=97=E5=A4=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 同步准入释放容量后,分别唤醒每个独立平台模型队列的 FIFO 队头;只有任务实际绑定用户组时才同时约束用户组队头,避免一个已饱和候选阻塞其他候选补位。\n\n调度器单次最多处理 256 个真实队头,不唤醒整个等待队列。新增 PostgreSQL 集成测试覆盖无用户组的双候选并行唤醒,以及共享用户组下仍保持组级 FIFO。\n\n验证:gofmt;Store/Runner 聚焦测试;一次性 PostgreSQL 集成测试;git diff --check。 --- apps/api/internal/runner/admission.go | 3 +- apps/api/internal/store/admission_queue.go | 14 +- .../store/admission_queue_integration_test.go | 146 ++++++++++++++++++ 3 files changed, 156 insertions(+), 7 deletions(-) diff --git a/apps/api/internal/runner/admission.go b/apps/api/internal/runner/admission.go index 9ea362d..4c7d6f1 100644 --- a/apps/api/internal/runner/admission.go +++ b/apps/api/internal/runner/admission.go @@ -34,6 +34,7 @@ const ( asyncWorkerMaxWaitSeconds = 24 * 60 * 60 acceptanceQueueLimit = 10000 acceptanceQueueMaxWait = 15 * 60 + synchronousAdmissionWakeMax = 256 ) func distributedAdmissionModelType(modelType string) bool { @@ -416,7 +417,7 @@ func (s *Service) dispatchWaitingSynchronousAdmissions(ctx context.Context) { return case <-s.admissionWake: } - taskIDs, err := s.store.ListWaitingTaskAdmissionIDs(ctx, 1) + taskIDs, err := s.store.ListWaitingTaskAdmissionIDs(ctx, synchronousAdmissionWakeMax) if err != nil { if s.logger != nil { s.logger.Warn("list waiting synchronous admissions failed", "error", err) diff --git a/apps/api/internal/store/admission_queue.go b/apps/api/internal/store/admission_queue.go index 5473471..f8c6e2c 100644 --- a/apps/api/internal/store/admission_queue.go +++ b/apps/api/internal/store/admission_queue.go @@ -713,10 +713,11 @@ LIMIT $1`, limit) return taskIDs, rows.Err() } -// ListWaitingTaskAdmissionIDs returns a bounded set of FIFO leaders across -// platform-model and user-group queues. It is used after capacity is released -// so API processes wake only plausible queue heads instead of every -// synchronous waiter. +// ListWaitingTaskAdmissionIDs returns one FIFO leader for every independent +// platform-model queue, additionally requiring the task to lead its user-group +// queue when one exists. It is used after capacity is released so API +// processes can refill every candidate without waking every synchronous +// waiter or allowing one saturated candidate to block another. func (s *Store) ListWaitingTaskAdmissionIDs(ctx context.Context, limit int) ([]string, error) { if limit <= 0 { limit = 256 @@ -724,6 +725,7 @@ func (s *Store) ListWaitingTaskAdmissionIDs(ctx context.Context, limit int) ([]s rows, err := s.pool.Query(ctx, ` WITH ranked AS ( SELECT task_id, + user_group_id, priority, enqueued_at, row_number() OVER ( @@ -740,8 +742,8 @@ WITH ranked AS ( ) SELECT task_id::text FROM ranked -WHERE platform_rank <= $1 - AND group_rank <= $1 +WHERE platform_rank = 1 + AND (user_group_id IS NULL OR group_rank = 1) ORDER BY priority ASC, enqueued_at ASC, task_id ASC LIMIT $1`, limit) if err != nil { diff --git a/apps/api/internal/store/admission_queue_integration_test.go b/apps/api/internal/store/admission_queue_integration_test.go index 648b8dc..a852b47 100644 --- a/apps/api/internal/store/admission_queue_integration_test.go +++ b/apps/api/internal/store/admission_queue_integration_test.go @@ -3,6 +3,7 @@ package store import ( "context" "errors" + "fmt" "os" "strings" "sync" @@ -754,6 +755,151 @@ SELECT } } +func TestWaitingAdmissionLeadersDoNotCrossBlockIndependentCandidates(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 task admission PostgreSQL integration tests") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + applyOIDCJITTestMigrations(t, ctx, databaseURL) + db, err := Connect(ctx, databaseURL) + if err != nil { + t.Fatalf("connect store: %v", err) + } + defer 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 := strings.ReplaceAll(uuid.NewString(), "-", "") + platform, err := db.CreatePlatform(ctx, CreatePlatformInput{ + Provider: "leader-test", + PlatformKey: "leader-test-" + suffix, + Name: "Admission Leader Test " + suffix, + AuthType: "none", + Status: "enabled", + }) + if err != nil { + t.Fatalf("create platform: %v", err) + } + modelIDs := make([]string, 0, 2) + for index := 0; index < 2; index++ { + modelName := fmt.Sprintf("leader-model-%d-%s", index, suffix) + var modelID 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(&modelID); err != nil { + t.Fatalf("create platform model %d: %v", index, err) + } + modelIDs = append(modelIDs, modelID) + } + group, err := db.CreateUserGroup(ctx, UserGroupInput{ + GroupKey: "leader-test-" + suffix, + Name: "Admission Leader Test " + suffix, + Status: "active", + }) + if err != nil { + t.Fatalf("create user group: %v", err) + } + taskIDs := make([]string, 0, 4) + t.Cleanup(func() { + cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cleanupCancel() + for _, taskID := range taskIDs { + _, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = $1::uuid`, taskID) + } + _, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_user_groups WHERE id = $1::uuid`, group.ID) + _, _ = db.pool.Exec(cleanupCtx, `DELETE FROM platform_models WHERE id = ANY($1::uuid[])`, modelIDs) + _, _ = db.pool.Exec(cleanupCtx, `DELETE FROM integration_platforms WHERE id = $1::uuid`, platform.ID) + }) + createTask := func(modelIndex int) GatewayTask { + t.Helper() + task, createErr := db.CreateTask(ctx, CreateTaskInput{ + Kind: "images.generate", + Model: fmt.Sprintf("leader-model-%d-%s", modelIndex, suffix), + RunMode: "production", + Request: map[string]any{"prompt": "admission leader test"}, + }, &auth.User{ID: "leader-test-user-" + suffix, Source: "gateway"}) + if createErr != nil { + t.Fatalf("create task: %v", createErr) + } + taskIDs = append(taskIDs, task.ID) + return task + } + tasks := []GatewayTask{createTask(0), createTask(0), createTask(1), createTask(1)} + insertWaiting := func(task GatewayTask, modelID string, groupID string, enqueuedAt time.Time) { + t.Helper() + if _, insertErr := db.pool.Exec(ctx, ` +INSERT INTO gateway_task_admissions ( + task_id, platform_id, platform_model_id, user_group_id, queue_key, mode, + status, priority, enqueued_at, wait_deadline_at, waiter_id, + waiter_lease_expires_at +) +VALUES ( + $1::uuid, $2::uuid, $3::uuid, NULLIF($4, '')::uuid, $5, 'sync', + 'waiting', 100, $6, now() + interval '10 minutes', $7, + now() + interval '15 seconds' +)`, + task.ID, + platform.ID, + modelID, + groupID, + "leader-test:"+modelID, + enqueuedAt, + "waiter-"+task.ID, + ); insertErr != nil { + t.Fatalf("insert waiting admission: %v", insertErr) + } + } + baseTime := time.Now().Add(-time.Minute) + insertWaiting(tasks[0], modelIDs[0], "", baseTime) + insertWaiting(tasks[1], modelIDs[0], "", baseTime.Add(time.Second)) + insertWaiting(tasks[2], modelIDs[1], "", baseTime.Add(2*time.Second)) + insertWaiting(tasks[3], modelIDs[1], "", baseTime.Add(3*time.Second)) + leaders, err := db.ListWaitingTaskAdmissionIDs(ctx, 256) + if err != nil { + t.Fatalf("list independent candidate leaders: %v", err) + } + leaderSet := make(map[string]bool, len(leaders)) + for _, taskID := range leaders { + leaderSet[taskID] = true + } + if !leaderSet[tasks[0].ID] || !leaderSet[tasks[2].ID] { + t.Fatalf("independent candidate leaders = %v, want %s and %s", leaders, tasks[0].ID, tasks[2].ID) + } + if leaderSet[tasks[1].ID] || leaderSet[tasks[3].ID] { + t.Fatalf("non-head independent candidate task was selected: %v", leaders) + } + + if _, err := db.pool.Exec(ctx, ` +DELETE FROM gateway_task_admissions +WHERE task_id = ANY($1::uuid[])`, taskIDs); err != nil { + t.Fatalf("clear independent admissions: %v", err) + } + insertWaiting(tasks[0], modelIDs[0], group.ID, baseTime) + insertWaiting(tasks[2], modelIDs[1], group.ID, baseTime.Add(time.Second)) + leaders, err = db.ListWaitingTaskAdmissionIDs(ctx, 256) + if err != nil { + t.Fatalf("list grouped candidate leaders: %v", err) + } + leaderSet = make(map[string]bool, len(leaders)) + for _, taskID := range leaders { + leaderSet[taskID] = true + } + if !leaderSet[tasks[0].ID] || leaderSet[tasks[2].ID] { + t.Fatalf("grouped candidate leaders = %v, want only group head %s", leaders, tasks[0].ID) + } +} + func TestWorkerCapacityAllocationAndFailover(t *testing.T) { databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) if databaseURL == "" {