feat(queue): 增加非文本模型分布式准入队列

使用 PostgreSQL 统一同步与异步非文本任务的并发准入、持久化等待和 Worker 容量分配,并将生产 API 与独立 Worker 角色拆分。

补充策略管理、共享契约、OpenAPI、Kubernetes 双节点 Worker 清单及跨节点验收脚本;未默认启用任何生产 queue_size 策略。

已在原基线完成 Go、前端、迁移、Shell、Kustomize 与长任务容量验收;合入最新主干后将重新执行发布门禁。
This commit is contained in:
2026-07-29 16:15:43 +08:00
parent 2e5a90731b
commit 9e4fc7362d
55 changed files with 4565 additions and 237 deletions
+1 -1
View File
@@ -38,7 +38,7 @@ func main() {
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop() defer stop()
db, err := store.Connect(ctx, cfg.DatabaseURL) db, err := store.ConnectWithMaxConns(ctx, cfg.DatabaseURL, cfg.DatabaseMaxConns)
if err != nil { if err != nil {
logger.Error("connect postgres failed", "error", err) logger.Error("connect postgres failed", "error", err)
os.Exit(1) os.Exit(1)
+61 -7
View File
@@ -6255,7 +6255,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 OpenAI-compatible Images 响应;设置 X-Async=true 时异步创建任务并返回 202。", "description": "默认同步返回 OpenAI-compatible Images 响应;非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时异步创建任务并返回 202。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -6319,6 +6319,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.OpenAIErrorEnvelope" "$ref": "#/definitions/httpapi.OpenAIErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.OpenAIErrorEnvelope"
}
} }
} }
} }
@@ -6330,7 +6336,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 OpenAI-compatible Images 响应;设置 X-Async=true 时异步创建任务并返回 202。", "description": "默认同步返回 OpenAI-compatible Images 响应;非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时异步创建任务并返回 202。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -6394,6 +6400,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.OpenAIErrorEnvelope" "$ref": "#/definitions/httpapi.OpenAIErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.OpenAIErrorEnvelope"
}
} }
} }
} }
@@ -7059,7 +7071,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 EasyAI GeneratedResponse;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。", "description": "默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -7147,6 +7159,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope" "$ref": "#/definitions/httpapi.ErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
} }
} }
} }
@@ -7921,7 +7939,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 EasyAI GeneratedResponse;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。", "description": "默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -8009,6 +8027,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope" "$ref": "#/definitions/httpapi.ErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
} }
} }
} }
@@ -8020,7 +8044,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 EasyAI GeneratedResponse;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。", "description": "默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -8108,6 +8132,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope" "$ref": "#/definitions/httpapi.ErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
} }
} }
} }
@@ -8462,7 +8492,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 EasyAI GeneratedResponse;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。", "description": "默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -8550,6 +8580,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope" "$ref": "#/definitions/httpapi.ErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
} }
} }
} }
@@ -8781,7 +8817,7 @@
"BearerAuth": [] "BearerAuth": []
} }
], ],
"description": "默认同步返回 EasyAI GeneratedResponse;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。", "description": "默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -8869,6 +8905,12 @@
"schema": { "schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope" "$ref": "#/definitions/httpapi.ErrorEnvelope"
} }
},
"504": {
"description": "Gateway Timeout",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
} }
} }
} }
@@ -14325,6 +14367,9 @@
"type": "string" "type": "string"
} }
}, },
"oldestQueueWaitSeconds": {
"type": "number"
},
"platformCooldownUntil": { "platformCooldownUntil": {
"type": "string" "type": "string"
}, },
@@ -14358,6 +14403,9 @@
"providerModelName": { "providerModelName": {
"type": "string" "type": "string"
}, },
"queueLimit": {
"type": "integer"
},
"queuedTasks": { "queuedTasks": {
"type": "number" "type": "number"
}, },
@@ -14376,6 +14424,12 @@
}, },
"tpm": { "tpm": {
"$ref": "#/definitions/store.RateLimitMetricStatus" "$ref": "#/definitions/store.RateLimitMetricStatus"
},
"waitingAsyncTasks": {
"type": "integer"
},
"waitingSyncTasks": {
"type": "integer"
} }
} }
}, },
+50 -12
View File
@@ -3174,6 +3174,8 @@ definitions:
items: items:
type: string type: string
type: array type: array
oldestQueueWaitSeconds:
type: number
platformCooldownUntil: platformCooldownUntil:
type: string type: string
platformDisabledReason: platformDisabledReason:
@@ -3196,6 +3198,8 @@ definitions:
type: string type: string
providerModelName: providerModelName:
type: string type: string
queueLimit:
type: integer
queuedTasks: queuedTasks:
type: number type: number
rateLimitPolicy: rateLimitPolicy:
@@ -3209,6 +3213,10 @@ definitions:
$ref: '#/definitions/store.RateLimitMetricStatus' $ref: '#/definitions/store.RateLimitMetricStatus'
tpm: tpm:
$ref: '#/definitions/store.RateLimitMetricStatus' $ref: '#/definitions/store.RateLimitMetricStatus'
waitingAsyncTasks:
type: integer
waitingSyncTasks:
type: integer
type: object type: object
store.Platform: store.Platform:
properties: properties:
@@ -8008,7 +8016,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 OpenAI-compatible Images 响应;设置 X-Async=true 时异步创建任务并返回 202。 description: 默认同步返回 OpenAI-compatible Images 响应;非文本模型并发饱和且配置队列时等待准入,超过最长等待返回
504 queue_timeout;设置 X-Async=true 时异步创建任务并返回 202。
parameters: parameters:
- description: OpenAI Images 请求 - description: OpenAI Images 请求
in: body in: body
@@ -8047,6 +8056,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.OpenAIErrorEnvelope' $ref: '#/definitions/httpapi.OpenAIErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.OpenAIErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建或编辑 OpenAI Images summary: 创建或编辑 OpenAI Images
@@ -8056,7 +8069,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 OpenAI-compatible Images 响应;设置 X-Async=true 时异步创建任务并返回 202。 description: 默认同步返回 OpenAI-compatible Images 响应;非文本模型并发饱和且配置队列时等待准入,超过最长等待返回
504 queue_timeout;设置 X-Async=true 时异步创建任务并返回 202。
parameters: parameters:
- description: OpenAI Images 请求 - description: OpenAI Images 请求
in: body in: body
@@ -8095,6 +8109,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.OpenAIErrorEnvelope' $ref: '#/definitions/httpapi.OpenAIErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.OpenAIErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建或编辑 OpenAI Images summary: 创建或编辑 OpenAI Images
@@ -8527,8 +8545,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 EasyAI GeneratedResponse设置 X-Async=true 时返回同时兼容 Gateway description: 默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504
与 server-main EasyAIClient 的异步提交结构。 queue_timeout;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。
parameters: parameters:
- description: true 时异步创建任务并返回 202 - description: true 时异步创建任务并返回 202
in: header in: header
@@ -8583,6 +8601,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.ErrorEnvelope' $ref: '#/definitions/httpapi.ErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建 EasyAI 媒体任务 summary: 创建 EasyAI 媒体任务
@@ -9088,8 +9110,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 EasyAI GeneratedResponse设置 X-Async=true 时返回同时兼容 Gateway description: 默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504
与 server-main EasyAIClient 的异步提交结构。 queue_timeout;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。
parameters: parameters:
- description: true 时异步创建任务并返回 202 - description: true 时异步创建任务并返回 202
in: header in: header
@@ -9144,6 +9166,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.ErrorEnvelope' $ref: '#/definitions/httpapi.ErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建 EasyAI 媒体任务 summary: 创建 EasyAI 媒体任务
@@ -9153,8 +9179,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 EasyAI GeneratedResponse设置 X-Async=true 时返回同时兼容 Gateway description: 默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504
与 server-main EasyAIClient 的异步提交结构。 queue_timeout;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。
parameters: parameters:
- description: true 时异步创建任务并返回 202 - description: true 时异步创建任务并返回 202
in: header in: header
@@ -9209,6 +9235,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.ErrorEnvelope' $ref: '#/definitions/httpapi.ErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建 EasyAI 媒体任务 summary: 创建 EasyAI 媒体任务
@@ -9437,8 +9467,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 EasyAI GeneratedResponse设置 X-Async=true 时返回同时兼容 Gateway description: 默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504
与 server-main EasyAIClient 的异步提交结构。 queue_timeout;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。
parameters: parameters:
- description: true 时异步创建任务并返回 202 - description: true 时异步创建任务并返回 202
in: header in: header
@@ -9493,6 +9523,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.ErrorEnvelope' $ref: '#/definitions/httpapi.ErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建 EasyAI 媒体任务 summary: 创建 EasyAI 媒体任务
@@ -9645,8 +9679,8 @@ paths:
post: post:
consumes: consumes:
- application/json - application/json
description: 默认同步返回 EasyAI GeneratedResponse设置 X-Async=true 时返回同时兼容 Gateway description: 默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504
与 server-main EasyAIClient 的异步提交结构。 queue_timeout;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。
parameters: parameters:
- description: true 时异步创建任务并返回 202 - description: true 时异步创建任务并返回 202
in: header in: header
@@ -9701,6 +9735,10 @@ paths:
description: Bad Gateway description: Bad Gateway
schema: schema:
$ref: '#/definitions/httpapi.ErrorEnvelope' $ref: '#/definitions/httpapi.ErrorEnvelope'
"504":
description: Gateway Timeout
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security: security:
- BearerAuth: [] - BearerAuth: []
summary: 创建 EasyAI 媒体任务 summary: 创建 EasyAI 媒体任务
+39
View File
@@ -57,6 +57,8 @@ type Config struct {
GlobalHTTPProxySource string GlobalHTTPProxySource string
LogLevel slog.Level LogLevel slog.Level
BillingEngineMode string BillingEngineMode string
ProcessRole string
DatabaseMaxConns int
AsyncQueueWorkerEnabled bool AsyncQueueWorkerEnabled bool
AsyncWorkerHardLimit int AsyncWorkerHardLimit int
AsyncWorkerRefreshIntervalSeconds int AsyncWorkerRefreshIntervalSeconds int
@@ -113,6 +115,8 @@ func Load() Config {
GlobalHTTPProxySource: globalProxy.Source, GlobalHTTPProxySource: globalProxy.Source,
LogLevel: logLevel(env("LOG_LEVEL", "info")), LogLevel: logLevel(env("LOG_LEVEL", "info")),
BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")), BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")),
ProcessRole: strings.ToLower(strings.TrimSpace(env("AI_GATEWAY_PROCESS_ROLE", "all"))),
DatabaseMaxConns: envInt("AI_GATEWAY_DATABASE_MAX_CONNS", 0),
AsyncQueueWorkerEnabled: env("AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED", "true") == "true", AsyncQueueWorkerEnabled: env("AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED", "true") == "true",
AsyncWorkerHardLimit: envIntValidated("AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT", 2048), AsyncWorkerHardLimit: envIntValidated("AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT", 2048),
AsyncWorkerRefreshIntervalSeconds: envIntValidated("AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS", 5), AsyncWorkerRefreshIntervalSeconds: envIntValidated("AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS", 5),
@@ -120,6 +124,14 @@ func Load() Config {
} }
func (c Config) Validate() error { func (c Config) Validate() error {
switch strings.ToLower(strings.TrimSpace(c.ProcessRole)) {
case "", "all", "api", "worker":
default:
return errors.New("AI_GATEWAY_PROCESS_ROLE must be all, api, or worker")
}
if c.DatabaseMaxConns < 0 || c.DatabaseMaxConns > 1000 {
return errors.New("AI_GATEWAY_DATABASE_MAX_CONNS must be between 1 and 1000 when configured")
}
switch strings.ToLower(strings.TrimSpace(c.BillingEngineMode)) { switch strings.ToLower(strings.TrimSpace(c.BillingEngineMode)) {
case "", "observe", "enforce", "hold": case "", "observe", "enforce", "hold":
default: default:
@@ -184,6 +196,33 @@ func (c Config) Validate() error {
return nil return nil
} }
func (c Config) EffectiveProcessRole() string {
role := strings.ToLower(strings.TrimSpace(c.ProcessRole))
if role == "" {
return "all"
}
return role
}
func (c Config) RunsPublicHTTP() bool {
return c.EffectiveProcessRole() != "worker"
}
func (c Config) RunsAsyncExecutionWorker() bool {
switch c.EffectiveProcessRole() {
case "api":
return false
case "worker":
return true
default:
return c.AsyncQueueWorkerEnabled
}
}
func (c Config) RunsBackgroundWorkers() bool {
return c.EffectiveProcessRole() != "api"
}
type GlobalHTTPProxyStatus struct { type GlobalHTTPProxyStatus struct {
HTTPProxy string HTTPProxy string
Source string Source string
+39
View File
@@ -94,6 +94,45 @@ func TestLoadAsyncQueueWorkerEnabled(t *testing.T) {
} }
} }
func TestProcessRolePrecedenceAndCompatibility(t *testing.T) {
t.Setenv("AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED", "false")
cfg := Load()
if got := cfg.EffectiveProcessRole(); got != "all" {
t.Fatalf("legacy configuration changed process role to %q", got)
}
if cfg.RunsAsyncExecutionWorker() {
t.Fatal("legacy async worker disable was ignored")
}
t.Setenv("AI_GATEWAY_PROCESS_ROLE", "worker")
cfg = Load()
if got := cfg.EffectiveProcessRole(); got != "worker" {
t.Fatalf("effective process role = %q, want worker", got)
}
if cfg.RunsPublicHTTP() || !cfg.RunsAsyncExecutionWorker() || !cfg.RunsBackgroundWorkers() {
t.Fatalf("worker role capabilities are inconsistent: %+v", cfg)
}
t.Setenv("AI_GATEWAY_PROCESS_ROLE", "api")
cfg = Load()
if !cfg.RunsPublicHTTP() || cfg.RunsAsyncExecutionWorker() || cfg.RunsBackgroundWorkers() {
t.Fatalf("api role capabilities are inconsistent: %+v", cfg)
}
}
func TestValidateProcessRoleAndDatabasePool(t *testing.T) {
cfg := Load()
cfg.ProcessRole = "invalid"
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "PROCESS_ROLE") {
t.Fatalf("Validate() error = %v, want invalid process role", err)
}
cfg.ProcessRole = "api"
cfg.DatabaseMaxConns = -1
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "DATABASE_MAX_CONNS") {
t.Fatalf("Validate() error = %v, want invalid database max conns", err)
}
}
func TestValidateTaskHistorySettings(t *testing.T) { func TestValidateTaskHistorySettings(t *testing.T) {
cfg := Load() cfg := Load()
cfg.TaskRetentionDays = 30 cfg.TaskRetentionDays = 30
@@ -205,6 +205,7 @@ func asyncAcceptancePlatformPayload(key, name string, concurrent int, status str
"status": status, "status": status,
"rateLimitPolicy": map[string]any{"rules": []any{ "rateLimitPolicy": map[string]any{"rules": []any{
map[string]any{"metric": "concurrent", "limit": concurrent, "leaseTtlSeconds": 120}, map[string]any{"metric": "concurrent", "limit": concurrent, "leaseTtlSeconds": 120},
map[string]any{"metric": "queue_size", "limit": 512, "maxWaitSeconds": 600},
}}, }},
} }
} }
@@ -222,6 +223,7 @@ func createAsyncAcceptanceModel(t *testing.T, baseURL, token, platformID, model,
if mode == "override" { if mode == "override" {
payload["rateLimitPolicy"] = map[string]any{"rules": []any{ payload["rateLimitPolicy"] = map[string]any{"rules": []any{
map[string]any{"metric": "concurrent", "limit": concurrent, "leaseTtlSeconds": ttl}, map[string]any{"metric": "concurrent", "limit": concurrent, "leaseTtlSeconds": ttl},
map[string]any{"metric": "queue_size", "limit": 512, "maxWaitSeconds": 600},
}} }}
} }
var response struct { var response struct {
@@ -2085,6 +2085,15 @@ func assertRuntimeRecoveryReleasesPendingRateReservations(t *testing.T, ctx cont
}}); err != nil { }}); err != nil {
t.Fatalf("reserve recovery rate limit: %v", err) t.Fatalf("reserve recovery rate limit: %v", err)
} }
if _, err := db.Pool().Exec(ctx, `
UPDATE gateway_tasks
SET status = 'failed',
error_code = 'integration_test',
error_message = 'terminal recovery fixture',
finished_at = now()
WHERE id = $1::uuid`, task.ID); err != nil {
t.Fatalf("mark recovery reservation task terminal: %v", err)
}
recovery, err := db.RecoverInterruptedRuntimeState(ctx) recovery, err := db.RecoverInterruptedRuntimeState(ctx)
if err != nil { if err != nil {
t.Fatalf("recover interrupted runtime state with pending reservation: %v", err) t.Fatalf("recover interrupted runtime state with pending reservation: %v", err)
+33 -18
View File
@@ -345,6 +345,10 @@ func (s *Server) createPlatform(w http.ResponseWriter, r *http.Request) {
input.Config = config input.Config = config
platform, err := s.store.CreatePlatform(r.Context(), input) platform, err := s.store.CreatePlatform(r.Context(), input)
if err != nil { if err != nil {
if store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return
}
s.logger.Error("create platform failed", "error", err) s.logger.Error("create platform failed", "error", err)
writeError(w, http.StatusInternalServerError, "create platform failed") writeError(w, http.StatusInternalServerError, "create platform failed")
return return
@@ -398,6 +402,10 @@ func (s *Server) updatePlatform(w http.ResponseWriter, r *http.Request) {
input.Config = config input.Config = config
platform, err := s.store.UpdatePlatform(r.Context(), r.PathValue("platformID"), input) platform, err := s.store.UpdatePlatform(r.Context(), r.PathValue("platformID"), input)
if err != nil { if err != nil {
if store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return
}
if store.IsNotFound(err) { if store.IsNotFound(err) {
writeError(w, http.StatusNotFound, "platform not found") writeError(w, http.StatusNotFound, "platform not found")
return return
@@ -471,7 +479,7 @@ func (s *Server) createPlatformModel(w http.ResponseWriter, r *http.Request) {
} }
model, err := s.store.CreatePlatformModel(r.Context(), input) model, err := s.store.CreatePlatformModel(r.Context(), input)
if err != nil { if err != nil {
if errors.Is(err, store.ErrInvalidPlatformModelConfiguration) { if errors.Is(err, store.ErrInvalidPlatformModelConfiguration) || store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter") writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return return
} }
@@ -519,7 +527,7 @@ func (s *Server) replacePlatformModels(w http.ResponseWriter, r *http.Request) {
models, err := s.store.ReplacePlatformModels(r.Context(), platformID, input.Models) models, err := s.store.ReplacePlatformModels(r.Context(), platformID, input.Models)
if err != nil { if err != nil {
if errors.Is(err, store.ErrInvalidPlatformModelConfiguration) { if errors.Is(err, store.ErrInvalidPlatformModelConfiguration) || store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter") writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return return
} }
@@ -1206,8 +1214,9 @@ func (s *Server) createTask(kind string, compatible bool) http.Handler {
return return
} }
if responsePlan.asyncMode { if responsePlan.asyncMode {
if err := s.runner.EnqueueAsyncTask(r.Context(), task); err != nil { if err := s.runner.SubmitAsyncTask(r.Context(), task); err != nil {
writeTaskError(http.StatusInternalServerError, err.Error(), nil, "enqueue_failed") applyRunErrorHeaders(w, err)
writeTaskError(statusFromRunError(err), runErrorMessage(err), runErrorDetails(err), runErrorCode(err))
return return
} }
writeTaskAccepted(w, task) writeTaskAccepted(w, task)
@@ -1307,7 +1316,7 @@ func openAIEmbeddingsDoc() {}
// openAIImagesDoc godoc // openAIImagesDoc godoc
// @Summary 创建或编辑 OpenAI Images // @Summary 创建或编辑 OpenAI Images
// @Description 默认同步返回 OpenAI-compatible Images 响应;设置 X-Async=true 时异步创建任务并返回 202。 // @Description 默认同步返回 OpenAI-compatible Images 响应;非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时异步创建任务并返回 202。
// @Tags images // @Tags images
// @Accept json // @Accept json
// @Produce json // @Produce json
@@ -1319,6 +1328,7 @@ func openAIEmbeddingsDoc() {}
// @Failure 400 {object} OpenAIErrorEnvelope // @Failure 400 {object} OpenAIErrorEnvelope
// @Failure 401 {object} OpenAIErrorEnvelope // @Failure 401 {object} OpenAIErrorEnvelope
// @Failure 429 {object} OpenAIErrorEnvelope // @Failure 429 {object} OpenAIErrorEnvelope
// @Failure 504 {object} OpenAIErrorEnvelope
// @Failure 502 {object} OpenAIErrorEnvelope // @Failure 502 {object} OpenAIErrorEnvelope
// @Router /api/v1/images/generations [post] // @Router /api/v1/images/generations [post]
// @Router /api/v1/images/edits [post] // @Router /api/v1/images/edits [post]
@@ -1326,7 +1336,7 @@ func openAIImagesDoc() {}
// easyAIMediaTasksDoc godoc // easyAIMediaTasksDoc godoc
// @Summary 创建 EasyAI 媒体任务 // @Summary 创建 EasyAI 媒体任务
// @Description 默认同步返回 EasyAI GeneratedResponse;设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。 // @Description 默认同步返回 EasyAI GeneratedResponse非文本模型并发饱和且配置队列时等待准入,超过最长等待返回 504 queue_timeout设置 X-Async=true 时返回同时兼容 Gateway 与 server-main EasyAIClient 的异步提交结构。
// @Tags media // @Tags media
// @Accept json // @Accept json
// @Produce json // @Produce json
@@ -1342,6 +1352,7 @@ func openAIImagesDoc() {}
// @Failure 403 {object} ErrorEnvelope // @Failure 403 {object} ErrorEnvelope
// @Failure 404 {object} ErrorEnvelope // @Failure 404 {object} ErrorEnvelope
// @Failure 429 {object} ErrorEnvelope // @Failure 429 {object} ErrorEnvelope
// @Failure 504 {object} ErrorEnvelope
// @Failure 502 {object} ErrorEnvelope // @Failure 502 {object} ErrorEnvelope
// @Router /api/v1/videos/generations [post] // @Router /api/v1/videos/generations [post]
// @Router /api/v1/song/generations [post] // @Router /api/v1/song/generations [post]
@@ -1351,19 +1362,10 @@ func openAIImagesDoc() {}
func easyAIMediaTasksDoc() {} func easyAIMediaTasksDoc() {}
func (s *Server) requestExecutionContext(r *http.Request) (context.Context, context.CancelFunc) { func (s *Server) requestExecutionContext(r *http.Request) (context.Context, context.CancelFunc) {
base := context.WithoutCancel(r.Context())
if s.ctx == nil { if s.ctx == nil {
return base, func() {} return runner.NewRequestExecutionContext(r.Context(), context.Background())
} }
ctx, cancel := context.WithCancel(base) return runner.NewRequestExecutionContext(r.Context(), s.ctx)
go func() {
select {
case <-s.ctx.Done():
cancel()
case <-ctx.Done():
}
}()
return ctx, cancel
} }
func requestStillConnected(r *http.Request) bool { func requestStillConnected(r *http.Request) bool {
@@ -1663,6 +1665,8 @@ func statusFromRunError(err error) int {
return http.StatusNotFound return http.StatusNotFound
case store.ModelCandidateErrorCode(err) == "platform_cooling_down" || store.ModelCandidateErrorCode(err) == "model_cooling_down": case store.ModelCandidateErrorCode(err) == "platform_cooling_down" || store.ModelCandidateErrorCode(err) == "model_cooling_down":
return http.StatusTooManyRequests return http.StatusTooManyRequests
case errors.Is(err, store.ErrQueueTimeout) || clients.ErrorCode(err) == "queue_timeout":
return http.StatusGatewayTimeout
case errors.Is(err, store.ErrNoModelCandidate): case errors.Is(err, store.ErrNoModelCandidate):
return http.StatusNotFound return http.StatusNotFound
case errors.Is(err, store.ErrRateLimited): case errors.Is(err, store.ErrRateLimited):
@@ -1710,7 +1714,11 @@ func runErrorDetails(err error) map[string]any {
} }
func applyRunErrorHeaders(w http.ResponseWriter, err error) { func applyRunErrorHeaders(w http.ResponseWriter, err error) {
if retryAfter := store.ModelCandidateRetryAfter(err); retryAfter > 0 { retryAfter := store.ModelCandidateRetryAfter(err)
if limitRetryAfter := store.RateLimitRetryAfter(err); limitRetryAfter > 0 {
retryAfter = limitRetryAfter
}
if retryAfter > 0 {
seconds := int((retryAfter + time.Second - 1) / time.Second) seconds := int((retryAfter + time.Second - 1) / time.Second)
if seconds < 1 { if seconds < 1 {
seconds = 1 seconds = 1
@@ -1783,6 +1791,13 @@ func rateLimitErrorDetail(err error) map[string]any {
"limit": limitErr.Limit, "limit": limitErr.Limit,
}, },
} }
if limitErr.Reason != "" {
detail["reason"] = limitErr.Reason
}
if limitErr.QueueLimit > 0 {
detail["queueDepth"] = limitErr.QueueDepth
detail["queueLimit"] = limitErr.QueueLimit
}
if limitErr.RetryAfter > 0 { if limitErr.RetryAfter > 0 {
detail["retryAfterMs"] = limitErr.RetryAfter.Milliseconds() detail["retryAfterMs"] = limitErr.RetryAfter.Milliseconds()
} }
@@ -259,6 +259,10 @@ func (s *Server) createUserGroup(w http.ResponseWriter, r *http.Request) {
} }
item, err := s.store.CreateUserGroup(r.Context(), input) item, err := s.store.CreateUserGroup(r.Context(), input)
if err != nil { if err != nil {
if store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return
}
if store.IsUniqueViolation(err) { if store.IsUniqueViolation(err) {
writeError(w, http.StatusConflict, "user group key already exists") writeError(w, http.StatusConflict, "user group key already exists")
return return
@@ -299,6 +303,10 @@ func (s *Server) updateUserGroup(w http.ResponseWriter, r *http.Request) {
} }
item, err := s.store.UpdateUserGroup(r.Context(), r.PathValue("groupID"), input) item, err := s.store.UpdateUserGroup(r.Context(), r.PathValue("groupID"), input)
if err != nil { if err != nil {
if store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return
}
if store.IsNotFound(err) { if store.IsNotFound(err) {
writeError(w, http.StatusNotFound, "user group not found") writeError(w, http.StatusNotFound, "user group not found")
return return
@@ -192,9 +192,10 @@ func (s *Server) createKelingOmniVideo(w http.ResponseWriter, r *http.Request) {
writeKelingCompatError(w, requestID, newKelingCompatError(http.StatusInternalServerError, 5000, "create compatibility metadata failed")) writeKelingCompatError(w, requestID, newKelingCompatError(http.StatusInternalServerError, 5000, "create compatibility metadata failed"))
return return
} }
if err := s.runner.EnqueueAsyncTask(r.Context(), task); err != nil { if err := s.runner.SubmitAsyncTask(r.Context(), task); err != nil {
s.logger.Error("enqueue Kling-compatible task failed", "taskId", task.ID, "error", err) s.logger.Error("enqueue Kling-compatible task failed", "taskId", task.ID, "error", err)
writeKelingCompatError(w, requestID, newKelingCompatError(http.StatusServiceUnavailable, 5001, "video task queue is unavailable")) applyRunErrorHeaders(w, err)
writeKelingCompatError(w, requestID, kelingCompatGatewayError(err))
return return
} }
task, createErr = s.waitForCompatibilitySubmission(r, task) task, createErr = s.waitForCompatibilitySubmission(r, task)
+3 -2
View File
@@ -162,8 +162,9 @@ func (s *Server) createKlingCompatTask(w http.ResponseWriter, r *http.Request, v
writeKlingCompatError(w, http.StatusInternalServerError, "create compatibility metadata failed", "task_create_failed") writeKlingCompatError(w, http.StatusInternalServerError, "create compatibility metadata failed", "task_create_failed")
return return
} }
if err := s.runner.EnqueueAsyncTask(r.Context(), task); err != nil { if err := s.runner.SubmitAsyncTask(r.Context(), task); err != nil {
writeKlingCompatError(w, http.StatusInternalServerError, err.Error(), "enqueue_failed") applyRunErrorHeaders(w, err)
writeKlingCompatError(w, statusFromRunError(err), runErrorMessage(err), runErrorCode(err))
return return
} }
task, err = s.waitForCompatibilitySubmission(r, task) task, err = s.waitForCompatibilitySubmission(r, task)
@@ -1,6 +1,8 @@
package httpapi package httpapi
import ( import (
"net/http"
"net/http/httptest"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -48,6 +50,40 @@ func TestRateLimitErrorDetailIncludesUserGroupAndExceededMetric(t *testing.T) {
} }
} }
func TestQueueFullAndTimeoutPublicContracts(t *testing.T) {
queueErr := &store.RateLimitExceededError{
ScopeType: "platform_model",
ScopeKey: "model-1",
Metric: "queue_size",
Reason: "queue_full",
Limit: 2,
Current: 2,
Projected: 3,
QueueDepth: 2,
QueueLimit: 2,
RetryAfter: 1500 * time.Millisecond,
Retryable: true,
}
recorder := httptest.NewRecorder()
applyRunErrorHeaders(recorder, queueErr)
if statusFromRunError(queueErr) != http.StatusTooManyRequests || runErrorCode(queueErr) != "rate_limit" {
t.Fatalf("queue full contract status=%d code=%s", statusFromRunError(queueErr), runErrorCode(queueErr))
}
if recorder.Header().Get("Retry-After") != "2" {
t.Fatalf("Retry-After = %q, want 2", recorder.Header().Get("Retry-After"))
}
details := runErrorDetails(queueErr)
detail, _ := details["rateLimit"].(map[string]any)
if detail["reason"] != "queue_full" || detail["queueDepth"] != 2 || detail["queueLimit"] != 2 {
t.Fatalf("queue full details = %+v", details)
}
timeoutErr := &store.QueueTimeoutError{TaskID: "task-1"}
if statusFromRunError(timeoutErr) != http.StatusGatewayTimeout || runErrorCode(timeoutErr) != "queue_timeout" {
t.Fatalf("queue timeout contract status=%d code=%s", statusFromRunError(timeoutErr), runErrorCode(timeoutErr))
}
}
func TestRunErrorMessageIncludesRateLimitSummary(t *testing.T) { func TestRunErrorMessageIncludesRateLimitSummary(t *testing.T) {
message := runErrorMessage(&store.RateLimitExceededError{ message := runErrorMessage(&store.RateLimitExceededError{
ScopeType: "user_group", ScopeType: "user_group",
@@ -276,6 +276,10 @@ func (s *Server) createRuntimePolicySet(w http.ResponseWriter, r *http.Request)
} }
item, err := s.store.CreateRuntimePolicySet(r.Context(), input) item, err := s.store.CreateRuntimePolicySet(r.Context(), input)
if err != nil { if err != nil {
if store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return
}
if store.IsUniqueViolation(err) { if store.IsUniqueViolation(err) {
writeError(w, http.StatusConflict, "runtime policy key already exists") writeError(w, http.StatusConflict, "runtime policy key already exists")
return return
@@ -316,6 +320,10 @@ func (s *Server) updateRuntimePolicySet(w http.ResponseWriter, r *http.Request)
} }
item, err := s.store.UpdateRuntimePolicySet(r.Context(), r.PathValue("policySetID"), input) item, err := s.store.UpdateRuntimePolicySet(r.Context(), r.PathValue("policySetID"), input)
if err != nil { if err != nil {
if store.IsInvalidRateLimitPolicy(err) {
writeError(w, http.StatusBadRequest, err.Error(), "invalid_parameter")
return
}
if store.IsNotFound(err) { if store.IsNotFound(err) {
writeError(w, http.StatusNotFound, "runtime policy set not found") writeError(w, http.StatusNotFound, "runtime policy set not found")
return return
+11 -5
View File
@@ -120,7 +120,8 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
return server.identityRuntime.LegacyJWTEnabled() return server.identityRuntime.LegacyJWTEnabled()
} }
server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey
if cfg.AsyncQueueWorkerEnabled { server.runner.StartAdmissionNotifier(ctx)
if cfg.RunsAsyncExecutionWorker() {
server.runner.StartAsyncQueueWorker(ctx) server.runner.StartAsyncQueueWorker(ctx)
} else { } else {
server.runner.StartAsyncQueueClient(ctx) server.runner.StartAsyncQueueClient(ctx)
@@ -128,10 +129,12 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
logger.Info("asynchronous queue worker disabled for this process") logger.Info("asynchronous queue worker disabled for this process")
} }
} }
server.runner.StartBillingSettlementWorker(ctx) if cfg.RunsBackgroundWorkers() {
server.runner.StartTaskHistoryWorkers(ctx) server.runner.StartBillingSettlementWorker(ctx)
server.startLocalTempAssetCleanup(ctx) server.runner.StartTaskHistoryWorkers(ctx)
server.startOIDCSessionCleanup(ctx) server.startLocalTempAssetCleanup(ctx)
server.startOIDCSessionCleanup(ctx)
}
mux := http.NewServeMux() mux := http.NewServeMux()
mux.HandleFunc("GET /healthz", server.health) mux.HandleFunc("GET /healthz", server.health)
@@ -139,6 +142,9 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
mux.HandleFunc("GET /api/v1/healthz", server.health) mux.HandleFunc("GET /api/v1/healthz", server.health)
mux.HandleFunc("GET /api/v1/readyz", server.ready) mux.HandleFunc("GET /api/v1/readyz", server.ready)
mux.Handle("GET /metrics", securityEventMetrics.DynamicHandler(db)) mux.Handle("GET /metrics", securityEventMetrics.DynamicHandler(db))
if !cfg.RunsPublicHTTP() {
return server.recover(mux)
}
mux.HandleFunc("GET /static/simulation/{asset}", serveSimulationAsset) mux.HandleFunc("GET /static/simulation/{asset}", serveSimulationAsset)
mux.HandleFunc("GET /static/generated/{asset}", server.serveGeneratedStaticAsset) mux.HandleFunc("GET /static/generated/{asset}", server.serveGeneratedStaticAsset)
mux.HandleFunc("GET /static/uploaded/{asset}", server.serveUploadedStaticAsset) mux.HandleFunc("GET /static/uploaded/{asset}", server.serveUploadedStaticAsset)
@@ -236,8 +236,8 @@ func (s *Server) createVolcesCompatibleTask(r *http.Request, user *auth.User, bo
return store.GatewayTask{}, err return store.GatewayTask{}, err
} }
task.CompatibilityProtocol = clients.ProtocolVolcesContents task.CompatibilityProtocol = clients.ProtocolVolcesContents
if err := s.runner.EnqueueAsyncTask(r.Context(), task); err != nil { if err := s.runner.SubmitAsyncTask(r.Context(), task); err != nil {
return store.GatewayTask{}, &clients.ClientError{Code: "enqueue_failed", Message: err.Error(), StatusCode: http.StatusServiceUnavailable, Retryable: true} return store.GatewayTask{}, err
} }
return s.waitForCompatibilitySubmission(r, task) return s.waitForCompatibilitySubmission(r, task)
} }
+432
View File
@@ -0,0 +1,432 @@
package runner
import (
"context"
"errors"
"strings"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
)
type taskAdmissionPlan struct {
Candidate store.RuntimeModelCandidate
Body map[string]any
ModelType string
GroupID string
Scopes []store.AdmissionScope
Eligible bool
}
func distributedAdmissionModelType(modelType string) bool {
switch strings.ToLower(strings.TrimSpace(modelType)) {
case "image_generate",
"image_edit",
"image_analysis",
"image_vectorize",
"video_generate",
"video_enhance",
"image_to_video",
"text_to_video",
"video_edit",
"video_reference",
"video_first_last_frame",
"video_understanding",
"omni_video",
"omni",
"audio_generate",
"audio_understanding",
"text_to_speech",
"voice_clone":
return true
default:
return false
}
}
func (s *Service) buildTaskAdmissionPlan(ctx context.Context, task store.GatewayTask, user *auth.User) (taskAdmissionPlan, error) {
restoredRequest, err := s.restoreTaskRequestReferences(ctx, task)
if err != nil {
return taskAdmissionPlan{}, err
}
body := normalizeRequest(task.Kind, restoredRequest)
modelType := modelTypeFromKind(task.Kind, body)
if !distributedAdmissionModelType(modelType) {
return taskAdmissionPlan{Body: body, ModelType: modelType}, nil
}
if err := validateRequest(task.Kind, body); err != nil {
return taskAdmissionPlan{}, parameterPreprocessClientError(err)
}
body, clonedVoice, err := s.resolveClonedVoiceBinding(ctx, user, task.Kind, body)
if err != nil {
return taskAdmissionPlan{}, err
}
runnerPolicy, err := s.store.GetActiveRunnerPolicy(ctx)
if err != nil {
return taskAdmissionPlan{}, err
}
cacheAffinityKeys := buildCacheAffinityKeys(task.Kind, modelType, body)
candidates, err := s.store.ListModelCandidates(ctx, task.Model, modelType, user, store.ListModelCandidatesOptions{
CacheAffinityKey: cacheAffinityKeys.Primary,
CacheAffinityKeys: cacheAffinityKeys.Lookup,
CacheAffinityPolicy: runnerPolicy.CacheAffinityPolicy,
})
if err == nil {
candidates, err = filterCandidatesByRequestedPlatform(candidates, body)
}
if err == nil {
candidates, err = filterCandidatesByClonedVoiceBinding(candidates, clonedVoice)
}
if err == nil {
candidates, _, err = filterRuntimeCandidatesByRequest(task.Kind, task.Model, modelType, body, candidates)
}
if err == nil {
candidates, _, err = filterRuntimeCandidatesByOutputTokens(task.Kind, task.Model, modelType, body, candidates)
}
if err != nil {
return taskAdmissionPlan{}, err
}
for _, candidate := range candidates {
available, availabilityErr := s.store.RuntimeCandidateAvailable(ctx, candidate.PlatformID, candidate.PlatformModelID)
if availabilityErr != nil {
return taskAdmissionPlan{}, availabilityErr
}
if !available {
continue
}
scopes, groupID := s.admissionScopes(ctx, user, candidate)
hasConcurrentLimit := false
for _, scope := range scopes {
if scope.ConcurrentLimit > 0 {
hasConcurrentLimit = true
break
}
}
if !hasConcurrentLimit {
return taskAdmissionPlan{
Candidate: candidate,
Body: body,
ModelType: modelType,
}, nil
}
if err := s.store.CheckRateLimits(ctx, s.rateLimitReservations(ctx, user, candidate, body)); err != nil {
return taskAdmissionPlan{}, err
}
return taskAdmissionPlan{
Candidate: candidate,
Body: body,
ModelType: modelType,
GroupID: groupID,
Scopes: scopes,
Eligible: true,
}, nil
}
return taskAdmissionPlan{}, store.ErrNoModelCandidate
}
func (s *Service) tryTaskAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string) (store.TaskAdmissionResult, error) {
return s.tryTaskAdmissionWithAdmittedHook(ctx, task, plan, waiterID, nil)
}
func (s *Service) tryTaskAdmissionWithAdmittedHook(
ctx context.Context,
task store.GatewayTask,
plan taskAdmissionPlan,
waiterID string,
onAdmitted func(pgx.Tx) error,
) (store.TaskAdmissionResult, error) {
mode := "sync"
if task.AsyncMode {
mode = "async"
}
input := store.TaskAdmissionInput{
TaskID: task.ID,
PlatformID: plan.Candidate.PlatformID,
PlatformModelID: plan.Candidate.PlatformModelID,
UserGroupID: plan.GroupID,
QueueKey: plan.Candidate.QueueKey,
Mode: mode,
Priority: task.Priority,
WaiterID: waiterID,
Scopes: plan.Scopes,
}
current, currentErr := s.store.GetTaskAdmission(ctx, task.ID)
if currentErr != nil && !errors.Is(currentErr, pgx.ErrNoRows) {
return store.TaskAdmissionResult{}, currentErr
}
if currentErr == nil && current.Status == "waiting" &&
(current.PlatformID != input.PlatformID || current.PlatformModelID != input.PlatformModelID || current.UserGroupID != input.UserGroupID) {
if _, rebindErr := s.store.RebindWaitingTaskAdmission(ctx, input); rebindErr != nil {
return store.TaskAdmissionResult{}, rebindErr
}
s.observeTaskAdmission("candidate_migrated")
}
var result store.TaskAdmissionResult
var err error
if onAdmitted == nil {
result, err = s.store.TryTaskAdmission(ctx, input)
} else {
result, err = s.store.TryTaskAdmissionWithAdmittedHook(ctx, input, onAdmitted)
}
if result.NewlyAdmitted {
s.observeTaskAdmission("admitted")
s.observeTaskAdmissionWait(time.Since(result.Admission.EnqueuedAt))
}
var limitErr *store.RateLimitExceededError
switch {
case errors.As(err, &limitErr) && limitErr.Reason == "queue_full":
s.observeTaskAdmission("queue_full")
case errors.Is(err, store.ErrQueueTimeout):
s.observeTaskAdmission("timeout")
}
return result, err
}
func (s *Service) waitForSynchronousAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan) (store.TaskAdmissionResult, error) {
waiterID := uuid.NewString()
return s.waitForTaskAdmission(ctx, task, plan, waiterID)
}
func (s *Service) waitForTaskAdmission(ctx context.Context, task store.GatewayTask, plan taskAdmissionPlan, waiterID string) (store.TaskAdmissionResult, error) {
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
for {
result, err := s.tryTaskAdmission(ctx, task, plan, waiterID)
if err != nil || result.Admitted {
return result, err
}
select {
case <-ctx.Done():
_ = s.store.DeleteTaskAdmission(context.WithoutCancel(ctx), task.ID)
s.observeTaskAdmission("cancelled")
return store.TaskAdmissionResult{}, ctx.Err()
case <-s.currentAdmissionWake():
case <-ticker.C:
}
}
}
func (s *Service) ensureCandidateAdmission(
ctx context.Context,
task store.GatewayTask,
user *auth.User,
body map[string]any,
candidate store.RuntimeModelCandidate,
) (store.TaskAdmissionResult, bool, error) {
scopes, groupID := s.admissionScopes(ctx, user, candidate)
hasConcurrentLimit := false
for _, scope := range scopes {
if scope.ConcurrentLimit > 0 {
hasConcurrentLimit = true
break
}
}
if !hasConcurrentLimit {
if err := s.store.DeleteTaskAdmission(context.WithoutCancel(ctx), task.ID); err != nil && !errors.Is(err, pgx.ErrNoRows) {
return store.TaskAdmissionResult{}, false, err
}
return store.TaskAdmissionResult{}, false, nil
}
if err := s.store.CheckRateLimits(ctx, s.rateLimitReservations(ctx, user, candidate, body)); err != nil {
return store.TaskAdmissionResult{}, true, err
}
plan := taskAdmissionPlan{
Candidate: candidate,
Body: body,
ModelType: candidate.ModelType,
GroupID: groupID,
Scopes: scopes,
Eligible: true,
}
waiterID := ""
if !task.AsyncMode {
waiterID = uuid.NewString()
}
result, err := s.waitForTaskAdmission(ctx, task, plan, waiterID)
return result, true, err
}
func (s *Service) observeTaskAdmission(event string) {
observer, ok := s.billingMetrics.(interface {
ObserveTaskAdmission(string)
})
if ok {
observer.ObserveTaskAdmission(event)
}
}
func (s *Service) observeTaskAdmissionWait(wait time.Duration) {
observer, ok := s.billingMetrics.(interface {
ObserveTaskAdmissionWait(time.Duration)
})
if ok {
observer.ObserveTaskAdmissionWait(wait)
}
}
func (s *Service) StartAdmissionNotifier(ctx context.Context) {
s.admissionListener.Do(func() {
go func() {
for ctx.Err() == nil {
err := s.store.ListenTaskAdmissionNotifications(ctx, func(string) {
s.broadcastAdmissionWake()
})
if ctx.Err() != nil {
return
}
if s.logger != nil {
s.logger.Warn("task admission LISTEN connection interrupted; periodic polling remains active", "error", err)
}
timer := time.NewTimer(time.Second)
select {
case <-ctx.Done():
timer.Stop()
return
case <-timer.C:
}
}
}()
})
}
func (s *Service) currentAdmissionWake() <-chan struct{} {
s.admissionWakeMu.Lock()
defer s.admissionWakeMu.Unlock()
return s.admissionWake
}
func (s *Service) broadcastAdmissionWake() {
s.admissionWakeMu.Lock()
close(s.admissionWake)
s.admissionWake = make(chan struct{})
s.admissionWakeMu.Unlock()
}
func (s *Service) SubmitAsyncTask(ctx context.Context, task store.GatewayTask) error {
if !task.AsyncMode {
return errors.New("only asynchronous tasks can be submitted to the async queue")
}
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
user := authUserFromTask(task)
plan, err := s.buildTaskAdmissionPlan(ctx, task, user)
if err != nil {
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
_, _ = s.store.FailQueuedTask(context.WithoutCancel(ctx), task.ID, clients.ErrorCode(err), err.Error())
return err
}
if !plan.Eligible {
if err := s.EnqueueAsyncTask(ctx, task); err != nil {
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
_, _ = s.store.FailQueuedTask(context.WithoutCancel(ctx), task.ID, "enqueue_failed", err.Error())
return err
}
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
return nil
}
result, err := s.tryTaskAdmissionWithAdmittedHook(ctx, task, plan, "", func(tx pgx.Tx) error {
return s.enqueueAsyncTaskTx(ctx, tx, task.ID, asyncTaskInsertOpts(task))
})
if err != nil {
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
_, _ = s.store.FailQueuedTask(context.WithoutCancel(ctx), task.ID, clients.ErrorCode(err), err.Error())
return err
}
if !result.Admitted {
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
return nil
}
if s.cancelAsyncSubmissionIfDisconnected(ctx, task.ID) {
return ctx.Err()
}
return nil
}
func (s *Service) cancelAsyncSubmissionIfDisconnected(ctx context.Context, taskID string) bool {
if ctx.Err() == nil {
return false
}
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second)
defer cancel()
if _, changed, err := s.store.CancelQueuedTask(cleanupCtx, taskID, "client disconnected before upstream submission"); err != nil {
if s.logger != nil {
s.logger.Warn("cancel disconnected asynchronous task failed", "taskID", taskID, "error", err)
}
} else if changed {
s.observeTaskAdmission("cancelled")
}
return true
}
func (s *Service) dispatchWaitingAsyncAdmissions(ctx context.Context) {
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-s.currentAdmissionWake():
case <-ticker.C:
}
taskIDs, err := s.store.ListWaitingAsyncAdmissionTaskIDs(ctx, 100)
if err != nil {
s.logger.Warn("list waiting async admissions failed", "error", err)
continue
}
for _, taskID := range taskIDs {
task, err := s.store.GetTask(ctx, taskID)
if err != nil {
if !errors.Is(err, pgx.ErrNoRows) {
s.logger.Warn("load waiting async task failed", "taskID", taskID, "error", err)
}
continue
}
if task.Status != "queued" || task.RiverJobID > 0 {
continue
}
if err := s.SubmitAsyncTask(ctx, task); err != nil {
s.logger.Warn("dispatch waiting async admission failed", "taskID", taskID, "error", err)
}
}
}
}
func (s *Service) reapExpiredTaskAdmissions(ctx context.Context) {
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
}
result, err := s.store.ReapExpiredTaskAdmissions(ctx, 500)
if err != nil {
if s.logger != nil {
s.logger.Warn("reap expired task admissions failed", "error", err)
}
continue
}
for range result.ExpiredWaiters {
s.observeTaskAdmission("expired")
}
for range result.ExpiredDeadlines {
s.observeTaskAdmission("timeout")
}
}
}
@@ -0,0 +1,21 @@
package runner
import "testing"
func TestDistributedAdmissionModelTypeBoundary(t *testing.T) {
for _, modelType := range []string{
"image_generate", "image_edit", "image_analysis", "image_vectorize",
"video_generate", "video_enhance", "image_to_video", "text_to_video",
"video_edit", "video_reference", "video_first_last_frame", "video_understanding", "omni_video", "omni",
"audio_generate", "audio_understanding", "text_to_speech", "voice_clone",
} {
if !distributedAdmissionModelType(modelType) {
t.Errorf("%s should use distributed admission", modelType)
}
}
for _, modelType := range []string{"text_generate", "embedding", "rerank", "", "unknown"} {
if distributedAdmissionModelType(modelType) {
t.Errorf("%s should preserve immediate execution/rate-limit behavior", modelType)
}
}
}
+69 -1
View File
@@ -81,6 +81,74 @@ func (s *Service) rateLimitReservations(ctx context.Context, user *auth.User, ca
return out return out
} }
func (s *Service) admissionScopes(ctx context.Context, user *auth.User, candidate store.RuntimeModelCandidate) ([]store.AdmissionScope, string) {
scopes := make([]store.AdmissionScope, 0, 2)
platformPolicy := effectiveRateLimitPolicy(candidate)
if scope, ok := admissionScopeFromPolicy(
"platform_model",
candidate.PlatformModelID,
firstNonEmptyString(candidate.DisplayName, candidate.ModelAlias, candidate.ModelName),
map[string]any{
"platformId": candidate.PlatformID,
"platformName": candidate.PlatformName,
"modelAlias": candidate.ModelAlias,
"modelName": candidate.ModelName,
},
platformPolicy,
); ok {
scopes = append(scopes, scope)
}
groupID := ""
if group, err := s.store.ResolveUserGroupPolicy(ctx, user); err == nil && group.ID != "" {
groupPolicy := store.NormalizeRateLimitPolicy(group.RateLimitPolicy)
if scope, ok := admissionScopeFromPolicy(
"user_group",
group.ID,
firstNonEmptyString(group.Name, group.GroupKey),
map[string]any{"groupKey": group.GroupKey, "name": group.Name},
groupPolicy,
); ok {
groupID = group.ID
scopes = append(scopes, scope)
}
}
return scopes, groupID
}
func admissionScopeFromPolicy(scopeType string, scopeKey string, scopeName string, metadata map[string]any, policy map[string]any) (store.AdmissionScope, bool) {
concurrentLimit, hasConcurrent := store.RateLimitPolicyMetric(policy, "concurrent")
queueRule, hasQueue := store.QueueRuleFromPolicy(policy)
if !hasConcurrent && !hasQueue {
return store.AdmissionScope{}, false
}
scope := store.AdmissionScope{
ScopeType: scopeType,
ScopeKey: scopeKey,
ScopeName: scopeName,
ScopeMetadata: metadata,
ConcurrentLimit: concurrentLimit,
Amount: 1,
LeaseTTLSeconds: 120,
Policy: policy,
}
if hasQueue {
scope.QueueLimit = queueRule.Limit
scope.MaxWaitSeconds = queueRule.MaxWaitSeconds
}
rules, _ := store.NormalizeRateLimitPolicy(policy)["rules"].([]any)
for _, rawRule := range rules {
rule, _ := rawRule.(map[string]any)
if strings.TrimSpace(stringFromMap(rule, "metric")) != "concurrent" {
continue
}
if ttl := int(floatFromAny(rule["leaseTtlSeconds"])); ttl > 0 {
scope.LeaseTTLSeconds = ttl
}
break
}
return scope, true
}
func effectiveRateLimitPolicy(candidate store.RuntimeModelCandidate) map[string]any { func effectiveRateLimitPolicy(candidate store.RuntimeModelCandidate) map[string]any {
return store.EffectiveRateLimitPolicy(store.EffectiveRateLimitPolicyInput{ return store.EffectiveRateLimitPolicy(store.EffectiveRateLimitPolicyInput{
BasePolicy: candidate.BaseRateLimitPolicy, BasePolicy: candidate.BaseRateLimitPolicy,
@@ -118,7 +186,7 @@ func reservationsFromPolicy(scopeType string, scopeKey string, scopeName string,
rule, _ := rawRule.(map[string]any) rule, _ := rawRule.(map[string]any)
metric := strings.TrimSpace(stringFromMap(rule, "metric")) metric := strings.TrimSpace(stringFromMap(rule, "metric"))
limit := floatFromAny(rule["limit"]) limit := floatFromAny(rule["limit"])
if metric == "" || limit <= 0 { if metric == "" || metric == "queue_size" || limit <= 0 {
continue continue
} }
amount := 1.0 amount := 1.0
+105 -42
View File
@@ -102,7 +102,7 @@ func (s *Service) startRiverQueue(ctx context.Context, workerEnabled bool) error
} }
controlClient, err := river.NewClient(driver, &river.Config{ controlClient, err := river.NewClient(driver, &river.Config{
ID: asyncWorkerID() + "-control", ID: s.workerInstanceID + "-control",
Logger: s.logger, Logger: s.logger,
TestOnly: s.cfg.AppEnv == "test", TestOnly: s.cfg.AppEnv == "test",
}) })
@@ -130,15 +130,18 @@ func (s *Service) startRiverQueue(ctx context.Context, workerEnabled bool) error
if err != nil { if err != nil {
return fmt.Errorf("calculate initial async worker capacity: %w", err) return fmt.Errorf("calculate initial async worker capacity: %w", err)
} }
executionClient, err := s.makeAsyncExecutionClient(snapshot.Capacity) var executionClient asyncExecutionClient
if err != nil { if snapshot.Capacity > 0 {
return err executionClient, err = s.makeAsyncExecutionClient(snapshot.Capacity)
} if err != nil {
if err := executionClient.Start(ctx); err != nil { return err
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) }
_ = executionClient.StopAndCancel(cleanupCtx) if err := executionClient.Start(ctx); err != nil {
cancel() cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
return err _ = executionClient.StopAndCancel(cleanupCtx)
cancel()
return err
}
} }
s.riverMu.Lock() s.riverMu.Lock()
s.riverControlClient = controlClient s.riverControlClient = controlClient
@@ -161,11 +164,15 @@ func (s *Service) startRiverQueue(ctx context.Context, workerEnabled bool) error
) )
if err := s.recoverAsyncRiverJobs(ctx); err != nil { if err := s.recoverAsyncRiverJobs(ctx); err != nil {
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
_ = executionClient.StopAndCancel(cleanupCtx) if executionClient != nil {
_ = executionClient.StopAndCancel(cleanupCtx)
}
cancel() cancel()
return err return err
} }
go s.refreshAsyncWorkerCapacity(ctx) go s.refreshAsyncWorkerCapacity(ctx)
go s.dispatchWaitingAsyncAdmissions(ctx)
go s.reapExpiredTaskAdmissions(ctx)
go s.stopAsyncWorkersOnShutdown(ctx) go s.stopAsyncWorkersOnShutdown(ctx)
return nil return nil
} }
@@ -176,7 +183,7 @@ func (s *Service) newRiverAsyncExecutionClient(capacity int) (*river.Client[pgx.
return nil, err return nil, err
} }
return river.NewClient(riverpgxv5.New(s.store.Pool()), &river.Config{ return river.NewClient(riverpgxv5.New(s.store.Pool()), &river.Config{
ID: fmt.Sprintf("%s-exec-%d-%d", asyncWorkerID(), capacity, time.Now().UnixNano()), ID: fmt.Sprintf("%s-exec-%d-%d", s.workerInstanceID, capacity, time.Now().UnixNano()),
JobTimeout: -1, JobTimeout: -1,
Logger: s.logger, Logger: s.logger,
CompletedJobRetentionPeriod: 24 * time.Hour, CompletedJobRetentionPeriod: 24 * time.Hour,
@@ -203,7 +210,26 @@ func (s *Service) loadAsyncWorkerCapacity(ctx context.Context) (store.AsyncWorke
if s.asyncCapacityLoader != nil { if s.asyncCapacityLoader != nil {
return s.asyncCapacityLoader(ctx, s.cfg.AsyncWorkerHardLimit) return s.asyncCapacityLoader(ctx, s.cfg.AsyncWorkerHardLimit)
} }
return s.store.AsyncWorkerCapacity(ctx, s.cfg.AsyncWorkerHardLimit) snapshot, err := s.store.AsyncWorkerCapacity(ctx, s.cfg.AsyncWorkerHardLimit)
if err != nil {
return store.AsyncWorkerCapacitySnapshot{}, err
}
allocation, err := s.store.RegisterWorkerInstance(ctx, store.WorkerRegistrationInput{
InstanceID: s.workerInstanceID,
PodUID: strings.TrimSpace(os.Getenv("POD_UID")),
PodName: strings.TrimSpace(os.Getenv("POD_NAME")),
Site: strings.TrimSpace(os.Getenv("EASYAI_SITE")),
Revision: strings.TrimSpace(os.Getenv("AI_GATEWAY_REVISION")),
DesiredCapacity: snapshot.Capacity,
})
if err != nil {
return store.AsyncWorkerCapacitySnapshot{}, err
}
snapshot.Capacity = allocation.Allocated
snapshot.GlobalCapacity = allocation.DesiredCapacity
snapshot.ActiveInstances = allocation.ActiveInstances
snapshot.InstanceID = allocation.InstanceID
return snapshot, nil
} }
func (s *Service) refreshAsyncWorkerCapacity(ctx context.Context) { func (s *Service) refreshAsyncWorkerCapacity(ctx context.Context) {
@@ -233,25 +259,30 @@ func (s *Service) resizeAsyncWorkerCapacity(ctx context.Context) {
if snapshot.Capacity == currentCapacity { if snapshot.Capacity == currentCapacity {
return return
} }
newClient, err := s.makeAsyncExecutionClient(snapshot.Capacity) var newClient asyncExecutionClient
if err != nil { if snapshot.Capacity > 0 {
s.observeAsyncWorkerResize("create_failed") newClient, err = s.makeAsyncExecutionClient(snapshot.Capacity)
s.logger.Warn("create replacement async worker client failed; keeping current client", if err != nil {
"error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity) s.observeAsyncWorkerResize("create_failed")
return s.logger.Warn("create replacement async worker client failed; keeping current client",
} "error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity)
if err := newClient.Start(ctx); err != nil { return
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) }
_ = newClient.StopAndCancel(cleanupCtx) if err := newClient.Start(ctx); err != nil {
cancel() cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
s.observeAsyncWorkerResize("start_failed") _ = newClient.StopAndCancel(cleanupCtx)
s.logger.Warn("start replacement async worker client failed; keeping current client", cancel()
"error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity) s.observeAsyncWorkerResize("start_failed")
return s.logger.Warn("start replacement async worker client failed; keeping current client",
"error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity)
return
}
} }
if ctx.Err() != nil { if ctx.Err() != nil {
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
_ = newClient.StopAndCancel(cleanupCtx) if newClient != nil {
_ = newClient.StopAndCancel(cleanupCtx)
}
cancel() cancel()
return return
} }
@@ -271,6 +302,8 @@ func (s *Service) resizeAsyncWorkerCapacity(ctx context.Context) {
"previousCapacity", currentCapacity, "previousCapacity", currentCapacity,
"capacity", snapshot.Capacity, "capacity", snapshot.Capacity,
"desiredCapacity", snapshot.Desired, "desiredCapacity", snapshot.Desired,
"activeInstances", snapshot.ActiveInstances,
"instanceID", snapshot.InstanceID,
"hardLimit", snapshot.HardLimit, "hardLimit", snapshot.HardLimit,
"modelDesired", snapshot.ModelDesired, "modelDesired", snapshot.ModelDesired,
"groupDesired", snapshot.GroupDesired, "groupDesired", snapshot.GroupDesired,
@@ -294,6 +327,13 @@ func (s *Service) drainAsyncWorkerClient(client asyncExecutionClient) {
func (s *Service) stopAsyncWorkersOnShutdown(ctx context.Context) { func (s *Service) stopAsyncWorkersOnShutdown(ctx context.Context) {
<-ctx.Done() <-ctx.Done()
if s.store != nil && strings.TrimSpace(s.workerInstanceID) != "" {
drainCtx, drainCancel := context.WithTimeout(context.Background(), 5*time.Second)
if err := s.store.MarkWorkerDraining(drainCtx, s.workerInstanceID); err != nil && !errors.Is(err, pgx.ErrNoRows) {
s.logger.Warn("mark async worker draining failed", "error", err)
}
drainCancel()
}
s.riverMu.Lock() s.riverMu.Lock()
s.riverControlClient = nil s.riverControlClient = nil
clients := make([]asyncExecutionClient, 0, 1+len(s.riverDrainingClients)) clients := make([]asyncExecutionClient, 0, 1+len(s.riverDrainingClients))
@@ -328,6 +368,12 @@ func (s *Service) observeAsyncWorkerCapacity(snapshot store.AsyncWorkerCapacityS
if ok { if ok {
observer.SetAsyncWorkerCapacity(snapshot.Capacity, snapshot.Desired, snapshot.HardLimit, snapshot.Capped) observer.SetAsyncWorkerCapacity(snapshot.Capacity, snapshot.Desired, snapshot.HardLimit, snapshot.Capped)
} }
distributedObserver, ok := s.billingMetrics.(interface {
SetDistributedWorkerCapacity(activeInstances, globalCapacity, allocatedCapacity int)
})
if ok {
distributedObserver.SetDistributedWorkerCapacity(snapshot.ActiveInstances, snapshot.GlobalCapacity, snapshot.Capacity)
}
} }
func (s *Service) observeAsyncWorkerResize(outcome string) { func (s *Service) observeAsyncWorkerResize(outcome string) {
@@ -340,16 +386,38 @@ func (s *Service) observeAsyncWorkerResize(outcome string) {
} }
func (s *Service) EnqueueAsyncTask(ctx context.Context, task store.GatewayTask) error { func (s *Service) EnqueueAsyncTask(ctx context.Context, task store.GatewayTask) error {
return s.enqueueAsyncTaskWithOptions(ctx, task.ID, asyncTaskInsertOpts(task))
}
func (s *Service) enqueueAsyncTaskWithOptions(ctx context.Context, taskID string, opts *river.InsertOpts) error {
tx, err := s.store.Pool().Begin(ctx)
if err != nil {
return err
}
defer tx.Rollback(ctx)
if err := s.enqueueAsyncTaskTx(ctx, tx, taskID, opts); err != nil {
return err
}
return tx.Commit(ctx)
}
func (s *Service) enqueueAsyncTaskTx(ctx context.Context, tx pgx.Tx, taskID string, opts *river.InsertOpts) error {
riverClient := s.asyncControlClient() riverClient := s.asyncControlClient()
if riverClient == nil { if riverClient == nil {
return errors.New("river async queue is not started") return errors.New("river async queue is not started")
} }
result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task)) result, err := riverClient.InsertTx(ctx, tx, asyncTaskArgs{TaskID: taskID}, opts)
if err != nil { if err != nil {
return err return err
} }
if result.Job != nil { if result.Job != nil {
return s.store.SetTaskRiverJobID(ctx, task.ID, result.Job.ID) if _, err := tx.Exec(ctx, `
UPDATE gateway_tasks
SET river_job_id = $2,
updated_at = now()
WHERE id = $1::uuid`, taskID, result.Job.ID); err != nil {
return err
}
} }
return nil return nil
} }
@@ -373,14 +441,13 @@ func (s *Service) recoverAsyncRiverJobs(ctx context.Context) error {
return err return err
} }
for _, item := range items { for _, item := range items {
result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item, time.Now())) if admission, admissionErr := s.store.GetTaskAdmission(ctx, item.ID); admissionErr == nil && admission.Status == "waiting" {
if err != nil { continue
return err } else if admissionErr != nil && !errors.Is(admissionErr, pgx.ErrNoRows) {
return admissionErr
} }
if result.Job != nil { if err := s.enqueueAsyncTaskWithOptions(ctx, item.ID, asyncTaskRecoveryInsertOpts(item, time.Now())); err != nil {
if err := s.store.SetTaskRiverJobID(ctx, item.ID, result.Job.ID); err != nil { return err
return err
}
} }
} }
if len(items) > 0 { if len(items) > 0 {
@@ -418,10 +485,6 @@ func asyncTaskRecoveryInsertOpts(item store.AsyncTaskQueueItem, now time.Time) *
if item.NextRunAt.After(now) { if item.NextRunAt.After(now) {
opts.ScheduledAt = item.NextRunAt opts.ScheduledAt = item.NextRunAt
} }
// A replacement process must not be blocked by a River row that the dead
// process left in running state. PostgreSQL execution leases still ensure
// that only one recovery job can call the upstream provider.
opts.UniqueOpts = river.UniqueOpts{}
return opts return opts
} }
@@ -0,0 +1,47 @@
package runner
import (
"context"
"sync/atomic"
)
type submissionCancellationGate struct {
submitted atomic.Bool
}
type submissionCancellationGateKey struct{}
func NewRequestExecutionContext(requestContext context.Context, serverContext context.Context) (context.Context, context.CancelFunc) {
base, cancel := context.WithCancel(context.WithoutCancel(requestContext))
gate := &submissionCancellationGate{}
ctx := context.WithValue(base, submissionCancellationGateKey{}, gate)
go func() {
select {
case <-serverContext.Done():
cancel()
case <-requestContext.Done():
if !gate.submitted.Load() {
cancel()
return
}
select {
case <-serverContext.Done():
cancel()
case <-base.Done():
}
case <-base.Done():
}
}()
return ctx, cancel
}
func markUpstreamSubmissionStarted(ctx context.Context) {
if gate, ok := ctx.Value(submissionCancellationGateKey{}).(*submissionCancellationGate); ok {
gate.submitted.Store(true)
}
}
func requestCancelledBeforeSubmission(ctx context.Context) bool {
gate, ok := ctx.Value(submissionCancellationGateKey{}).(*submissionCancellationGate)
return ok && ctx.Err() != nil && !gate.submitted.Load()
}
@@ -0,0 +1,55 @@
package runner
import (
"context"
"testing"
"time"
)
func TestRequestExecutionContextCancelsBeforeSubmission(t *testing.T) {
requestCtx, cancelRequest := context.WithCancel(context.Background())
serverCtx, cancelServer := context.WithCancel(context.Background())
defer cancelServer()
ctx, cancel := NewRequestExecutionContext(requestCtx, serverCtx)
defer cancel()
cancelRequest()
select {
case <-ctx.Done():
case <-time.After(time.Second):
t.Fatal("request context did not cancel before upstream submission")
}
}
func TestRequestExecutionContextSurvivesDisconnectAfterSubmission(t *testing.T) {
requestCtx, cancelRequest := context.WithCancel(context.Background())
serverCtx, cancelServer := context.WithCancel(context.Background())
ctx, cancel := NewRequestExecutionContext(requestCtx, serverCtx)
defer cancel()
markUpstreamSubmissionStarted(ctx)
cancelRequest()
select {
case <-ctx.Done():
t.Fatal("client disconnect propagated after upstream submission")
case <-time.After(50 * time.Millisecond):
}
cancelServer()
select {
case <-ctx.Done():
case <-time.After(time.Second):
t.Fatal("server shutdown did not cancel detached execution")
}
}
func TestRequestExecutionContextPreservesRequestValues(t *testing.T) {
type contextKey struct{}
requestCtx := context.WithValue(context.Background(), contextKey{}, "trace-value")
ctx, cancel := NewRequestExecutionContext(requestCtx, context.Background())
defer cancel()
if got := ctx.Value(contextKey{}); got != "trace-value" {
t.Fatalf("request context value = %v, want trace-value", got)
}
}
+203 -11
View File
@@ -36,7 +36,11 @@ type Service struct {
riverExecutionClient asyncExecutionClient riverExecutionClient asyncExecutionClient
riverDrainingClients map[asyncExecutionClient]struct{} riverDrainingClients map[asyncExecutionClient]struct{}
riverWorkerCapacity int riverWorkerCapacity int
workerInstanceID string
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
admissionWakeMu sync.Mutex
admissionWake chan struct{}
admissionListener sync.Once
asyncClientFactory func(int) (asyncExecutionClient, error) asyncClientFactory func(int) (asyncExecutionClient, error)
billingMetrics billingMetricsObserver billingMetrics billingMetricsObserver
} }
@@ -135,7 +139,9 @@ func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...b
"universal": clients.UniversalClient{HTTPClient: httpClients.none, ScriptExecutor: scriptExecutor}, "universal": clients.UniversalClient{HTTPClient: httpClients.none, ScriptExecutor: scriptExecutor},
"simulation": clients.SimulationClient{}, "simulation": clients.SimulationClient{},
}, },
httpClients: httpClients, httpClients: httpClients,
workerInstanceID: asyncWorkerID(),
admissionWake: make(chan struct{}),
} }
if len(observers) > 0 { if len(observers) > 0 {
service.billingMetrics = observers[0] service.billingMetrics = observers[0]
@@ -172,7 +178,15 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask, user *auth.User, onDelta clients.StreamDelta, executionToken string) (Result, error) { func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask, user *auth.User, onDelta clients.StreamDelta, executionToken string) (Result, error) {
wasRunning := task.Status == "running" wasRunning := task.Status == "running"
claimed, err := s.store.ClaimTaskExecution(ctx, task.ID, executionToken, taskExecutionLeaseTTL) preliminaryModelType := modelTypeFromKind(task.Kind, normalizeRequest(task.Kind, task.Request))
distributedAdmission := distributedAdmissionModelType(preliminaryModelType)
var claimed store.GatewayTask
var err error
if distributedAdmission && !wasRunning {
claimed, err = s.store.ClaimTaskPreparation(ctx, task.ID, executionToken, taskExecutionLeaseTTL)
} else {
claimed, err = s.store.ClaimTaskExecution(ctx, task.ID, executionToken, taskExecutionLeaseTTL)
}
if err != nil { if err != nil {
return Result{}, err return Result{}, err
} }
@@ -181,6 +195,21 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask,
defer stopExecution() defer stopExecution()
go s.renewTaskExecutionLease(executionCtx, stopExecution, task.ID, task.ExecutionToken) go s.renewTaskExecutionLease(executionCtx, stopExecution, task.ID, task.ExecutionToken)
ctx = executionCtx ctx = executionCtx
defer func() {
if !requestCancelledBeforeSubmission(ctx) {
return
}
cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second)
defer cancel()
if _, _, cleanupErr := s.store.CancelTaskBeforeUpstreamSubmission(
cleanupCtx,
task.ID,
task.ExecutionToken,
"client disconnected before upstream submission",
); cleanupErr != nil && s.logger != nil {
s.logger.Warn("cancel disconnected task before upstream submission failed", "taskID", task.ID, "error", cleanupErr)
}
}()
executeStartedAt := time.Now() executeStartedAt := time.Now()
restoredRequest, err := s.restoreTaskRequestReferences(ctx, task) restoredRequest, err := s.restoreTaskRequestReferences(ctx, task)
if err != nil { if err != nil {
@@ -189,13 +218,16 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask,
body := normalizeRequest(task.Kind, restoredRequest) body := normalizeRequest(task.Kind, restoredRequest)
responseExecution := responseExecutionContext{} responseExecution := responseExecutionContext{}
modelType := modelTypeFromKind(task.Kind, body) modelType := modelTypeFromKind(task.Kind, body)
if err := s.store.MarkTaskRunning(ctx, task.ID, task.ExecutionToken, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil { distributedAdmission = distributedAdmissionModelType(modelType)
return Result{}, err if !distributedAdmission {
} if err := s.store.MarkTaskRunning(ctx, task.ID, task.ExecutionToken, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil {
if !wasRunning {
if err := s.emit(ctx, task.ID, "task.running", "running", "starting", 0.12, "task pulled from queue and started", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil {
return Result{}, err return Result{}, err
} }
if !wasRunning {
if err := s.emit(ctx, task.ID, "task.running", "running", "starting", 0.12, "task pulled from queue and started", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil {
return Result{}, err
}
}
} }
if err := validateRequest(task.Kind, body); err != nil { if err := validateRequest(task.Kind, body); err != nil {
validationErr := parameterPreprocessClientError(err) validationErr := parameterPreprocessClientError(err)
@@ -477,6 +509,90 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask,
} }
} }
} }
var admittedLeases []store.ConcurrencyLease
admittedPlatformModelID := ""
if distributedAdmission && len(candidates) > 0 {
admissionScopes, groupID := s.admissionScopes(ctx, user, candidates[0])
hasConcurrentLimit := false
for _, scope := range admissionScopes {
if scope.ConcurrentLimit > 0 {
hasConcurrentLimit = true
break
}
}
if hasConcurrentLimit {
if err := s.store.CheckRateLimits(ctx, s.rateLimitReservations(ctx, user, candidates[0], body)); err != nil {
failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, clients.ErrorCode(err), err.Error(), task.RunMode == "simulation", err)
if finishErr != nil {
return Result{}, finishErr
}
return Result{Task: failed, Output: failed.Result}, err
}
plan := taskAdmissionPlan{
Candidate: candidates[0],
Body: body,
ModelType: modelType,
GroupID: groupID,
Scopes: admissionScopes,
Eligible: true,
}
var admissionResult store.TaskAdmissionResult
var admissionErr error
if task.AsyncMode {
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}
}
} else {
admissionResult, admissionErr = s.waitForSynchronousAdmission(ctx, task, plan)
}
if admissionErr != nil {
if requestCancelledBeforeSubmission(ctx) {
cancelled, changed, cancelErr := s.store.CancelTaskBeforeUpstreamSubmission(
context.WithoutCancel(ctx),
task.ID,
task.ExecutionToken,
"client disconnected before upstream submission",
)
if cancelErr != nil {
return Result{}, cancelErr
}
if changed {
return Result{Task: cancelled, Output: cancelled.Result}, admissionErr
}
}
code := clients.ErrorCode(admissionErr)
if errors.Is(admissionErr, context.Canceled) {
code = "client_disconnected"
}
failed, finishErr := s.failTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, code, admissionErr.Error(), task.RunMode == "simulation", admissionErr)
if finishErr != nil && !errors.Is(finishErr, store.ErrTaskExecutionLeaseLost) {
return Result{}, finishErr
}
return Result{Task: failed, Output: failed.Result}, admissionErr
}
admittedLeases = admissionResult.Leases
if len(admittedLeases) == 0 {
admittedLeases, err = s.store.ActiveTaskAdmissionLeases(ctx, task.ID)
if err != nil {
return Result{}, err
}
}
admittedPlatformModelID = admissionResult.Admission.PlatformModelID
}
if err := s.store.MarkTaskRunning(ctx, task.ID, task.ExecutionToken, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil {
return Result{}, err
}
if !wasRunning {
if err := s.emit(ctx, task.ID, "task.running", "running", "starting", 0.12, "task admitted and started", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil {
return Result{}, err
}
}
}
if distributedAdmission {
defer s.store.DeleteTaskAdmission(context.WithoutCancel(ctx), task.ID)
}
firstCandidateBody := body firstCandidateBody := body
normalizedModelType := modelType normalizedModelType := modelType
attemptNo := task.AttemptCount attemptNo := task.AttemptCount
@@ -579,6 +695,33 @@ candidatesLoop:
if !available { if !available {
continue continue
} }
if distributedAdmission && candidate.PlatformModelID != admittedPlatformModelID {
admissionResult, candidateLimited, admissionErr := s.ensureCandidateAdmission(ctx, task, user, body, candidate)
if admissionErr != nil {
lastErr = admissionErr
lastCandidate = candidate
if errors.Is(admissionErr, store.ErrQueueTimeout) {
break
}
if candidateLimited && errors.Is(admissionErr, store.ErrRateLimited) {
continue
}
return Result{}, admissionErr
}
if candidateLimited {
admittedPlatformModelID = admissionResult.Admission.PlatformModelID
admittedLeases = admissionResult.Leases
if len(admittedLeases) == 0 {
admittedLeases, err = s.store.ActiveTaskAdmissionLeases(ctx, task.ID)
if err != nil {
return Result{}, err
}
}
} else {
admittedPlatformModelID = ""
admittedLeases = nil
}
}
platformsVisited++ platformsVisited++
lastCandidate = candidate lastCandidate = candidate
clientAttempts := clientAttemptsForCandidate(candidate) clientAttempts := clientAttemptsForCandidate(candidate)
@@ -614,7 +757,21 @@ candidatesLoop:
} }
candidateBody := preprocessing.Body candidateBody := preprocessing.Body
candidatePricing := pricingByCandidate[pricingCandidateKey(candidate)] candidatePricing := pricingByCandidate[pricingCandidateKey(candidate)]
response, err := s.runCandidate(ctx, task, user, candidateBody, preprocessing.Log, candidate, candidatePricing, nextAttemptNo, candidateDelta, responseExecution, runnerPolicy.CacheAffinityPolicy, cacheAffinityKeys.Record) response, err := s.runCandidate(ctx, task, user, candidateBody, preprocessing.Log, candidate, candidatePricing, nextAttemptNo, candidateDelta, responseExecution, runnerPolicy.CacheAffinityPolicy, cacheAffinityKeys.Record, admittedPlatformModelID, admittedLeases)
if err != nil && requestCancelledBeforeSubmission(ctx) {
cancelled, changed, cancelErr := s.store.CancelTaskBeforeUpstreamSubmission(
context.WithoutCancel(ctx),
task.ID,
task.ExecutionToken,
"client disconnected before upstream submission",
)
if cancelErr != nil {
return Result{}, cancelErr
}
if changed {
return Result{Task: cancelled, Output: cancelled.Result}, err
}
}
if err != nil && isVolcesRemoteTaskCancellation(candidate, err) { if err != nil && isVolcesRemoteTaskCancellation(candidate, err) {
cancelled, changed, cancelErr := s.store.CancelSubmittedTask(ctx, task.ID, task.ExecutionToken, "任务已由火山引擎取消") cancelled, changed, cancelErr := s.store.CancelSubmittedTask(ctx, task.ID, task.ExecutionToken, "任务已由火山引擎取消")
if cancelErr != nil { if cancelErr != nil {
@@ -981,7 +1138,22 @@ func billingItemsFixedTotal(items []any) fixedAmount {
return total return total
} }
func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user *auth.User, body map[string]any, preprocessing parameterPreprocessingLog, candidate store.RuntimeModelCandidate, pricing resolvedPricing, attemptNo int, onDelta clients.StreamDelta, responseExecution responseExecutionContext, cacheAffinityPolicy map[string]any, cacheAffinityRecordKeys []string) (clients.Response, error) { func (s *Service) runCandidate(
ctx context.Context,
task store.GatewayTask,
user *auth.User,
body map[string]any,
preprocessing parameterPreprocessingLog,
candidate store.RuntimeModelCandidate,
pricing resolvedPricing,
attemptNo int,
onDelta clients.StreamDelta,
responseExecution responseExecutionContext,
cacheAffinityPolicy map[string]any,
cacheAffinityRecordKeys []string,
admittedPlatformModelID string,
admittedLeases []store.ConcurrencyLease,
) (clients.Response, error) {
var err error var err error
candidate, err = candidateWithEnvironmentCredentials(candidate) candidate, err = candidateWithEnvironmentCredentials(candidate)
if err != nil { if err != nil {
@@ -990,12 +1162,24 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
simulated := isSimulation(task, candidate) simulated := isSimulation(task, candidate)
baseAttemptMetrics := mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), parameterPreprocessingMetrics(preprocessing)) baseAttemptMetrics := mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), parameterPreprocessingMetrics(preprocessing))
reservations := s.rateLimitReservations(ctx, user, candidate, body) reservations := s.rateLimitReservations(ctx, user, candidate, body)
if admittedPlatformModelID == candidate.PlatformModelID && len(admittedLeases) > 0 {
filtered := make([]store.RateLimitReservation, 0, len(reservations))
for _, reservation := range reservations {
if reservation.Metric != "concurrent" {
filtered = append(filtered, reservation)
}
}
reservations = filtered
}
limitResult, err := s.store.ReserveRateLimits(ctx, task.ID, "", reservations) limitResult, err := s.store.ReserveRateLimits(ctx, task.ID, "", reservations)
if err != nil { if err != nil {
retryable := store.RateLimitRetryable(err) retryable := store.RateLimitRetryable(err)
clientErr := &clients.ClientError{Code: "rate_limit", Message: err.Error(), Retryable: retryable} clientErr := &clients.ClientError{Code: "rate_limit", Message: err.Error(), Retryable: retryable}
return clients.Response{}, &localRateLimitError{clientErr: clientErr, cause: err, retryAfter: localRateLimitRetryAfter(err)} return clients.Response{}, &localRateLimitError{clientErr: clientErr, cause: err, retryAfter: localRateLimitRetryAfter(err)}
} }
if admittedPlatformModelID == candidate.PlatformModelID {
limitResult.Leases = append(limitResult.Leases, admittedLeases...)
}
rateReservationsFinalized := false rateReservationsFinalized := false
defer func() { defer func() {
if !rateReservationsFinalized { if !rateReservationsFinalized {
@@ -1124,6 +1308,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
submissionStatus := "not_submitted" submissionStatus := "not_submitted"
if strings.TrimSpace(task.RemoteTaskID) != "" { if strings.TrimSpace(task.RemoteTaskID) != "" {
submissionStatus = "response_received" submissionStatus = "response_received"
markUpstreamSubmissionStarted(ctx)
if err := s.store.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, submissionStatus); err != nil { if err := s.store.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, submissionStatus); err != nil {
return clients.Response{}, fmt.Errorf("restore upstream submission status: %w", err) return clients.Response{}, fmt.Errorf("restore upstream submission status: %w", err)
} }
@@ -1181,7 +1366,11 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
return setSubmissionStatus("response_received") return setSubmissionStatus("response_received")
}, },
OnUpstreamSubmissionStarted: func() error { OnUpstreamSubmissionStarted: func() error {
return setSubmissionStatus("submitting") if err := setSubmissionStatus("submitting"); err != nil {
return err
}
markUpstreamSubmissionStarted(ctx)
return nil
}, },
OnUpstreamResponseReceived: func() error { OnUpstreamResponseReceived: func() error {
return setSubmissionStatus("response_received") return setSubmissionStatus("response_received")
@@ -1907,7 +2096,10 @@ func canonicalModelType(value string) string {
func isKnownModelType(value string) bool { func isKnownModelType(value string) bool {
switch value { switch value {
case "text_generate", "text_embedding", "text_rerank", "image_generate", "image_edit", "image_vectorize", "video_generate", "video_enhance", "image_to_video", "text_to_video", "video_edit", "video_reference", "video_first_last_frame", "omni_video", "omni", "audio_generate", "text_to_speech", "voice_clone": case "text_generate", "text_embedding", "text_rerank",
"image_generate", "image_edit", "image_analysis", "image_vectorize",
"video_generate", "video_enhance", "image_to_video", "text_to_video", "video_edit", "video_reference", "video_first_last_frame", "video_understanding",
"omni_video", "omni", "audio_generate", "audio_understanding", "text_to_speech", "voice_clone":
return true return true
default: default:
return false return false
+111
View File
@@ -48,6 +48,9 @@ type Metrics struct {
asyncWorkerDesiredCapacity atomic.Int64 asyncWorkerDesiredCapacity atomic.Int64
asyncWorkerHardLimit atomic.Int64 asyncWorkerHardLimit atomic.Int64
asyncWorkerCapacityCapped atomic.Int64 asyncWorkerCapacityCapped atomic.Int64
asyncWorkerActiveInstances atomic.Int64
asyncWorkerGlobalCapacity atomic.Int64
asyncWorkerAllocated atomic.Int64
asyncWorkerResizeSuccess atomic.Uint64 asyncWorkerResizeSuccess atomic.Uint64
asyncWorkerRefreshFailed atomic.Uint64 asyncWorkerRefreshFailed atomic.Uint64
asyncWorkerCreateFailed atomic.Uint64 asyncWorkerCreateFailed atomic.Uint64
@@ -58,6 +61,14 @@ type Metrics struct {
taskEventDuplicate atomic.Uint64 taskEventDuplicate atomic.Uint64
taskEventUnknownType atomic.Uint64 taskEventUnknownType atomic.Uint64
taskEventBudgetExceeded atomic.Uint64 taskEventBudgetExceeded atomic.Uint64
taskAdmissionAdmitted atomic.Uint64
taskAdmissionQueueFull atomic.Uint64
taskAdmissionTimeout atomic.Uint64
taskAdmissionCancelled atomic.Uint64
taskAdmissionExpired atomic.Uint64
taskAdmissionMigrated atomic.Uint64
taskAdmissionWaitBuckets [11]atomic.Uint64
taskAdmissionWaitMicros atomic.Uint64
} }
var processingDurationBounds = [...]time.Duration{ var processingDurationBounds = [...]time.Duration{
@@ -65,6 +76,19 @@ var processingDurationBounds = [...]time.Duration{
500 * time.Millisecond, time.Second, 3 * time.Second, 500 * time.Millisecond, time.Second, 3 * time.Second,
} }
var taskAdmissionWaitBounds = [...]time.Duration{
time.Second,
5 * time.Second,
15 * time.Second,
30 * time.Second,
time.Minute,
2 * time.Minute,
5 * time.Minute,
10 * time.Minute,
30 * time.Minute,
time.Hour,
}
func (m *Metrics) ObserveReceipt(outcome string, duration time.Duration, sessionsDeleted int64) { func (m *Metrics) ObserveReceipt(outcome string, duration time.Duration, sessionsDeleted int64) {
switch outcome { switch outcome {
case "accepted": case "accepted":
@@ -148,6 +172,42 @@ func (m *Metrics) SetAsyncWorkerCapacity(current, desired, hardLimit int, capped
m.asyncWorkerCapacityCapped.Store(cappedValue) m.asyncWorkerCapacityCapped.Store(cappedValue)
} }
func (m *Metrics) SetDistributedWorkerCapacity(activeInstances, globalCapacity, allocatedCapacity int) {
m.asyncWorkerActiveInstances.Store(int64(activeInstances))
m.asyncWorkerGlobalCapacity.Store(int64(globalCapacity))
m.asyncWorkerAllocated.Store(int64(allocatedCapacity))
}
func (m *Metrics) ObserveTaskAdmission(event string) {
switch event {
case "admitted":
m.taskAdmissionAdmitted.Add(1)
case "queue_full":
m.taskAdmissionQueueFull.Add(1)
case "timeout":
m.taskAdmissionTimeout.Add(1)
case "cancelled":
m.taskAdmissionCancelled.Add(1)
case "expired":
m.taskAdmissionExpired.Add(1)
case "candidate_migrated":
m.taskAdmissionMigrated.Add(1)
}
}
func (m *Metrics) ObserveTaskAdmissionWait(wait time.Duration) {
if wait < 0 {
wait = 0
}
for index, bound := range taskAdmissionWaitBounds {
if wait <= bound {
m.taskAdmissionWaitBuckets[index].Add(1)
}
}
m.taskAdmissionWaitBuckets[len(m.taskAdmissionWaitBuckets)-1].Add(1)
m.taskAdmissionWaitMicros.Add(uint64(wait / time.Microsecond))
}
func (m *Metrics) ObserveAsyncWorkerResize(outcome string) { func (m *Metrics) ObserveAsyncWorkerResize(outcome string) {
switch outcome { switch outcome {
case "success": case "success":
@@ -208,6 +268,17 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
return return
} }
} }
admission := store.TaskAdmissionMetricsSnapshot{}
if admissionProvider, ok := provider.(interface {
TaskAdmissionMetrics(context.Context) (store.TaskAdmissionMetricsSnapshot, error)
}); ok {
var err error
admission, err = admissionProvider.TaskAdmissionMetrics(r.Context())
if err != nil {
http.Error(w, "metrics unavailable", http.StatusServiceUnavailable)
return
}
}
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8") w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
outcomeCounters(w, "easyai_gateway_ssf_receipts_total", "Received SETs by bounded outcome.", []outcomeValue{ outcomeCounters(w, "easyai_gateway_ssf_receipts_total", "Received SETs by bounded outcome.", []outcomeValue{
{"accepted", m.accepted.Load()}, {"rejected", m.rejected.Load()}, {"duplicate", m.duplicate.Load()}, {"accepted", m.accepted.Load()}, {"rejected", m.rejected.Load()}, {"duplicate", m.duplicate.Load()},
@@ -276,6 +347,9 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
plainGauge(w, "easyai_gateway_async_worker_desired_capacity", "Uncapped asynchronous worker capacity derived from model and user-group policies.", m.asyncWorkerDesiredCapacity.Load()) plainGauge(w, "easyai_gateway_async_worker_desired_capacity", "Uncapped asynchronous worker capacity derived from model and user-group policies.", m.asyncWorkerDesiredCapacity.Load())
plainGauge(w, "easyai_gateway_async_worker_hard_limit", "Per-process asynchronous worker safety limit.", m.asyncWorkerHardLimit.Load()) plainGauge(w, "easyai_gateway_async_worker_hard_limit", "Per-process asynchronous worker safety limit.", m.asyncWorkerHardLimit.Load())
plainGauge(w, "easyai_gateway_async_worker_capacity_capped", "Whether desired asynchronous worker capacity is capped by the hard limit.", m.asyncWorkerCapacityCapped.Load()) plainGauge(w, "easyai_gateway_async_worker_capacity_capped", "Whether desired asynchronous worker capacity is capped by the hard limit.", m.asyncWorkerCapacityCapped.Load())
plainGauge(w, "easyai_gateway_worker_active_instances", "Active distributed worker instances with a fresh heartbeat.", m.asyncWorkerActiveInstances.Load())
plainGauge(w, "easyai_gateway_worker_global_capacity", "Global asynchronous execution capacity before instance allocation.", m.asyncWorkerGlobalCapacity.Load())
plainGauge(w, "easyai_gateway_worker_allocated_capacity", "Asynchronous execution capacity allocated to this worker instance.", m.asyncWorkerAllocated.Load())
outcomeCounters(w, "easyai_gateway_async_worker_resizes_total", "Asynchronous worker resize attempts by bounded outcome.", []outcomeValue{ outcomeCounters(w, "easyai_gateway_async_worker_resizes_total", "Asynchronous worker resize attempts by bounded outcome.", []outcomeValue{
{"success", m.asyncWorkerResizeSuccess.Load()}, {"success", m.asyncWorkerResizeSuccess.Load()},
{"refresh_failed", m.asyncWorkerRefreshFailed.Load()}, {"refresh_failed", m.asyncWorkerRefreshFailed.Load()},
@@ -292,6 +366,21 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
{"unknown_type", m.taskEventUnknownType.Load()}, {"unknown_type", m.taskEventUnknownType.Load()},
{"budget_exceeded", m.taskEventBudgetExceeded.Load()}, {"budget_exceeded", m.taskEventBudgetExceeded.Load()},
}) })
plainGauge(w, "easyai_gateway_task_admission_queue_depth", "Current persistent non-text task admission queue depth.", int64(admission.QueueDepth))
plainGauge(w, "easyai_gateway_task_admission_waiting_sync", "Current synchronous admission waiters.", int64(admission.WaitingSync))
plainGauge(w, "easyai_gateway_task_admission_waiting_async", "Current asynchronous admission waiters.", int64(admission.WaitingAsync))
plainFloatGauge(w, "easyai_gateway_task_admission_oldest_wait_seconds", "Age of the oldest persistent admission waiter.", admission.OldestWaitSeconds)
plainGauge(w, "easyai_gateway_task_admission_expired_waiter_backlog", "Expired synchronous waiter leases awaiting reclamation.", int64(admission.ExpiredWaiterBacklog))
plainGauge(w, "easyai_gateway_task_admission_expired_deadline_backlog", "Expired queue deadlines awaiting reclamation.", int64(admission.ExpiredDeadlineBacklog))
outcomeCounters(w, "easyai_gateway_task_admissions_total", "Task admission transitions by bounded outcome.", []outcomeValue{
{"admitted", m.taskAdmissionAdmitted.Load()},
{"queue_full", m.taskAdmissionQueueFull.Load()},
{"timeout", m.taskAdmissionTimeout.Load()},
{"cancelled", m.taskAdmissionCancelled.Load()},
{"expired", m.taskAdmissionExpired.Load()},
{"candidate_migrated", m.taskAdmissionMigrated.Load()},
})
taskAdmissionWaitHistogram(w, m)
}) })
} }
@@ -333,3 +422,25 @@ func plainCounter(w http.ResponseWriter, name, help string, value uint64) {
func plainGauge(w http.ResponseWriter, name, help string, value int64) { func plainGauge(w http.ResponseWriter, name, help string, value int64) {
fmt.Fprintf(w, "# HELP %s %s\n# TYPE %s gauge\n%s %d\n", name, help, name, name, value) fmt.Fprintf(w, "# HELP %s %s\n# TYPE %s gauge\n%s %d\n", name, help, name, name, value)
} }
func plainFloatGauge(w http.ResponseWriter, name, help string, value float64) {
fmt.Fprintf(w, "# HELP %s %s\n# TYPE %s gauge\n%s %.6f\n", name, help, name, name, value)
}
func taskAdmissionWaitHistogram(w http.ResponseWriter, metrics *Metrics) {
const name = "easyai_gateway_task_admission_wait_seconds"
fmt.Fprintf(w, "# HELP %s Time spent waiting for persistent task admission.\n# TYPE %s histogram\n", name, name)
for index, bound := range taskAdmissionWaitBounds {
fmt.Fprintf(
w,
"%s_bucket{le=\"%g\"} %d\n",
name,
bound.Seconds(),
metrics.taskAdmissionWaitBuckets[index].Load(),
)
}
count := metrics.taskAdmissionWaitBuckets[len(metrics.taskAdmissionWaitBuckets)-1].Load()
fmt.Fprintf(w, "%s_bucket{le=\"+Inf\"} %d\n", name, count)
fmt.Fprintf(w, "%s_sum %.6f\n", name, float64(metrics.taskAdmissionWaitMicros.Load())/1_000_000)
fmt.Fprintf(w, "%s_count %d\n", name, count)
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,499 @@
package store
import (
"context"
"errors"
"os"
"strings"
"sync"
"testing"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
)
func TestDistributedTaskAdmissionQueueAcrossStores(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(), 60*time.Second)
defer cancel()
applyOIDCJITTestMigrations(t, ctx, databaseURL)
first, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect first store: %v", err)
}
defer first.Close()
second, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect second store: %v", err)
}
defer second.Close()
var databaseName string
if err := first.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 := first.CreatePlatform(ctx, CreatePlatformInput{
Provider: "queue-test",
PlatformKey: "queue-test-" + suffix,
Name: "Queue Test " + suffix,
AuthType: "none",
Status: "enabled",
})
if err != nil {
t.Fatalf("create platform: %v", err)
}
var platformModelID string
if err := first.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, "queue-model-"+suffix).Scan(&platformModelID); err != nil {
t.Fatalf("create platform model: %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 {
_, _ = first.pool.Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = $1::uuid`, taskID)
}
_, _ = first.pool.Exec(cleanupCtx, `DELETE FROM platform_models WHERE id = $1::uuid`, platformModelID)
_, _ = first.pool.Exec(cleanupCtx, `DELETE FROM integration_platforms WHERE id = $1::uuid`, platform.ID)
_, _ = first.pool.Exec(cleanupCtx, `DELETE FROM gateway_worker_instances WHERE instance_id LIKE $1`, "queue-test-"+suffix+"%")
})
createTask := func(async bool) GatewayTask {
t.Helper()
task, createErr := first.CreateTask(ctx, CreateTaskInput{
Kind: "images.generate",
Model: "queue-model-" + suffix,
RunMode: "production",
Async: async,
Request: map[string]any{"prompt": "queue admission test"},
}, &auth.User{ID: "queue-test-user-" + suffix, Source: "gateway"})
if createErr != nil {
t.Fatalf("create task: %v", createErr)
}
taskIDs = append(taskIDs, task.ID)
return task
}
scope := AdmissionScope{
ScopeType: "platform_model",
ScopeKey: platformModelID,
ScopeName: "queue-model",
ConcurrentLimit: 1,
Amount: 1,
LeaseTTLSeconds: 137,
QueueLimit: 2,
MaxWaitSeconds: 600,
}
inputFor := func(task GatewayTask, priority int, waiterID string) TaskAdmissionInput {
mode := "sync"
if task.AsyncMode {
mode = "async"
}
return TaskAdmissionInput{
TaskID: task.ID,
PlatformID: platform.ID,
PlatformModelID: platformModelID,
QueueKey: "queue-test:" + suffix,
Mode: mode,
Priority: priority,
WaiterID: waiterID,
Scopes: []AdmissionScope{scope},
}
}
running := createTask(false)
result, err := first.TryTaskAdmission(ctx, inputFor(running, 100, "waiter-running"))
if err != nil || !result.Admitted || len(result.Leases) != 1 {
t.Fatalf("first task admission = %+v, err=%v", result, err)
}
reloadedLeases, err := second.ActiveTaskAdmissionLeases(ctx, running.ID)
if err != nil || len(reloadedLeases) != 1 || reloadedLeases[0].TTL != 137*time.Second {
t.Fatalf("reloaded admission leases = %+v, err=%v, want one 137s lease", reloadedLeases, err)
}
lowerPriority := createTask(false)
result, err = second.TryTaskAdmission(ctx, inputFor(lowerPriority, 20, "waiter-low"))
if err != nil || result.Admitted {
t.Fatalf("second task should wait: result=%+v err=%v", result, err)
}
higherPriorityAsync := createTask(true)
result, err = first.TryTaskAdmission(ctx, inputFor(higherPriorityAsync, 10, ""))
if err != nil || result.Admitted {
t.Fatalf("third task should wait: result=%+v err=%v", result, err)
}
rejected := createTask(true)
_, err = second.TryTaskAdmission(ctx, inputFor(rejected, 30, ""))
var limitErr *RateLimitExceededError
if !errors.As(err, &limitErr) || limitErr.Metric != "queue_size" || limitErr.Reason != "queue_full" ||
limitErr.QueueDepth != 2 || limitErr.QueueLimit != 2 {
t.Fatalf("fourth task error = %#v, want queue_full depth 2/2", err)
}
var queueCounters int
if err := first.pool.QueryRow(ctx, `
SELECT COUNT(*)
FROM gateway_rate_limit_counters
WHERE metric = 'queue_size'
AND scope_key = $1`, platformModelID).Scan(&queueCounters); err != nil {
t.Fatalf("count queue counters: %v", err)
}
if queueCounters != 0 {
t.Fatalf("queue_size wrote %d fixed-window counters", queueCounters)
}
var attempts int
if err := first.pool.QueryRow(ctx, `
SELECT COUNT(*)
FROM gateway_task_attempts
WHERE task_id = ANY($1::uuid[])`, taskIDs).Scan(&attempts); err != nil {
t.Fatalf("count attempts: %v", err)
}
if attempts != 0 {
t.Fatalf("waiting admissions created %d task attempts", attempts)
}
if err := first.DeleteTaskAdmission(ctx, running.ID); err != nil {
t.Fatalf("release running admission: %v", err)
}
result, err = second.TryTaskAdmission(ctx, inputFor(lowerPriority, 20, "waiter-low"))
if err != nil {
t.Fatalf("try lower-priority task after release: %v", err)
}
if result.Admitted {
t.Fatal("lower-priority task bypassed higher-priority asynchronous waiter")
}
result, err = first.TryTaskAdmission(ctx, inputFor(higherPriorityAsync, 10, ""))
if err != nil || !result.Admitted {
t.Fatalf("higher-priority async task was not admitted first: result=%+v err=%v", result, err)
}
for _, taskID := range taskIDs {
_ = first.DeleteTaskAdmission(ctx, taskID)
}
concurrentTasks := []GatewayTask{createTask(false), createTask(true), createTask(false), createTask(true)}
type concurrentResult struct {
result TaskAdmissionResult
err error
}
start := make(chan struct{})
results := make(chan concurrentResult, len(concurrentTasks))
var waitGroup sync.WaitGroup
for index, task := range concurrentTasks {
waitGroup.Add(1)
go func(index int, task GatewayTask) {
defer waitGroup.Done()
<-start
target := first
if index%2 == 1 {
target = second
}
waiterID := ""
if !task.AsyncMode {
waiterID = "concurrent-waiter-" + task.ID
}
admission, admissionErr := target.TryTaskAdmission(ctx, inputFor(task, 100, waiterID))
results <- concurrentResult{result: admission, err: admissionErr}
}(index, task)
}
close(start)
waitGroup.Wait()
close(results)
admittedCount := 0
waitingCount := 0
queueFullCount := 0
for item := range results {
if item.err == nil && item.result.Admitted {
admittedCount++
continue
}
if item.err == nil {
waitingCount++
continue
}
if errors.As(item.err, &limitErr) && limitErr.Metric == "queue_size" {
queueFullCount++
continue
}
t.Fatalf("unexpected concurrent admission result: %+v err=%v", item.result, item.err)
}
if admittedCount != 1 || waitingCount != 2 || queueFullCount != 1 {
t.Fatalf("concurrent outcomes admitted=%d waiting=%d queue_full=%d, want 1/2/1", admittedCount, waitingCount, queueFullCount)
}
var activeLeases, waitingAdmissions int
if err := first.pool.QueryRow(ctx, `
SELECT COUNT(*)
FROM gateway_concurrency_leases
WHERE scope_type = 'platform_model'
AND scope_key = $1
AND released_at IS NULL
AND expires_at > now()`, platformModelID).Scan(&activeLeases); err != nil {
t.Fatalf("count active leases: %v", err)
}
if err := first.pool.QueryRow(ctx, `
SELECT COUNT(*)
FROM gateway_task_admissions
WHERE platform_model_id = $1::uuid
AND status = 'waiting'`, platformModelID).Scan(&waitingAdmissions); err != nil {
t.Fatalf("count waiting admissions: %v", err)
}
if activeLeases != 1 || waitingAdmissions != 2 {
t.Fatalf("database bounds active=%d waiting=%d, want 1/2", activeLeases, waitingAdmissions)
}
var deadlineTaskID string
if err := first.pool.QueryRow(ctx, `
SELECT task_id::text
FROM gateway_task_admissions
WHERE platform_model_id = $1::uuid
AND status = 'waiting'
ORDER BY task_id
LIMIT 1`, platformModelID).Scan(&deadlineTaskID); err != nil {
t.Fatalf("select deadline waiter: %v", err)
}
if _, err := first.pool.Exec(ctx, `
UPDATE gateway_task_admissions
SET enqueued_at = now() - interval '2 seconds',
wait_deadline_at = now() - interval '1 second'
WHERE task_id = $1::uuid`, deadlineTaskID); err != nil {
t.Fatalf("expire queue deadline: %v", err)
}
reaped, err := second.ReapExpiredTaskAdmissions(ctx, 10)
if err != nil || reaped.ExpiredDeadlines < 1 {
t.Fatalf("deadline reap = %+v, err=%v", reaped, err)
}
var taskStatus, taskErrorCode string
if err := first.pool.QueryRow(ctx, `
SELECT status, COALESCE(error_code, '')
FROM gateway_tasks
WHERE id = $1::uuid`, deadlineTaskID).Scan(&taskStatus, &taskErrorCode); err != nil {
t.Fatalf("read deadline task: %v", err)
}
if taskStatus != "failed" || taskErrorCode != "queue_timeout" {
t.Fatalf("deadline task status=%s code=%s, want failed/queue_timeout", taskStatus, taskErrorCode)
}
for _, taskID := range taskIDs {
_ = first.DeleteTaskAdmission(ctx, taskID)
}
leaseHolder := createTask(false)
if result, err := first.TryTaskAdmission(ctx, inputFor(leaseHolder, 100, "lease-holder")); err != nil || !result.Admitted {
t.Fatalf("admit lease holder: result=%+v err=%v", result, err)
}
disconnectedWaiter := createTask(false)
if result, err := second.TryTaskAdmission(ctx, inputFor(disconnectedWaiter, 100, "disconnected-waiter")); err != nil || result.Admitted {
t.Fatalf("queue disconnected waiter: result=%+v err=%v", result, err)
}
if _, err := first.pool.Exec(ctx, `
UPDATE gateway_task_admissions
SET waiter_lease_expires_at = now() - interval '1 second'
WHERE task_id = $1::uuid`, disconnectedWaiter.ID); err != nil {
t.Fatalf("expire waiter lease: %v", err)
}
reaped, err = second.ReapExpiredTaskAdmissions(ctx, 10)
if err != nil || reaped.ExpiredWaiters < 1 {
t.Fatalf("waiter reap = %+v, err=%v", reaped, err)
}
if err := first.pool.QueryRow(ctx, `
SELECT status, COALESCE(error_code, '')
FROM gateway_tasks
WHERE id = $1::uuid`, disconnectedWaiter.ID).Scan(&taskStatus, &taskErrorCode); err != nil {
t.Fatalf("read disconnected task: %v", err)
}
if taskStatus != "cancelled" || taskErrorCode != "client_disconnected" {
t.Fatalf("disconnected task status=%s code=%s, want cancelled/client_disconnected", taskStatus, taskErrorCode)
}
if err := first.DeleteTaskAdmission(ctx, leaseHolder.ID); err != nil {
t.Fatalf("release lease holder admission: %v", err)
}
preSubmission := createTask(false)
claimed, err := first.ClaimTaskPreparation(ctx, preSubmission.ID, uuid.NewString(), time.Minute)
if err != nil {
t.Fatalf("claim pre-submission task: %v", err)
}
result, err = first.TryTaskAdmission(ctx, inputFor(claimed, 100, "pre-submission"))
if err != nil || !result.Admitted {
t.Fatalf("admit pre-submission task: result=%+v err=%v", result, err)
}
var attemptID string
if err := first.pool.QueryRow(ctx, `
INSERT INTO gateway_task_attempts (
task_id, attempt_no, platform_id, platform_model_id, queue_key, status,
upstream_submission_status
)
VALUES ($1::uuid, 1, $2::uuid, $3::uuid, $4, 'running', 'not_submitted')
RETURNING id::text`,
preSubmission.ID, platform.ID, platformModelID, "queue-test:"+suffix,
).Scan(&attemptID); err != nil {
t.Fatalf("create pre-submission attempt: %v", err)
}
if _, err := first.pool.Exec(ctx, `
INSERT INTO gateway_task_param_preprocessing_logs (
task_id, attempt_id, attempt_no, model_type
)
VALUES ($1::uuid, $2::uuid, 1, 'image_generate')`, preSubmission.ID, attemptID); err != nil {
t.Fatalf("create pre-submission log: %v", err)
}
cancelled, changed, err := first.CancelTaskBeforeUpstreamSubmission(
ctx,
preSubmission.ID,
claimed.ExecutionToken,
"client disconnected before upstream submission",
)
if err != nil || !changed || cancelled.Status != "cancelled" || cancelled.ErrorCode != "client_disconnected" {
t.Fatalf("cancel pre-submission task = %+v changed=%v err=%v", cancelled, changed, err)
}
var remainingAdmissions, remainingLeases, remainingAttempts, remainingPreprocessing int
if err := first.pool.QueryRow(ctx, `
SELECT
(SELECT COUNT(*) FROM gateway_task_admissions WHERE task_id = $1::uuid),
(SELECT COUNT(*) FROM gateway_concurrency_leases WHERE task_id = $1::uuid AND released_at IS NULL),
(SELECT COUNT(*) FROM gateway_task_attempts WHERE task_id = $1::uuid),
(SELECT COUNT(*) FROM gateway_task_param_preprocessing_logs WHERE task_id = $1::uuid)`,
preSubmission.ID,
).Scan(&remainingAdmissions, &remainingLeases, &remainingAttempts, &remainingPreprocessing); err != nil {
t.Fatalf("read pre-submission cleanup: %v", err)
}
if remainingAdmissions != 0 || remainingLeases != 0 || remainingAttempts != 0 || remainingPreprocessing != 0 {
t.Fatalf(
"pre-submission cleanup admissions=%d leases=%d attempts=%d preprocessing=%d",
remainingAdmissions,
remainingLeases,
remainingAttempts,
remainingPreprocessing,
)
}
postSubmission := createTask(false)
claimed, err = first.ClaimTaskPreparation(ctx, postSubmission.ID, uuid.NewString(), time.Minute)
if err != nil {
t.Fatalf("claim post-submission task: %v", err)
}
if _, err := first.pool.Exec(ctx, `
INSERT INTO gateway_task_attempts (
task_id, attempt_no, platform_id, platform_model_id, queue_key, status,
upstream_submission_status
)
VALUES ($1::uuid, 1, $2::uuid, $3::uuid, $4, 'running', 'submitting')`,
postSubmission.ID, platform.ID, platformModelID, "queue-test:"+suffix,
); err != nil {
t.Fatalf("create submitted attempt: %v", err)
}
_, changed, err = first.CancelTaskBeforeUpstreamSubmission(
ctx,
postSubmission.ID,
claimed.ExecutionToken,
"client disconnected after upstream submission",
)
if err != nil || changed {
t.Fatalf("post-submission cancellation changed=%v err=%v, want unchanged", changed, err)
}
atomicTask := createTask(true)
hookFailure := errors.New("synthetic admitted hook failure")
_, err = first.TryTaskAdmissionWithAdmittedHook(ctx, inputFor(atomicTask, 100, ""), func(pgx.Tx) error {
return hookFailure
})
if !errors.Is(err, hookFailure) {
t.Fatalf("admitted hook failure = %v, want synthetic failure", err)
}
var atomicAdmissions, atomicLeases int
if err := first.pool.QueryRow(ctx, `
SELECT
(SELECT COUNT(*) FROM gateway_task_admissions WHERE task_id = $1::uuid),
(SELECT COUNT(*) FROM gateway_concurrency_leases WHERE task_id = $1::uuid AND released_at IS NULL)`,
atomicTask.ID,
).Scan(&atomicAdmissions, &atomicLeases); err != nil {
t.Fatalf("read rolled back admitted hook state: %v", err)
}
if atomicAdmissions != 0 || atomicLeases != 0 {
t.Fatalf("failed admitted hook left admissions=%d leases=%d", atomicAdmissions, atomicLeases)
}
const syntheticRiverJobID int64 = 987654321
result, err = first.TryTaskAdmissionWithAdmittedHook(ctx, inputFor(atomicTask, 100, ""), func(tx pgx.Tx) error {
_, hookErr := tx.Exec(ctx, `
UPDATE gateway_tasks
SET river_job_id = $2
WHERE id = $1::uuid`, atomicTask.ID, syntheticRiverJobID)
return hookErr
})
if err != nil || !result.Admitted {
t.Fatalf("atomic admitted hook result=%+v err=%v", result, err)
}
var riverJobID int64
if err := first.pool.QueryRow(ctx, `
SELECT river_job_id
FROM gateway_tasks
WHERE id = $1::uuid`, atomicTask.ID).Scan(&riverJobID); err != nil {
t.Fatalf("read atomic River job marker: %v", err)
}
if riverJobID != syntheticRiverJobID {
t.Fatalf("atomic River job marker=%d, want %d", riverJobID, syntheticRiverJobID)
}
if err := first.DeleteTaskAdmission(ctx, atomicTask.ID); err != nil {
t.Fatalf("release atomic admitted hook task: %v", err)
}
}
func TestWorkerCapacityAllocationAndFailover(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 worker registry 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)
}
prefix := "queue-test-" + strings.ReplaceAll(uuid.NewString(), "-", "")
firstID := prefix + "-a"
secondID := prefix + "-b"
defer db.pool.Exec(context.Background(), `DELETE FROM gateway_worker_instances WHERE instance_id = ANY($1::text[])`, []string{firstID, secondID})
first, err := db.RegisterWorkerInstance(ctx, WorkerRegistrationInput{InstanceID: firstID, DesiredCapacity: 5})
if err != nil || first.Allocated != 5 || first.ActiveInstances != 1 {
t.Fatalf("first allocation = %+v, err=%v", first, err)
}
second, err := db.RegisterWorkerInstance(ctx, WorkerRegistrationInput{InstanceID: secondID, DesiredCapacity: 5})
if err != nil || second.Allocated != 2 || second.ActiveInstances != 2 {
t.Fatalf("second allocation = %+v, err=%v", second, err)
}
first, err = db.RegisterWorkerInstance(ctx, WorkerRegistrationInput{InstanceID: firstID, DesiredCapacity: 5})
if err != nil || first.Allocated != 3 || first.ActiveInstances != 2 {
t.Fatalf("rebalanced first allocation = %+v, err=%v", first, err)
}
if _, err := db.pool.Exec(ctx, `
UPDATE gateway_worker_instances
SET heartbeat_at = now() - interval '16 seconds'
WHERE instance_id = $1`, secondID); err != nil {
t.Fatalf("expire second heartbeat: %v", err)
}
first, err = db.RegisterWorkerInstance(ctx, WorkerRegistrationInput{InstanceID: firstID, DesiredCapacity: 5})
if err != nil || first.Allocated != 5 || first.ActiveInstances != 1 {
t.Fatalf("single-worker failover allocation = %+v, err=%v", first, err)
}
}
@@ -7,6 +7,7 @@ import (
type AsyncWorkerCapacitySnapshot struct { type AsyncWorkerCapacitySnapshot struct {
Capacity int Capacity int
GlobalCapacity int
Desired int Desired int
HardLimit int HardLimit int
Capped bool Capped bool
@@ -16,6 +17,8 @@ type AsyncWorkerCapacitySnapshot struct {
UnlimitedGroups int UnlimitedGroups int
ModelDesired int ModelDesired int
GroupDesired int GroupDesired int
ActiveInstances int
InstanceID string
} }
func (s *Store) AsyncWorkerCapacity(ctx context.Context, hardLimit int) (AsyncWorkerCapacitySnapshot, error) { func (s *Store) AsyncWorkerCapacity(ctx context.Context, hardLimit int) (AsyncWorkerCapacitySnapshot, error) {
@@ -131,6 +134,7 @@ func asyncWorkerCapacityFromPolicySets(modelPolicies []map[string]any, groupPoli
snapshot.Capacity = hardLimit snapshot.Capacity = hardLimit
snapshot.Capped = true snapshot.Capped = true
} }
snapshot.GlobalCapacity = snapshot.Capacity
return snapshot return snapshot
} }
+10
View File
@@ -212,6 +212,11 @@ WHERE id = $1::uuid AND deleted_at IS NULL`, id)
func (s *Store) CreateUserGroup(ctx context.Context, input UserGroupInput) (UserGroup, error) { func (s *Store) CreateUserGroup(ctx context.Context, input UserGroupInput) (UserGroup, error) {
input = normalizeUserGroupInput(input) input = normalizeUserGroupInput(input)
normalizedPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return UserGroup{}, err
}
input.RateLimitPolicy = normalizedPolicy
rechargeDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.RechargeDiscountPolicy)) rechargeDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.RechargeDiscountPolicy))
billingDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.BillingDiscountPolicy)) billingDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.BillingDiscountPolicy))
rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy)) rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy))
@@ -231,6 +236,11 @@ RETURNING `+userGroupColumns,
func (s *Store) UpdateUserGroup(ctx context.Context, id string, input UserGroupInput) (UserGroup, error) { func (s *Store) UpdateUserGroup(ctx context.Context, id string, input UserGroupInput) (UserGroup, error) {
input = normalizeUserGroupInput(input) input = normalizeUserGroupInput(input)
normalizedPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return UserGroup{}, err
}
input.RateLimitPolicy = normalizedPolicy
rechargeDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.RechargeDiscountPolicy)) rechargeDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.RechargeDiscountPolicy))
billingDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.BillingDiscountPolicy)) billingDiscountPolicy, _ := json.Marshal(emptyObjectIfNil(input.BillingDiscountPolicy))
rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy)) rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy))
@@ -93,6 +93,19 @@ WHERE resource_type = 'platform_model'
func (s *Store) createPlatformModel(ctx context.Context, q platformModelQuerier, input CreatePlatformModelInput) (PlatformModel, error) { func (s *Store) createPlatformModel(ctx context.Context, q platformModelQuerier, input CreatePlatformModelInput) (PlatformModel, error) {
input.ModelName = strings.TrimSpace(input.ModelName) input.ModelName = strings.TrimSpace(input.ModelName)
input.ProviderModelName = strings.TrimSpace(input.ProviderModelName) input.ProviderModelName = strings.TrimSpace(input.ProviderModelName)
normalizedRateLimitPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return PlatformModel{}, err
}
input.RateLimitPolicy = normalizedRateLimitPolicy
if rawPolicy, ok := input.RuntimePolicyOverride["rateLimitPolicy"].(map[string]any); ok {
normalizedRuntimeOverridePolicy, normalizeErr := NormalizeAndValidateRateLimitPolicy(rawPolicy)
if normalizeErr != nil {
return PlatformModel{}, normalizeErr
}
input.RuntimePolicyOverride = cloneObject(input.RuntimePolicyOverride)
input.RuntimePolicyOverride["rateLimitPolicy"] = normalizedRuntimeOverridePolicy
}
base, err := s.lookupBaseModel(ctx, q, input.BaseModelID, input.CanonicalModelKey, input.ModelName) base, err := s.lookupBaseModel(ctx, q, input.BaseModelID, input.CanonicalModelKey, input.ModelName)
if err != nil && !IsNotFound(err) { if err != nil && !IsNotFound(err) {
return PlatformModel{}, err return PlatformModel{}, err
+27 -5
View File
@@ -76,7 +76,11 @@ var (
) )
func Connect(ctx context.Context, databaseURL string) (*Store, error) { func Connect(ctx context.Context, databaseURL string) (*Store, error) {
config, err := postgresPoolConfig(databaseURL) return ConnectWithMaxConns(ctx, databaseURL, 0)
}
func ConnectWithMaxConns(ctx context.Context, databaseURL string, maxConns int) (*Store, error) {
config, err := postgresPoolConfig(databaseURL, maxConns)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -91,13 +95,16 @@ func Connect(ctx context.Context, databaseURL string) (*Store, error) {
return &Store{pool: pool}, nil return &Store{pool: pool}, nil
} }
func postgresPoolConfig(databaseURL string) (*pgxpool.Config, error) { func postgresPoolConfig(databaseURL string, maxConns ...int) (*pgxpool.Config, error) {
config, err := pgxpool.ParseConfig(databaseURL) config, err := pgxpool.ParseConfig(databaseURL)
if err != nil { if err != nil {
return nil, err return nil, err
} }
config.ConnConfig.ConnectTimeout = postgresConnectTimeout config.ConnConfig.ConnectTimeout = postgresConnectTimeout
config.ConnConfig.RuntimeParams["application_name"] = postgresApplicationName config.ConnConfig.RuntimeParams["application_name"] = postgresApplicationName
if len(maxConns) > 0 && maxConns[0] > 0 {
config.MaxConns = int32(maxConns[0])
}
return config, nil return config, nil
} }
@@ -523,6 +530,8 @@ type GatewayTask struct {
AsyncMode bool `json:"asyncMode"` AsyncMode bool `json:"asyncMode"`
RiverJobID int64 `json:"riverJobId,omitempty"` RiverJobID int64 `json:"riverJobId,omitempty"`
Status string `json:"status"` Status string `json:"status"`
QueueKey string `json:"-"`
Priority int `json:"-"`
Cancellable *bool `json:"cancellable,omitempty"` Cancellable *bool `json:"cancellable,omitempty"`
Submitted *bool `json:"submitted,omitempty"` Submitted *bool `json:"submitted,omitempty"`
Message string `json:"message,omitempty"` Message string `json:"message,omitempty"`
@@ -570,7 +579,8 @@ COALESCE(api_key_id, ''), COALESCE(api_key_name, ''), COALESCE(api_key_prefix, '
COALESCE(user_group_id::text, ''), COALESCE(user_group_key, ''), model, COALESCE(user_group_id::text, ''), COALESCE(user_group_key, ''), model,
COALESCE(model_type, ''), COALESCE(requested_model, ''), COALESCE(resolved_model, ''), COALESCE(request_id, ''), COALESCE(model_type, ''), COALESCE(requested_model, ''), COALESCE(resolved_model, ''), COALESCE(request_id, ''),
COALESCE(conversation_id::text, ''), COALESCE(new_message_count, 0), COALESCE(conversation_id::text, ''), COALESCE(new_message_count, 0),
request, COALESCE(async_mode, false), COALESCE(river_job_id, 0), status, COALESCE(attempt_count, 0), request, COALESCE(async_mode, false), COALESCE(river_job_id, 0), status,
COALESCE(queue_key, 'default'), COALESCE(priority, 100), COALESCE(attempt_count, 0),
COALESCE(remote_task_id, ''), COALESCE(remote_task_payload, '{}'::jsonb), COALESCE(remote_task_id, ''), COALESCE(remote_task_payload, '{}'::jsonb),
COALESCE(result, '{}'::jsonb), COALESCE(billings, '[]'::jsonb), COALESCE(result, '{}'::jsonb), COALESCE(billings, '[]'::jsonb),
COALESCE(usage, '{}'::jsonb), COALESCE(metrics, '{}'::jsonb), COALESCE(billing_summary, '{}'::jsonb), COALESCE(usage, '{}'::jsonb), COALESCE(metrics, '{}'::jsonb), COALESCE(billing_summary, '{}'::jsonb),
@@ -715,6 +725,11 @@ ORDER BY COALESCE(dynamic_priority, priority) ASC, priority ASC, created_at DESC
} }
func (s *Store) CreatePlatform(ctx context.Context, input CreatePlatformInput) (Platform, error) { func (s *Store) CreatePlatform(ctx context.Context, input CreatePlatformInput) (Platform, error) {
normalizedPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return Platform{}, err
}
input.RateLimitPolicy = normalizedPolicy
credentials, _ := json.Marshal(emptyObjectIfNil(input.Credentials)) credentials, _ := json.Marshal(emptyObjectIfNil(input.Credentials))
config, _ := json.Marshal(emptyObjectIfNil(input.Config)) config, _ := json.Marshal(emptyObjectIfNil(input.Config))
retryPolicy, _ := json.Marshal(emptyObjectIfNil(input.RetryPolicy)) retryPolicy, _ := json.Marshal(emptyObjectIfNil(input.RetryPolicy))
@@ -734,7 +749,7 @@ func (s *Store) CreatePlatform(ctx context.Context, input CreatePlatformInput) (
var retryPolicyBytes []byte var retryPolicyBytes []byte
var rateLimitPolicyBytes []byte var rateLimitPolicyBytes []byte
var dynamicPriority sql.NullInt64 var dynamicPriority sql.NullInt64
err := s.pool.QueryRow(ctx, ` err = s.pool.QueryRow(ctx, `
INSERT INTO integration_platforms ( INSERT INTO integration_platforms (
provider, platform_key, name, internal_name, base_url, auth_type, credentials, config, provider, platform_key, name, internal_name, base_url, auth_type, credentials, config,
default_pricing_mode, default_discount_factor, pricing_rule_set_id, default_pricing_mode, default_discount_factor, pricing_rule_set_id,
@@ -789,6 +804,11 @@ RETURNING id::text, provider, platform_key, name, COALESCE(internal_name, ''), C
} }
func (s *Store) UpdatePlatform(ctx context.Context, id string, input CreatePlatformInput) (Platform, error) { func (s *Store) UpdatePlatform(ctx context.Context, id string, input CreatePlatformInput) (Platform, error) {
normalizedPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return Platform{}, err
}
input.RateLimitPolicy = normalizedPolicy
var credentials any var credentials any
if input.Credentials != nil { if input.Credentials != nil {
credentialsBytes, _ := json.Marshal(emptyObjectIfNil(input.Credentials)) credentialsBytes, _ := json.Marshal(emptyObjectIfNil(input.Credentials))
@@ -812,7 +832,7 @@ func (s *Store) UpdatePlatform(ctx context.Context, id string, input CreatePlatf
var retryPolicyBytes []byte var retryPolicyBytes []byte
var rateLimitPolicyBytes []byte var rateLimitPolicyBytes []byte
var dynamicPriority sql.NullInt64 var dynamicPriority sql.NullInt64
err := s.pool.QueryRow(ctx, ` err = s.pool.QueryRow(ctx, `
UPDATE integration_platforms UPDATE integration_platforms
SET provider = $2, SET provider = $2,
platform_key = COALESCE(NULLIF($3, ''), platform_key), platform_key = COALESCE(NULLIF($3, ''), platform_key),
@@ -2138,6 +2158,8 @@ func scanGatewayTask(scanner taskScanner) (GatewayTask, error) {
&task.AsyncMode, &task.AsyncMode,
&task.RiverJobID, &task.RiverJobID,
&task.Status, &task.Status,
&task.QueueKey,
&task.Priority,
&task.AttemptCount, &task.AttemptCount,
&task.RemoteTaskID, &task.RemoteTaskID,
&remoteTaskPayloadBytes, &remoteTaskPayloadBytes,
+149 -34
View File
@@ -1,13 +1,24 @@
package store package store
import ( import (
"errors"
"fmt"
"math" "math"
"strings" "strings"
) )
var ErrInvalidRateLimitPolicy = errors.New("invalid rate limit policy")
func IsInvalidRateLimitPolicy(err error) bool {
return errors.Is(err, ErrInvalidRateLimitPolicy)
}
const ( const (
RateLimitPolicyModeInherit = "inherit" RateLimitPolicyModeInherit = "inherit"
RateLimitPolicyModeOverride = "override" RateLimitPolicyModeOverride = "override"
DefaultQueueMaxWaitSeconds = 600
MaxQueueSize = 10000
MaxQueueWaitSeconds = 3600
) )
type EffectiveRateLimitPolicyInput struct { type EffectiveRateLimitPolicyInput struct {
@@ -78,46 +89,150 @@ func NormalizeRateLimitPolicy(policy map[string]any) map[string]any {
} }
out := clonePolicy(policy) out := clonePolicy(policy)
rules, _ := out["rules"].([]any) rules, _ := out["rules"].([]any)
if len(rules) > 0 { if len(rules) == 0 {
return out legacyScopes := []map[string]any{policy}
} for _, key := range []string{"platformLimits", "modelLimits", "platform_limits", "model_limits"} {
if nested, ok := policy[key].(map[string]any); ok {
legacyScopes := []map[string]any{policy} legacyScopes = append(legacyScopes, nested)
for _, key := range []string{"platformLimits", "modelLimits", "platform_limits", "model_limits"} { }
if nested, ok := policy[key].(map[string]any); ok { }
legacyScopes = append(legacyScopes, nested) if limit, ok := lowestPositiveLegacyLimit(legacyScopes,
"max_concurrent_requests", "maxConcurrentRequests", "concurrent"); ok {
rules = append(rules, map[string]any{
"metric": "concurrent",
"limit": limit,
"leaseTtlSeconds": 120,
})
}
if limit, ok := lowestPositiveLegacyLimit(legacyScopes,
"max_request_per_minute", "maxRequestsPerMinute", "rpm"); ok {
rules = append(rules, map[string]any{
"metric": "rpm",
"limit": limit,
"windowSeconds": 60,
})
}
if limit, ok := lowestPositiveLegacyLimit(legacyScopes,
"max_tokens_per_minute", "maxTokensPerMinute", "tpm_total"); ok {
rules = append(rules, map[string]any{
"metric": "tpm_total",
"limit": limit,
"windowSeconds": 60,
})
} }
} }
if limit, ok := lowestPositiveLegacyLimit(legacyScopes, normalizedRules := make([]any, 0, len(rules))
"max_concurrent_requests", "maxConcurrentRequests", "concurrent"); ok { for _, rawRule := range rules {
rules = append(rules, map[string]any{ rule, ok := rawRule.(map[string]any)
"metric": "concurrent", if !ok {
"limit": limit, normalizedRules = append(normalizedRules, rawRule)
"leaseTtlSeconds": 120, continue
}) }
} normalizedRule := clonePolicy(rule)
if limit, ok := lowestPositiveLegacyLimit(legacyScopes, if strings.TrimSpace(stringValue(normalizedRule["metric"])) == "queue_size" {
"max_request_per_minute", "maxRequestsPerMinute", "rpm"); ok { if floatValue(normalizedRule["limit"]) <= 0 {
rules = append(rules, map[string]any{ continue
"metric": "rpm", }
"limit": limit, if floatValue(normalizedRule["maxWaitSeconds"]) <= 0 {
"windowSeconds": 60, normalizedRule["maxWaitSeconds"] = DefaultQueueMaxWaitSeconds
}) }
} } else {
if limit, ok := lowestPositiveLegacyLimit(legacyScopes, delete(normalizedRule, "maxWaitSeconds")
"max_tokens_per_minute", "maxTokensPerMinute", "tpm_total"); ok { }
rules = append(rules, map[string]any{ normalizedRules = append(normalizedRules, normalizedRule)
"metric": "tpm_total",
"limit": limit,
"windowSeconds": 60,
})
}
if len(rules) > 0 {
out["rules"] = rules
} }
out["rules"] = normalizedRules
return out return out
} }
type QueuePolicyRule struct {
Limit int
MaxWaitSeconds int
Policy map[string]any
}
func QueueRuleFromPolicy(policy map[string]any) (QueuePolicyRule, bool) {
normalized := NormalizeRateLimitPolicy(policy)
rules, _ := normalized["rules"].([]any)
for _, rawRule := range rules {
rule, _ := rawRule.(map[string]any)
if strings.TrimSpace(stringValue(rule["metric"])) != "queue_size" {
continue
}
limit := int(math.Floor(floatValue(rule["limit"])))
if limit <= 0 {
return QueuePolicyRule{}, false
}
maxWaitSeconds := int(math.Floor(floatValue(rule["maxWaitSeconds"])))
if maxWaitSeconds <= 0 {
maxWaitSeconds = DefaultQueueMaxWaitSeconds
}
return QueuePolicyRule{
Limit: limit,
MaxWaitSeconds: maxWaitSeconds,
Policy: normalized,
}, true
}
return QueuePolicyRule{}, false
}
// NormalizeAndValidateRateLimitPolicy validates only the canonical queue
// extension. Existing rate-limit metrics intentionally keep their historical
// permissive parsing contract.
func NormalizeAndValidateRateLimitPolicy(policy map[string]any) (map[string]any, error) {
if policy == nil {
return nil, nil
}
rulesValue, rulesPresent := policy["rules"]
if rulesPresent {
if _, ok := rulesValue.([]any); !ok {
return nil, fmt.Errorf("%w: rateLimitPolicy.rules must be an array", ErrInvalidRateLimitPolicy)
}
}
rules, _ := rulesValue.([]any)
for index, rawRule := range rules {
rule, ok := rawRule.(map[string]any)
if !ok {
return nil, fmt.Errorf("%w: rateLimitPolicy.rules[%d] must be an object", ErrInvalidRateLimitPolicy, index)
}
metric := strings.TrimSpace(stringValue(rule["metric"]))
maxWait, hasMaxWait := numericRuleValue(rule, "maxWaitSeconds")
if metric != "queue_size" {
if hasMaxWait {
return nil, fmt.Errorf("%w: rateLimitPolicy.rules[%d].maxWaitSeconds is only valid for queue_size", ErrInvalidRateLimitPolicy, index)
}
continue
}
limit, hasLimit := numericRuleValue(rule, "limit")
if !hasLimit {
return nil, fmt.Errorf("%w: rateLimitPolicy.rules[%d].limit must be an integer", ErrInvalidRateLimitPolicy, index)
}
if limit != math.Trunc(limit) || limit < 0 || limit > MaxQueueSize {
return nil, fmt.Errorf("%w: queue_size limit must be 0 or an integer between 1 and %d", ErrInvalidRateLimitPolicy, MaxQueueSize)
}
if limit == 0 {
continue
}
if hasMaxWait && (maxWait != math.Trunc(maxWait) || maxWait < 1 || maxWait > MaxQueueWaitSeconds) {
return nil, fmt.Errorf("%w: queue_size maxWaitSeconds must be an integer between 1 and %d", ErrInvalidRateLimitPolicy, MaxQueueWaitSeconds)
}
}
return NormalizeRateLimitPolicy(policy), nil
}
func numericRuleValue(rule map[string]any, key string) (float64, bool) {
raw, ok := rule[key]
if !ok || raw == nil {
return 0, false
}
switch raw.(type) {
case float64, float32, int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64:
return floatValue(raw), true
default:
return 0, false
}
}
func ConcurrentPolicyCapacity(policy map[string]any) (int, bool) { func ConcurrentPolicyCapacity(policy map[string]any) (int, bool) {
limit, ok := RateLimitPolicyMetric(policy, "concurrent") limit, ok := RateLimitPolicyMetric(policy, "concurrent")
if !ok || limit <= 0 { if !ok || limit <= 0 {
@@ -2,6 +2,66 @@ package store
import "testing" import "testing"
func TestNormalizeAndValidateQueuePolicy(t *testing.T) {
t.Run("defaults maximum wait and preserves extensions", func(t *testing.T) {
policy, err := NormalizeAndValidateRateLimitPolicy(map[string]any{
"strategy": "strict",
"rules": []any{map[string]any{
"metric": "queue_size",
"limit": 12,
"source": "admin",
}},
})
if err != nil {
t.Fatalf("NormalizeAndValidateRateLimitPolicy() error = %v", err)
}
queue, ok := QueueRuleFromPolicy(policy)
if !ok || queue.Limit != 12 || queue.MaxWaitSeconds != 600 {
t.Fatalf("queue rule = %+v, %v", queue, ok)
}
rules := policy["rules"].([]any)
rule := rules[0].(map[string]any)
if rule["source"] != "admin" {
t.Fatalf("unknown extension was lost: %+v", rule)
}
})
t.Run("normalizes zero to a removed queue rule", func(t *testing.T) {
policy, err := NormalizeAndValidateRateLimitPolicy(map[string]any{
"rules": []any{
map[string]any{"metric": "rpm", "limit": 30},
map[string]any{"metric": "queue_size", "limit": 0},
},
})
if err != nil {
t.Fatalf("NormalizeAndValidateRateLimitPolicy() error = %v", err)
}
if _, ok := QueueRuleFromPolicy(policy); ok {
t.Fatalf("zero queue rule remained enabled: %+v", policy)
}
rules := policy["rules"].([]any)
if len(rules) != 1 {
t.Fatalf("rules length = %d, want 1", len(rules))
}
})
for _, test := range []struct {
name string
policy map[string]any
}{
{name: "fractional queue", policy: map[string]any{"rules": []any{map[string]any{"metric": "queue_size", "limit": 1.5}}}},
{name: "queue too large", policy: map[string]any{"rules": []any{map[string]any{"metric": "queue_size", "limit": 10001}}}},
{name: "wait too large", policy: map[string]any{"rules": []any{map[string]any{"metric": "queue_size", "limit": 1, "maxWaitSeconds": 3601}}}},
{name: "wait on rpm", policy: map[string]any{"rules": []any{map[string]any{"metric": "rpm", "limit": 1, "maxWaitSeconds": 60}}}},
} {
t.Run(test.name, func(t *testing.T) {
if _, err := NormalizeAndValidateRateLimitPolicy(test.policy); err == nil {
t.Fatalf("invalid policy was accepted: %+v", test.policy)
}
})
}
}
func TestEffectiveRateLimitPolicyPrecedence(t *testing.T) { func TestEffectiveRateLimitPolicyPrecedence(t *testing.T) {
policy := func(limit float64) map[string]any { policy := func(limit float64) map[string]any {
return map[string]any{"rules": []any{map[string]any{"metric": "concurrent", "limit": limit}}} return map[string]any{"rules": []any{map[string]any{"metric": "concurrent", "limit": limit}}}
@@ -42,6 +42,10 @@ type ModelRateLimitStatus struct {
ModelCooldownUntil string `json:"modelCooldownUntil,omitempty"` ModelCooldownUntil string `json:"modelCooldownUntil,omitempty"`
Concurrent RateLimitMetricStatus `json:"concurrent"` Concurrent RateLimitMetricStatus `json:"concurrent"`
QueuedTasks float64 `json:"queuedTasks"` QueuedTasks float64 `json:"queuedTasks"`
WaitingSyncTasks int `json:"waitingSyncTasks,omitempty"`
WaitingAsyncTasks int `json:"waitingAsyncTasks,omitempty"`
OldestQueueWaitSeconds float64 `json:"oldestQueueWaitSeconds,omitempty"`
QueueLimit int `json:"queueLimit,omitempty"`
RPM RateLimitMetricStatus `json:"rpm"` RPM RateLimitMetricStatus `json:"rpm"`
TPM RateLimitMetricStatus `json:"tpm"` TPM RateLimitMetricStatus `json:"tpm"`
LoadRatio float64 `json:"loadRatio"` LoadRatio float64 `json:"loadRatio"`
@@ -151,6 +155,9 @@ func (s *Store) ListModelRateLimitStatuses(ctx context.Context) ([]ModelRateLimi
COALESCE(to_char(m.cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''), COALESCE(to_char(m.cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''),
COALESCE(con.active, 0)::float8, COALESCE(con.active, 0)::float8,
COALESCE(queued.waiting, 0)::float8, COALESCE(queued.waiting, 0)::float8,
COALESCE(admission.waiting_sync, 0)::int,
COALESCE(admission.waiting_async, 0)::int,
COALESCE(admission.oldest_wait_seconds, 0)::float8,
COALESCE(rpm.used_value, 0)::float8, COALESCE(rpm.reserved_value, 0)::float8, COALESCE(rpm.reset_at::text, ''), COALESCE(rpm.used_value, 0)::float8, COALESCE(rpm.reserved_value, 0)::float8, COALESCE(rpm.reset_at::text, ''),
COALESCE(tpm.used_value, 0)::float8, COALESCE(tpm.reserved_value, 0)::float8, COALESCE(tpm.reset_at::text, '') COALESCE(tpm.used_value, 0)::float8, COALESCE(tpm.reserved_value, 0)::float8, COALESCE(tpm.reset_at::text, '')
FROM platform_models m FROM platform_models m
@@ -192,6 +199,15 @@ LEFT JOIN (
) queued_sources ) queued_sources
GROUP BY queued_sources.platform_model_id GROUP BY queued_sources.platform_model_id
) queued ON queued.platform_model_id = m.id::text ) queued ON queued.platform_model_id = m.id::text
LEFT JOIN (
SELECT platform_model_id::text AS platform_model_id,
COUNT(*) FILTER (WHERE mode = 'sync') AS waiting_sync,
COUNT(*) FILTER (WHERE mode = 'async') AS waiting_async,
EXTRACT(EPOCH FROM now() - MIN(enqueued_at)) AS oldest_wait_seconds
FROM gateway_task_admissions
WHERE status = 'waiting'
GROUP BY platform_model_id
) admission ON admission.platform_model_id = m.id::text
LEFT JOIN ( LEFT JOIN (
SELECT DISTINCT ON (scope_key) scope_key, used_value, reserved_value, reset_at SELECT DISTINCT ON (scope_key) scope_key, used_value, reserved_value, reset_at
FROM gateway_rate_limit_counters FROM gateway_rate_limit_counters
@@ -231,6 +247,9 @@ ORDER BY p.priority ASC, m.model_name ASC`)
var modelCooldownUntil string var modelCooldownUntil string
var concurrentCurrent float64 var concurrentCurrent float64
var queuedTasks float64 var queuedTasks float64
var waitingSyncTasks int
var waitingAsyncTasks int
var oldestQueueWaitSeconds float64
var rpmUsed float64 var rpmUsed float64
var rpmReserved float64 var rpmReserved float64
var rpmResetAt string var rpmResetAt string
@@ -263,6 +282,9 @@ ORDER BY p.priority ASC, m.model_name ASC`)
&modelCooldownUntil, &modelCooldownUntil,
&concurrentCurrent, &concurrentCurrent,
&queuedTasks, &queuedTasks,
&waitingSyncTasks,
&waitingAsyncTasks,
&oldestQueueWaitSeconds,
&rpmUsed, &rpmUsed,
&rpmReserved, &rpmReserved,
&rpmResetAt, &rpmResetAt,
@@ -287,6 +309,12 @@ ORDER BY p.priority ASC, m.model_name ASC`)
item.ModelCooldownUntil = modelCooldownUntil item.ModelCooldownUntil = modelCooldownUntil
item.RateLimitPolicy = policy item.RateLimitPolicy = policy
item.QueuedTasks = queuedTasks item.QueuedTasks = queuedTasks
item.WaitingSyncTasks = waitingSyncTasks
item.WaitingAsyncTasks = waitingAsyncTasks
item.OldestQueueWaitSeconds = oldestQueueWaitSeconds
if queueRule, ok := QueueRuleFromPolicy(policy); ok {
item.QueueLimit = queueRule.Limit
}
item.Concurrent = metricStatus(concurrentCurrent, concurrentCurrent, 0, rateLimitForMetric(policy, "concurrent"), "") item.Concurrent = metricStatus(concurrentCurrent, concurrentCurrent, 0, rateLimitForMetric(policy, "concurrent"), "")
item.RPM = metricStatus(rpmUsed+rpmReserved, rpmUsed, rpmReserved, rateLimitForMetric(policy, "rpm"), rpmResetAt) item.RPM = metricStatus(rpmUsed+rpmReserved, rpmUsed, rpmReserved, rateLimitForMetric(policy, "rpm"), rpmResetAt)
item.TPM = metricStatus(tpmUsed+tpmReserved, tpmUsed, tpmReserved, tpmLimit(policy), tpmResetAt) item.TPM = metricStatus(tpmUsed+tpmReserved, tpmUsed, tpmReserved, tpmLimit(policy), tpmResetAt)
+117 -7
View File
@@ -20,6 +20,84 @@ type RuntimeRecoveryResult struct {
var ErrConcurrencyLeaseLost = errors.New("concurrency lease lost") var ErrConcurrencyLeaseLost = errors.New("concurrency lease lost")
// CheckRateLimits performs a non-consuming admission preflight for fixed-window
// limits. The same limits are checked and reserved again when execution starts.
func (s *Store) CheckRateLimits(ctx context.Context, reservations []RateLimitReservation) error {
for _, reservation := range reservations {
if reservation.Limit <= 0 || reservation.Amount <= 0 || reservation.Metric == "concurrent" || reservation.Metric == "queue_size" {
continue
}
if reservation.Metric == "" || reservation.Amount > reservation.Limit {
return &RateLimitExceededError{
ScopeType: reservation.ScopeType,
ScopeKey: reservation.ScopeKey,
ScopeName: reservation.ScopeName,
ScopeMetadata: reservation.ScopeMetadata,
Metric: reservation.Metric,
Limit: reservation.Limit,
Amount: reservation.Amount,
Projected: reservation.Amount,
WindowSeconds: reservation.WindowSeconds,
Policy: reservation.Policy,
Message: fmt.Sprintf("rate limit exceeded: %s request amount %.0f is greater than limit %.0f", reservation.Metric, reservation.Amount, reservation.Limit),
Retryable: false,
}
}
windowSeconds := reservation.WindowSeconds
if windowSeconds <= 0 {
windowSeconds = 60
}
var used float64
var reserved float64
var resetAt time.Time
err := s.pool.QueryRow(ctx, `
WITH bounds AS (
SELECT to_timestamp(floor(extract(epoch FROM now()) / $4::int) * $4::int) AS window_start,
to_timestamp(floor(extract(epoch FROM now()) / $4::int) * $4::int) + ($4::int * interval '1 second') AS reset_at
)
SELECT COALESCE(counters.used_value, 0)::float8,
COALESCE(counters.reserved_value, 0)::float8,
bounds.reset_at
FROM bounds
LEFT JOIN gateway_rate_limit_counters counters
ON counters.scope_type = $1
AND counters.scope_key = $2
AND counters.metric = $3
AND counters.window_start = bounds.window_start`,
reservation.ScopeType,
reservation.ScopeKey,
reservation.Metric,
windowSeconds,
).Scan(&used, &reserved, &resetAt)
if err != nil {
return err
}
current := used + reserved
if current+reservation.Amount > reservation.Limit {
return &RateLimitExceededError{
ScopeType: reservation.ScopeType,
ScopeKey: reservation.ScopeKey,
ScopeName: reservation.ScopeName,
ScopeMetadata: reservation.ScopeMetadata,
Metric: reservation.Metric,
Limit: reservation.Limit,
Amount: reservation.Amount,
Current: current,
Used: used,
Reserved: reserved,
Projected: current + reservation.Amount,
WindowSeconds: windowSeconds,
ResetAt: resetAt,
Policy: reservation.Policy,
Message: fmt.Sprintf("rate limit exceeded: %s window has no remaining capacity", reservation.Metric),
RetryAfter: retryAfterUntil(resetAt),
Retryable: true,
}
}
}
return nil
}
func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID string, reservations []RateLimitReservation) (RateLimitResult, error) { func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID string, reservations []RateLimitReservation) (RateLimitResult, error) {
tx, err := s.pool.Begin(ctx) tx, err := s.pool.Begin(ctx)
if err != nil { if err != nil {
@@ -30,6 +108,9 @@ func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID
lockKeys := make([]string, 0) lockKeys := make([]string, 0)
lockKeySet := make(map[string]struct{}) lockKeySet := make(map[string]struct{})
for _, reservation := range reservations { for _, reservation := range reservations {
if reservation.Metric == "queue_size" {
continue
}
if reservation.Metric != "concurrent" || reservation.Limit <= 0 || reservation.Amount <= 0 { if reservation.Metric != "concurrent" || reservation.Limit <= 0 || reservation.Amount <= 0 {
continue continue
} }
@@ -377,15 +458,28 @@ func (s *Store) RecoverInterruptedRuntimeState(ctx context.Context) (RuntimeReco
return RuntimeRecoveryResult{}, err return RuntimeRecoveryResult{}, err
} }
defer tx.Rollback(ctx) defer tx.Rollback(ctx)
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-runtime-recovery', 0))`); err != nil {
return RuntimeRecoveryResult{}, err
}
result := RuntimeRecoveryResult{} result := RuntimeRecoveryResult{}
rows, err := tx.Query(ctx, ` rows, err := tx.Query(ctx, `
UPDATE gateway_rate_limit_reservations UPDATE gateway_rate_limit_reservations reservation
SET status = 'released', SET status = 'released',
reason = 'server_restarted', reason = 'execution_lease_expired',
finalized_at = now(), finalized_at = now(),
updated_at = now() updated_at = now()
WHERE status = 'reserved' FROM gateway_tasks task
WHERE reservation.task_id = task.id
AND reservation.status = 'reserved'
AND (
task.status IN ('succeeded', 'failed', 'cancelled')
OR (
task.status IN ('queued', 'running')
AND task.execution_lease_expires_at IS NOT NULL
AND task.execution_lease_expires_at <= now()
)
)
RETURNING scope_type, scope_key, metric, window_start, reserved_amount::float8`) RETURNING scope_type, scope_key, metric, window_start, reserved_amount::float8`)
if err != nil { if err != nil {
return RuntimeRecoveryResult{}, err return RuntimeRecoveryResult{}, err
@@ -415,7 +509,7 @@ RETURNING scope_type, scope_key, metric, window_start, reserved_amount::float8`)
UPDATE gateway_concurrency_leases UPDATE gateway_concurrency_leases
SET released_at = now() SET released_at = now()
WHERE released_at IS NULL WHERE released_at IS NULL
AND expires_at > now()`) AND expires_at <= now()`)
if err != nil { if err != nil {
return RuntimeRecoveryResult{}, err return RuntimeRecoveryResult{}, err
} }
@@ -425,10 +519,22 @@ WHERE released_at IS NULL
UPDATE gateway_task_attempts UPDATE gateway_task_attempts
SET status = 'failed', SET status = 'failed',
retryable = true, retryable = true,
error_code = 'server_restarted', error_code = 'execution_lease_expired',
error_message = 'attempt interrupted by service restart', error_message = 'attempt execution lease expired',
finished_at = now() finished_at = now()
WHERE status = 'running'`) WHERE status = 'running'
AND EXISTS (
SELECT 1
FROM gateway_tasks task
WHERE task.id = gateway_task_attempts.task_id
AND (
task.status IN ('succeeded', 'failed', 'cancelled')
OR (
task.execution_lease_expires_at IS NOT NULL
AND task.execution_lease_expires_at <= now()
)
)
)`)
if err != nil { if err != nil {
return RuntimeRecoveryResult{}, err return RuntimeRecoveryResult{}, err
} }
@@ -448,6 +554,8 @@ SET status = 'queued',
updated_at = now() updated_at = now()
WHERE async_mode = true WHERE async_mode = true
AND status = 'running' AND status = 'running'
AND execution_lease_expires_at IS NOT NULL
AND execution_lease_expires_at <= now()
RETURNING id::text`) RETURNING id::text`)
if err != nil { if err != nil {
return RuntimeRecoveryResult{}, err return RuntimeRecoveryResult{}, err
@@ -516,6 +624,8 @@ SET status = 'failed',
updated_at = now() updated_at = now()
WHERE async_mode = false WHERE async_mode = false
AND status = 'running' AND status = 'running'
AND execution_lease_expires_at IS NOT NULL
AND execution_lease_expires_at <= now()
RETURNING id::text`) RETURNING id::text`)
if err != nil { if err != nil {
return RuntimeRecoveryResult{}, err return RuntimeRecoveryResult{}, err
@@ -76,6 +76,11 @@ func (s *Store) ListRuntimePolicySets(ctx context.Context) ([]RuntimePolicySet,
func (s *Store) CreateRuntimePolicySet(ctx context.Context, input RuntimePolicySetInput) (RuntimePolicySet, error) { func (s *Store) CreateRuntimePolicySet(ctx context.Context, input RuntimePolicySetInput) (RuntimePolicySet, error) {
input = normalizeRuntimePolicyInput(input) input = normalizeRuntimePolicyInput(input)
normalizedPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return RuntimePolicySet{}, err
}
input.RateLimitPolicy = normalizedPolicy
rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy)) rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy))
retryPolicy, _ := json.Marshal(emptyObjectIfNil(input.RetryPolicy)) retryPolicy, _ := json.Marshal(emptyObjectIfNil(input.RetryPolicy))
autoDisablePolicy, _ := json.Marshal(emptyObjectIfNil(input.AutoDisablePolicy)) autoDisablePolicy, _ := json.Marshal(emptyObjectIfNil(input.AutoDisablePolicy))
@@ -94,6 +99,11 @@ RETURNING `+runtimePolicyColumns,
func (s *Store) UpdateRuntimePolicySet(ctx context.Context, id string, input RuntimePolicySetInput) (RuntimePolicySet, error) { func (s *Store) UpdateRuntimePolicySet(ctx context.Context, id string, input RuntimePolicySetInput) (RuntimePolicySet, error) {
input = normalizeRuntimePolicyInput(input) input = normalizeRuntimePolicyInput(input)
normalizedPolicy, err := NormalizeAndValidateRateLimitPolicy(input.RateLimitPolicy)
if err != nil {
return RuntimePolicySet{}, err
}
input.RateLimitPolicy = normalizedPolicy
rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy)) rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.RateLimitPolicy))
retryPolicy, _ := json.Marshal(emptyObjectIfNil(input.RetryPolicy)) retryPolicy, _ := json.Marshal(emptyObjectIfNil(input.RetryPolicy))
autoDisablePolicy, _ := json.Marshal(emptyObjectIfNil(input.AutoDisablePolicy)) autoDisablePolicy, _ := json.Marshal(emptyObjectIfNil(input.AutoDisablePolicy))
+7
View File
@@ -78,6 +78,9 @@ type RateLimitExceededError struct {
Message string Message string
RetryAfter time.Duration RetryAfter time.Duration
Retryable bool Retryable bool
Reason string
QueueDepth int
QueueLimit int
} }
func (e *RateLimitExceededError) Error() string { func (e *RateLimitExceededError) Error() string {
@@ -90,6 +93,10 @@ func (e *RateLimitExceededError) Error() string {
return ErrRateLimited.Error() return ErrRateLimited.Error()
} }
func (e *RateLimitExceededError) ErrorCode() string {
return "rate_limit"
}
func (e *RateLimitExceededError) Unwrap() error { func (e *RateLimitExceededError) Unwrap() error {
return ErrRateLimited return ErrRateLimited
} }
+196 -12
View File
@@ -389,6 +389,82 @@ RETURNING `+gatewayTaskColumns, taskID, executionToken, int(leaseTTL/time.Second
return task, nil return task, nil
} }
func (s *Store) ClaimTaskPreparation(ctx context.Context, taskID string, executionToken string, leaseTTL time.Duration) (GatewayTask, error) {
if leaseTTL <= 0 {
leaseTTL = 5 * time.Minute
}
task, err := scanGatewayTask(s.pool.QueryRow(ctx, `
UPDATE gateway_tasks
SET execution_token = $2::uuid,
execution_lease_expires_at = now() + ($3::int * interval '1 second'),
locked_at = now(),
heartbeat_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status = 'queued'
AND next_run_at <= now()
AND (
execution_token IS NULL
OR execution_lease_expires_at IS NULL
OR execution_lease_expires_at <= now()
)
RETURNING `+gatewayTaskColumns, taskID, executionToken, int(leaseTTL/time.Second)))
if errors.Is(err, pgx.ErrNoRows) {
return GatewayTask{}, ErrTaskExecutionLeaseUnavailable
}
return task, err
}
func (s *Store) ReleaseTaskPreparation(ctx context.Context, taskID string, executionToken string) error {
tag, err := s.pool.Exec(ctx, `
UPDATE gateway_tasks
SET execution_token = NULL,
execution_lease_expires_at = NULL,
locked_at = NULL,
heartbeat_at = NULL,
updated_at = now()
WHERE id = $1::uuid
AND status = 'queued'
AND execution_token = $2::uuid`, taskID, executionToken)
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
return nil
}
func (s *Store) FailQueuedTask(ctx context.Context, taskID string, code string, message string) (GatewayTask, error) {
tag, err := s.pool.Exec(ctx, `
UPDATE gateway_tasks
SET status = 'failed',
error = NULL,
error_code = NULLIF($2, ''),
error_message = NULLIF($3, ''),
billing_status = CASE
WHEN run_mode <> 'production' OR gateway_user_id IS NULL THEN 'not_required'
ELSE 'released'
END,
billing_updated_at = now(),
locked_by = NULL,
locked_at = NULL,
heartbeat_at = NULL,
execution_token = NULL,
execution_lease_expires_at = NULL,
finished_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status = 'queued'`, taskID, strings.TrimSpace(code), truncateUTF8Bytes(message, 2048))
if err != nil {
return GatewayTask{}, err
}
if tag.RowsAffected() != 1 {
return GatewayTask{}, ErrTaskExecutionFinished
}
return s.GetTask(ctx, taskID)
}
func (s *Store) RenewTaskExecutionLease(ctx context.Context, taskID string, executionToken string, leaseTTL time.Duration) error { func (s *Store) RenewTaskExecutionLease(ctx context.Context, taskID string, executionToken string, leaseTTL time.Duration) error {
if leaseTTL <= 0 { if leaseTTL <= 0 {
leaseTTL = 5 * time.Minute leaseTTL = 5 * time.Minute
@@ -399,7 +475,7 @@ SET execution_lease_expires_at = now() + ($3::int * interval '1 second'),
heartbeat_at = now(), heartbeat_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid WHERE id = $1::uuid
AND status = 'running' AND status IN ('queued', 'running')
AND execution_token = $2::uuid`, taskID, executionToken, int(leaseTTL/time.Second)) AND execution_token = $2::uuid`, taskID, executionToken, int(leaseTTL/time.Second))
if err != nil { if err != nil {
return err return err
@@ -407,7 +483,7 @@ WHERE id = $1::uuid
if tag.RowsAffected() != 1 { if tag.RowsAffected() != 1 {
var stillRunning bool var stillRunning bool
if err := s.pool.QueryRow(ctx, ` if err := s.pool.QueryRow(ctx, `
SELECT COALESCE((SELECT status = 'running' FROM gateway_tasks WHERE id = $1::uuid), false)`, taskID).Scan(&stillRunning); err != nil { SELECT COALESCE((SELECT status IN ('queued', 'running') FROM gateway_tasks WHERE id = $1::uuid), false)`, taskID).Scan(&stillRunning); err != nil {
return err return err
} }
if !stillRunning { if !stillRunning {
@@ -428,7 +504,7 @@ SET status = 'running',
heartbeat_at = now(), heartbeat_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid WHERE id = $1::uuid
AND status = 'running' AND status IN ('queued', 'running')
AND execution_token = $2::uuid`, taskID, executionToken, modelType) AND execution_token = $2::uuid`, taskID, executionToken, modelType)
if err != nil { if err != nil {
return err return err
@@ -492,7 +568,7 @@ SET status = 'queued',
error_message = NULL, error_message = NULL,
updated_at = now() updated_at = now()
WHERE id = $1::uuid WHERE id = $1::uuid
AND status = 'running' AND status IN ('queued', 'running')
AND execution_token = $2::uuid AND execution_token = $2::uuid
RETURNING `+gatewayTaskColumns, taskID, executionToken, nextRunAt, strings.TrimSpace(queueKey))) RETURNING `+gatewayTaskColumns, taskID, executionToken, nextRunAt, strings.TrimSpace(queueKey)))
} }
@@ -566,7 +642,14 @@ func (s *Store) CancelQueuedTask(ctx context.Context, taskID string, message str
message = "任务已取消" message = "任务已取消"
} }
message = truncateUTF8Bytes(message, 2048) message = truncateUTF8Bytes(message, 2048)
tag, err := s.pool.Exec(ctx, ` var task GatewayTask
changed := false
err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil {
return err
}
var err error
task, err = scanGatewayTask(tx.QueryRow(ctx, `
UPDATE gateway_tasks UPDATE gateway_tasks
SET status = 'cancelled', SET status = 'cancelled',
error = NULL, error = NULL,
@@ -576,22 +659,123 @@ SET status = 'cancelled',
locked_by = NULL, locked_by = NULL,
locked_at = NULL, locked_at = NULL,
heartbeat_at = NULL, heartbeat_at = NULL,
execution_token = NULL,
execution_lease_expires_at = NULL,
finished_at = now(), finished_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid WHERE id = $1::uuid
AND status = 'queued' AND status = 'queued'
AND COALESCE(remote_task_id, '') = ''`, taskID, message) AND COALESCE(remote_task_id, '') = ''
RETURNING `+gatewayTaskColumns, taskID, message))
if errors.Is(err, pgx.ErrNoRows) {
return nil
}
if err != nil {
return err
}
changed = true
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
WHERE task_id = $1::uuid
AND released_at IS NULL`, taskID); err != nil {
return err
}
if _, err := tx.Exec(ctx, `DELETE FROM gateway_task_admissions WHERE task_id = $1::uuid`, taskID); err != nil {
return err
}
return notifyTaskAdmissionTx(ctx, tx, taskID)
})
if err != nil { if err != nil {
return GatewayTask{}, false, err return GatewayTask{}, false, err
} }
if tag.RowsAffected() == 0 { return task, changed, nil
return GatewayTask{}, false, nil }
func (s *Store) CancelTaskBeforeUpstreamSubmission(
ctx context.Context,
taskID string,
executionToken string,
message string,
) (GatewayTask, bool, error) {
message = strings.TrimSpace(message)
if message == "" {
message = "client disconnected before upstream submission"
} }
task, err := s.GetTask(ctx, taskID) message = truncateUTF8Bytes(message, 2048)
var task GatewayTask
changed := false
err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, "task-admission:"+taskID); err != nil {
return err
}
var err error
task, err = scanGatewayTask(tx.QueryRow(ctx, `
UPDATE gateway_tasks task
SET status = 'cancelled',
error = NULL,
error_code = 'client_disconnected',
error_message = NULLIF($3::text, ''),
billing_status = CASE
WHEN run_mode <> 'production' OR gateway_user_id IS NULL THEN 'not_required'
ELSE 'released'
END,
billing_updated_at = now(),
locked_by = NULL,
locked_at = NULL,
heartbeat_at = NULL,
execution_token = NULL,
execution_lease_expires_at = NULL,
finished_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status IN ('queued', 'running')
AND execution_token = NULLIF($2, '')::uuid
AND COALESCE(remote_task_id, '') = ''
AND NOT EXISTS (
SELECT 1
FROM gateway_task_attempts attempt
WHERE attempt.task_id = task.id
AND attempt.upstream_submission_status <> 'not_submitted'
)
RETURNING `+gatewayTaskColumns, taskID, strings.TrimSpace(executionToken), message))
if errors.Is(err, pgx.ErrNoRows) {
return nil
}
if err != nil {
return err
}
changed = true
if _, err := tx.Exec(ctx, `
DELETE FROM gateway_task_param_preprocessing_logs log
USING gateway_task_attempts attempt
WHERE log.attempt_id = attempt.id
AND attempt.task_id = $1::uuid
AND attempt.upstream_submission_status = 'not_submitted'`, taskID); err != nil {
return err
}
if _, err := tx.Exec(ctx, `
DELETE FROM gateway_task_attempts
WHERE task_id = $1::uuid
AND upstream_submission_status = 'not_submitted'`, taskID); err != nil {
return err
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
WHERE task_id = $1::uuid
AND released_at IS NULL`, taskID); err != nil {
return err
}
if _, err := tx.Exec(ctx, `DELETE FROM gateway_task_admissions WHERE task_id = $1::uuid`, taskID); err != nil {
return err
}
return notifyTaskAdmissionTx(ctx, tx, taskID)
})
if err != nil { if err != nil {
return GatewayTask{}, true, err return GatewayTask{}, false, err
} }
return task, true, nil return task, changed, nil
} }
// CancelSubmittedTask records a confirmed upstream cancellation. Callers must // CancelSubmittedTask records a confirmed upstream cancellation. Callers must
@@ -1506,7 +1690,7 @@ func (s *Store) FinishTaskFailure(ctx context.Context, input FinishTaskFailureIn
finished_at = now(), finished_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid WHERE id = $1::uuid
AND status = 'running' AND status IN ('queued', 'running')
AND execution_token = $10::uuid`, AND execution_token = $10::uuid`,
input.TaskID, input.TaskID,
message, message,
+162
View File
@@ -0,0 +1,162 @@
package store
import (
"context"
"errors"
"strings"
"time"
"github.com/jackc/pgx/v5"
)
const workerHeartbeatStaleAfter = 15 * time.Second
type WorkerRegistrationInput struct {
InstanceID string
PodUID string
PodName string
Site string
Revision string
DesiredCapacity int
}
type WorkerAllocation struct {
InstanceID string
DesiredCapacity int
Allocated int
ActiveInstances int
HeartbeatAt time.Time
}
func (s *Store) RegisterWorkerInstance(ctx context.Context, input WorkerRegistrationInput) (WorkerAllocation, error) {
input.InstanceID = strings.TrimSpace(input.InstanceID)
if input.InstanceID == "" {
return WorkerAllocation{}, errors.New("worker instance ID is required")
}
if input.DesiredCapacity < 0 {
return WorkerAllocation{}, errors.New("worker desired capacity cannot be negative")
}
tx, err := s.pool.Begin(ctx)
if err != nil {
return WorkerAllocation{}, err
}
defer tx.Rollback(ctx)
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended('gateway-worker-capacity', 0))`); err != nil {
return WorkerAllocation{}, err
}
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_worker_instances (
instance_id, pod_uid, pod_name, site, revision, status,
desired_capacity, allocated_capacity, started_at, heartbeat_at, updated_at
)
VALUES ($1, $2, $3, $4, $5, 'active', $6, 0, now(), now(), now())
ON CONFLICT (instance_id) DO UPDATE
SET pod_uid = EXCLUDED.pod_uid,
pod_name = EXCLUDED.pod_name,
site = EXCLUDED.site,
revision = EXCLUDED.revision,
status = 'active',
desired_capacity = EXCLUDED.desired_capacity,
heartbeat_at = now(),
updated_at = now()`,
input.InstanceID,
strings.TrimSpace(input.PodUID),
strings.TrimSpace(input.PodName),
strings.TrimSpace(input.Site),
strings.TrimSpace(input.Revision),
input.DesiredCapacity,
); err != nil {
return WorkerAllocation{}, err
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_worker_instances
SET allocated_capacity = 0,
updated_at = now()
WHERE status <> 'active'
OR heartbeat_at <= now() - $1::interval`, workerHeartbeatStaleAfter.String()); err != nil {
return WorkerAllocation{}, err
}
rows, err := tx.Query(ctx, `
SELECT instance_id
FROM gateway_worker_instances
WHERE status = 'active'
AND heartbeat_at > now() - $1::interval
ORDER BY instance_id ASC
FOR UPDATE`, workerHeartbeatStaleAfter.String())
if err != nil {
return WorkerAllocation{}, err
}
activeIDs := make([]string, 0)
for rows.Next() {
var instanceID string
if err := rows.Scan(&instanceID); err != nil {
rows.Close()
return WorkerAllocation{}, err
}
activeIDs = append(activeIDs, instanceID)
}
if err := rows.Err(); err != nil {
rows.Close()
return WorkerAllocation{}, err
}
rows.Close()
if len(activeIDs) == 0 {
return WorkerAllocation{}, errors.New("worker registration was not active after heartbeat")
}
base := input.DesiredCapacity / len(activeIDs)
remainder := input.DesiredCapacity % len(activeIDs)
allocated := 0
for index, instanceID := range activeIDs {
capacity := base
if index < remainder {
capacity++
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_worker_instances
SET desired_capacity = $2,
allocated_capacity = $3,
updated_at = now()
WHERE instance_id = $1`, instanceID, input.DesiredCapacity, capacity); err != nil {
return WorkerAllocation{}, err
}
if instanceID == input.InstanceID {
allocated = capacity
}
}
var heartbeatAt time.Time
if err := tx.QueryRow(ctx, `
SELECT heartbeat_at
FROM gateway_worker_instances
WHERE instance_id = $1`, input.InstanceID).Scan(&heartbeatAt); err != nil {
return WorkerAllocation{}, err
}
if err := tx.Commit(ctx); err != nil {
return WorkerAllocation{}, err
}
return WorkerAllocation{
InstanceID: input.InstanceID,
DesiredCapacity: input.DesiredCapacity,
Allocated: allocated,
ActiveInstances: len(activeIDs),
HeartbeatAt: heartbeatAt,
}, nil
}
func (s *Store) MarkWorkerDraining(ctx context.Context, instanceID string) error {
result, err := s.pool.Exec(ctx, `
UPDATE gateway_worker_instances
SET status = 'draining',
allocated_capacity = 0,
heartbeat_at = now(),
updated_at = now()
WHERE instance_id = $1`, strings.TrimSpace(instanceID))
if err != nil {
return err
}
if result.RowsAffected() == 0 {
return pgx.ErrNoRows
}
return nil
}
@@ -0,0 +1,75 @@
-- Shared admission state for non-text synchronous and asynchronous tasks.
CREATE TABLE IF NOT EXISTS gateway_task_admissions (
task_id uuid PRIMARY KEY REFERENCES gateway_tasks(id) ON DELETE CASCADE,
platform_id uuid NOT NULL REFERENCES integration_platforms(id) ON DELETE RESTRICT,
platform_model_id uuid NOT NULL REFERENCES platform_models(id) ON DELETE RESTRICT,
user_group_id uuid REFERENCES gateway_user_groups(id) ON DELETE SET NULL,
queue_key text NOT NULL,
mode text NOT NULL,
status text NOT NULL DEFAULT 'waiting',
priority integer NOT NULL DEFAULT 100,
enqueued_at timestamptz NOT NULL DEFAULT now(),
wait_deadline_at timestamptz NOT NULL,
admitted_at timestamptz,
waiter_id text,
waiter_lease_expires_at timestamptz,
reselect_requested_at timestamptz,
created_at timestamptz NOT NULL DEFAULT now(),
updated_at timestamptz NOT NULL DEFAULT now(),
CONSTRAINT gateway_task_admissions_mode_check
CHECK (mode IN ('sync', 'async')),
CONSTRAINT gateway_task_admissions_status_check
CHECK (status IN ('waiting', 'admitted')),
CONSTRAINT gateway_task_admissions_deadline_check
CHECK (wait_deadline_at >= enqueued_at),
CONSTRAINT gateway_task_admissions_sync_waiter_check
CHECK (
mode = 'async'
OR status = 'admitted'
OR (
NULLIF(waiter_id, '') IS NOT NULL
AND waiter_lease_expires_at IS NOT NULL
)
)
);
CREATE INDEX IF NOT EXISTS idx_task_admissions_platform_fifo
ON gateway_task_admissions(platform_model_id, priority, enqueued_at, task_id)
WHERE status = 'waiting';
CREATE INDEX IF NOT EXISTS idx_task_admissions_user_group_fifo
ON gateway_task_admissions(user_group_id, priority, enqueued_at, task_id)
WHERE status = 'waiting' AND user_group_id IS NOT NULL;
CREATE INDEX IF NOT EXISTS idx_task_admissions_waiter_expiry
ON gateway_task_admissions(waiter_lease_expires_at, task_id)
WHERE status = 'waiting' AND mode = 'sync';
CREATE INDEX IF NOT EXISTS idx_task_admissions_deadline
ON gateway_task_admissions(wait_deadline_at, task_id)
WHERE status = 'waiting';
CREATE INDEX IF NOT EXISTS idx_task_admissions_reselection
ON gateway_task_admissions(reselect_requested_at, priority, enqueued_at, task_id)
WHERE status = 'waiting' AND reselect_requested_at IS NOT NULL;
CREATE TABLE IF NOT EXISTS gateway_worker_instances (
instance_id text PRIMARY KEY,
pod_uid text NOT NULL DEFAULT '',
pod_name text NOT NULL DEFAULT '',
site text NOT NULL DEFAULT '',
revision text NOT NULL DEFAULT '',
status text NOT NULL DEFAULT 'active',
desired_capacity integer NOT NULL DEFAULT 0,
allocated_capacity integer NOT NULL DEFAULT 0,
started_at timestamptz NOT NULL DEFAULT now(),
heartbeat_at timestamptz NOT NULL DEFAULT now(),
updated_at timestamptz NOT NULL DEFAULT now(),
CONSTRAINT gateway_worker_instances_status_check
CHECK (status IN ('active', 'draining')),
CONSTRAINT gateway_worker_instances_capacity_check
CHECK (desired_capacity >= 0 AND allocated_capacity >= 0)
);
CREATE INDEX IF NOT EXISTS idx_worker_instances_active_heartbeat
ON gateway_worker_instances(status, heartbeat_at, instance_id);
@@ -498,6 +498,8 @@ export function PlatformManagementPanel(props: {
<Label>RPS / <Input value={form.rpsLimit} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, rpsLimit: event.target.value })} /></Label> <Label>RPS / <Input value={form.rpsLimit} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, rpsLimit: event.target.value })} /></Label>
<Label>TPM / Token<Input value={form.tpmLimit} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, tpmLimit: event.target.value })} /></Label> <Label>TPM / Token<Input value={form.tpmLimit} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, tpmLimit: event.target.value })} /></Label>
<Label><Input value={form.concurrencyLimit} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, concurrencyLimit: event.target.value })} /></Label> <Label><Input value={form.concurrencyLimit} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, concurrencyLimit: event.target.value })} /></Label>
<Label><Input value={form.queueSize} placeholder="0 或留空关闭排队" inputMode="numeric" onChange={(event) => setForm({ ...form, queueSize: event.target.value })} /></Label>
<Label><Input value={form.queueMaxWaitSeconds} placeholder="默认 600" inputMode="numeric" disabled={!Number(form.queueSize)} onChange={(event) => setForm({ ...form, queueMaxWaitSeconds: event.target.value })} /></Label>
<div className="platformTogglePair"> <div className="platformTogglePair">
<ToggleField checked={form.supportBase64Input} label="支持 Base64 输入" onChange={(checked) => setForm({ ...form, supportBase64Input: checked })} /> <ToggleField checked={form.supportBase64Input} label="支持 Base64 输入" onChange={(checked) => setForm({ ...form, supportBase64Input: checked })} />
<ToggleField checked={form.supportUrlInput} label="支持 URL 输入" onChange={(checked) => setForm({ ...form, supportUrlInput: checked })} /> <ToggleField checked={form.supportUrlInput} label="支持 URL 输入" onChange={(checked) => setForm({ ...form, supportUrlInput: checked })} />
@@ -567,6 +569,8 @@ function ModelBindingPolicy(props: { form: PlatformWizardForm; onChange: (value:
<Label> RPS<Input value={form.modelRpsLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelRpsLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label> <Label> RPS<Input value={form.modelRpsLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelRpsLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
<Label> TPM<Input value={form.modelTpmLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelTpmLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label> <Label> TPM<Input value={form.modelTpmLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelTpmLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
<Label><Input value={form.modelConcurrencyLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelConcurrencyLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label> <Label><Input value={form.modelConcurrencyLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelConcurrencyLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
<Label><Input value={form.modelQueueSize} placeholder="0 或留空关闭排队" inputMode="numeric" onChange={(event) => onChange({ ...form, modelQueueSize: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
<Label><Input value={form.modelQueueMaxWaitSeconds} placeholder="默认 600" inputMode="numeric" disabled={!Number(form.modelQueueSize)} onChange={(event) => onChange({ ...form, modelQueueMaxWaitSeconds: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
</> </>
)} )}
</div> </div>
@@ -1151,9 +1155,11 @@ function platformToForm(
const config = platform.config ?? {}; const config = platform.config ?? {};
const retryPolicy = platform.retryPolicy ?? {}; const retryPolicy = platform.retryPolicy ?? {};
const rateLimitPolicy = platform.rateLimitPolicy ?? {}; const rateLimitPolicy = platform.rateLimitPolicy ?? {};
const rateLimitParts = splitRateLimitPolicy(rateLimitPolicy);
const networkProxy = readNetworkProxyConfig(config); const networkProxy = readNetworkProxyConfig(config);
const currentModels = platformModels.filter((model) => model.platformId === platform.id); const currentModels = platformModels.filter((model) => model.platformId === platform.id);
const modelRateLimitPolicy = currentModels.find((model) => model.rateLimitPolicyMode === 'override')?.rateLimitPolicy ?? {}; const modelRateLimitPolicy = currentModels.find((model) => model.rateLimitPolicyMode === 'override')?.rateLimitPolicy ?? {};
const modelRateLimitParts = splitRateLimitPolicy(modelRateLimitPolicy);
return { return {
...createEmptyPlatformForm(platform.provider, defaults), ...createEmptyPlatformForm(platform.provider, defaults),
provider: platform.provider, provider: platform.provider,
@@ -1176,6 +1182,11 @@ function platformToForm(
rpsLimit: readLimit(rateLimitPolicy, 'rps'), rpsLimit: readLimit(rateLimitPolicy, 'rps'),
tpmLimit: readLimit(rateLimitPolicy, 'tpm_total'), tpmLimit: readLimit(rateLimitPolicy, 'tpm_total'),
concurrencyLimit: readLimit(rateLimitPolicy, 'concurrent'), concurrencyLimit: readLimit(rateLimitPolicy, 'concurrent'),
queueSize: readLimit(rateLimitPolicy, 'queue_size'),
queueMaxWaitSeconds: readRateLimitValue(rateLimitPolicy, 'queue_size', 'maxWaitSeconds', '600'),
rateLimitPolicyExtra: rateLimitParts.extra,
rateLimitRuleExtras: rateLimitParts.ruleExtras,
rateLimitPreservedRules: rateLimitParts.preservedRules,
testMode: readBoolean(config, 'testMode', false), testMode: readBoolean(config, 'testMode', false),
proxyMode: networkProxy.proxyMode, proxyMode: networkProxy.proxyMode,
httpProxy: networkProxy.httpProxy, httpProxy: networkProxy.httpProxy,
@@ -1190,6 +1201,11 @@ function platformToForm(
modelRpsLimit: readLimit(modelRateLimitPolicy, 'rps'), modelRpsLimit: readLimit(modelRateLimitPolicy, 'rps'),
modelTpmLimit: readLimit(modelRateLimitPolicy, 'tpm_total'), modelTpmLimit: readLimit(modelRateLimitPolicy, 'tpm_total'),
modelConcurrencyLimit: readLimit(modelRateLimitPolicy, 'concurrent'), modelConcurrencyLimit: readLimit(modelRateLimitPolicy, 'concurrent'),
modelQueueSize: readLimit(modelRateLimitPolicy, 'queue_size'),
modelQueueMaxWaitSeconds: readRateLimitValue(modelRateLimitPolicy, 'queue_size', 'maxWaitSeconds', '600'),
modelRateLimitPolicyExtra: modelRateLimitParts.extra,
modelRateLimitRuleExtras: modelRateLimitParts.ruleExtras,
modelRateLimitPreservedRules: modelRateLimitParts.preservedRules,
selectionMode: 'partial', selectionMode: 'partial',
}; };
} }
@@ -1283,6 +1299,36 @@ function readLimit(policy: Record<string, unknown>, metric: string) {
return typeof rule?.limit === 'number' ? String(rule.limit) : ''; return typeof rule?.limit === 'number' ? String(rule.limit) : '';
} }
function readRateLimitValue(policy: Record<string, unknown>, metric: string, key: string, fallback = '') {
const rules = Array.isArray(policy.rules) ? policy.rules : [];
const rule = rules.find((item): item is Record<string, unknown> => typeof item === 'object' && item !== null && item.metric === metric);
const value = rule?.[key];
return typeof value === 'number' && Number.isFinite(value) ? String(value) : fallback;
}
function splitRateLimitPolicy(policy: Record<string, unknown>) {
const editableMetrics = new Set(['rpm', 'rps', 'tpm_total', 'concurrent', 'queue_size']);
const extra = { ...policy };
delete extra.rules;
const ruleExtras: Record<string, Record<string, unknown>> = {};
const preservedRules: Array<Record<string, unknown>> = [];
const rules = Array.isArray(policy.rules) ? policy.rules : [];
rules.forEach((rawRule) => {
const rule = readRecord(rawRule);
const metric = typeof rule.metric === 'string' ? rule.metric.trim() : '';
if (!editableMetrics.has(metric)) {
preservedRules.push({ ...rule });
return;
}
const ruleExtra = { ...rule };
delete ruleExtra.metric;
delete ruleExtra.limit;
delete ruleExtra.maxWaitSeconds;
ruleExtras[metric] = ruleExtra;
});
return { extra, ruleExtras, preservedRules };
}
function numberText(value: unknown, fallback: string) { function numberText(value: unknown, fallback: string) {
return typeof value === 'number' && Number.isFinite(value) && value > 0 ? String(value) : fallback; return typeof value === 'number' && Number.isFinite(value) && value > 0 ? String(value) : fallback;
} }
+21 -6
View File
@@ -356,11 +356,12 @@ export function filterAndSortRealtimeLoad(
const cooling = cooldownRemainingMs(status.modelCooldownUntil, options.now) > 0 || const cooling = cooldownRemainingMs(status.modelCooldownUntil, options.now) > 0 ||
cooldownRemainingMs(status.platformCooldownUntil, options.now) > 0; cooldownRemainingMs(status.platformCooldownUntil, options.now) > 0;
const disabled = !status.enabled || status.platformStatus !== 'enabled'; const disabled = !status.enabled || status.platformStatus !== 'enabled';
const queuedTasks = admissionQueueDepth(status);
const matchesRuntime = !options.runtimeState || const matchesRuntime = !options.runtimeState ||
(options.runtimeState === 'high' && status.loadRatio >= 0.8) || (options.runtimeState === 'high' && status.loadRatio >= 0.8) ||
(options.runtimeState === 'queued' && status.queuedTasks > 0) || (options.runtimeState === 'queued' && queuedTasks > 0) ||
(options.runtimeState === 'attention' && (cooling || disabled)) || (options.runtimeState === 'attention' && (cooling || disabled)) ||
(options.runtimeState === 'normal' && !cooling && !disabled && status.loadRatio < 0.8 && status.queuedTasks <= 0); (options.runtimeState === 'normal' && !cooling && !disabled && status.loadRatio < 0.8 && queuedTasks <= 0);
return (!keyword || text.includes(keyword)) && return (!keyword || text.includes(keyword)) &&
(!options.platformId || status.platformId === options.platformId) && (!options.platformId || status.platformId === options.platformId) &&
(!options.modelType || status.modelType.includes(options.modelType)) && (!options.modelType || status.modelType.includes(options.modelType)) &&
@@ -375,7 +376,7 @@ export function filterAndSortRealtimeLoad(
} }
function realtimeLoadSortValue(status: ModelRateLimitStatus, field: RealtimeLoadSortField) { function realtimeLoadSortValue(status: ModelRateLimitStatus, field: RealtimeLoadSortField) {
if (field === 'queued') return status.queuedTasks; if (field === 'queued') return admissionQueueDepth(status);
if (field === 'concurrent') return status.concurrent.currentValue; if (field === 'concurrent') return status.concurrent.currentValue;
if (field === 'tpm') return status.tpm.currentValue; if (field === 'tpm') return status.tpm.currentValue;
if (field === 'rpm') return status.rpm.currentValue; if (field === 'rpm') return status.rpm.currentValue;
@@ -383,6 +384,13 @@ function realtimeLoadSortValue(status: ModelRateLimitStatus, field: RealtimeLoad
return status.loadRatio; return status.loadRatio;
} }
function admissionQueueDepth(status: ModelRateLimitStatus) {
if (status.waitingSyncTasks !== undefined || status.waitingAsyncTasks !== undefined) {
return (status.waitingSyncTasks ?? 0) + (status.waitingAsyncTasks ?? 0);
}
return status.queuedTasks ?? 0;
}
function platformDisplayName(platform: IntegrationPlatform) { function platformDisplayName(platform: IntegrationPlatform) {
return platform.internalName?.trim() || platform.name; return platform.internalName?.trim() || platform.name;
} }
@@ -627,11 +635,18 @@ function metricCell(metric: ModelRateLimitStatus['rpm'], includeReserved = false
} }
function concurrencyMetricCell(status: ModelRateLimitStatus) { function concurrencyMetricCell(status: ModelRateLimitStatus) {
const queuedTasks = status.queuedTasks ?? 0; const hasAdmissionBreakdown = status.waitingSyncTasks !== undefined || status.waitingAsyncTasks !== undefined;
const queuedTasks = hasAdmissionBreakdown
? (status.waitingSyncTasks ?? 0) + (status.waitingAsyncTasks ?? 0)
: (status.queuedTasks ?? 0);
const limitText = status.concurrent.limited ? formatLimit(status.concurrent.limitValue) : '不限'; const limitText = status.concurrent.limited ? formatLimit(status.concurrent.limitValue) : '不限';
const queueText = status.queueLimit
? `${formatLimit(queuedTasks)} / ${formatLimit(status.queueLimit)}`
: formatLimit(queuedTasks);
return ( return (
<span className="rateMetricCell" title="正在执行 / 并发上限 / 排队任务"> <span className="rateMetricCell" title={`正在执行 / 并发上限 / 排队任务${status.oldestQueueWaitSeconds ? `;最老等待 ${Math.round(status.oldestQueueWaitSeconds)}` : ''}`}>
<strong>{formatLimit(status.concurrent.currentValue)} / {limitText} / {formatLimit(queuedTasks)}</strong> <strong>{formatLimit(status.concurrent.currentValue)} / {limitText} / {queueText}</strong>
{hasAdmissionBreakdown && <small> {status.waitingSyncTasks ?? 0} / {status.waitingAsyncTasks ?? 0}</small>}
</span> </span>
); );
} }
@@ -0,0 +1,55 @@
import type { RuntimePolicySet } from '@easyai-ai-gateway/contracts';
import { describe, expect, it } from 'vitest';
import { formToPayload, policyToForm } from './RuntimePoliciesPanel';
describe('runtime queue policy editor', () => {
it('preserves unknown policy extensions while editing queue values', () => {
const policy = {
id: 'runtime-1',
policyKey: 'runtime-image',
name: 'Image runtime',
status: 'active',
rateLimitPolicy: {
extensionMode: 'legacy',
rules: [
{ metric: 'concurrent', limit: 2, leaseTtlSeconds: 90, vendorHint: 'keep' },
{ metric: 'queue_size', limit: 4, maxWaitSeconds: 120, vendorQueueMode: 'shared' },
{ metric: 'rps', limit: 5, windowSeconds: 1, vendorWindow: true },
],
},
retryPolicy: {},
autoDisablePolicy: {},
degradePolicy: {},
metadata: {},
createdAt: '2026-07-29T00:00:00Z',
updatedAt: '2026-07-29T00:00:00Z',
} as RuntimePolicySet;
const form = policyToForm(policy);
const payload = formToPayload({
...form,
queueSize: '6',
queueMaxWaitSeconds: '300',
});
expect(payload.rateLimitPolicy).toEqual({
extensionMode: 'legacy',
rules: [
{ metric: 'rps', limit: 5, windowSeconds: 1, vendorWindow: true },
{
vendorHint: 'keep',
metric: 'concurrent',
limit: 2,
windowSeconds: 60,
leaseTtlSeconds: 90,
},
{
vendorQueueMode: 'shared',
metric: 'queue_size',
limit: 6,
maxWaitSeconds: 300,
},
],
});
});
});
@@ -26,13 +26,18 @@ type RunnerActionRule = {
cooldownSeconds: string; cooldownSeconds: string;
}; };
type RuntimePolicyForm = { export type RuntimePolicyForm = {
policyKey: string; policyKey: string;
name: string; name: string;
description: string; description: string;
rpm: string; rpm: string;
tpm: string; tpm: string;
concurrency: string; concurrency: string;
queueSize: string;
queueMaxWaitSeconds: string;
rateLimitPolicyExtra: Record<string, unknown>;
rateLimitRuleExtras: Record<string, Record<string, unknown>>;
rateLimitPreservedRules: Array<Record<string, unknown>>;
retryEnabled: boolean; retryEnabled: boolean;
retryMaxAttempts: string; retryMaxAttempts: string;
retryAllowKeywords: string[]; retryAllowKeywords: string[];
@@ -275,11 +280,13 @@ export function RuntimePoliciesPanel(props: {
<Label className="spanTwo"><Input value={form.description} onChange={(event) => setForm({ ...form, description: event.target.value })} /></Label> <Label className="spanTwo"><Input value={form.description} onChange={(event) => setForm({ ...form, description: event.target.value })} /></Label>
<section className="runtimePolicySection spanTwo"> <section className="runtimePolicySection spanTwo">
<header><strong></strong><span>TPM / RPM / </span></header> <header><strong></strong><span>TPM / RPM / / </span></header>
<div className="runtimePolicyRows"> <div className="runtimePolicyRows">
<Label>RPM / <Input value={form.rpm} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, rpm: event.target.value })} /></Label> <Label>RPM / <Input value={form.rpm} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, rpm: event.target.value })} /></Label>
<Label>TPM / Token<Input value={form.tpm} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, tpm: event.target.value })} /></Label> <Label>TPM / Token<Input value={form.tpm} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, tpm: event.target.value })} /></Label>
<Label><Input value={form.concurrency} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, concurrency: event.target.value })} /></Label> <Label><Input value={form.concurrency} placeholder="不填或负数为不限" inputMode="numeric" onChange={(event) => setForm({ ...form, concurrency: event.target.value })} /></Label>
<Label><Input value={form.queueSize} placeholder="0 或留空关闭排队" inputMode="numeric" onChange={(event) => setForm({ ...form, queueSize: event.target.value })} /></Label>
<Label><Input value={form.queueMaxWaitSeconds} placeholder="默认 600" inputMode="numeric" disabled={!Number(form.queueSize)} onChange={(event) => setForm({ ...form, queueMaxWaitSeconds: event.target.value })} /></Label>
</div> </div>
</section> </section>
@@ -615,6 +622,11 @@ function createDefaultForm(policyKey = 'default-runtime-v1'): RuntimePolicyForm
rpm: '120', rpm: '120',
tpm: '240000', tpm: '240000',
concurrency: '6', concurrency: '6',
queueSize: '',
queueMaxWaitSeconds: '600',
rateLimitPolicyExtra: {},
rateLimitRuleExtras: {},
rateLimitPreservedRules: [],
retryEnabled: true, retryEnabled: true,
retryMaxAttempts: '2', retryMaxAttempts: '2',
retryAllowKeywords: ['rate_limit', 'timeout', 'server_error', 'network', '429', '5xx'], retryAllowKeywords: ['rate_limit', 'timeout', 'server_error', 'network', '429', '5xx'],
@@ -949,8 +961,9 @@ function defaultRunnerActionTarget(action: RunnerFailoverAction): RunnerFailover
return action === 'cooldown_and_next' ? 'model' : 'platform'; return action === 'cooldown_and_next' ? 'model' : 'platform';
} }
function policyToForm(policy: RuntimePolicySet): RuntimePolicyForm { export function policyToForm(policy: RuntimePolicySet): RuntimePolicyForm {
const rateRules = Array.isArray(policy.rateLimitPolicy?.rules) ? policy.rateLimitPolicy.rules : []; const rateRules = Array.isArray(policy.rateLimitPolicy?.rules) ? policy.rateLimitPolicy.rules : [];
const rateLimitParts = splitRuntimeRateLimitPolicy(readObject(policy.rateLimitPolicy));
const retry = readObject(policy.retryPolicy); const retry = readObject(policy.retryPolicy);
const disable = readObject(policy.autoDisablePolicy); const disable = readObject(policy.autoDisablePolicy);
const degrade = readObject(policy.degradePolicy); const degrade = readObject(policy.degradePolicy);
@@ -961,6 +974,11 @@ function policyToForm(policy: RuntimePolicySet): RuntimePolicyForm {
rpm: formRateLimitText(readRateLimit(rateRules, 'rpm')), rpm: formRateLimitText(readRateLimit(rateRules, 'rpm')),
tpm: formRateLimitText(readRateLimit(rateRules, 'tpm_total')), tpm: formRateLimitText(readRateLimit(rateRules, 'tpm_total')),
concurrency: formRateLimitText(readRateLimit(rateRules, 'concurrent')), concurrency: formRateLimitText(readRateLimit(rateRules, 'concurrent')),
queueSize: formRateLimitText(readRateLimit(rateRules, 'queue_size')),
queueMaxWaitSeconds: formRateLimitPropertyText(readRateLimitRule(rateRules, 'queue_size'), 'maxWaitSeconds', '600'),
rateLimitPolicyExtra: rateLimitParts.extra,
rateLimitRuleExtras: rateLimitParts.ruleExtras,
rateLimitPreservedRules: rateLimitParts.preservedRules,
retryEnabled: readBool(retry.enabled, true), retryEnabled: readBool(retry.enabled, true),
retryMaxAttempts: String(readNumber(retry.maxAttempts, 2)), retryMaxAttempts: String(readNumber(retry.maxAttempts, 2)),
retryAllowKeywords: tagsFromValue(retry.allowKeywords), retryAllowKeywords: tagsFromValue(retry.allowKeywords),
@@ -976,12 +994,15 @@ function policyToForm(policy: RuntimePolicySet): RuntimePolicyForm {
}; };
} }
function formToPayload(form: RuntimePolicyForm): RuntimePolicySetUpsertRequest { export function formToPayload(form: RuntimePolicyForm): RuntimePolicySetUpsertRequest {
return { return {
policyKey: form.policyKey.trim(), policyKey: form.policyKey.trim(),
name: form.name.trim(), name: form.name.trim(),
description: form.description.trim() || undefined, description: form.description.trim() || undefined,
rateLimitPolicy: { rules: rateLimitRules(form) }, rateLimitPolicy: {
...form.rateLimitPolicyExtra,
rules: rateLimitRules(form),
},
retryPolicy: { retryPolicy: {
enabled: form.retryEnabled, enabled: form.retryEnabled,
maxAttempts: positiveInt(form.retryMaxAttempts, 2), maxAttempts: positiveInt(form.retryMaxAttempts, 2),
@@ -1005,16 +1026,54 @@ function formToPayload(form: RuntimePolicyForm): RuntimePolicySetUpsertRequest {
function rateLimitRules(form: RuntimePolicyForm) { function rateLimitRules(form: RuntimePolicyForm) {
return [ return [
limitRule('rpm', form.rpm), ...form.rateLimitPreservedRules,
limitRule('tpm_total', form.tpm), limitRule('rpm', form.rpm, 60, undefined, form.rateLimitRuleExtras.rpm),
limitRule('concurrent', form.concurrency, 60, 120), limitRule('tpm_total', form.tpm, 60, undefined, form.rateLimitRuleExtras.tpm_total),
limitRule('concurrent', form.concurrency, 60, 120, form.rateLimitRuleExtras.concurrent),
queueLimitRule(form.queueSize, form.queueMaxWaitSeconds, form.rateLimitRuleExtras.queue_size),
].filter((rule): rule is NonNullable<ReturnType<typeof limitRule>> => Boolean(rule)); ].filter((rule): rule is NonNullable<ReturnType<typeof limitRule>> => Boolean(rule));
} }
function limitRule(metric: string, value: string, windowSeconds = 60, leaseTtlSeconds?: number) { function queueLimitRule(queueSize: string, maxWaitSeconds: string, extra: Record<string, unknown> = {}) {
if (!queueSize.trim() || Number(queueSize) === 0) return undefined;
const limit = boundedPositiveInt(queueSize, '排队长度', 10000);
const maxWait = maxWaitSeconds.trim() ? boundedPositiveInt(maxWaitSeconds, '最长等待', 3600) : 600;
return { ...extra, metric: 'queue_size', limit, maxWaitSeconds: maxWait };
}
function limitRule(metric: string, value: string, windowSeconds = 60, leaseTtlSeconds?: number, extra: Record<string, unknown> = {}) {
const limit = Number(value); const limit = Number(value);
if (!Number.isFinite(limit) || limit <= 0) return undefined; if (!Number.isFinite(limit) || limit <= 0) return undefined;
return { metric, limit, windowSeconds, leaseTtlSeconds }; return {
...extra,
metric,
limit,
windowSeconds: typeof extra.windowSeconds === 'number' ? extra.windowSeconds : windowSeconds,
leaseTtlSeconds: typeof extra.leaseTtlSeconds === 'number' ? extra.leaseTtlSeconds : leaseTtlSeconds,
};
}
function splitRuntimeRateLimitPolicy(policy: Record<string, unknown>) {
const editableMetrics = new Set(['rpm', 'tpm_total', 'concurrent', 'queue_size']);
const extra = { ...policy };
delete extra.rules;
const ruleExtras: Record<string, Record<string, unknown>> = {};
const preservedRules: Array<Record<string, unknown>> = [];
const rules = Array.isArray(policy.rules) ? policy.rules : [];
rules.forEach((rawRule) => {
const rule = readObject(rawRule);
const metric = typeof rule.metric === 'string' ? rule.metric.trim() : '';
if (!editableMetrics.has(metric)) {
preservedRules.push({ ...rule });
return;
}
const ruleExtra = { ...rule };
delete ruleExtra.metric;
delete ruleExtra.limit;
delete ruleExtra.maxWaitSeconds;
ruleExtras[metric] = ruleExtra;
});
return { extra, ruleExtras, preservedRules };
} }
function isDefaultPolicy(policy: RuntimePolicySet) { function isDefaultPolicy(policy: RuntimePolicySet) {
@@ -1026,7 +1085,9 @@ function rateLimitSummary(policy: RuntimePolicySet) {
const rpm = rateLimitText(readRateLimit(rules, 'rpm')); const rpm = rateLimitText(readRateLimit(rules, 'rpm'));
const tpm = rateLimitText(readRateLimit(rules, 'tpm_total')); const tpm = rateLimitText(readRateLimit(rules, 'tpm_total'));
const concurrent = rateLimitText(readRateLimit(rules, 'concurrent')); const concurrent = rateLimitText(readRateLimit(rules, 'concurrent'));
return `RPM ${rpm} / TPM ${tpm} / 并发 ${concurrent}`; const queue = readRateLimit(rules, 'queue_size');
const queueText = queue ? ` / 排队 ${rateLimitText(queue)} / 最长 ${formRateLimitPropertyText(readRateLimitRule(rules, 'queue_size'), 'maxWaitSeconds', '600')}` : '';
return `RPM ${rpm} / TPM ${tpm} / 并发 ${concurrent}${queueText}`;
} }
function retrySummary(policy: RuntimePolicySet) { function retrySummary(policy: RuntimePolicySet) {
@@ -1045,11 +1106,20 @@ function degradeSummary(policy: RuntimePolicySet) {
} }
function readRateLimit(rules: unknown[], metric: string) { function readRateLimit(rules: unknown[], metric: string) {
const rule = rules.find((item) => readObject(item).metric === metric); const rule = readRateLimitRule(rules, metric);
const limit = Number(readObject(rule).limit); const limit = Number(readObject(rule).limit);
return Number.isFinite(limit) ? limit : undefined; return Number.isFinite(limit) ? limit : undefined;
} }
function readRateLimitRule(rules: unknown[], metric: string) {
return rules.find((item) => readObject(item).metric === metric);
}
function formRateLimitPropertyText(rule: unknown, key: string, fallback = '') {
const value = Number(readObject(rule)[key]);
return Number.isFinite(value) && value > 0 ? String(value) : fallback;
}
function formRateLimitText(value: number | undefined) { function formRateLimitText(value: number | undefined) {
return value === undefined ? '' : String(value); return value === undefined ? '' : String(value);
} }
@@ -1090,6 +1160,14 @@ function positiveInt(value: string, fallback: number) {
return Number.isFinite(parsed) && parsed > 0 ? parsed : fallback; return Number.isFinite(parsed) && parsed > 0 ? parsed : fallback;
} }
function boundedPositiveInt(value: string, label: string, max: number) {
const parsed = Number(value);
if (!Number.isInteger(parsed) || parsed < 1 || parsed > max) {
throw new Error(`${label}必须是 1${max} 的整数`);
}
return parsed;
}
function readNumber(value: unknown, fallback: number) { function readNumber(value: unknown, fallback: number) {
const parsed = Number(value); const parsed = Number(value);
return Number.isFinite(parsed) ? parsed : fallback; return Number.isFinite(parsed) ? parsed : fallback;
@@ -34,6 +34,31 @@ describe('user group policy editors', () => {
expect(() => rateLimitPolicyFromForm(form.rules, form.extra)).toThrow('不能重复'); expect(() => rateLimitPolicyFromForm(form.rules, form.extra)).toThrow('不能重复');
}); });
it('normalizes zero queues and round-trips queue wait settings with extensions', () => {
const source = {
rules: [
{ metric: 'queue_size' as const, limit: 8, maxWaitSeconds: 75, source: 'custom' },
],
};
const form = rateLimitPolicyToForm(source);
expect(rateLimitPolicyFromForm(form.rules, form.extra)).toEqual(source);
expect(rateLimitPolicySummary(source)).toEqual(['排队 8 / 75秒']);
form.rules[0].limit = '0';
expect(rateLimitPolicyFromForm(form.rules, form.extra)).toBeUndefined();
});
it('defaults queue wait to 600 seconds and validates upper bounds', () => {
const form = rateLimitPolicyToForm({
rules: [{ metric: 'queue_size', limit: 2 }],
});
expect(rateLimitPolicyFromForm(form.rules, form.extra)).toEqual({
rules: [{ metric: 'queue_size', limit: 2, maxWaitSeconds: 600 }],
});
form.rules[0].maxWaitSeconds = '3601';
expect(() => rateLimitPolicyFromForm(form.rules, form.extra)).toThrow('13600');
});
it('round-trips quota values without changing their types', () => { it('round-trips quota values without changing their types', () => {
const source = { const source = {
monthlyResource: 10000, monthlyResource: 10000,
@@ -8,6 +8,7 @@ export type UserGroupRateLimitRuleForm = {
limit: string; limit: string;
windowSeconds: string; windowSeconds: string;
leaseTtlSeconds: string; leaseTtlSeconds: string;
maxWaitSeconds: string;
consume: string; consume: string;
extra: Record<string, unknown>; extra: Record<string, unknown>;
}; };
@@ -59,18 +60,34 @@ export function UserGroupRateLimitEditor(props: {
<span role="columnheader"></span> <span role="columnheader"></span>
<span role="columnheader"></span> <span role="columnheader"></span>
<span role="columnheader"> TTL</span> <span role="columnheader"> TTL</span>
<span role="columnheader"></span>
<span role="columnheader"></span> <span role="columnheader"></span>
<span role="columnheader"></span> <span role="columnheader"></span>
</div> </div>
{props.value.map((rule) => ( {props.value.map((rule) => (
<div className="userGroupPolicyRow" role="row" key={rule.id}> <div className="userGroupPolicyRow" role="row" key={rule.id}>
<Select aria-label="限流指标" size="sm" value={rule.metric} onChange={(event) => update(rule.id, { metric: event.target.value as RateLimitMetric })}> <Select
aria-label="限流指标"
size="sm"
value={rule.metric}
onChange={(event) => {
const metric = event.target.value as RateLimitMetric;
update(rule.id, {
metric,
windowSeconds: metric === 'queue_size' || metric === 'concurrent' ? '' : (rule.windowSeconds || '60'),
leaseTtlSeconds: metric === 'concurrent' ? (rule.leaseTtlSeconds || '120') : '',
maxWaitSeconds: metric === 'queue_size' ? (rule.maxWaitSeconds || '600') : '',
consume: metric === 'queue_size' ? '' : rule.consume,
});
}}
>
{rateLimitMetrics.map((item) => <option value={item.value} key={item.value}>{item.label}</option>)} {rateLimitMetrics.map((item) => <option value={item.value} key={item.value}>{item.label}</option>)}
</Select> </Select>
<Input aria-label={`${metricLabel(rule.metric)}上限`} size="sm" inputMode="decimal" value={rule.limit} placeholder="不限" onChange={(event) => update(rule.id, { limit: event.target.value })} /> <Input aria-label={`${metricLabel(rule.metric)}上限`} size="sm" inputMode="decimal" value={rule.limit} placeholder="不限" onChange={(event) => update(rule.id, { limit: event.target.value })} />
<Input aria-label={`${metricLabel(rule.metric)}窗口秒数`} size="sm" inputMode="numeric" value={rule.windowSeconds} placeholder={rule.metric === 'concurrent' ? '-' : '60'} onChange={(event) => update(rule.id, { windowSeconds: event.target.value })} /> <Input aria-label={`${metricLabel(rule.metric)}窗口秒数`} size="sm" inputMode="numeric" value={rule.windowSeconds} placeholder={rule.metric === 'queue_size' || rule.metric === 'concurrent' ? '-' : '60'} disabled={rule.metric === 'queue_size' || rule.metric === 'concurrent'} onChange={(event) => update(rule.id, { windowSeconds: event.target.value })} />
<Input aria-label={`${metricLabel(rule.metric)}租约 TTL`} size="sm" inputMode="numeric" value={rule.leaseTtlSeconds} placeholder={rule.metric === 'concurrent' ? '120' : '-'} onChange={(event) => update(rule.id, { leaseTtlSeconds: event.target.value })} /> <Input aria-label={`${metricLabel(rule.metric)}租约 TTL`} size="sm" inputMode="numeric" value={rule.leaseTtlSeconds} placeholder={rule.metric === 'concurrent' ? '120' : '-'} disabled={rule.metric !== 'concurrent'} onChange={(event) => update(rule.id, { leaseTtlSeconds: event.target.value })} />
<Select aria-label={`${metricLabel(rule.metric)}消费方式`} size="sm" value={rule.consume} onChange={(event) => update(rule.id, { consume: event.target.value })}> <Input aria-label={`${metricLabel(rule.metric)}最长等待`} size="sm" inputMode="numeric" value={rule.maxWaitSeconds} placeholder={rule.metric === 'queue_size' ? '600' : '-'} disabled={rule.metric !== 'queue_size'} onChange={(event) => update(rule.id, { maxWaitSeconds: event.target.value })} />
<Select aria-label={`${metricLabel(rule.metric)}消费方式`} size="sm" value={rule.consume} disabled={rule.metric === 'queue_size'} onChange={(event) => update(rule.id, { consume: event.target.value })}>
<option value=""></option> <option value=""></option>
<option value="fixed_window"></option> <option value="fixed_window"></option>
<option value="reserve_then_reconcile"></option> <option value="reserve_then_reconcile"></option>
@@ -172,6 +189,7 @@ export function rateLimitPolicyToForm(policy?: RateLimitPolicy) {
delete ruleExtra.limit; delete ruleExtra.limit;
delete ruleExtra.windowSeconds; delete ruleExtra.windowSeconds;
delete ruleExtra.leaseTtlSeconds; delete ruleExtra.leaseTtlSeconds;
delete ruleExtra.maxWaitSeconds;
delete ruleExtra.consume; delete ruleExtra.consume;
return { return {
id: nextPolicyEditorRowId('rate'), id: nextPolicyEditorRowId('rate'),
@@ -179,6 +197,7 @@ export function rateLimitPolicyToForm(policy?: RateLimitPolicy) {
limit: numberText(rule.limit), limit: numberText(rule.limit),
windowSeconds: numberText(rule.windowSeconds), windowSeconds: numberText(rule.windowSeconds),
leaseTtlSeconds: numberText(rule.leaseTtlSeconds), leaseTtlSeconds: numberText(rule.leaseTtlSeconds),
maxWaitSeconds: rule.metric === 'queue_size' ? numberText(rule.maxWaitSeconds) || '600' : '',
consume: typeof rule.consume === 'string' ? rule.consume : '', consume: typeof rule.consume === 'string' ? rule.consume : '',
extra: ruleExtra, extra: ruleExtra,
}; };
@@ -193,17 +212,24 @@ export function rateLimitPolicyFromForm(
const metrics = new Set<string>(); const metrics = new Set<string>();
const normalizedRules = rules.flatMap((rule) => { const normalizedRules = rules.flatMap((rule) => {
if (!rule.limit.trim()) return []; if (!rule.limit.trim()) return [];
if (rule.metric === 'queue_size' && Number(rule.limit) === 0) return [];
if (metrics.has(rule.metric)) throw new Error(`限流指标 ${metricLabel(rule.metric)} 不能重复`); if (metrics.has(rule.metric)) throw new Error(`限流指标 ${metricLabel(rule.metric)} 不能重复`);
metrics.add(rule.metric); metrics.add(rule.metric);
const limit = positiveNumber(rule.limit, `${metricLabel(rule.metric)}上限`); const limit = rule.metric === 'queue_size'
const windowSeconds = optionalPositiveInteger(rule.windowSeconds, `${metricLabel(rule.metric)}窗口`); ? boundedPositiveInteger(rule.limit, `${metricLabel(rule.metric)}上限`, 10000)
const leaseTtlSeconds = optionalPositiveInteger(rule.leaseTtlSeconds, `${metricLabel(rule.metric)}租约 TTL`); : positiveNumber(rule.limit, `${metricLabel(rule.metric)}上限`);
const windowSeconds = rule.metric === 'queue_size' ? undefined : optionalPositiveInteger(rule.windowSeconds, `${metricLabel(rule.metric)}窗口`);
const leaseTtlSeconds = rule.metric === 'queue_size' ? undefined : optionalPositiveInteger(rule.leaseTtlSeconds, `${metricLabel(rule.metric)}租约 TTL`);
const maxWaitSeconds = rule.metric === 'queue_size'
? boundedPositiveInteger(rule.maxWaitSeconds || '600', `${metricLabel(rule.metric)}最长等待`, 3600)
: undefined;
return [{ return [{
...rule.extra, ...rule.extra,
metric: rule.metric, metric: rule.metric,
limit, limit,
...(windowSeconds ? { windowSeconds } : {}), ...(windowSeconds ? { windowSeconds } : {}),
...(leaseTtlSeconds ? { leaseTtlSeconds } : {}), ...(leaseTtlSeconds ? { leaseTtlSeconds } : {}),
...(maxWaitSeconds ? { maxWaitSeconds } : {}),
...(rule.consume ? { consume: rule.consume } : {}), ...(rule.consume ? { consume: rule.consume } : {}),
}]; }];
}); });
@@ -242,7 +268,9 @@ export function rateLimitPolicySummary(policy?: RateLimitPolicy) {
const rules = Array.isArray(policy?.rules) ? policy.rules : []; const rules = Array.isArray(policy?.rules) ? policy.rules : [];
const parts = rules const parts = rules
.filter((rule) => Number.isFinite(Number(rule.limit)) && Number(rule.limit) > 0) .filter((rule) => Number.isFinite(Number(rule.limit)) && Number(rule.limit) > 0)
.map((rule) => `${metricShortLabel(rule.metric)} ${compactNumber(Number(rule.limit))}`); .map((rule) => rule.metric === 'queue_size'
? `${metricShortLabel(rule.metric)} ${compactNumber(Number(rule.limit))} / ${Number(rule.maxWaitSeconds) || 600}`
: `${metricShortLabel(rule.metric)} ${compactNumber(Number(rule.limit))}`);
if (parts.length) return parts; if (parts.length) return parts;
return Object.entries(policy ?? {}) return Object.entries(policy ?? {})
.filter(([key]) => key !== 'rules') .filter(([key]) => key !== 'rules')
@@ -261,6 +289,7 @@ function createRateLimitRuleForm(metric: RateLimitMetric): UserGroupRateLimitRul
limit: '', limit: '',
windowSeconds: metric === 'concurrent' ? '' : '60', windowSeconds: metric === 'concurrent' ? '' : '60',
leaseTtlSeconds: metric === 'concurrent' ? '120' : '', leaseTtlSeconds: metric === 'concurrent' ? '120' : '',
maxWaitSeconds: metric === 'queue_size' ? '600' : '',
consume: '', consume: '',
extra: {}, extra: {},
}; };
@@ -321,6 +350,14 @@ function optionalPositiveInteger(value: string, label: string) {
return parsed; return parsed;
} }
function boundedPositiveInteger(value: string, label: string, max: number) {
const parsed = Number(value);
if (!Number.isInteger(parsed) || parsed < 1 || parsed > max) {
throw new Error(`${label}必须是 1${max} 的整数`);
}
return parsed;
}
function serializeQuotaValue(value: unknown): Pick<UserGroupQuotaEntryForm, 'valueType' | 'value'> { function serializeQuotaValue(value: unknown): Pick<UserGroupQuotaEntryForm, 'valueType' | 'value'> {
if (typeof value === 'number') return { valueType: 'number', value: String(value) }; if (typeof value === 'number') return { valueType: 'number', value: String(value) };
if (typeof value === 'boolean') return { valueType: 'boolean', value: String(value) }; if (typeof value === 'boolean') return { valueType: 'boolean', value: String(value) };
+75 -1
View File
@@ -1,6 +1,6 @@
import type { BaseModelCatalogItem } from '@easyai-ai-gateway/contracts'; import type { BaseModelCatalogItem } from '@easyai-ai-gateway/contracts';
import { describe, expect, it } from 'vitest'; import { describe, expect, it } from 'vitest';
import { createEmptyPlatformForm, platformModelPayloads } from './platform-form'; import { createEmptyPlatformForm, platformModelPayloads, platformPayload, validatePlatformForm } from './platform-form';
const baseModel = { const baseModel = {
id: 'model-1', id: 'model-1',
@@ -53,4 +53,78 @@ describe('platform model rate limit payload', () => {
expect(payload.rateLimitPolicyMode).toBe('override'); expect(payload.rateLimitPolicyMode).toBe('override');
expect(payload.rateLimitPolicy).toEqual({ rules: [] }); expect(payload.rateLimitPolicy).toEqual({ rules: [] });
}); });
it('serializes queue length and defaults maximum wait to 600 seconds', () => {
const form = {
...selectedForm(),
name: 'Queue platform',
baseUrl: 'https://example.test',
testMode: true,
queueSize: '12',
queueMaxWaitSeconds: '',
modelOverrideRateLimit: true,
modelRateLimitOverrideTouched: true,
modelQueueSize: '4',
modelQueueMaxWaitSeconds: '90',
};
expect(platformPayload(form).rateLimitPolicy).toEqual({
rules: [{ metric: 'queue_size', limit: 12, maxWaitSeconds: 600 }],
});
expect(platformModelPayloads([baseModel], form)[0].rateLimitPolicy).toEqual({
rules: [{ metric: 'queue_size', limit: 4, maxWaitSeconds: 90 }],
});
});
it('rejects queue values outside the shared policy bounds', () => {
const form = {
...selectedForm(),
name: 'Queue platform',
baseUrl: 'https://example.test',
testMode: true,
queueSize: '10001',
};
expect(validatePlatformForm(form, 1)).toContain('110000');
expect(validatePlatformForm({ ...form, queueSize: '2', queueMaxWaitSeconds: '3601' }, 1)).toContain('13600');
});
it('preserves unknown policy fields and non-editable rules', () => {
const form = {
...selectedForm(),
name: 'Compatible platform',
baseUrl: 'https://example.test',
testMode: true,
concurrencyLimit: '3',
queueSize: '6',
rateLimitPolicyExtra: { extensionMode: 'legacy' },
rateLimitRuleExtras: {
concurrent: { consume: 'lease', vendorHint: 'keep' },
queue_size: { vendorQueueMode: 'shared' },
},
rateLimitPreservedRules: [
{ metric: 'tpm_input', limit: 1234, vendorWindow: true },
],
};
expect(platformPayload(form).rateLimitPolicy).toEqual({
extensionMode: 'legacy',
rules: [
{ metric: 'tpm_input', limit: 1234, vendorWindow: true },
{
consume: 'lease',
vendorHint: 'keep',
metric: 'concurrent',
limit: 3,
windowSeconds: 60,
leaseTtlSeconds: 120,
},
{
vendorQueueMode: 'shared',
metric: 'queue_size',
limit: 6,
maxWaitSeconds: 600,
},
],
});
});
}); });
+81 -9
View File
@@ -29,6 +29,11 @@ export interface PlatformWizardForm {
rpsLimit: string; rpsLimit: string;
tpmLimit: string; tpmLimit: string;
concurrencyLimit: string; concurrencyLimit: string;
queueSize: string;
queueMaxWaitSeconds: string;
rateLimitPolicyExtra: Record<string, unknown>;
rateLimitRuleExtras: Record<string, Record<string, unknown>>;
rateLimitPreservedRules: Array<Record<string, unknown>>;
testMode: boolean; testMode: boolean;
proxyMode: 'none' | 'global' | 'custom'; proxyMode: 'none' | 'global' | 'custom';
httpProxy: string; httpProxy: string;
@@ -46,6 +51,11 @@ export interface PlatformWizardForm {
modelRpsLimit: string; modelRpsLimit: string;
modelTpmLimit: string; modelTpmLimit: string;
modelConcurrencyLimit: string; modelConcurrencyLimit: string;
modelQueueSize: string;
modelQueueMaxWaitSeconds: string;
modelRateLimitPolicyExtra: Record<string, unknown>;
modelRateLimitRuleExtras: Record<string, Record<string, unknown>>;
modelRateLimitPreservedRules: Array<Record<string, unknown>>;
selectionMode: 'all' | 'partial'; selectionMode: 'all' | 'partial';
selectedModelIds: string[]; selectedModelIds: string[];
} }
@@ -89,6 +99,11 @@ export function createEmptyPlatformForm(provider = '', defaults?: ProviderConnec
rpsLimit: '', rpsLimit: '',
tpmLimit: '', tpmLimit: '',
concurrencyLimit: '', concurrencyLimit: '',
queueSize: '',
queueMaxWaitSeconds: '600',
rateLimitPolicyExtra: {},
rateLimitRuleExtras: {},
rateLimitPreservedRules: [],
testMode: false, testMode: false,
proxyMode: 'none', proxyMode: 'none',
httpProxy: '', httpProxy: '',
@@ -106,6 +121,11 @@ export function createEmptyPlatformForm(provider = '', defaults?: ProviderConnec
modelRpsLimit: '', modelRpsLimit: '',
modelTpmLimit: '', modelTpmLimit: '',
modelConcurrencyLimit: '', modelConcurrencyLimit: '',
modelQueueSize: '',
modelQueueMaxWaitSeconds: '600',
modelRateLimitPolicyExtra: {},
modelRateLimitRuleExtras: {},
modelRateLimitPreservedRules: [],
selectionMode: 'partial', selectionMode: 'partial',
selectedModelIds: [], selectedModelIds: [],
}; };
@@ -194,6 +214,11 @@ export function platformModelPayloads(models: BaseModelCatalogItem[], form: Plat
rpsLimit: form.modelRpsLimit, rpsLimit: form.modelRpsLimit,
tpmLimit: form.modelTpmLimit, tpmLimit: form.modelTpmLimit,
concurrencyLimit: form.modelConcurrencyLimit, concurrencyLimit: form.modelConcurrencyLimit,
queueSize: form.modelQueueSize,
queueMaxWaitSeconds: form.modelQueueMaxWaitSeconds,
rateLimitPolicyExtra: form.modelRateLimitPolicyExtra,
rateLimitRuleExtras: form.modelRateLimitRuleExtras,
rateLimitPreservedRules: form.modelRateLimitPreservedRules,
}) : undefined, }) : undefined,
rateLimitPolicyMode: form.modelRateLimitOverrideTouched rateLimitPolicyMode: form.modelRateLimitOverrideTouched
? (form.modelOverrideRateLimit ? 'override' : 'inherit') ? (form.modelOverrideRateLimit ? 'override' : 'inherit')
@@ -326,6 +351,12 @@ export function validatePlatformForm(form: PlatformWizardForm, selectedCount: nu
return '请填写 AccessKey 和 SecretKey,或开启测试模式。'; return '请填写 AccessKey 和 SecretKey,或开启测试模式。';
} }
if (form.proxyMode === 'custom' && !form.httpProxy.trim()) return '请填写自定义 HTTP 代理地址。'; if (form.proxyMode === 'custom' && !form.httpProxy.trim()) return '请填写自定义 HTTP 代理地址。';
const platformQueueError = validateQueueRule(form.queueSize, form.queueMaxWaitSeconds, '平台');
if (platformQueueError) return platformQueueError;
if (form.modelOverrideRateLimit) {
const modelQueueError = validateQueueRule(form.modelQueueSize, form.modelQueueMaxWaitSeconds, '模型');
if (modelQueueError) return modelQueueError;
}
if (selectedCount === 0) return '请至少添加一个模型。'; if (selectedCount === 0) return '请至少添加一个模型。';
return ''; return '';
} }
@@ -388,14 +419,52 @@ function credentialPreviewValue(preview: Record<string, unknown> | undefined, ..
return ''; return '';
} }
function rateLimitPolicyPayload(form: Pick<PlatformWizardForm, 'rpmLimit' | 'rpsLimit' | 'tpmLimit' | 'concurrencyLimit'>) { function rateLimitPolicyPayload(form: Pick<
PlatformWizardForm,
| 'rpmLimit'
| 'rpsLimit'
| 'tpmLimit'
| 'concurrencyLimit'
| 'queueSize'
| 'queueMaxWaitSeconds'
| 'rateLimitPolicyExtra'
| 'rateLimitRuleExtras'
| 'rateLimitPreservedRules'
>) {
const rules = [ const rules = [
limitRule('rpm', form.rpmLimit), ...form.rateLimitPreservedRules,
limitRule('rps', form.rpsLimit, 1), limitRule('rpm', form.rpmLimit, 60, form.rateLimitRuleExtras.rpm),
limitRule('tpm_total', form.tpmLimit), limitRule('rps', form.rpsLimit, 1, form.rateLimitRuleExtras.rps),
limitRule('concurrent', form.concurrencyLimit), limitRule('tpm_total', form.tpmLimit, 60, form.rateLimitRuleExtras.tpm_total),
limitRule('concurrent', form.concurrencyLimit, 60, form.rateLimitRuleExtras.concurrent),
queueLimitRule(form.queueSize, form.queueMaxWaitSeconds, form.rateLimitRuleExtras.queue_size),
].filter((rule): rule is NonNullable<ReturnType<typeof limitRule>> => Boolean(rule)); ].filter((rule): rule is NonNullable<ReturnType<typeof limitRule>> => Boolean(rule));
return { rules }; return { ...form.rateLimitPolicyExtra, rules };
}
function queueLimitRule(queueSize: string, maxWaitSeconds: string, extra: Record<string, unknown> = {}) {
const limit = Number(queueSize);
if (!Number.isFinite(limit) || limit <= 0) return undefined;
const parsedMaxWait = Number(maxWaitSeconds);
return {
...extra,
metric: 'queue_size',
limit,
maxWaitSeconds: Number.isInteger(parsedMaxWait) && parsedMaxWait > 0 ? parsedMaxWait : 600,
};
}
function validateQueueRule(queueSize: string, maxWaitSeconds: string, scopeLabel: string) {
if (!queueSize.trim() || Number(queueSize) === 0) return '';
const size = Number(queueSize);
if (!Number.isInteger(size) || size < 1 || size > 10000) {
return `${scopeLabel}排队长度必须是 1–10000 的整数,0 或留空表示关闭排队。`;
}
const wait = maxWaitSeconds.trim() ? Number(maxWaitSeconds) : 600;
if (!Number.isInteger(wait) || wait < 1 || wait > 3600) {
return `${scopeLabel}最长等待必须是 1–3600 秒的整数。`;
}
return '';
} }
function networkProxyPayload(form: PlatformWizardForm) { function networkProxyPayload(form: PlatformWizardForm) {
@@ -408,14 +477,17 @@ function networkProxyPayload(form: PlatformWizardForm) {
return { mode: form.proxyMode }; return { mode: form.proxyMode };
} }
function limitRule(metric: string, value: string, windowSeconds = 60) { function limitRule(metric: string, value: string, windowSeconds = 60, extra: Record<string, unknown> = {}) {
const limit = positiveNumber(value, 0); const limit = positiveNumber(value, 0);
if (!limit) return undefined; if (!limit) return undefined;
return { return {
...extra,
metric, metric,
limit, limit,
windowSeconds, windowSeconds: typeof extra.windowSeconds === 'number' ? extra.windowSeconds : windowSeconds,
leaseTtlSeconds: metric === 'concurrent' ? 120 : 70, leaseTtlSeconds: typeof extra.leaseTtlSeconds === 'number'
? extra.leaseTtlSeconds
: (metric === 'concurrent' ? 120 : 70),
}; };
} }
+1 -1
View File
@@ -855,7 +855,7 @@
} }
.userGroupRateLimitTable .userGroupPolicyRow { .userGroupRateLimitTable .userGroupPolicyRow {
grid-template-columns: minmax(170px, 1.15fr) minmax(100px, 0.65fr) minmax(112px, 0.65fr) minmax(120px, 0.72fr) minmax(156px, 0.9fr) 42px; grid-template-columns: minmax(170px, 1.15fr) minmax(100px, 0.65fr) minmax(112px, 0.65fr) minmax(120px, 0.72fr) minmax(120px, 0.72fr) minmax(156px, 0.9fr) 42px;
} }
.userGroupQuotaTable .userGroupPolicyRow { .userGroupQuotaTable .userGroupPolicyRow {
@@ -84,6 +84,16 @@ verify_site() {
wait_for_url "$site Web" "http://$address:$web_port/" 'EasyAI AI Gateway' wait_for_url "$site Web" "http://$address:$web_port/" 'EasyAI AI Gateway'
} }
rollout_worker_site() {
local site=$1
local api_image=$2
"${kubectl[@]}" set image "deployment/easyai-worker-$site" -n "$NAMESPACE" \
"worker=$api_image"
"${kubectl[@]}" rollout status "deployment/easyai-worker-$site" -n "$NAMESPACE" --timeout=300s
[[ $("${kubectl[@]}" get deployment "easyai-worker-$site" -n "$NAMESPACE" \
-o jsonpath='{.spec.template.spec.containers[?(@.name=="worker")].env[?(@.name=="AI_GATEWAY_PROCESS_ROLE")].value}') == worker ]]
}
rollout_site() { rollout_site() {
local site=$1 local site=$1
local api_image=$2 local api_image=$2
@@ -236,11 +246,20 @@ activate_manifest() {
run_migrator "$source_sha" "$api_image" run_migrator "$source_sha" "$api_image"
fi fi
if ! rollout_site hongkong "$api_image" "$web_image" "$api_changed" "$web_changed" || if { [[ $api_changed != true ]] ||
! rollout_site ningbo "$api_image" "$web_image" "$api_changed" "$web_changed" || { rollout_worker_site hongkong "$api_image" &&
! wait_for_url 'public readiness' "$PUBLIC_BASE_URL/api/v1/readyz" '"ok":true'; then rollout_worker_site ningbo "$api_image"; }; } &&
rollout_site hongkong "$api_image" "$web_image" "$api_changed" "$web_changed" &&
rollout_site ningbo "$api_image" "$web_image" "$api_changed" "$web_changed" &&
wait_for_url 'public readiness' "$PUBLIC_BASE_URL/api/v1/readyz" '"ok":true'; then
:
else
if [[ -n $previous_api && -n $previous_web ]]; then if [[ -n $previous_api && -n $previous_web ]]; then
echo '[cluster-release] restoring previous application images' >&2 echo '[cluster-release] restoring previous application images' >&2
if [[ $api_changed == true ]]; then
rollout_worker_site hongkong "$previous_api" || true
rollout_worker_site ningbo "$previous_api" || true
fi
rollout_site hongkong "$previous_api" "$previous_web" "$api_changed" "$web_changed" || true rollout_site hongkong "$previous_api" "$previous_web" "$api_changed" "$web_changed" || true
rollout_site ningbo "$previous_api" "$previous_web" "$api_changed" "$web_changed" || true rollout_site ningbo "$previous_api" "$previous_web" "$api_changed" "$web_changed" || true
fi fi
@@ -266,6 +285,10 @@ case ${1:-} in
node "$manifest_tool" validate "$2" >/dev/null node "$manifest_tool" validate "$2" >/dev/null
api_image=$(manifest_get "$2" images.api) api_image=$(manifest_get "$2" images.api)
web_image=$(manifest_get "$2" images.web) web_image=$(manifest_get "$2" images.web)
for site in ningbo hongkong; do
[[ $("${kubectl[@]}" get deployment "easyai-worker-$site" -n "$NAMESPACE" \
-o jsonpath='{.spec.template.spec.containers[?(@.name=="worker")].image}') == "$api_image" ]]
done
[[ $("${kubectl[@]}" get deployment easyai-api-ningbo -n "$NAMESPACE" \ [[ $("${kubectl[@]}" get deployment easyai-api-ningbo -n "$NAMESPACE" \
-o jsonpath='{.spec.template.spec.containers[?(@.name=="api")].image}') == "$api_image" ]] -o jsonpath='{.spec.template.spec.containers[?(@.name=="api")].image}') == "$api_image" ]]
[[ $("${kubectl[@]}" get deployment easyai-web-ningbo -n "$NAMESPACE" \ [[ $("${kubectl[@]}" get deployment easyai-web-ningbo -n "$NAMESPACE" \
@@ -48,6 +48,12 @@ spec:
- secretRef: - secretRef:
name: easyai-ai-gateway-runtime name: easyai-ai-gateway-runtime
env: env:
- name: AI_GATEWAY_PROCESS_ROLE
value: api
- name: AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED
value: "false"
- name: AI_GATEWAY_DATABASE_MAX_CONNS
value: "16"
- name: POD_NAMESPACE - name: POD_NAMESPACE
valueFrom: valueFrom:
fieldRef: fieldRef:
@@ -56,6 +62,12 @@ spec:
valueFrom: valueFrom:
fieldRef: fieldRef:
fieldPath: metadata.name fieldPath: metadata.name
- name: POD_UID
valueFrom:
fieldRef:
fieldPath: metadata.uid
- name: EASYAI_SITE
value: ningbo
ports: ports:
- name: http - name: http
containerPort: 8088 containerPort: 8088
@@ -154,6 +166,12 @@ spec:
- secretRef: - secretRef:
name: easyai-ai-gateway-runtime name: easyai-ai-gateway-runtime
env: env:
- name: AI_GATEWAY_PROCESS_ROLE
value: api
- name: AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED
value: "false"
- name: AI_GATEWAY_DATABASE_MAX_CONNS
value: "16"
- name: POD_NAMESPACE - name: POD_NAMESPACE
valueFrom: valueFrom:
fieldRef: fieldRef:
@@ -162,6 +180,12 @@ spec:
valueFrom: valueFrom:
fieldRef: fieldRef:
fieldPath: metadata.name fieldPath: metadata.name
- name: POD_UID
valueFrom:
fieldRef:
fieldPath: metadata.uid
- name: EASYAI_SITE
value: hongkong
ports: ports:
- name: http - name: http
containerPort: 8088 containerPort: 8088
@@ -212,6 +236,242 @@ spec:
--- ---
apiVersion: apps/v1 apiVersion: apps/v1
kind: Deployment kind: Deployment
metadata:
name: easyai-worker-ningbo
labels:
app.kubernetes.io/name: easyai-worker
app.kubernetes.io/component: worker
app.kubernetes.io/part-of: easyai-ai-gateway
spec:
replicas: 1
strategy:
type: RollingUpdate
rollingUpdate:
maxSurge: 1
maxUnavailable: 0
selector:
matchLabels:
app.kubernetes.io/name: easyai-worker
easyai.io/site: ningbo
template:
metadata:
labels:
app.kubernetes.io/name: easyai-worker
app.kubernetes.io/component: worker
app.kubernetes.io/part-of: easyai-ai-gateway
easyai.io/site: ningbo
spec:
serviceAccountName: easyai-ai-gateway
terminationGracePeriodSeconds: 90
nodeSelector:
easyai.io/site: ningbo
easyai.io/workload: "true"
securityContext:
runAsNonRoot: true
runAsUser: 10001
runAsGroup: 10001
fsGroup: 10001
fsGroupChangePolicy: OnRootMismatch
seccompProfile:
type: RuntimeDefault
containers:
- name: worker
image: easyai-api
imagePullPolicy: IfNotPresent
envFrom:
- configMapRef:
name: easyai-ai-gateway-config
- secretRef:
name: easyai-ai-gateway-runtime
env:
- name: AI_GATEWAY_PROCESS_ROLE
value: worker
- name: AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED
value: "true"
- name: AI_GATEWAY_DATABASE_MAX_CONNS
value: "24"
- name: POD_NAMESPACE
valueFrom:
fieldRef:
fieldPath: metadata.namespace
- name: POD_NAME
valueFrom:
fieldRef:
fieldPath: metadata.name
- name: POD_UID
valueFrom:
fieldRef:
fieldPath: metadata.uid
- name: EASYAI_SITE
value: ningbo
ports:
- name: health
containerPort: 8088
readinessProbe:
httpGet:
path: /readyz
port: health
initialDelaySeconds: 3
periodSeconds: 5
timeoutSeconds: 3
failureThreshold: 6
livenessProbe:
httpGet:
path: /healthz
port: health
initialDelaySeconds: 15
periodSeconds: 10
timeoutSeconds: 3
failureThreshold: 6
lifecycle:
preStop:
exec:
command: ["/bin/sh", "-c", "sleep 10"]
resources:
requests:
cpu: 250m
memory: 512Mi
limits:
cpu: "2"
memory: 2Gi
securityContext:
allowPrivilegeEscalation: false
capabilities:
drop: ["ALL"]
readOnlyRootFilesystem: true
volumeMounts:
- name: application-data
mountPath: /app/data
- name: temporary-files
mountPath: /tmp
volumes:
- name: application-data
emptyDir:
sizeLimit: 5Gi
- name: temporary-files
emptyDir:
sizeLimit: 1Gi
---
apiVersion: apps/v1
kind: Deployment
metadata:
name: easyai-worker-hongkong
labels:
app.kubernetes.io/name: easyai-worker
app.kubernetes.io/component: worker
app.kubernetes.io/part-of: easyai-ai-gateway
spec:
replicas: 1
strategy:
type: RollingUpdate
rollingUpdate:
maxSurge: 1
maxUnavailable: 0
selector:
matchLabels:
app.kubernetes.io/name: easyai-worker
easyai.io/site: hongkong
template:
metadata:
labels:
app.kubernetes.io/name: easyai-worker
app.kubernetes.io/component: worker
app.kubernetes.io/part-of: easyai-ai-gateway
easyai.io/site: hongkong
spec:
serviceAccountName: easyai-ai-gateway
terminationGracePeriodSeconds: 90
nodeSelector:
easyai.io/site: hongkong
easyai.io/workload: "true"
securityContext:
runAsNonRoot: true
runAsUser: 10001
runAsGroup: 10001
fsGroup: 10001
fsGroupChangePolicy: OnRootMismatch
seccompProfile:
type: RuntimeDefault
containers:
- name: worker
image: easyai-api
imagePullPolicy: IfNotPresent
envFrom:
- configMapRef:
name: easyai-ai-gateway-config
- secretRef:
name: easyai-ai-gateway-runtime
env:
- name: AI_GATEWAY_PROCESS_ROLE
value: worker
- name: AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED
value: "true"
- name: AI_GATEWAY_DATABASE_MAX_CONNS
value: "24"
- name: POD_NAMESPACE
valueFrom:
fieldRef:
fieldPath: metadata.namespace
- name: POD_NAME
valueFrom:
fieldRef:
fieldPath: metadata.name
- name: POD_UID
valueFrom:
fieldRef:
fieldPath: metadata.uid
- name: EASYAI_SITE
value: hongkong
ports:
- name: health
containerPort: 8088
readinessProbe:
httpGet:
path: /readyz
port: health
initialDelaySeconds: 3
periodSeconds: 5
timeoutSeconds: 3
failureThreshold: 6
livenessProbe:
httpGet:
path: /healthz
port: health
initialDelaySeconds: 15
periodSeconds: 10
timeoutSeconds: 3
failureThreshold: 6
lifecycle:
preStop:
exec:
command: ["/bin/sh", "-c", "sleep 10"]
resources:
requests:
cpu: 250m
memory: 512Mi
limits:
cpu: "2"
memory: 2Gi
securityContext:
allowPrivilegeEscalation: false
capabilities:
drop: ["ALL"]
readOnlyRootFilesystem: true
volumeMounts:
- name: application-data
mountPath: /app/data
- name: temporary-files
mountPath: /tmp
volumes:
- name: application-data
emptyDir:
sizeLimit: 5Gi
- name: temporary-files
emptyDir:
sizeLimit: 1Gi
---
apiVersion: apps/v1
kind: Deployment
metadata: metadata:
name: easyai-web-ningbo name: easyai-web-ningbo
labels: labels:
@@ -529,3 +789,13 @@ spec:
selector: selector:
matchLabels: matchLabels:
app.kubernetes.io/name: easyai-web app.kubernetes.io/name: easyai-web
---
apiVersion: policy/v1
kind: PodDisruptionBudget
metadata:
name: easyai-worker
spec:
minAvailable: 1
selector:
matchLabels:
app.kubernetes.io/name: easyai-worker
@@ -7,7 +7,7 @@ spec:
matchExpressions: matchExpressions:
- key: app.kubernetes.io/name - key: app.kubernetes.io/name
operator: In operator: In
values: [easyai-api, easyai-web] values: [easyai-api, easyai-worker, easyai-web]
policyTypes: [Ingress] policyTypes: [Ingress]
ingress: ingress:
- from: - from:
+5
View File
@@ -711,6 +711,7 @@ export interface RateLimitRule {
limit: number; limit: number;
windowSeconds?: number; windowSeconds?: number;
leaseTtlSeconds?: number; leaseTtlSeconds?: number;
maxWaitSeconds?: number;
consume?: 'fixed_window' | 'reserve_then_reconcile' | string; consume?: 'fixed_window' | 'reserve_then_reconcile' | string;
} }
@@ -950,6 +951,10 @@ export interface ModelRateLimitStatus {
modelCooldownUntil?: string; modelCooldownUntil?: string;
concurrent: RateLimitMetricStatus; concurrent: RateLimitMetricStatus;
queuedTasks: number; queuedTasks: number;
waitingSyncTasks?: number;
waitingAsyncTasks?: number;
oldestQueueWaitSeconds?: number;
queueLimit?: number;
rpm: RateLimitMetricStatus; rpm: RateLimitMetricStatus;
tpm: RateLimitMetricStatus; tpm: RateLimitMetricStatus;
loadRatio: number; loadRatio: number;
+34 -6
View File
@@ -23,6 +23,11 @@ node "$cluster_root/scripts/release-manifest.mjs" validate "$release_manifest" >
api_image=$(node "$cluster_root/scripts/release-manifest.mjs" get "$release_manifest" images.api) api_image=$(node "$cluster_root/scripts/release-manifest.mjs" get "$release_manifest" images.api)
web_image=$(node "$cluster_root/scripts/release-manifest.mjs" get "$release_manifest" images.web) web_image=$(node "$cluster_root/scripts/release-manifest.mjs" get "$release_manifest" images.web)
source_sha=$(node "$cluster_root/scripts/release-manifest.mjs" get "$release_manifest" sourceSha) source_sha=$(node "$cluster_root/scripts/release-manifest.mjs" get "$release_manifest" sourceSha)
activate_workers=${AI_GATEWAY_ACTIVATE_WORKERS:-true}
[[ $activate_workers == true || $activate_workers == false ]] || {
echo 'AI_GATEWAY_ACTIVATE_WORKERS must be true or false' >&2
exit 64
}
[[ $api_image =~ ^registry\.cn-shanghai\.aliyuncs\.com/easyaigc/ai-gateway@sha256:[0-9a-f]{64}$ ]] [[ $api_image =~ ^registry\.cn-shanghai\.aliyuncs\.com/easyaigc/ai-gateway@sha256:[0-9a-f]{64}$ ]]
[[ $web_image =~ ^registry\.cn-shanghai\.aliyuncs\.com/easyaigc/ai-gateway-web@sha256:[0-9a-f]{64}$ ]] [[ $web_image =~ ^registry\.cn-shanghai\.aliyuncs\.com/easyaigc/ai-gateway-web@sha256:[0-9a-f]{64}$ ]]
@@ -87,11 +92,28 @@ if [[ $mode == prepare ]]; then
sed -i.bak 's/^ replicas: 1$/ replicas: 0/' "$working_directory/production/application.yaml" sed -i.bak 's/^ replicas: 1$/ replicas: 0/' "$working_directory/production/application.yaml"
rm -f -- "$working_directory/production/application.yaml.bak" rm -f -- "$working_directory/production/application.yaml.bak"
fi fi
if [[ ${AI_GATEWAY_ACTIVATE_WORKERS:-true} == false ]]; then if [[ $activate_workers == false ]]; then
sed -i.bak \ awk '
's/^ AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED: "true"$/ AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED: "false"/' \ /^---$/ {
"$working_directory/production/application-config.yaml" document_kind = ""
rm -f -- "$working_directory/production/application-config.yaml.bak" worker_deployment = 0
print
next
}
/^kind: / {
document_kind = $2
}
document_kind == "Deployment" && /^ name: easyai-worker-/ {
worker_deployment = 1
}
worker_deployment && /^ replicas: 1$/ {
sub(/1$/, "0")
}
{ print }
' "$working_directory/production/application.yaml" \
>"$working_directory/production/application.yaml.next"
mv "$working_directory/production/application.yaml.next" \
"$working_directory/production/application.yaml"
fi fi
kubectl kustomize "$working_directory/production" >"$working_directory/easyai-production.yaml" kubectl kustomize "$working_directory/production" >"$working_directory/easyai-production.yaml"
if rg -q 'sha256:0{64}|image: (easyai-api|easyai-web)$' "$working_directory/easyai-production.yaml"; then if rg -q 'sha256:0{64}|image: (easyai-api|easyai-web)$' "$working_directory/easyai-production.yaml"; then
@@ -136,10 +158,12 @@ for file in "${platform_files[@]}"; do
cluster_scp "$working_directory/$file" "$CLUSTER_NINGBO_HOST:/root/$file" cluster_scp "$working_directory/$file" "$CLUSTER_NINGBO_HOST:/root/$file"
done done
cluster_ssh "$CLUSTER_NINGBO_HOST" bash -s -- "$source_sha" "$mode" <<'REMOTE' cluster_ssh "$CLUSTER_NINGBO_HOST" bash -s -- \
"$source_sha" "$mode" "$activate_workers" <<'REMOTE'
set -euo pipefail set -euo pipefail
source_sha=$1 source_sha=$1
mode=$2 mode=$2
activate_workers=$3
kubectl='k3s kubectl' kubectl='k3s kubectl'
runtime_environment= runtime_environment=
cleanup_bootstrap_files() { cleanup_bootstrap_files() {
@@ -234,6 +258,10 @@ $kubectl apply --server-side --dry-run=server --force-conflicts \
$kubectl apply -f /root/easyai-production.yaml >/dev/null $kubectl apply -f /root/easyai-production.yaml >/dev/null
$kubectl annotate namespace easyai easyai.io/release="$source_sha" --overwrite >/dev/null $kubectl annotate namespace easyai easyai.io/release="$source_sha" --overwrite >/dev/null
if [[ $mode == activate ]]; then if [[ $mode == activate ]]; then
if [[ $activate_workers == true ]]; then
$kubectl rollout status deployment/easyai-worker-hongkong -n easyai --timeout=300s
$kubectl rollout status deployment/easyai-worker-ningbo -n easyai --timeout=300s
fi
$kubectl rollout status deployment/easyai-api-hongkong -n easyai --timeout=300s $kubectl rollout status deployment/easyai-api-hongkong -n easyai --timeout=300s
$kubectl rollout status deployment/easyai-web-hongkong -n easyai --timeout=300s $kubectl rollout status deployment/easyai-web-hongkong -n easyai --timeout=300s
$kubectl rollout status deployment/easyai-api-ningbo -n easyai --timeout=300s $kubectl rollout status deployment/easyai-api-ningbo -n easyai --timeout=300s
+11 -9
View File
@@ -114,7 +114,9 @@ rollback_before_commit() {
if [[ $traffic_committed == false ]]; then if [[ $traffic_committed == false ]]; then
echo '[cutover] rolling back to the frozen Docker database' >&2 echo '[cutover] rolling back to the frozen Docker database' >&2
k3s kubectl scale deployment -n easyai \ k3s kubectl scale deployment -n easyai \
easyai-api-ningbo easyai-api-hongkong easyai-web-ningbo easyai-web-hongkong \ easyai-api-ningbo easyai-api-hongkong \
easyai-worker-ningbo easyai-worker-hongkong \
easyai-web-ningbo easyai-web-hongkong \
--replicas=0 >/dev/null 2>&1 || true --replicas=0 >/dev/null 2>&1 || true
docker start easyai-ai-gateway-api-1 easyai-ai-gateway-web-1 >/dev/null 2>&1 || true docker start easyai-ai-gateway-api-1 easyai-ai-gateway-web-1 >/dev/null 2>&1 || true
if [[ -f $old_site_backup ]]; then if [[ -f $old_site_backup ]]; then
@@ -312,7 +314,9 @@ set -euo pipefail
backup_directory=$(cat /root/.easyai-pre-cutover-backup) backup_directory=$(cat /root/.easyai-pre-cutover-backup)
[[ $backup_directory == /var/backups/easyai-ai-gateway/pre-cutover-* ]] [[ $backup_directory == /var/backups/easyai-ai-gateway/pre-cutover-* ]]
k3s kubectl scale deployment -n easyai \ k3s kubectl scale deployment -n easyai \
easyai-api-ningbo easyai-api-hongkong easyai-web-ningbo easyai-web-hongkong \ easyai-api-ningbo easyai-api-hongkong \
easyai-worker-ningbo easyai-worker-hongkong \
easyai-web-ningbo easyai-web-hongkong \
--replicas=0 >/dev/null 2>&1 || true --replicas=0 >/dev/null 2>&1 || true
docker start easyai-ai-gateway-api-1 easyai-ai-gateway-web-1 >/dev/null docker start easyai-ai-gateway-api-1 easyai-ai-gateway-web-1 >/dev/null
install -m 0644 "$backup_directory/nginx-site.conf" \ install -m 0644 "$backup_directory/nginx-site.conf" \
@@ -325,7 +329,7 @@ REMOTE
} }
trap rollback_local_before_commit ERR trap rollback_local_before_commit ERR
echo '[cutover] starting both application sites with River workers frozen' echo '[cutover] starting both application sites with dedicated workers frozen'
AI_GATEWAY_ACTIVATE_WORKERS=false \ AI_GATEWAY_ACTIVATE_WORKERS=false \
"$script_dir/bootstrap-platform.sh" activate "$release_manifest" "$script_dir/bootstrap-platform.sh" activate "$release_manifest"
@@ -352,12 +356,10 @@ echo '[cutover] traffic committed; enabling Hong Kong worker, then Ningbo worker
cluster_ssh "$CLUSTER_NINGBO_HOST" 'bash -s' <<'REMOTE' cluster_ssh "$CLUSTER_NINGBO_HOST" 'bash -s' <<'REMOTE'
set -euo pipefail set -euo pipefail
kubectl='k3s kubectl' kubectl='k3s kubectl'
$kubectl set env deployment/easyai-api-hongkong -n easyai \ $kubectl scale deployment/easyai-worker-hongkong -n easyai --replicas=1
AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED=true $kubectl rollout status deployment/easyai-worker-hongkong -n easyai --timeout=300s
$kubectl rollout status deployment/easyai-api-hongkong -n easyai --timeout=300s $kubectl scale deployment/easyai-worker-ningbo -n easyai --replicas=1
$kubectl set env deployment/easyai-api-ningbo -n easyai \ $kubectl rollout status deployment/easyai-worker-ningbo -n easyai --timeout=300s
AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED=true
$kubectl rollout status deployment/easyai-api-ningbo -n easyai --timeout=300s
backup_resource=$(cat <<'EOF' | $kubectl create -f - -o name backup_resource=$(cat <<'EOF' | $kubectl create -f - -o name
apiVersion: postgresql.cnpg.io/v1 apiVersion: postgresql.cnpg.io/v1
kind: Backup kind: Backup
+17 -19
View File
@@ -13,43 +13,41 @@ require_commands base64 jq node
restore_workers() { restore_workers() {
cluster_ssh "$CLUSTER_NINGBO_HOST" 'bash -s' <<'REMOTE' || true cluster_ssh "$CLUSTER_NINGBO_HOST" 'bash -s' <<'REMOTE' || true
set -euo pipefail set -euo pipefail
k3s kubectl set env deployment/easyai-api-ningbo -n easyai \ k3s kubectl scale deployment/easyai-worker-ningbo -n easyai --replicas=1 >/dev/null
AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED=true >/dev/null k3s kubectl scale deployment/easyai-worker-hongkong -n easyai --replicas=1 >/dev/null
k3s kubectl set env deployment/easyai-api-hongkong -n easyai \ k3s kubectl rollout status deployment/easyai-worker-ningbo -n easyai --timeout=300s
AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED=true >/dev/null k3s kubectl rollout status deployment/easyai-worker-hongkong -n easyai --timeout=300s
k3s kubectl rollout status deployment/easyai-api-ningbo -n easyai --timeout=300s
k3s kubectl rollout status deployment/easyai-api-hongkong -n easyai --timeout=300s
REMOTE REMOTE
} }
trap restore_workers EXIT HUP INT TERM trap restore_workers EXIT HUP INT TERM
cluster_ssh "$CLUSTER_NINGBO_HOST" 'bash -s' <<'REMOTE' cluster_ssh "$CLUSTER_NINGBO_HOST" 'bash -s' <<'REMOTE'
set -euo pipefail set -euo pipefail
k3s kubectl set env deployment/easyai-api-hongkong -n easyai \ k3s kubectl scale deployment/easyai-worker-hongkong -n easyai --replicas=1 >/dev/null
AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED=true >/dev/null k3s kubectl rollout status deployment/easyai-worker-hongkong -n easyai --timeout=300s
k3s kubectl rollout status deployment/easyai-api-hongkong -n easyai --timeout=300s k3s kubectl scale deployment/easyai-worker-ningbo -n easyai --replicas=0 >/dev/null
k3s kubectl set env deployment/easyai-api-ningbo -n easyai \
AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED=false >/dev/null
k3s kubectl rollout status deployment/easyai-api-ningbo -n easyai --timeout=300s
for _ in $(seq 1 60); do for _ in $(seq 1 60); do
pods=$(k3s kubectl get pods -n easyai \ pods=$(k3s kubectl get pods -n easyai \
-l app.kubernetes.io/name=easyai-api,easyai.io/site=ningbo \ -l app.kubernetes.io/name=easyai-worker,easyai.io/site=hongkong \
-o jsonpath='{range .items[*]}{.metadata.name}{"\n"}{end}') -o jsonpath='{range .items[*]}{.metadata.name}{"\n"}{end}')
pod_count=$(wc -w <<<"$pods" | tr -d ' ') pod_count=$(wc -w <<<"$pods" | tr -d ' ')
if [[ $pod_count == 1 ]]; then if [[ $pod_count == 1 ]]; then
pod=$pods pod=$pods
worker_enabled=$(k3s kubectl get pod "$pod" -n easyai \ process_role=$(k3s kubectl get pod "$pod" -n easyai \
-o jsonpath='{.spec.containers[0].env[?(@.name=="AI_GATEWAY_ASYNC_QUEUE_WORKER_ENABLED")].value}') -o jsonpath='{.spec.containers[0].env[?(@.name=="AI_GATEWAY_PROCESS_ROLE")].value}')
pod_ready=$(k3s kubectl get pod "$pod" -n easyai \ pod_ready=$(k3s kubectl get pod "$pod" -n easyai \
-o jsonpath='{.status.containerStatuses[0].ready}') -o jsonpath='{.status.containerStatuses[0].ready}')
if [[ $worker_enabled == false && $pod_ready == true ]]; then if [[ $process_role == worker && $pod_ready == true ]]; then
break break
fi fi
fi fi
sleep 1 sleep 1
done done
[[ ${pod_count:-0} == 1 && ${worker_enabled:-} == false && ${pod_ready:-} == true ]] [[ ${pod_count:-0} == 1 && ${process_role:-} == worker && ${pod_ready:-} == true ]]
k3s kubectl logs "$pod" -n easyai | grep -q 'asynchronous queue worker disabled for this process' [[ $(k3s kubectl get deployment easyai-api-ningbo -n easyai \
-o jsonpath='{.spec.template.spec.containers[0].env[?(@.name=="AI_GATEWAY_PROCESS_ROLE")].value}') == api ]]
[[ $(k3s kubectl get deployment easyai-worker-ningbo -n easyai \
-o jsonpath='{.status.readyReplicas}') == 0 ]]
curl -fsS http://10.77.0.1:31088/api/v1/readyz | grep -q '"ok":true' curl -fsS http://10.77.0.1:31088/api/v1/readyz | grep -q '"ok":true'
curl -fsS http://10.77.0.2:31089/api/v1/readyz | grep -q '"ok":true' curl -fsS http://10.77.0.2:31089/api/v1/readyz | grep -q '"ok":true'
REMOTE REMOTE
@@ -98,7 +96,7 @@ query() {
} }
attempted_by=$(query -c \ attempted_by=$(query -c \
"SELECT attempted_by::text FROM river_job WHERE id=(SELECT river_job_id FROM gateway_tasks WHERE id='$task_id'::uuid);") "SELECT attempted_by::text FROM river_job WHERE id=(SELECT river_job_id FROM gateway_tasks WHERE id='$task_id'::uuid);")
[[ $attempted_by == *easyai-api-hongkong* ]] || { [[ $attempted_by == *easyai-worker-hongkong* ]] || {
echo "River attempted_by does not identify Hong Kong: $attempted_by" >&2 echo "River attempted_by does not identify Hong Kong: $attempted_by" >&2
exit 1 exit 1
} }
+13 -1
View File
@@ -95,6 +95,18 @@ kubectl='k3s kubectl'
--field-selector=status.phase=Running --no-headers | wc -l) -eq 2 ]] --field-selector=status.phase=Running --no-headers | wc -l) -eq 2 ]]
[[ $($kubectl get pods -n easyai -l app.kubernetes.io/name=easyai-web \ [[ $($kubectl get pods -n easyai -l app.kubernetes.io/name=easyai-web \
--field-selector=status.phase=Running --no-headers | wc -l) -eq 2 ]] --field-selector=status.phase=Running --no-headers | wc -l) -eq 2 ]]
[[ $($kubectl get pods -n easyai -l app.kubernetes.io/name=easyai-worker \
--field-selector=status.phase=Running --no-headers | wc -l) -eq 2 ]]
for site in ningbo hongkong; do
[[ $($kubectl get deployment "easyai-api-$site" -n easyai \
-o jsonpath='{.status.readyReplicas}') -eq 1 ]]
[[ $($kubectl get deployment "easyai-worker-$site" -n easyai \
-o jsonpath='{.status.readyReplicas}') -eq 1 ]]
[[ $($kubectl get deployment "easyai-api-$site" -n easyai \
-o jsonpath='{.spec.template.spec.containers[0].env[?(@.name=="AI_GATEWAY_PROCESS_ROLE")].value}') == api ]]
[[ $($kubectl get deployment "easyai-worker-$site" -n easyai \
-o jsonpath='{.spec.template.spec.containers[0].env[?(@.name=="AI_GATEWAY_PROCESS_ROLE")].value}') == worker ]]
done
images=$($kubectl get deployments -n easyai -o json | images=$($kubectl get deployments -n easyai -o json |
jq -r '.items[] | select(.metadata.name | startswith("easyai-")) | jq -r '.items[] | select(.metadata.name | startswith("easyai-")) |
.spec.template.spec.containers[].image') .spec.template.spec.containers[].image')
@@ -112,7 +124,7 @@ printf 'applications=ready\n'
REMOTE REMOTE
) )
[[ $application_state == applications=ready ]] [[ $application_state == applications=ready ]]
echo 'application_acceptance=PASS api=2 web=2 digests=pinned' echo 'application_acceptance=PASS api=2 worker=2 web=2 roles=isolated digests=pinned'
for edge_ip in "${CLUSTER_NINGBO_HOST#*@}" "${CLUSTER_HONGKONG_HOST#*@}"; do for edge_ip in "${CLUSTER_NINGBO_HOST#*@}" "${CLUSTER_HONGKONG_HOST#*@}"; do
health_body=$(curl -fsS --max-time 10 --resolve "ai.51easyai.com:443:$edge_ip" \ health_body=$(curl -fsS --max-time 10 --resolve "ai.51easyai.com:443:$edge_ip" \