fix(admission): 按候选唤醒同步队列头
同步准入释放容量后,分别唤醒每个独立平台模型队列的 FIFO 队头;只有任务实际绑定用户组时才同时约束用户组队头,避免一个已饱和候选阻塞其他候选补位。\n\n调度器单次最多处理 256 个真实队头,不唤醒整个等待队列。新增 PostgreSQL 集成测试覆盖无用户组的双候选并行唤醒,以及共享用户组下仍保持组级 FIFO。\n\n验证:gofmt;Store/Runner 聚焦测试;一次性 PostgreSQL 集成测试;git diff --check。
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 == "" {
|
||||
|
||||
Reference in New Issue
Block a user