fix(admission): 按候选唤醒同步队列头

同步准入释放容量后,分别唤醒每个独立平台模型队列的 FIFO 队头;只有任务实际绑定用户组时才同时约束用户组队头,避免一个已饱和候选阻塞其他候选补位。\n\n调度器单次最多处理 256 个真实队头,不唤醒整个等待队列。新增 PostgreSQL 集成测试覆盖无用户组的双候选并行唤醒,以及共享用户组下仍保持组级 FIFO。\n\n验证:gofmt;Store/Runner 聚焦测试;一次性 PostgreSQL 集成测试;git diff --check。
This commit is contained in:
2026-07-31 02:18:22 +08:00
parent 3976cfb64d
commit e0f841e8fb
3 changed files with 156 additions and 7 deletions
@@ -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 == "" {