perf(queue): 按策略动态扩缩异步 Worker
将 River 执行容量按平台模型与用户组的有效并发策略动态调整,并为长任务续租、并发租约原子抢占和限流退避补充保护。\n\n统一平台模型限流继承语义,兼容历史 platformLimits/modelLimits,并为三个迁移图像模型建立独立并发租约。补充管理端显式继承/覆盖配置、指标、单元测试及隔离 PostgreSQL 验收。\n\n验证:go test ./... -count=1;pnpm lint;pnpm test;pnpm build;隔离 PostgreSQL 并发原子性/续租测试;128 任务与三分钟长任务动态 Worker 验收。
This commit is contained in:
@@ -29,6 +29,11 @@ IDENTITY_MODE=hybrid
|
|||||||
# - hold: reject new production generation before any upstream request; existing settlements continue.
|
# - hold: reject new production generation before any upstream request; existing settlements continue.
|
||||||
BILLING_ENGINE_MODE=observe
|
BILLING_ENGINE_MODE=observe
|
||||||
|
|
||||||
|
# River 执行 worker 会按平台模型和活跃用户组中更严格的 concurrent 总量每 5 秒重算。
|
||||||
|
# hard limit 是单进程安全边界;平台模型/用户组 concurrency lease 才是业务并发真值。
|
||||||
|
AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT=2048
|
||||||
|
AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS=5
|
||||||
|
|
||||||
# Unified identity business settings are managed in System Settings > Unified
|
# Unified identity business settings are managed in System Settings > Unified
|
||||||
# Identity. Deployment only supplies the SecretStore and infrastructure timing.
|
# Identity. Deployment only supplies the SecretStore and infrastructure timing.
|
||||||
AI_GATEWAY_PUBLIC_BASE_URL=http://localhost:8088
|
AI_GATEWAY_PUBLIC_BASE_URL=http://localhost:8088
|
||||||
|
|||||||
@@ -159,6 +159,8 @@ AI_GATEWAY_DATABASE_URL=postgresql://easyai:easyai2025@localhost:5432/easyai_ai_
|
|||||||
|
|
||||||
如果现有 `easyai-pgvector` 没有把 `5432` 映射到宿主机,就需要补端口映射,或者把 AI Gateway 后端容器化后接入同一个 `easyai` Docker network。
|
如果现有 `easyai-pgvector` 没有把 `5432` 映射到宿主机,就需要补端口映射,或者把 AI Gateway 后端容器化后接入同一个 `easyai` Docker network。
|
||||||
|
|
||||||
|
异步队列 worker 不使用固定并发。服务分别汇总启用平台模型和活跃用户组的有效 `concurrent` 策略,采用两者中更严格的容量,默认每 5 秒在线调整 River 执行容量。策略解析兼容历史 `platformLimits/modelLimits.max_concurrent_requests`,运行时统一转换为 `rules`。`AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT` 默认 `2048`,仅作为单进程资源安全边界;平台模型和用户组的 PostgreSQL concurrency lease 才是业务并发真值。可通过 `AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS` 调整刷新周期。
|
||||||
|
|
||||||
## 迁移原则
|
## 迁移原则
|
||||||
|
|
||||||
1. 新服务先并行运行,不直接删除 `easyai-server-main` 内现有模块。
|
1. 新服务先并行运行,不直接删除 `easyai-server-main` 内现有模块。
|
||||||
|
|||||||
@@ -12848,6 +12848,13 @@
|
|||||||
"type": "object",
|
"type": "object",
|
||||||
"additionalProperties": {}
|
"additionalProperties": {}
|
||||||
},
|
},
|
||||||
|
"rateLimitPolicyMode": {
|
||||||
|
"type": "string",
|
||||||
|
"enum": [
|
||||||
|
"inherit",
|
||||||
|
"override"
|
||||||
|
]
|
||||||
|
},
|
||||||
"retryPolicy": {
|
"retryPolicy": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"additionalProperties": {}
|
"additionalProperties": {}
|
||||||
@@ -13838,6 +13845,13 @@
|
|||||||
"type": "object",
|
"type": "object",
|
||||||
"additionalProperties": {}
|
"additionalProperties": {}
|
||||||
},
|
},
|
||||||
|
"rateLimitPolicyMode": {
|
||||||
|
"type": "string",
|
||||||
|
"enum": [
|
||||||
|
"inherit",
|
||||||
|
"override"
|
||||||
|
]
|
||||||
|
},
|
||||||
"retryPolicy": {
|
"retryPolicy": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"additionalProperties": {}
|
"additionalProperties": {}
|
||||||
|
|||||||
@@ -2439,6 +2439,11 @@ definitions:
|
|||||||
rateLimitPolicy:
|
rateLimitPolicy:
|
||||||
additionalProperties: {}
|
additionalProperties: {}
|
||||||
type: object
|
type: object
|
||||||
|
rateLimitPolicyMode:
|
||||||
|
enum:
|
||||||
|
- inherit
|
||||||
|
- override
|
||||||
|
type: string
|
||||||
retryPolicy:
|
retryPolicy:
|
||||||
additionalProperties: {}
|
additionalProperties: {}
|
||||||
type: object
|
type: object
|
||||||
@@ -3108,6 +3113,11 @@ definitions:
|
|||||||
rateLimitPolicy:
|
rateLimitPolicy:
|
||||||
additionalProperties: {}
|
additionalProperties: {}
|
||||||
type: object
|
type: object
|
||||||
|
rateLimitPolicyMode:
|
||||||
|
enum:
|
||||||
|
- inherit
|
||||||
|
- override
|
||||||
|
type: string
|
||||||
retryPolicy:
|
retryPolicy:
|
||||||
additionalProperties: {}
|
additionalProperties: {}
|
||||||
type: object
|
type: object
|
||||||
|
|||||||
@@ -48,6 +48,8 @@ type Config struct {
|
|||||||
GlobalHTTPProxySource string
|
GlobalHTTPProxySource string
|
||||||
LogLevel slog.Level
|
LogLevel slog.Level
|
||||||
BillingEngineMode string
|
BillingEngineMode string
|
||||||
|
AsyncWorkerHardLimit int
|
||||||
|
AsyncWorkerRefreshIntervalSeconds int
|
||||||
}
|
}
|
||||||
|
|
||||||
func Load() Config {
|
func Load() Config {
|
||||||
@@ -85,13 +87,15 @@ func Load() Config {
|
|||||||
TaskProgressCallbackURL: env("TASK_PROGRESS_CALLBACK_URL",
|
TaskProgressCallbackURL: env("TASK_PROGRESS_CALLBACK_URL",
|
||||||
strings.TrimRight(env("SERVER_MAIN_BASE_URL", "http://localhost:3000"), "/")+"/internal/platform/task-progress-callbacks",
|
strings.TrimRight(env("SERVER_MAIN_BASE_URL", "http://localhost:3000"), "/")+"/internal/platform/task-progress-callbacks",
|
||||||
),
|
),
|
||||||
TaskProgressCallbackTimeoutMS: env("TASK_PROGRESS_CALLBACK_TIMEOUT_MS", "5000"),
|
TaskProgressCallbackTimeoutMS: env("TASK_PROGRESS_CALLBACK_TIMEOUT_MS", "5000"),
|
||||||
TaskProgressCallbackMaxAttempts: env("TASK_PROGRESS_CALLBACK_MAX_ATTEMPTS", "10"),
|
TaskProgressCallbackMaxAttempts: env("TASK_PROGRESS_CALLBACK_MAX_ATTEMPTS", "10"),
|
||||||
CORSAllowedOrigin: env("CORS_ALLOWED_ORIGIN", "http://localhost:5178,http://127.0.0.1:5178"),
|
CORSAllowedOrigin: env("CORS_ALLOWED_ORIGIN", "http://localhost:5178,http://127.0.0.1:5178"),
|
||||||
GlobalHTTPProxy: globalProxy.HTTPProxy,
|
GlobalHTTPProxy: globalProxy.HTTPProxy,
|
||||||
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")),
|
||||||
|
AsyncWorkerHardLimit: envIntValidated("AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT", 2048),
|
||||||
|
AsyncWorkerRefreshIntervalSeconds: envIntValidated("AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS", 5),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -101,6 +105,12 @@ func (c Config) Validate() error {
|
|||||||
default:
|
default:
|
||||||
return errors.New("BILLING_ENGINE_MODE must be observe, enforce, or hold")
|
return errors.New("BILLING_ENGINE_MODE must be observe, enforce, or hold")
|
||||||
}
|
}
|
||||||
|
if c.AsyncWorkerHardLimit < 1 || c.AsyncWorkerHardLimit > 10000 {
|
||||||
|
return errors.New("AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT must be between 1 and 10000")
|
||||||
|
}
|
||||||
|
if c.AsyncWorkerRefreshIntervalSeconds < 1 {
|
||||||
|
return errors.New("AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS must be positive")
|
||||||
|
}
|
||||||
switch strings.ToLower(strings.TrimSpace(c.IdentitySecretStore)) {
|
switch strings.ToLower(strings.TrimSpace(c.IdentitySecretStore)) {
|
||||||
case "":
|
case "":
|
||||||
case "file":
|
case "file":
|
||||||
@@ -210,6 +220,18 @@ func envInt(key string, fallback int) int {
|
|||||||
return parsed
|
return parsed
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func envIntValidated(key string, fallback int) int {
|
||||||
|
value := envValue(key)
|
||||||
|
if value == "" {
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
parsed, err := strconv.Atoi(value)
|
||||||
|
if err != nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return parsed
|
||||||
|
}
|
||||||
|
|
||||||
func logLevel(value string) slog.Level {
|
func logLevel(value string) slog.Level {
|
||||||
switch strings.ToLower(value) {
|
switch strings.ToLower(value) {
|
||||||
case "debug":
|
case "debug":
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ func TestLoadIdentitySecretStoreUsesNewEnvironmentNamesAndIgnoresLegacyBusinessV
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateIdentityFileSecretStoreRequiresDirectory(t *testing.T) {
|
func TestValidateIdentityFileSecretStoreRequiresDirectory(t *testing.T) {
|
||||||
cfg := Config{IdentitySecretStore: "file"}
|
cfg := Config{IdentitySecretStore: "file", AsyncWorkerHardLimit: 2048, AsyncWorkerRefreshIntervalSeconds: 5}
|
||||||
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "IDENTITY_SECRET_DIR") {
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "IDENTITY_SECRET_DIR") {
|
||||||
t.Fatalf("Validate() error = %v, want missing identity secret directory", err)
|
t.Fatalf("Validate() error = %v, want missing identity secret directory", err)
|
||||||
}
|
}
|
||||||
@@ -37,7 +37,12 @@ func TestValidateIdentityFileSecretStoreRequiresDirectory(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateIdentityKubernetesSecretStore(t *testing.T) {
|
func TestValidateIdentityKubernetesSecretStore(t *testing.T) {
|
||||||
cfg := Config{IdentitySecretStore: "kubernetes", IdentityKubernetesSecretName: "easyai-gateway-identity"}
|
cfg := Config{
|
||||||
|
IdentitySecretStore: "kubernetes",
|
||||||
|
IdentityKubernetesSecretName: "easyai-gateway-identity",
|
||||||
|
AsyncWorkerHardLimit: 2048,
|
||||||
|
AsyncWorkerRefreshIntervalSeconds: 5,
|
||||||
|
}
|
||||||
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "namespace") {
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "namespace") {
|
||||||
t.Fatalf("Validate() error = %v, want missing namespace", err)
|
t.Fatalf("Validate() error = %v, want missing namespace", err)
|
||||||
}
|
}
|
||||||
@@ -52,8 +57,27 @@ func TestValidateIdentitySecurityEventTiming(t *testing.T) {
|
|||||||
IdentitySecurityEventHeartbeatIntervalSeconds: 60,
|
IdentitySecurityEventHeartbeatIntervalSeconds: 60,
|
||||||
IdentitySecurityEventStaleAfterSeconds: 60,
|
IdentitySecurityEventStaleAfterSeconds: 60,
|
||||||
IdentitySecurityEventClockSkewSeconds: 60,
|
IdentitySecurityEventClockSkewSeconds: 60,
|
||||||
|
AsyncWorkerHardLimit: 2048,
|
||||||
|
AsyncWorkerRefreshIntervalSeconds: 5,
|
||||||
}
|
}
|
||||||
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "heartbeat") {
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "heartbeat") {
|
||||||
t.Fatalf("Validate() error = %v, want invalid stale threshold", err)
|
t.Fatalf("Validate() error = %v, want invalid stale threshold", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestValidateAsyncWorkerSettings(t *testing.T) {
|
||||||
|
cfg := Config{AsyncWorkerHardLimit: 10001, AsyncWorkerRefreshIntervalSeconds: 5}
|
||||||
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "HARD_LIMIT") {
|
||||||
|
t.Fatalf("Validate() error = %v, want invalid hard limit", err)
|
||||||
|
}
|
||||||
|
cfg.AsyncWorkerHardLimit = 2048
|
||||||
|
cfg.AsyncWorkerRefreshIntervalSeconds = 0
|
||||||
|
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "REFRESH_INTERVAL") {
|
||||||
|
t.Fatalf("Validate() error = %v, want invalid refresh interval", err)
|
||||||
|
}
|
||||||
|
t.Setenv("AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT", "not-an-integer")
|
||||||
|
loaded := Load()
|
||||||
|
if err := loaded.Validate(); err == nil || !strings.Contains(err.Error(), "HARD_LIMIT") {
|
||||||
|
t.Fatalf("Validate() error = %v, want invalid non-integer hard limit", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,479 @@
|
|||||||
|
package httpapi
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAsyncWorkerDynamicConcurrencyAcceptance(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 dynamic worker acceptance tests")
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 7*time.Minute)
|
||||||
|
defer cancel()
|
||||||
|
applyMigration(t, ctx, databaseURL)
|
||||||
|
|
||||||
|
db, err := store.Connect(ctx, databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect store: %v", err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
pool, err := pgxpool.New(ctx, databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect acceptance pool: %v", err)
|
||||||
|
}
|
||||||
|
defer pool.Close()
|
||||||
|
if _, err := pool.Exec(ctx, `UPDATE integration_platforms SET status = 'disabled' WHERE deleted_at IS NULL`); err != nil {
|
||||||
|
t.Fatalf("isolate acceptance platforms: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
serverCtx, cancelServer := context.WithCancel(ctx)
|
||||||
|
defer cancelServer()
|
||||||
|
server := httptest.NewServer(NewServerWithContext(serverCtx, config.Config{
|
||||||
|
AppEnv: "test",
|
||||||
|
HTTPAddr: ":0",
|
||||||
|
DatabaseURL: databaseURL,
|
||||||
|
IdentityMode: "hybrid",
|
||||||
|
JWTSecret: "test-secret",
|
||||||
|
BillingEngineMode: "observe",
|
||||||
|
CORSAllowedOrigin: "*",
|
||||||
|
AsyncWorkerHardLimit: 256,
|
||||||
|
AsyncWorkerRefreshIntervalSeconds: 1,
|
||||||
|
}, db, slog.New(slog.NewTextHandler(io.Discard, nil))))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
adminToken := createAsyncAcceptanceAdmin(t, ctx, pool, server.URL)
|
||||||
|
if _, err := pool.Exec(ctx, `
|
||||||
|
UPDATE gateway_user_groups
|
||||||
|
SET rate_limit_policy = '{"rules":[{"metric":"concurrent","limit":256,"leaseTtlSeconds":120}]}'::jsonb
|
||||||
|
WHERE status = 'active'`); err != nil {
|
||||||
|
t.Fatalf("raise acceptance user-group concurrency: %v", err)
|
||||||
|
}
|
||||||
|
suffix := strconv.FormatInt(time.Now().UnixNano(), 10)
|
||||||
|
|
||||||
|
t.Run("超过旧64上限的高并发", func(t *testing.T) {
|
||||||
|
model := "worker-burst-" + suffix
|
||||||
|
platform := createAsyncAcceptancePlatform(t, server.URL, adminToken, "burst-"+suffix, "Burst Simulation", 128)
|
||||||
|
defer updateAsyncAcceptancePlatform(t, server.URL, adminToken, platform, 128, "disabled")
|
||||||
|
modelID := createAsyncAcceptanceModel(t, server.URL, adminToken, platform.ID, model, "video_generate", "override", 96, 120)
|
||||||
|
|
||||||
|
waitForAsyncWorkerMetric(t, server.URL, "easyai_gateway_async_worker_capacity", 96, 15*time.Second)
|
||||||
|
startedAt := time.Now()
|
||||||
|
taskIDs := submitAsyncSimulationTasks(t, server.URL, adminToken, "/api/v1/videos/generations", model, 128, 15*time.Second)
|
||||||
|
peakRunning, peakLeases := waitForAcceptanceTasks(t, ctx, pool, taskIDs, modelID, 60*time.Second)
|
||||||
|
elapsed := time.Since(startedAt)
|
||||||
|
if peakRunning < 80 || peakLeases < 80 {
|
||||||
|
t.Fatalf("running peak=%d lease peak=%d, want both >=80", peakRunning, peakLeases)
|
||||||
|
}
|
||||||
|
if peakLeases > 96 {
|
||||||
|
t.Fatalf("lease peak=%d exceeded model concurrent limit 96", peakLeases)
|
||||||
|
}
|
||||||
|
assertAcceptanceAttempts(t, ctx, pool, taskIDs, len(taskIDs))
|
||||||
|
if elapsed >= 60*time.Second {
|
||||||
|
t.Fatalf("128 tasks completed in %s, want <60s", elapsed)
|
||||||
|
}
|
||||||
|
t.Logf("高并发证据: tasks=128 duration=%s running_peak=%d lease_peak=%d capacity=96", elapsed.Round(time.Millisecond), peakRunning, peakLeases)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("三分钟视频任务在线扩缩容与续租", func(t *testing.T) {
|
||||||
|
model := "worker-video-long-" + suffix
|
||||||
|
platform := createAsyncAcceptancePlatform(t, server.URL, adminToken, "video-long-"+suffix, "Long Video Simulation", 1)
|
||||||
|
modelID := createAsyncAcceptanceModel(t, server.URL, adminToken, platform.ID, model, "video_generate", "inherit", 0, 120)
|
||||||
|
waitForAsyncWorkerMetric(t, server.URL, "easyai_gateway_async_worker_capacity", 1, 15*time.Second)
|
||||||
|
|
||||||
|
longTaskIDs := submitAsyncSimulationTasks(t, server.URL, adminToken, "/api/v1/videos/generations", model, 3, 180*time.Second)
|
||||||
|
waitForTaskCount(t, ctx, pool, longTaskIDs, "running", 1, 10*time.Second)
|
||||||
|
updateStartedAt := time.Now()
|
||||||
|
updateAsyncAcceptancePlatform(t, server.URL, adminToken, platform, 3, "enabled")
|
||||||
|
waitForAsyncWorkerMetric(t, server.URL, "easyai_gateway_async_worker_desired_capacity", 3, 15*time.Second)
|
||||||
|
waitForAsyncWorkerMetric(t, server.URL, "easyai_gateway_async_worker_capacity", 3, 15*time.Second)
|
||||||
|
waitForActiveLeaseCount(t, ctx, pool, modelID, 3, 15*time.Second)
|
||||||
|
waitForTaskCount(t, ctx, pool, longTaskIDs, "running", 3, 5*time.Second)
|
||||||
|
if time.Since(updateStartedAt) > 15*time.Second {
|
||||||
|
t.Fatalf("capacity expansion took %s, want <=15s", time.Since(updateStartedAt))
|
||||||
|
}
|
||||||
|
|
||||||
|
initialExpiry := activeLeaseExpiry(t, ctx, pool, modelID, 3)
|
||||||
|
updateAsyncAcceptancePlatform(t, server.URL, adminToken, platform, 1, "enabled")
|
||||||
|
waitForAsyncWorkerMetric(t, server.URL, "easyai_gateway_async_worker_capacity", 1, 15*time.Second)
|
||||||
|
shortTaskIDs := submitAsyncSimulationTasks(t, server.URL, adminToken, "/api/v1/videos/generations", model, 1, time.Second)
|
||||||
|
time.Sleep(3 * time.Second)
|
||||||
|
if attempts := taskAttemptCount(t, ctx, pool, shortTaskIDs); attempts != 0 {
|
||||||
|
t.Fatalf("short task created %d upstream attempts while three long leases exceeded reduced limit", attempts)
|
||||||
|
}
|
||||||
|
waitForTaskCount(t, ctx, pool, longTaskIDs, "running", 3, 5*time.Second)
|
||||||
|
|
||||||
|
time.Sleep(130 * time.Second)
|
||||||
|
renewedExpiry := activeLeaseExpiry(t, ctx, pool, modelID, 3)
|
||||||
|
if !renewedExpiry.After(initialExpiry.Add(30*time.Second)) || !renewedExpiry.After(time.Now()) {
|
||||||
|
t.Fatalf("lease was not renewed: initial=%s renewed=%s now=%s", initialExpiry, renewedExpiry, time.Now())
|
||||||
|
}
|
||||||
|
|
||||||
|
waitForTaskCount(t, ctx, pool, longTaskIDs, "succeeded", 3, 75*time.Second)
|
||||||
|
shortStartedAt := time.Now()
|
||||||
|
waitForTaskAttempts(t, ctx, pool, shortTaskIDs, 1, 5*time.Second)
|
||||||
|
waitForTaskCount(t, ctx, pool, shortTaskIDs, "succeeded", 1, 10*time.Second)
|
||||||
|
if time.Since(shortStartedAt) > 5*time.Second {
|
||||||
|
t.Fatalf("short task started after %s once long tasks completed, want <=5s", time.Since(shortStartedAt))
|
||||||
|
}
|
||||||
|
allTaskIDs := append(append([]string{}, longTaskIDs...), shortTaskIDs...)
|
||||||
|
assertAcceptanceAttempts(t, ctx, pool, allTaskIDs, 4)
|
||||||
|
assertSimulationOutputsAndEvents(t, ctx, pool, allTaskIDs)
|
||||||
|
t.Logf("长任务证据: initial_capacity=1 expanded_capacity=3 reduced_capacity=1 initial_expiry=%s renewed_expiry=%s tasks=4", initialExpiry.UTC().Format(time.RFC3339Nano), renewedExpiry.UTC().Format(time.RFC3339Nano))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
type asyncAcceptancePlatform struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Provider string `json:"provider"`
|
||||||
|
PlatformKey string `json:"platformKey"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
BaseURL string `json:"baseUrl"`
|
||||||
|
AuthType string `json:"authType"`
|
||||||
|
Priority int `json:"priority"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func createAsyncAcceptanceAdmin(t *testing.T, ctx context.Context, pool *pgxpool.Pool, baseURL string) string {
|
||||||
|
t.Helper()
|
||||||
|
suffix := strconv.FormatInt(time.Now().UnixNano(), 10)
|
||||||
|
username := "async_acceptance_" + suffix
|
||||||
|
password := "password123"
|
||||||
|
doJSON(t, baseURL, http.MethodPost, "/api/v1/auth/register", "", map[string]any{
|
||||||
|
"username": username,
|
||||||
|
"email": username + "@example.com",
|
||||||
|
"password": password,
|
||||||
|
}, http.StatusCreated, &map[string]any{})
|
||||||
|
if _, err := pool.Exec(ctx, `UPDATE gateway_users SET roles = '["admin"]'::jsonb WHERE username = $1`, username); err != nil {
|
||||||
|
t.Fatalf("promote acceptance user: %v", err)
|
||||||
|
}
|
||||||
|
var login struct {
|
||||||
|
AccessToken string `json:"accessToken"`
|
||||||
|
}
|
||||||
|
doJSON(t, baseURL, http.MethodPost, "/api/v1/auth/login", "", map[string]any{
|
||||||
|
"account": username, "password": password,
|
||||||
|
}, http.StatusOK, &login)
|
||||||
|
if login.AccessToken == "" {
|
||||||
|
t.Fatal("acceptance login returned no access token")
|
||||||
|
}
|
||||||
|
return login.AccessToken
|
||||||
|
}
|
||||||
|
|
||||||
|
func createAsyncAcceptancePlatform(t *testing.T, baseURL, token, key, name string, concurrent int) asyncAcceptancePlatform {
|
||||||
|
t.Helper()
|
||||||
|
var platform asyncAcceptancePlatform
|
||||||
|
doJSON(t, baseURL, http.MethodPost, "/api/admin/platforms", token, asyncAcceptancePlatformPayload(key, name, concurrent, "enabled"), http.StatusCreated, &platform)
|
||||||
|
if platform.ID == "" {
|
||||||
|
t.Fatal("acceptance platform returned no id")
|
||||||
|
}
|
||||||
|
return platform
|
||||||
|
}
|
||||||
|
|
||||||
|
func updateAsyncAcceptancePlatform(t *testing.T, baseURL, token string, platform asyncAcceptancePlatform, concurrent int, status string) {
|
||||||
|
t.Helper()
|
||||||
|
doJSON(t, baseURL, http.MethodPatch, "/api/admin/platforms/"+platform.ID, token,
|
||||||
|
asyncAcceptancePlatformPayload(platform.PlatformKey, platform.Name, concurrent, status), http.StatusOK, &asyncAcceptancePlatform{})
|
||||||
|
}
|
||||||
|
|
||||||
|
func asyncAcceptancePlatformPayload(key, name string, concurrent int, status string) map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"provider": "openai",
|
||||||
|
"platformKey": key,
|
||||||
|
"name": name,
|
||||||
|
"baseUrl": "https://api.openai.com/v1",
|
||||||
|
"authType": "bearer",
|
||||||
|
"credentials": map[string]any{"mode": "simulation"},
|
||||||
|
"config": map[string]any{"testMode": true},
|
||||||
|
"priority": 1,
|
||||||
|
"status": status,
|
||||||
|
"rateLimitPolicy": map[string]any{"rules": []any{
|
||||||
|
map[string]any{"metric": "concurrent", "limit": concurrent, "leaseTtlSeconds": 120},
|
||||||
|
}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func createAsyncAcceptanceModel(t *testing.T, baseURL, token, platformID, model, modelType, mode string, concurrent, ttl int) string {
|
||||||
|
t.Helper()
|
||||||
|
payload := map[string]any{
|
||||||
|
"modelName": model,
|
||||||
|
"providerModelName": model,
|
||||||
|
"modelAlias": model,
|
||||||
|
"modelType": []string{modelType},
|
||||||
|
"displayName": model,
|
||||||
|
"rateLimitPolicyMode": mode,
|
||||||
|
}
|
||||||
|
if mode == "override" {
|
||||||
|
payload["rateLimitPolicy"] = map[string]any{"rules": []any{
|
||||||
|
map[string]any{"metric": "concurrent", "limit": concurrent, "leaseTtlSeconds": ttl},
|
||||||
|
}}
|
||||||
|
}
|
||||||
|
var response struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
}
|
||||||
|
doJSON(t, baseURL, http.MethodPost, "/api/admin/platforms/"+platformID+"/models", token, payload, http.StatusCreated, &response)
|
||||||
|
if response.ID == "" {
|
||||||
|
t.Fatal("acceptance model returned no id")
|
||||||
|
}
|
||||||
|
return response.ID
|
||||||
|
}
|
||||||
|
|
||||||
|
func submitAsyncSimulationTasks(t *testing.T, baseURL, token, path, model string, count int, duration time.Duration) []string {
|
||||||
|
t.Helper()
|
||||||
|
ids := make([]string, count)
|
||||||
|
errs := make(chan error, count)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for index := 0; index < count; index++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(index int) {
|
||||||
|
defer wg.Done()
|
||||||
|
payload := map[string]any{
|
||||||
|
"model": model,
|
||||||
|
"runMode": "simulation",
|
||||||
|
"simulation": true,
|
||||||
|
"simulationDurationMs": duration.Milliseconds(),
|
||||||
|
"prompt": fmt.Sprintf("async acceptance %d", index),
|
||||||
|
"input": fmt.Sprintf("async acceptance %d", index),
|
||||||
|
}
|
||||||
|
raw, _ := json.Marshal(payload)
|
||||||
|
req, err := http.NewRequest(http.MethodPost, baseURL+path, bytes.NewReader(raw))
|
||||||
|
if err != nil {
|
||||||
|
errs <- err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
req.Header.Set("X-Async", "true")
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
errs <- err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
responseBody, _ := io.ReadAll(resp.Body)
|
||||||
|
if resp.StatusCode != http.StatusAccepted {
|
||||||
|
errs <- fmt.Errorf("submit task %d status=%d body=%s", index, resp.StatusCode, responseBody)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var response struct {
|
||||||
|
TaskID string `json:"taskId"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(responseBody, &response); err != nil || response.TaskID == "" {
|
||||||
|
errs <- fmt.Errorf("decode task %d: %w body=%s", index, err, responseBody)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ids[index] = response.TaskID
|
||||||
|
}(index)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(errs)
|
||||||
|
for err := range errs {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForAcceptanceTasks(t *testing.T, ctx context.Context, pool *pgxpool.Pool, taskIDs []string, modelID string, timeout time.Duration) (int64, int64) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(timeout)
|
||||||
|
var peakRunning, peakLeases int64
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
var running, succeeded, failed, leases int64
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT COUNT(*) FILTER (WHERE status = 'running'),
|
||||||
|
COUNT(*) FILTER (WHERE status = 'succeeded'),
|
||||||
|
COUNT(*) FILTER (WHERE status IN ('failed', 'cancelled', 'manual_review'))
|
||||||
|
FROM gateway_tasks
|
||||||
|
WHERE id = ANY($1::uuid[])`, taskIDs).Scan(&running, &succeeded, &failed); err != nil {
|
||||||
|
t.Fatalf("read acceptance tasks: %v", err)
|
||||||
|
}
|
||||||
|
if err := 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()`, modelID).Scan(&leases); err != nil {
|
||||||
|
t.Fatalf("read acceptance leases: %v", err)
|
||||||
|
}
|
||||||
|
peakRunning = maxInt64(peakRunning, running)
|
||||||
|
peakLeases = maxInt64(peakLeases, leases)
|
||||||
|
if failed > 0 {
|
||||||
|
t.Fatalf("acceptance tasks entered failed/manual status: %d", failed)
|
||||||
|
}
|
||||||
|
if succeeded == int64(len(taskIDs)) {
|
||||||
|
return peakRunning, peakLeases
|
||||||
|
}
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("tasks did not finish within %s", timeout)
|
||||||
|
return 0, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForTaskCount(t *testing.T, ctx context.Context, pool *pgxpool.Pool, taskIDs []string, status string, count int, timeout time.Duration) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(timeout)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
var got int
|
||||||
|
if err := pool.QueryRow(ctx, `SELECT COUNT(*) FROM gateway_tasks WHERE id = ANY($1::uuid[]) AND status = $2`, taskIDs, status).Scan(&got); err != nil {
|
||||||
|
t.Fatalf("read task count: %v", err)
|
||||||
|
}
|
||||||
|
if got == count {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("task status %s count did not reach %d within %s", status, count, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func activeLeaseExpiry(t *testing.T, ctx context.Context, pool *pgxpool.Pool, modelID string, wantCount int) time.Time {
|
||||||
|
t.Helper()
|
||||||
|
var count int
|
||||||
|
var earliest time.Time
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT COUNT(*), MIN(expires_at)
|
||||||
|
FROM gateway_concurrency_leases
|
||||||
|
WHERE scope_type = 'platform_model'
|
||||||
|
AND scope_key = $1
|
||||||
|
AND released_at IS NULL
|
||||||
|
AND expires_at > now()`, modelID).Scan(&count, &earliest); err != nil {
|
||||||
|
t.Fatalf("read active lease expiry: %v", err)
|
||||||
|
}
|
||||||
|
if count != wantCount {
|
||||||
|
t.Fatalf("active lease count=%d, want=%d", count, wantCount)
|
||||||
|
}
|
||||||
|
return earliest
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForActiveLeaseCount(t *testing.T, ctx context.Context, pool *pgxpool.Pool, modelID string, count int, timeout time.Duration) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(timeout)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
var got int
|
||||||
|
if err := 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()`, modelID).Scan(&got); err != nil {
|
||||||
|
t.Fatalf("read active lease count: %v", err)
|
||||||
|
}
|
||||||
|
if got == count {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("active lease count did not reach %d within %s", count, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func taskAttemptCount(t *testing.T, ctx context.Context, pool *pgxpool.Pool, taskIDs []string) int {
|
||||||
|
t.Helper()
|
||||||
|
var count int
|
||||||
|
if err := pool.QueryRow(ctx, `SELECT COUNT(*) FROM gateway_task_attempts WHERE task_id = ANY($1::uuid[])`, taskIDs).Scan(&count); err != nil {
|
||||||
|
t.Fatalf("count task attempts: %v", err)
|
||||||
|
}
|
||||||
|
return count
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForTaskAttempts(t *testing.T, ctx context.Context, pool *pgxpool.Pool, taskIDs []string, count int, timeout time.Duration) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(timeout)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if taskAttemptCount(t, ctx, pool, taskIDs) == count {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("attempt count did not reach %d within %s", count, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertAcceptanceAttempts(t *testing.T, ctx context.Context, pool *pgxpool.Pool, taskIDs []string, want int) {
|
||||||
|
t.Helper()
|
||||||
|
var attempts, duplicateTasks, manualReview int
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT COUNT(*),
|
||||||
|
COUNT(DISTINCT task_id) FILTER (WHERE per_task.attempts > 1),
|
||||||
|
COUNT(*) FILTER (WHERE status = 'manual_review')
|
||||||
|
FROM (
|
||||||
|
SELECT task_id, COUNT(*) AS attempts, MAX(status) AS status
|
||||||
|
FROM gateway_task_attempts
|
||||||
|
WHERE task_id = ANY($1::uuid[])
|
||||||
|
GROUP BY task_id
|
||||||
|
) per_task`, taskIDs).Scan(&attempts, &duplicateTasks, &manualReview); err != nil {
|
||||||
|
t.Fatalf("read attempt acceptance evidence: %v", err)
|
||||||
|
}
|
||||||
|
if attempts != want || duplicateTasks != 0 || manualReview != 0 {
|
||||||
|
t.Fatalf("attempts=%d duplicate_tasks=%d manual_review=%d, want %d/0/0", attempts, duplicateTasks, manualReview, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertSimulationOutputsAndEvents(t *testing.T, ctx context.Context, pool *pgxpool.Pool, taskIDs []string) {
|
||||||
|
t.Helper()
|
||||||
|
var complete int
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT COUNT(*)
|
||||||
|
FROM gateway_tasks task
|
||||||
|
WHERE task.id = ANY($1::uuid[])
|
||||||
|
AND task.status = 'succeeded'
|
||||||
|
AND task.finished_at IS NOT NULL
|
||||||
|
AND task.result IS NOT NULL
|
||||||
|
AND task.result <> '{}'::jsonb
|
||||||
|
AND task.metrics IS NOT NULL
|
||||||
|
AND jsonb_typeof(task.billings) = 'array'
|
||||||
|
AND jsonb_array_length(task.billings) > 0
|
||||||
|
AND EXISTS (SELECT 1 FROM gateway_task_events event WHERE event.task_id = task.id AND event.event_type = 'task.completed')`,
|
||||||
|
taskIDs).Scan(&complete); err != nil {
|
||||||
|
t.Fatalf("read simulation output evidence: %v", err)
|
||||||
|
}
|
||||||
|
if complete != len(taskIDs) {
|
||||||
|
t.Fatalf("complete simulation outputs/events=%d, want=%d", complete, len(taskIDs))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForAsyncWorkerMetric(t *testing.T, baseURL, metric string, want int64, timeout time.Duration) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(timeout)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
resp, err := http.Get(baseURL + "/metrics")
|
||||||
|
if err == nil {
|
||||||
|
raw, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode == http.StatusOK {
|
||||||
|
for _, line := range strings.Split(string(raw), "\n") {
|
||||||
|
fields := strings.Fields(line)
|
||||||
|
if len(fields) == 2 && fields[0] == metric {
|
||||||
|
value, _ := strconv.ParseInt(fields[1], 10, 64)
|
||||||
|
if value == want {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("metric %s did not reach %d within %s", metric, want, timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
func maxInt64(left, right int64) int64 {
|
||||||
|
if right > left {
|
||||||
|
return right
|
||||||
|
}
|
||||||
|
return left
|
||||||
|
}
|
||||||
@@ -487,18 +487,20 @@ func discountTitle(label string, discount float64) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func effectiveModelRateLimits(model store.PlatformModel, platform store.Platform, runtimePolicyMap map[string]store.RuntimePolicySet) ModelCatalogRateLimits {
|
func effectiveModelRateLimits(model store.PlatformModel, platform store.Platform, runtimePolicyMap map[string]store.RuntimePolicySet) ModelCatalogRateLimits {
|
||||||
overridePolicyRaw, hasOverridePolicy := model.RuntimePolicyOverride["rateLimitPolicy"]
|
|
||||||
overridePolicy := objectValue(overridePolicyRaw)
|
|
||||||
runtimePolicy := map[string]any(nil)
|
runtimePolicy := map[string]any(nil)
|
||||||
if model.RuntimePolicySetID != "" {
|
if model.RuntimePolicySetID != "" {
|
||||||
runtimePolicy = runtimePolicyMap[model.RuntimePolicySetID].RateLimitPolicy
|
runtimePolicy = runtimePolicyMap[model.RuntimePolicySetID].RateLimitPolicy
|
||||||
}
|
}
|
||||||
policies := []rateLimitPolicySource{
|
policy := store.EffectiveRateLimitPolicy(store.EffectiveRateLimitPolicyInput{
|
||||||
{policy: overridePolicy, authoritative: hasOverridePolicy},
|
BasePolicy: model.BaseRateLimitPolicy,
|
||||||
{policy: model.RateLimitPolicy, authoritative: len(model.RateLimitPolicy) > 0},
|
PlatformPolicy: platform.RateLimitPolicy,
|
||||||
{policy: runtimePolicy, authoritative: strings.TrimSpace(model.RuntimePolicySetID) != ""},
|
RuntimePolicy: runtimePolicy,
|
||||||
{policy: platform.RateLimitPolicy},
|
RuntimePolicyExplicit: model.RuntimePolicySetID != "",
|
||||||
}
|
RuntimePolicyOverride: model.RuntimePolicyOverride,
|
||||||
|
ModelPolicy: model.RateLimitPolicy,
|
||||||
|
ModelPolicyMode: model.RateLimitPolicyMode,
|
||||||
|
})
|
||||||
|
policies := []rateLimitPolicySource{{policy: policy, authoritative: policy != nil}}
|
||||||
limits := ModelCatalogRateLimits{
|
limits := ModelCatalogRateLimits{
|
||||||
RPM: firstRateLimit(policies, "rpm"),
|
RPM: firstRateLimit(policies, "rpm"),
|
||||||
TPM: firstRateLimit(policies, "tpm_total"),
|
TPM: firstRateLimit(policies, "tpm_total"),
|
||||||
|
|||||||
@@ -74,7 +74,7 @@ func (s *Service) rateLimitReservations(ctx context.Context, user *auth.User, ca
|
|||||||
"groupKey": group.GroupKey,
|
"groupKey": group.GroupKey,
|
||||||
"name": group.Name,
|
"name": group.Name,
|
||||||
},
|
},
|
||||||
group.RateLimitPolicy,
|
store.NormalizeRateLimitPolicy(group.RateLimitPolicy),
|
||||||
body,
|
body,
|
||||||
)...)
|
)...)
|
||||||
}
|
}
|
||||||
@@ -82,27 +82,15 @@ func (s *Service) rateLimitReservations(ctx context.Context, user *auth.User, ca
|
|||||||
}
|
}
|
||||||
|
|
||||||
func effectiveRateLimitPolicy(candidate store.RuntimeModelCandidate) map[string]any {
|
func effectiveRateLimitPolicy(candidate store.RuntimeModelCandidate) map[string]any {
|
||||||
policy := candidate.PlatformRateLimitPolicy
|
return store.EffectiveRateLimitPolicy(store.EffectiveRateLimitPolicyInput{
|
||||||
if strings.TrimSpace(candidate.RuntimePolicySetID) != "" {
|
BasePolicy: candidate.BaseRateLimitPolicy,
|
||||||
policy = candidate.RuntimeRateLimitPolicy
|
PlatformPolicy: candidate.PlatformRateLimitPolicy,
|
||||||
} else if hasRules(candidate.RuntimeRateLimitPolicy) {
|
RuntimePolicy: candidate.RuntimeRateLimitPolicy,
|
||||||
policy = mergeMap(policy, candidate.RuntimeRateLimitPolicy)
|
RuntimePolicyExplicit: candidate.RuntimePolicyExplicit,
|
||||||
}
|
RuntimePolicyOverride: candidate.RateLimitRuntimeOverride,
|
||||||
if _, hasOverride := candidate.RuntimePolicyOverride["rateLimitPolicy"]; hasOverride {
|
ModelPolicy: candidate.ModelRateLimitPolicy,
|
||||||
nested, _ := candidate.RuntimePolicyOverride["rateLimitPolicy"].(map[string]any)
|
ModelPolicyMode: candidate.ModelRateLimitPolicyMode,
|
||||||
if len(nested) == 0 {
|
})
|
||||||
policy = nil
|
|
||||||
} else {
|
|
||||||
policy = mergeMap(policy, nested)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if hasRules(candidate.ModelRateLimitPolicy) {
|
|
||||||
policy = mergeMap(policy, candidate.ModelRateLimitPolicy)
|
|
||||||
}
|
|
||||||
if hasRules(policy) {
|
|
||||||
return policy
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func effectiveRetryPolicy(candidate store.RuntimeModelCandidate) map[string]any {
|
func effectiveRetryPolicy(candidate store.RuntimeModelCandidate) map[string]any {
|
||||||
|
|||||||
@@ -102,8 +102,10 @@ func TestEffectiveRateLimitPolicyTreatsEmptyRuntimePolicyAsUnlimited(t *testing.
|
|||||||
PlatformRateLimitPolicy: map[string]any{"rules": []any{
|
PlatformRateLimitPolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"metric": "rpm", "limit": 500},
|
map[string]any{"metric": "rpm", "limit": 500},
|
||||||
}},
|
}},
|
||||||
RuntimePolicySetID: "runtime-policy-1",
|
RuntimePolicySetID: "runtime-policy-1",
|
||||||
RuntimeRateLimitPolicy: map[string]any{"rules": []any{}},
|
RuntimePolicyExplicit: true,
|
||||||
|
RuntimeRateLimitPolicy: map[string]any{"rules": []any{}},
|
||||||
|
ModelRateLimitPolicyMode: "inherit",
|
||||||
})
|
})
|
||||||
|
|
||||||
if hasRules(policy) {
|
if hasRules(policy) {
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
"github.com/riverqueue/river"
|
"github.com/riverqueue/river"
|
||||||
"github.com/riverqueue/river/riverdriver/riverpgxv5"
|
"github.com/riverqueue/river/riverdriver/riverpgxv5"
|
||||||
"github.com/riverqueue/river/rivermigrate"
|
"github.com/riverqueue/river/rivermigrate"
|
||||||
@@ -23,6 +24,12 @@ type asyncTaskArgs struct {
|
|||||||
TaskID string `json:"task_id" river:"unique"`
|
TaskID string `json:"task_id" river:"unique"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type asyncExecutionClient interface {
|
||||||
|
Start(context.Context) error
|
||||||
|
Stop(context.Context) error
|
||||||
|
StopAndCancel(context.Context) error
|
||||||
|
}
|
||||||
|
|
||||||
func (asyncTaskArgs) Kind() string { return "gateway_task_run" }
|
func (asyncTaskArgs) Kind() string { return "gateway_task_run" }
|
||||||
|
|
||||||
type asyncTaskWorker struct {
|
type asyncTaskWorker struct {
|
||||||
@@ -87,21 +94,70 @@ func (s *Service) startRiverQueue(ctx context.Context) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
workers := river.NewWorkers()
|
controlClient, err := river.NewClient(driver, &river.Config{
|
||||||
if err := river.AddWorkerSafely(workers, &asyncTaskWorker{service: s}); err != nil {
|
ID: asyncWorkerID() + "-control",
|
||||||
|
Logger: s.logger,
|
||||||
|
TestOnly: s.cfg.AppEnv == "test",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
riverClient, err := river.NewClient(driver, &river.Config{
|
snapshot, err := s.loadAsyncWorkerCapacity(ctx)
|
||||||
ID: asyncWorkerID(),
|
if err != nil {
|
||||||
|
return fmt.Errorf("calculate initial async worker capacity: %w", err)
|
||||||
|
}
|
||||||
|
executionClient, err := s.makeAsyncExecutionClient(snapshot.Capacity)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := executionClient.Start(ctx); err != nil {
|
||||||
|
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
_ = executionClient.StopAndCancel(cleanupCtx)
|
||||||
|
cancel()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.riverMu.Lock()
|
||||||
|
s.riverControlClient = controlClient
|
||||||
|
s.riverExecutionClient = executionClient
|
||||||
|
s.riverWorkerCapacity = snapshot.Capacity
|
||||||
|
s.riverDrainingClients = make(map[asyncExecutionClient]struct{})
|
||||||
|
s.riverMu.Unlock()
|
||||||
|
s.observeAsyncWorkerCapacity(snapshot)
|
||||||
|
s.logger.Info("async worker capacity initialized",
|
||||||
|
"capacity", snapshot.Capacity,
|
||||||
|
"desiredCapacity", snapshot.Desired,
|
||||||
|
"hardLimit", snapshot.HardLimit,
|
||||||
|
"enabledModels", snapshot.EnabledModels,
|
||||||
|
"unlimitedModels", snapshot.UnlimitedModels,
|
||||||
|
"modelDesired", snapshot.ModelDesired,
|
||||||
|
"enabledGroups", snapshot.EnabledGroups,
|
||||||
|
"unlimitedGroups", snapshot.UnlimitedGroups,
|
||||||
|
"groupDesired", snapshot.GroupDesired,
|
||||||
|
"capped", snapshot.Capped,
|
||||||
|
)
|
||||||
|
if err := s.recoverAsyncRiverJobs(ctx); err != nil {
|
||||||
|
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
_ = executionClient.StopAndCancel(cleanupCtx)
|
||||||
|
cancel()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
go s.refreshAsyncWorkerCapacity(ctx)
|
||||||
|
go s.stopAsyncWorkersOnShutdown(ctx)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) newRiverAsyncExecutionClient(capacity int) (*river.Client[pgx.Tx], error) {
|
||||||
|
workers := river.NewWorkers()
|
||||||
|
if err := river.AddWorkerSafely(workers, &asyncTaskWorker{service: s}); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return river.NewClient(riverpgxv5.New(s.store.Pool()), &river.Config{
|
||||||
|
ID: fmt.Sprintf("%s-exec-%d-%d", asyncWorkerID(), capacity, time.Now().UnixNano()),
|
||||||
JobTimeout: -1,
|
JobTimeout: -1,
|
||||||
Logger: s.logger,
|
Logger: s.logger,
|
||||||
CompletedJobRetentionPeriod: 24 * time.Hour,
|
CompletedJobRetentionPeriod: 24 * time.Hour,
|
||||||
Queues: map[string]river.QueueConfig{
|
Queues: map[string]river.QueueConfig{
|
||||||
// Image providers may hold a worker while polling for several
|
asyncTaskQueueName: {MaxWorkers: capacity},
|
||||||
// minutes. Keep enough workers available for production bursts so
|
|
||||||
// unrelated models do not remain queued behind long-running media
|
|
||||||
// tasks.
|
|
||||||
asyncTaskQueueName: {MaxWorkers: 96},
|
|
||||||
},
|
},
|
||||||
// Provider-backed media jobs commonly poll for 10-20 minutes. River may
|
// Provider-backed media jobs commonly poll for 10-20 minutes. River may
|
||||||
// execute a still-running job again once this window elapses, so keep the
|
// execute a still-running job again once this window elapses, so keep the
|
||||||
@@ -110,32 +166,160 @@ func (s *Service) startRiverQueue(ctx context.Context) error {
|
|||||||
TestOnly: s.cfg.AppEnv == "test",
|
TestOnly: s.cfg.AppEnv == "test",
|
||||||
Workers: workers,
|
Workers: workers,
|
||||||
})
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) makeAsyncExecutionClient(capacity int) (asyncExecutionClient, error) {
|
||||||
|
if s.asyncClientFactory != nil {
|
||||||
|
return s.asyncClientFactory(capacity)
|
||||||
|
}
|
||||||
|
return s.newRiverAsyncExecutionClient(capacity)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) loadAsyncWorkerCapacity(ctx context.Context) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||||
|
if s.asyncCapacityLoader != nil {
|
||||||
|
return s.asyncCapacityLoader(ctx, s.cfg.AsyncWorkerHardLimit)
|
||||||
|
}
|
||||||
|
return s.store.AsyncWorkerCapacity(ctx, s.cfg.AsyncWorkerHardLimit)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) refreshAsyncWorkerCapacity(ctx context.Context) {
|
||||||
|
ticker := time.NewTicker(time.Duration(s.cfg.AsyncWorkerRefreshIntervalSeconds) * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
s.resizeAsyncWorkerCapacity(ctx)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) resizeAsyncWorkerCapacity(ctx context.Context) {
|
||||||
|
snapshot, err := s.loadAsyncWorkerCapacity(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
s.observeAsyncWorkerResize("refresh_failed")
|
||||||
|
s.logger.Warn("refresh async worker capacity failed; keeping current client", "error", err)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
s.riverClient = riverClient
|
s.riverMu.RLock()
|
||||||
if err := riverClient.Start(ctx); err != nil {
|
currentCapacity := s.riverWorkerCapacity
|
||||||
return err
|
s.riverMu.RUnlock()
|
||||||
|
s.observeAsyncWorkerCapacity(snapshot)
|
||||||
|
if snapshot.Capacity == currentCapacity {
|
||||||
|
return
|
||||||
}
|
}
|
||||||
if err := s.recoverAsyncRiverJobs(ctx); err != nil {
|
newClient, err := s.makeAsyncExecutionClient(snapshot.Capacity)
|
||||||
return err
|
if err != nil {
|
||||||
|
s.observeAsyncWorkerResize("create_failed")
|
||||||
|
s.logger.Warn("create replacement async worker client failed; keeping current client",
|
||||||
|
"error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
go func() {
|
if err := newClient.Start(ctx); err != nil {
|
||||||
<-ctx.Done()
|
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
_ = newClient.StopAndCancel(cleanupCtx)
|
||||||
|
cancel()
|
||||||
|
s.observeAsyncWorkerResize("start_failed")
|
||||||
|
s.logger.Warn("start replacement async worker client failed; keeping current client",
|
||||||
|
"error", err, "currentCapacity", currentCapacity, "desiredCapacity", snapshot.Capacity)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
_ = newClient.StopAndCancel(cleanupCtx)
|
||||||
|
cancel()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.riverMu.Lock()
|
||||||
|
oldClient := s.riverExecutionClient
|
||||||
|
s.riverExecutionClient = newClient
|
||||||
|
s.riverWorkerCapacity = snapshot.Capacity
|
||||||
|
if oldClient != nil {
|
||||||
|
if s.riverDrainingClients == nil {
|
||||||
|
s.riverDrainingClients = make(map[asyncExecutionClient]struct{})
|
||||||
|
}
|
||||||
|
s.riverDrainingClients[oldClient] = struct{}{}
|
||||||
|
}
|
||||||
|
s.riverMu.Unlock()
|
||||||
|
s.observeAsyncWorkerResize("success")
|
||||||
|
s.logger.Info("async worker capacity resized",
|
||||||
|
"previousCapacity", currentCapacity,
|
||||||
|
"capacity", snapshot.Capacity,
|
||||||
|
"desiredCapacity", snapshot.Desired,
|
||||||
|
"hardLimit", snapshot.HardLimit,
|
||||||
|
"modelDesired", snapshot.ModelDesired,
|
||||||
|
"groupDesired", snapshot.GroupDesired,
|
||||||
|
"capped", snapshot.Capped,
|
||||||
|
)
|
||||||
|
if oldClient != nil {
|
||||||
|
go s.drainAsyncWorkerClient(oldClient)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) drainAsyncWorkerClient(client asyncExecutionClient) {
|
||||||
|
stopCtx, cancel := context.WithTimeout(context.Background(), time.Hour)
|
||||||
|
defer cancel()
|
||||||
|
if err := client.Stop(stopCtx); err != nil {
|
||||||
|
s.logger.Warn("gracefully drain previous async worker client failed", "error", err)
|
||||||
|
}
|
||||||
|
s.riverMu.Lock()
|
||||||
|
delete(s.riverDrainingClients, client)
|
||||||
|
s.riverMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) stopAsyncWorkersOnShutdown(ctx context.Context) {
|
||||||
|
<-ctx.Done()
|
||||||
|
s.riverMu.Lock()
|
||||||
|
clients := make([]asyncExecutionClient, 0, 1+len(s.riverDrainingClients))
|
||||||
|
if s.riverExecutionClient != nil {
|
||||||
|
clients = append(clients, s.riverExecutionClient)
|
||||||
|
}
|
||||||
|
for client := range s.riverDrainingClients {
|
||||||
|
clients = append(clients, client)
|
||||||
|
}
|
||||||
|
s.riverExecutionClient = nil
|
||||||
|
s.riverDrainingClients = nil
|
||||||
|
s.riverMu.Unlock()
|
||||||
|
for _, client := range clients {
|
||||||
stopCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
stopCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
defer cancel()
|
if err := client.StopAndCancel(stopCtx); err != nil {
|
||||||
if err := riverClient.StopAndCancel(stopCtx); err != nil {
|
|
||||||
s.logger.Warn("stop river async queue failed", "error", err)
|
s.logger.Warn("stop river async queue failed", "error", err)
|
||||||
}
|
}
|
||||||
}()
|
cancel()
|
||||||
return nil
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) asyncControlClient() *river.Client[pgx.Tx] {
|
||||||
|
s.riverMu.RLock()
|
||||||
|
defer s.riverMu.RUnlock()
|
||||||
|
return s.riverControlClient
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) observeAsyncWorkerCapacity(snapshot store.AsyncWorkerCapacitySnapshot) {
|
||||||
|
observer, ok := s.billingMetrics.(interface {
|
||||||
|
SetAsyncWorkerCapacity(current, desired, hardLimit int, capped bool)
|
||||||
|
})
|
||||||
|
if ok {
|
||||||
|
observer.SetAsyncWorkerCapacity(snapshot.Capacity, snapshot.Desired, snapshot.HardLimit, snapshot.Capped)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) observeAsyncWorkerResize(outcome string) {
|
||||||
|
observer, ok := s.billingMetrics.(interface {
|
||||||
|
ObserveAsyncWorkerResize(string)
|
||||||
|
})
|
||||||
|
if ok {
|
||||||
|
observer.ObserveAsyncWorkerResize(outcome)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Service) EnqueueAsyncTask(ctx context.Context, task store.GatewayTask) error {
|
func (s *Service) EnqueueAsyncTask(ctx context.Context, task store.GatewayTask) error {
|
||||||
if s.riverClient == nil {
|
riverClient := s.asyncControlClient()
|
||||||
|
if riverClient == nil {
|
||||||
return errors.New("river async queue is not started")
|
return errors.New("river async queue is not started")
|
||||||
}
|
}
|
||||||
result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task))
|
result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: task.ID}, asyncTaskInsertOpts(task))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -155,12 +339,16 @@ func (s *Service) RunAsyncTask(ctx context.Context, task store.GatewayTask, user
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *Service) recoverAsyncRiverJobs(ctx context.Context) error {
|
func (s *Service) recoverAsyncRiverJobs(ctx context.Context) error {
|
||||||
|
riverClient := s.asyncControlClient()
|
||||||
|
if riverClient == nil {
|
||||||
|
return errors.New("river async queue is not started")
|
||||||
|
}
|
||||||
items, err := s.store.ListRecoverableAsyncTasks(ctx, 1000)
|
items, err := s.store.ListRecoverableAsyncTasks(ctx, 1000)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
for _, item := range items {
|
for _, item := range items {
|
||||||
result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item, time.Now()))
|
result, err := riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item, time.Now()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,189 @@
|
|||||||
|
package runner
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeAsyncExecutionClient struct {
|
||||||
|
startErr error
|
||||||
|
stopGate <-chan struct{}
|
||||||
|
started atomic.Int64
|
||||||
|
stopped atomic.Int64
|
||||||
|
stopCancelled atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeAsyncExecutionClient) Start(context.Context) error {
|
||||||
|
c.started.Add(1)
|
||||||
|
return c.startErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeAsyncExecutionClient) Stop(ctx context.Context) error {
|
||||||
|
c.stopped.Add(1)
|
||||||
|
if c.stopGate == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-c.stopGate:
|
||||||
|
return nil
|
||||||
|
case <-ctx.Done():
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeAsyncExecutionClient) StopAndCancel(context.Context) error {
|
||||||
|
c.stopCancelled.Add(1)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResizeAsyncWorkerCapacityStartsReplacementBeforeGracefulDrain(t *testing.T) {
|
||||||
|
oldClient := &fakeAsyncExecutionClient{}
|
||||||
|
drainGate := make(chan struct{})
|
||||||
|
oldClient.stopGate = drainGate
|
||||||
|
newClient := &fakeAsyncExecutionClient{}
|
||||||
|
service := asyncWorkerManagerTestService(1)
|
||||||
|
service.riverExecutionClient = oldClient
|
||||||
|
service.asyncCapacityLoader = fixedAsyncCapacity(3)
|
||||||
|
service.asyncClientFactory = func(capacity int) (asyncExecutionClient, error) {
|
||||||
|
if capacity != 3 {
|
||||||
|
t.Fatalf("factory capacity=%d, want=3", capacity)
|
||||||
|
}
|
||||||
|
return newClient, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
service.resizeAsyncWorkerCapacity(context.Background())
|
||||||
|
if newClient.started.Load() != 1 {
|
||||||
|
t.Fatalf("replacement start count=%d, want=1", newClient.started.Load())
|
||||||
|
}
|
||||||
|
if service.riverExecutionClient != newClient || service.riverWorkerCapacity != 3 {
|
||||||
|
t.Fatalf("replacement was not installed: capacity=%d client=%T", service.riverWorkerCapacity, service.riverExecutionClient)
|
||||||
|
}
|
||||||
|
waitForAtomicValue(t, &oldClient.stopped, 1)
|
||||||
|
if oldClient.stopCancelled.Load() != 0 {
|
||||||
|
t.Fatal("graceful drain cancelled an already running old task")
|
||||||
|
}
|
||||||
|
close(drainGate)
|
||||||
|
waitForDrainingClients(t, service, 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResizeAsyncWorkerCapacityFailuresKeepCurrentClient(t *testing.T) {
|
||||||
|
t.Run("refresh failure", func(t *testing.T) {
|
||||||
|
oldClient := &fakeAsyncExecutionClient{}
|
||||||
|
service := asyncWorkerManagerTestService(4)
|
||||||
|
service.riverExecutionClient = oldClient
|
||||||
|
service.asyncCapacityLoader = func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||||
|
return store.AsyncWorkerCapacitySnapshot{}, errors.New("database unavailable")
|
||||||
|
}
|
||||||
|
service.asyncClientFactory = func(int) (asyncExecutionClient, error) {
|
||||||
|
t.Fatal("factory must not run after refresh failure")
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
service.resizeAsyncWorkerCapacity(context.Background())
|
||||||
|
if service.riverExecutionClient != oldClient || service.riverWorkerCapacity != 4 {
|
||||||
|
t.Fatal("refresh failure replaced the current client")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("create failure", func(t *testing.T) {
|
||||||
|
oldClient := &fakeAsyncExecutionClient{}
|
||||||
|
service := asyncWorkerManagerTestService(4)
|
||||||
|
service.riverExecutionClient = oldClient
|
||||||
|
service.asyncCapacityLoader = fixedAsyncCapacity(8)
|
||||||
|
service.asyncClientFactory = func(int) (asyncExecutionClient, error) {
|
||||||
|
return nil, errors.New("factory failed")
|
||||||
|
}
|
||||||
|
service.resizeAsyncWorkerCapacity(context.Background())
|
||||||
|
if service.riverExecutionClient != oldClient || service.riverWorkerCapacity != 4 {
|
||||||
|
t.Fatal("create failure replaced the current client")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("start failure", func(t *testing.T) {
|
||||||
|
oldClient := &fakeAsyncExecutionClient{}
|
||||||
|
failedClient := &fakeAsyncExecutionClient{startErr: errors.New("start failed")}
|
||||||
|
service := asyncWorkerManagerTestService(4)
|
||||||
|
service.riverExecutionClient = oldClient
|
||||||
|
service.asyncCapacityLoader = fixedAsyncCapacity(8)
|
||||||
|
service.asyncClientFactory = func(int) (asyncExecutionClient, error) {
|
||||||
|
return failedClient, nil
|
||||||
|
}
|
||||||
|
service.resizeAsyncWorkerCapacity(context.Background())
|
||||||
|
if service.riverExecutionClient != oldClient || service.riverWorkerCapacity != 4 {
|
||||||
|
t.Fatal("start failure replaced the current client")
|
||||||
|
}
|
||||||
|
if failedClient.stopCancelled.Load() != 1 {
|
||||||
|
t.Fatal("failed replacement client was not cleaned up")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStopAsyncWorkersIncludesCurrentAndDrainingClients(t *testing.T) {
|
||||||
|
current := &fakeAsyncExecutionClient{}
|
||||||
|
draining := &fakeAsyncExecutionClient{}
|
||||||
|
service := asyncWorkerManagerTestService(2)
|
||||||
|
service.riverExecutionClient = current
|
||||||
|
service.riverDrainingClients[draining] = struct{}{}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
service.stopAsyncWorkersOnShutdown(ctx)
|
||||||
|
if current.stopCancelled.Load() != 1 || draining.stopCancelled.Load() != 1 {
|
||||||
|
t.Fatalf("shutdown cancellations current=%d draining=%d, want 1/1", current.stopCancelled.Load(), draining.stopCancelled.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func asyncWorkerManagerTestService(capacity int) *Service {
|
||||||
|
return &Service{
|
||||||
|
cfg: config.Config{
|
||||||
|
AsyncWorkerHardLimit: 2048,
|
||||||
|
AsyncWorkerRefreshIntervalSeconds: 5,
|
||||||
|
},
|
||||||
|
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
||||||
|
riverDrainingClients: make(map[asyncExecutionClient]struct{}),
|
||||||
|
riverWorkerCapacity: capacity,
|
||||||
|
riverExecutionClient: nil,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func fixedAsyncCapacity(capacity int) func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||||
|
return func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error) {
|
||||||
|
return store.AsyncWorkerCapacitySnapshot{Capacity: capacity, Desired: capacity, HardLimit: 2048}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForAtomicValue(t *testing.T, value *atomic.Int64, want int64) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if value.Load() == want {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("atomic value=%d, want=%d", value.Load(), want)
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForDrainingClients(t *testing.T, service *Service, want int) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(time.Second)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
service.riverMu.RLock()
|
||||||
|
count := len(service.riverDrainingClients)
|
||||||
|
service.riverMu.RUnlock()
|
||||||
|
if count == want {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(time.Millisecond)
|
||||||
|
}
|
||||||
|
service.riverMu.RLock()
|
||||||
|
count := len(service.riverDrainingClients)
|
||||||
|
service.riverMu.RUnlock()
|
||||||
|
t.Fatalf("draining client count=%d, want=%d", count, want)
|
||||||
|
}
|
||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
@@ -23,14 +24,20 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type Service struct {
|
type Service struct {
|
||||||
cfg config.Config
|
cfg config.Config
|
||||||
store *store.Store
|
store *store.Store
|
||||||
logger *slog.Logger
|
logger *slog.Logger
|
||||||
clients map[string]clients.Client
|
clients map[string]clients.Client
|
||||||
scriptExecutor *scriptengine.Executor
|
scriptExecutor *scriptengine.Executor
|
||||||
httpClients *httpClientCache
|
httpClients *httpClientCache
|
||||||
riverClient *river.Client[pgx.Tx]
|
riverMu sync.RWMutex
|
||||||
billingMetrics billingMetricsObserver
|
riverControlClient *river.Client[pgx.Tx]
|
||||||
|
riverExecutionClient asyncExecutionClient
|
||||||
|
riverDrainingClients map[asyncExecutionClient]struct{}
|
||||||
|
riverWorkerCapacity int
|
||||||
|
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
|
||||||
|
asyncClientFactory func(int) (asyncExecutionClient, error)
|
||||||
|
billingMetrics billingMetricsObserver
|
||||||
}
|
}
|
||||||
|
|
||||||
type billingMetricsObserver interface {
|
type billingMetricsObserver interface {
|
||||||
@@ -75,6 +82,12 @@ func (e *TaskQueuedError) Is(target error) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...billingMetricsObserver) *Service {
|
func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...billingMetricsObserver) *Service {
|
||||||
|
if cfg.AsyncWorkerHardLimit == 0 {
|
||||||
|
cfg.AsyncWorkerHardLimit = 2048
|
||||||
|
}
|
||||||
|
if cfg.AsyncWorkerRefreshIntervalSeconds == 0 {
|
||||||
|
cfg.AsyncWorkerRefreshIntervalSeconds = 5
|
||||||
|
}
|
||||||
httpClients := newHTTPClientCache()
|
httpClients := newHTTPClientCache()
|
||||||
scriptExecutor := &scriptengine.Executor{Logger: logger}
|
scriptExecutor := &scriptengine.Executor{Logger: logger}
|
||||||
service := &Service{
|
service := &Service{
|
||||||
@@ -891,7 +904,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
|||||||
_ = s.store.ReleaseRateLimitReservations(context.WithoutCancel(ctx), limitResult.Reservations, "attempt_failed")
|
_ = s.store.ReleaseRateLimitReservations(context.WithoutCancel(ctx), limitResult.Reservations, "attempt_failed")
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.LeaseIDs)
|
defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.Leases)
|
||||||
|
|
||||||
attemptID, err := s.store.CreateTaskAttempt(ctx, store.CreateTaskAttemptInput{
|
attemptID, err := s.store.CreateTaskAttempt(ctx, store.CreateTaskAttemptInput{
|
||||||
TaskID: task.ID,
|
TaskID: task.ID,
|
||||||
@@ -1031,7 +1044,8 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
var submissionWire *clients.WireResponse
|
var submissionWire *clients.WireResponse
|
||||||
response, err := client.Run(ctx, clients.Request{
|
runCtx, stopLeaseRenewal := s.startConcurrencyLeaseRenewal(ctx, task.ID, limitResult.Leases)
|
||||||
|
response, err := client.Run(runCtx, clients.Request{
|
||||||
Kind: task.Kind,
|
Kind: task.Kind,
|
||||||
ModelType: candidate.ModelType,
|
ModelType: candidate.ModelType,
|
||||||
Model: task.Model,
|
Model: task.Model,
|
||||||
@@ -1083,6 +1097,13 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
|
|||||||
UpstreamPreviousResponseID: responseExecution.UpstreamPreviousResponseID,
|
UpstreamPreviousResponseID: responseExecution.UpstreamPreviousResponseID,
|
||||||
PreviousResponseTurns: responseExecution.PreviousTurns,
|
PreviousResponseTurns: responseExecution.PreviousTurns,
|
||||||
})
|
})
|
||||||
|
if leaseErr := stopLeaseRenewal(); leaseErr != nil {
|
||||||
|
err = &clients.ClientError{
|
||||||
|
Code: "concurrency_lease_lost",
|
||||||
|
Message: leaseErr.Error(),
|
||||||
|
Retryable: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
callFinishedAt := time.Now()
|
callFinishedAt := time.Now()
|
||||||
if err == nil {
|
if err == nil {
|
||||||
if markErr := setSubmissionStatus("response_received"); markErr != nil {
|
if markErr := setSubmissionStatus("response_received"); markErr != nil {
|
||||||
@@ -1500,6 +1521,7 @@ func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.Gateway
|
|||||||
if delay <= 0 {
|
if delay <= 0 {
|
||||||
delay = 5 * time.Second
|
delay = 5 * time.Second
|
||||||
}
|
}
|
||||||
|
delay += taskRetryJitter(task.ID)
|
||||||
queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, delay, candidate.QueueKey)
|
queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, delay, candidate.QueueKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return store.GatewayTask{}, 0, err
|
return store.GatewayTask{}, 0, err
|
||||||
@@ -1515,6 +1537,81 @@ func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.Gateway
|
|||||||
return queued, delay, nil
|
return queued, delay, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func taskRetryJitter(taskID string) time.Duration {
|
||||||
|
var sum uint32
|
||||||
|
for _, value := range []byte(taskID) {
|
||||||
|
sum = sum*33 + uint32(value)
|
||||||
|
}
|
||||||
|
return time.Duration(sum%251) * time.Millisecond
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) startConcurrencyLeaseRenewal(ctx context.Context, taskID string, leases []store.ConcurrencyLease) (context.Context, func() error) {
|
||||||
|
if len(leases) == 0 {
|
||||||
|
return ctx, func() error { return nil }
|
||||||
|
}
|
||||||
|
interval := 30 * time.Second
|
||||||
|
for _, lease := range leases {
|
||||||
|
ttl := lease.TTL
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = 120 * time.Second
|
||||||
|
}
|
||||||
|
candidate := ttl / 3
|
||||||
|
if candidate < time.Second {
|
||||||
|
candidate = time.Second
|
||||||
|
}
|
||||||
|
if candidate < interval {
|
||||||
|
interval = candidate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
runCtx, cancelRun := context.WithCancel(ctx)
|
||||||
|
renewCtx, cancelRenew := context.WithCancel(ctx)
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
ticker := time.NewTicker(interval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-renewCtx.Done():
|
||||||
|
done <- nil
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
if err := s.store.RenewConcurrencyLeases(renewCtx, leases); err != nil {
|
||||||
|
if renewCtx.Err() != nil {
|
||||||
|
done <- nil
|
||||||
|
return
|
||||||
|
}
|
||||||
|
outcome := "failure"
|
||||||
|
if errors.Is(err, store.ErrConcurrencyLeaseLost) {
|
||||||
|
outcome = "lost"
|
||||||
|
}
|
||||||
|
s.observeConcurrencyLeaseRenewal(outcome)
|
||||||
|
s.logger.Error("concurrency lease renewal failed; cancelling upstream execution",
|
||||||
|
"taskID", taskID, "leaseCount", len(leases), "outcome", outcome, "error", err)
|
||||||
|
done <- err
|
||||||
|
cancelRun()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.observeConcurrencyLeaseRenewal("success")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
return runCtx, func() error {
|
||||||
|
cancelRenew()
|
||||||
|
renewalErr := <-done
|
||||||
|
cancelRun()
|
||||||
|
return renewalErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Service) observeConcurrencyLeaseRenewal(outcome string) {
|
||||||
|
observer, ok := s.billingMetrics.(interface {
|
||||||
|
ObserveConcurrencyLeaseRenewal(string)
|
||||||
|
})
|
||||||
|
if ok {
|
||||||
|
observer.ObserveConcurrencyLeaseRenewal(outcome)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Service) requeueInterruptedAsyncTask(ctx context.Context, task store.GatewayTask) (store.GatewayTask, error) {
|
func (s *Service) requeueInterruptedAsyncTask(ctx context.Context, task store.GatewayTask) (store.GatewayTask, error) {
|
||||||
queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, 0, "")
|
queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, 0, "")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -58,10 +58,11 @@ func (s *Service) CancelTask(ctx context.Context, taskID string, user *auth.User
|
|||||||
return taskCancelUnavailable(task, "任务已开始执行,当前阶段不可取消,请继续查询结果"), nil
|
return taskCancelUnavailable(task, "任务已开始执行,当前阶段不可取消,请继续查询结果"), nil
|
||||||
}
|
}
|
||||||
if task.RiverJobID > 0 {
|
if task.RiverJobID > 0 {
|
||||||
if s.riverClient == nil {
|
riverClient := s.asyncControlClient()
|
||||||
|
if riverClient == nil {
|
||||||
return taskCancelUnavailable(task, "任务取消队列未就绪,请继续查询结果"), nil
|
return taskCancelUnavailable(task, "任务取消队列未就绪,请继续查询结果"), nil
|
||||||
}
|
}
|
||||||
job, err := s.riverClient.JobGet(ctx, task.RiverJobID)
|
job, err := riverClient.JobGet(ctx, task.RiverJobID)
|
||||||
if errors.Is(err, rivertype.ErrNotFound) {
|
if errors.Is(err, rivertype.ErrNotFound) {
|
||||||
return taskCancelUnavailable(task, "任务已不在本地排队队列,可能已提交上游,当前不可取消,请继续查询结果"), nil
|
return taskCancelUnavailable(task, "任务已不在本地排队队列,可能已提交上游,当前不可取消,请继续查询结果"), nil
|
||||||
}
|
}
|
||||||
@@ -71,7 +72,7 @@ func (s *Service) CancelTask(ctx context.Context, taskID string, user *auth.User
|
|||||||
if job == nil || !riverJobStateCancellable(job.State) {
|
if job == nil || !riverJobStateCancellable(job.State) {
|
||||||
return taskCancelUnavailable(task, "任务已不在可取消队列状态,请继续查询结果"), nil
|
return taskCancelUnavailable(task, "任务已不在可取消队列状态,请继续查询结果"), nil
|
||||||
}
|
}
|
||||||
if _, err := s.riverClient.JobDelete(ctx, task.RiverJobID); err != nil {
|
if _, err := riverClient.JobDelete(ctx, task.RiverJobID); err != nil {
|
||||||
if errors.Is(err, rivertype.ErrJobRunning) || errors.Is(err, rivertype.ErrNotFound) {
|
if errors.Is(err, rivertype.ErrJobRunning) || errors.Is(err, rivertype.ErrNotFound) {
|
||||||
return taskCancelUnavailable(task, "任务已被工作进程领取,当前不可取消,请继续查询结果"), nil
|
return taskCancelUnavailable(task, "任务已被工作进程领取,当前不可取消,请继续查询结果"), nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -44,6 +44,17 @@ type Metrics struct {
|
|||||||
billingEstimateFailed atomic.Uint64
|
billingEstimateFailed atomic.Uint64
|
||||||
billingIdempotentReplay atomic.Uint64
|
billingIdempotentReplay atomic.Uint64
|
||||||
billingPricingUnavailable atomic.Uint64
|
billingPricingUnavailable atomic.Uint64
|
||||||
|
asyncWorkerCapacity atomic.Int64
|
||||||
|
asyncWorkerDesiredCapacity atomic.Int64
|
||||||
|
asyncWorkerHardLimit atomic.Int64
|
||||||
|
asyncWorkerCapacityCapped atomic.Int64
|
||||||
|
asyncWorkerResizeSuccess atomic.Uint64
|
||||||
|
asyncWorkerRefreshFailed atomic.Uint64
|
||||||
|
asyncWorkerCreateFailed atomic.Uint64
|
||||||
|
asyncWorkerStartFailed atomic.Uint64
|
||||||
|
leaseRenewalSuccess atomic.Uint64
|
||||||
|
leaseRenewalFailure atomic.Uint64
|
||||||
|
leaseRenewalLost atomic.Uint64
|
||||||
}
|
}
|
||||||
|
|
||||||
var processingDurationBounds = [...]time.Duration{
|
var processingDurationBounds = [...]time.Duration{
|
||||||
@@ -123,6 +134,41 @@ func (m *Metrics) ObserveBillingEvent(event string) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *Metrics) SetAsyncWorkerCapacity(current, desired, hardLimit int, capped bool) {
|
||||||
|
m.asyncWorkerCapacity.Store(int64(current))
|
||||||
|
m.asyncWorkerDesiredCapacity.Store(int64(desired))
|
||||||
|
m.asyncWorkerHardLimit.Store(int64(hardLimit))
|
||||||
|
cappedValue := int64(0)
|
||||||
|
if capped {
|
||||||
|
cappedValue = 1
|
||||||
|
}
|
||||||
|
m.asyncWorkerCapacityCapped.Store(cappedValue)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Metrics) ObserveAsyncWorkerResize(outcome string) {
|
||||||
|
switch outcome {
|
||||||
|
case "success":
|
||||||
|
m.asyncWorkerResizeSuccess.Add(1)
|
||||||
|
case "refresh_failed":
|
||||||
|
m.asyncWorkerRefreshFailed.Add(1)
|
||||||
|
case "create_failed":
|
||||||
|
m.asyncWorkerCreateFailed.Add(1)
|
||||||
|
case "start_failed":
|
||||||
|
m.asyncWorkerStartFailed.Add(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Metrics) ObserveConcurrencyLeaseRenewal(outcome string) {
|
||||||
|
switch outcome {
|
||||||
|
case "success":
|
||||||
|
m.leaseRenewalSuccess.Add(1)
|
||||||
|
case "lost":
|
||||||
|
m.leaseRenewalLost.Add(1)
|
||||||
|
default:
|
||||||
|
m.leaseRenewalFailure.Add(1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience string, enabled bool) http.Handler {
|
func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience string, enabled bool) http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
mode := "disabled"
|
mode := "disabled"
|
||||||
@@ -212,6 +258,21 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
|
|||||||
plainCounter(w, "easyai_gateway_billing_estimate_failures_total", "Pricing estimate requests that failed.", m.billingEstimateFailed.Load())
|
plainCounter(w, "easyai_gateway_billing_estimate_failures_total", "Pricing estimate requests that failed.", m.billingEstimateFailed.Load())
|
||||||
plainCounter(w, "easyai_gateway_billing_idempotent_replays_total", "Generation requests replayed idempotently.", m.billingIdempotentReplay.Load())
|
plainCounter(w, "easyai_gateway_billing_idempotent_replays_total", "Generation requests replayed idempotently.", m.billingIdempotentReplay.Load())
|
||||||
plainCounter(w, "easyai_gateway_billing_pricing_unavailable_total", "Pricing requests rejected because no effective price was available.", m.billingPricingUnavailable.Load())
|
plainCounter(w, "easyai_gateway_billing_pricing_unavailable_total", "Pricing requests rejected because no effective price was available.", m.billingPricingUnavailable.Load())
|
||||||
|
plainGauge(w, "easyai_gateway_async_worker_capacity", "Current River asynchronous worker capacity.", m.asyncWorkerCapacity.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_capacity_capped", "Whether desired asynchronous worker capacity is capped by the hard limit.", m.asyncWorkerCapacityCapped.Load())
|
||||||
|
outcomeCounters(w, "easyai_gateway_async_worker_resizes_total", "Asynchronous worker resize attempts by bounded outcome.", []outcomeValue{
|
||||||
|
{"success", m.asyncWorkerResizeSuccess.Load()},
|
||||||
|
{"refresh_failed", m.asyncWorkerRefreshFailed.Load()},
|
||||||
|
{"create_failed", m.asyncWorkerCreateFailed.Load()},
|
||||||
|
{"start_failed", m.asyncWorkerStartFailed.Load()},
|
||||||
|
})
|
||||||
|
outcomeCounters(w, "easyai_gateway_concurrency_lease_renewals_total", "Concurrency lease renewals by bounded outcome.", []outcomeValue{
|
||||||
|
{"success", m.leaseRenewalSuccess.Load()},
|
||||||
|
{"failure", m.leaseRenewalFailure.Load()},
|
||||||
|
{"lost", m.leaseRenewalLost.Load()},
|
||||||
|
})
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -249,3 +310,7 @@ func outcomeCounters(w http.ResponseWriter, name, help string, values []outcomeV
|
|||||||
func plainCounter(w http.ResponseWriter, name, help string, value uint64) {
|
func plainCounter(w http.ResponseWriter, name, help string, value uint64) {
|
||||||
fmt.Fprintf(w, "# HELP %s %s\n# TYPE %s counter\n%s %d\n", name, help, name, name, value)
|
fmt.Fprintf(w, "# HELP %s %s\n# TYPE %s counter\n%s %d\n", name, help, name, name, value)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -26,6 +26,10 @@ func TestMetricsExposeBoundedOutcomesAndState(t *testing.T) {
|
|||||||
metrics.ObserveHeartbeat("accepted")
|
metrics.ObserveHeartbeat("accepted")
|
||||||
metrics.ObserveIntrospection("failed")
|
metrics.ObserveIntrospection("failed")
|
||||||
metrics.ObserveJWKSRefreshFailure("ssf")
|
metrics.ObserveJWKSRefreshFailure("ssf")
|
||||||
|
metrics.SetAsyncWorkerCapacity(96, 128, 96, true)
|
||||||
|
metrics.ObserveAsyncWorkerResize("success")
|
||||||
|
metrics.ObserveConcurrencyLeaseRenewal("success")
|
||||||
|
metrics.ObserveConcurrencyLeaseRenewal("lost")
|
||||||
|
|
||||||
recorder := httptest.NewRecorder()
|
recorder := httptest.NewRecorder()
|
||||||
metrics.Handler(metricsSnapshot{mode: "push_healthy", last: time.Now().Add(-time.Minute)}, "issuer", "audience", true).
|
metrics.Handler(metricsSnapshot{mode: "push_healthy", last: time.Now().Add(-time.Minute)}, "issuer", "audience", true).
|
||||||
@@ -42,6 +46,12 @@ func TestMetricsExposeBoundedOutcomesAndState(t *testing.T) {
|
|||||||
`easyai_gateway_ssf_mode{mode="push_healthy"} 1`,
|
`easyai_gateway_ssf_mode{mode="push_healthy"} 1`,
|
||||||
`easyai_gateway_oidc_introspection_total{outcome="failed"} 1`,
|
`easyai_gateway_oidc_introspection_total{outcome="failed"} 1`,
|
||||||
`easyai_gateway_jwks_refresh_failures_total{outcome="ssf"} 1`,
|
`easyai_gateway_jwks_refresh_failures_total{outcome="ssf"} 1`,
|
||||||
|
`easyai_gateway_async_worker_capacity 96`,
|
||||||
|
`easyai_gateway_async_worker_desired_capacity 128`,
|
||||||
|
`easyai_gateway_async_worker_capacity_capped 1`,
|
||||||
|
`easyai_gateway_async_worker_resizes_total{outcome="success"} 1`,
|
||||||
|
`easyai_gateway_concurrency_lease_renewals_total{outcome="success"} 1`,
|
||||||
|
`easyai_gateway_concurrency_lease_renewals_total{outcome="lost"} 1`,
|
||||||
} {
|
} {
|
||||||
if !strings.Contains(recorder.Body.String(), expected) {
|
if !strings.Contains(recorder.Body.String(), expected) {
|
||||||
t.Fatalf("missing metric %q in:\n%s", expected, recorder.Body.String())
|
t.Fatalf("missing metric %q in:\n%s", expected, recorder.Body.String())
|
||||||
|
|||||||
@@ -0,0 +1,152 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
type AsyncWorkerCapacitySnapshot struct {
|
||||||
|
Capacity int
|
||||||
|
Desired int
|
||||||
|
HardLimit int
|
||||||
|
Capped bool
|
||||||
|
EnabledModels int
|
||||||
|
UnlimitedModels int
|
||||||
|
EnabledGroups int
|
||||||
|
UnlimitedGroups int
|
||||||
|
ModelDesired int
|
||||||
|
GroupDesired int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) AsyncWorkerCapacity(ctx context.Context, hardLimit int) (AsyncWorkerCapacitySnapshot, error) {
|
||||||
|
if hardLimit < 1 {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, fmt.Errorf("async worker hard limit must be positive")
|
||||||
|
}
|
||||||
|
rows, err := s.pool.Query(ctx, `
|
||||||
|
SELECT COALESCE(b.default_rate_limit_policy, '{}'::jsonb),
|
||||||
|
p.rate_limit_policy,
|
||||||
|
COALESCE(rp.rate_limit_policy, '{}'::jsonb),
|
||||||
|
(m.runtime_policy_set_id IS NOT NULL),
|
||||||
|
COALESCE(m.runtime_policy_override, '{}'::jsonb),
|
||||||
|
m.rate_limit_policy,
|
||||||
|
m.rate_limit_policy_mode
|
||||||
|
FROM platform_models m
|
||||||
|
JOIN integration_platforms p ON p.id = m.platform_id
|
||||||
|
LEFT JOIN base_model_catalog b ON b.id = m.base_model_id
|
||||||
|
LEFT JOIN model_runtime_policy_sets rp ON rp.id = COALESCE(m.runtime_policy_set_id, b.runtime_policy_set_id)
|
||||||
|
WHERE p.status = 'enabled'
|
||||||
|
AND p.deleted_at IS NULL
|
||||||
|
AND m.enabled = true`)
|
||||||
|
if err != nil {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
|
||||||
|
modelPolicies := make([]map[string]any, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var basePolicyBytes, platformPolicyBytes, runtimePolicyBytes []byte
|
||||||
|
var runtimeOverrideBytes, modelPolicyBytes []byte
|
||||||
|
var runtimeExplicit bool
|
||||||
|
var modelPolicyMode string
|
||||||
|
if err := rows.Scan(
|
||||||
|
&basePolicyBytes,
|
||||||
|
&platformPolicyBytes,
|
||||||
|
&runtimePolicyBytes,
|
||||||
|
&runtimeExplicit,
|
||||||
|
&runtimeOverrideBytes,
|
||||||
|
&modelPolicyBytes,
|
||||||
|
&modelPolicyMode,
|
||||||
|
); err != nil {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, err
|
||||||
|
}
|
||||||
|
modelPolicies = append(modelPolicies, EffectiveRateLimitPolicy(EffectiveRateLimitPolicyInput{
|
||||||
|
BasePolicy: decodeObject(basePolicyBytes),
|
||||||
|
PlatformPolicy: decodeObject(platformPolicyBytes),
|
||||||
|
RuntimePolicy: decodeObject(runtimePolicyBytes),
|
||||||
|
RuntimePolicyExplicit: runtimeExplicit,
|
||||||
|
RuntimePolicyOverride: decodeObject(runtimeOverrideBytes),
|
||||||
|
ModelPolicy: decodeObject(modelPolicyBytes),
|
||||||
|
ModelPolicyMode: modelPolicyMode,
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, err
|
||||||
|
}
|
||||||
|
groupRows, err := s.pool.Query(ctx, `
|
||||||
|
SELECT rate_limit_policy
|
||||||
|
FROM gateway_user_groups
|
||||||
|
WHERE status = 'active'`)
|
||||||
|
if err != nil {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, err
|
||||||
|
}
|
||||||
|
defer groupRows.Close()
|
||||||
|
groupPolicies := make([]map[string]any, 0)
|
||||||
|
for groupRows.Next() {
|
||||||
|
var policyBytes []byte
|
||||||
|
if err := groupRows.Scan(&policyBytes); err != nil {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, err
|
||||||
|
}
|
||||||
|
groupPolicies = append(groupPolicies, NormalizeRateLimitPolicy(decodeObject(policyBytes)))
|
||||||
|
}
|
||||||
|
if err := groupRows.Err(); err != nil {
|
||||||
|
return AsyncWorkerCapacitySnapshot{}, err
|
||||||
|
}
|
||||||
|
return asyncWorkerCapacityFromPolicySets(modelPolicies, groupPolicies, hardLimit), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func asyncWorkerCapacityFromPolicies(policies []map[string]any, hardLimit int) AsyncWorkerCapacitySnapshot {
|
||||||
|
return asyncWorkerCapacityFromPolicySets(policies, nil, hardLimit)
|
||||||
|
}
|
||||||
|
|
||||||
|
func asyncWorkerCapacityFromPolicySets(modelPolicies []map[string]any, groupPolicies []map[string]any, hardLimit int) AsyncWorkerCapacitySnapshot {
|
||||||
|
snapshot := AsyncWorkerCapacitySnapshot{
|
||||||
|
HardLimit: hardLimit,
|
||||||
|
EnabledModels: len(modelPolicies),
|
||||||
|
EnabledGroups: len(groupPolicies),
|
||||||
|
}
|
||||||
|
modelDesired, modelFinite, unlimitedModels := concurrentPolicySetCapacity(modelPolicies)
|
||||||
|
groupDesired, groupFinite, unlimitedGroups := concurrentPolicySetCapacity(groupPolicies)
|
||||||
|
snapshot.ModelDesired = modelDesired
|
||||||
|
snapshot.GroupDesired = groupDesired
|
||||||
|
snapshot.UnlimitedModels = unlimitedModels
|
||||||
|
snapshot.UnlimitedGroups = unlimitedGroups
|
||||||
|
|
||||||
|
desired := hardLimit
|
||||||
|
switch {
|
||||||
|
case snapshot.EnabledModels == 0:
|
||||||
|
desired = 1
|
||||||
|
case modelFinite && groupFinite:
|
||||||
|
desired = min(modelDesired, groupDesired)
|
||||||
|
case modelFinite:
|
||||||
|
desired = modelDesired
|
||||||
|
case groupFinite:
|
||||||
|
desired = groupDesired
|
||||||
|
}
|
||||||
|
if desired < 1 {
|
||||||
|
desired = 1
|
||||||
|
}
|
||||||
|
snapshot.Desired = desired
|
||||||
|
snapshot.Capacity = desired
|
||||||
|
if snapshot.Capacity > hardLimit {
|
||||||
|
snapshot.Capacity = hardLimit
|
||||||
|
snapshot.Capped = true
|
||||||
|
}
|
||||||
|
return snapshot
|
||||||
|
}
|
||||||
|
|
||||||
|
func concurrentPolicySetCapacity(policies []map[string]any) (total int, finite bool, unlimited int) {
|
||||||
|
if len(policies) == 0 {
|
||||||
|
return 0, false, 0
|
||||||
|
}
|
||||||
|
finite = true
|
||||||
|
for _, policy := range policies {
|
||||||
|
capacity, policyFinite := ConcurrentPolicyCapacity(policy)
|
||||||
|
if !policyFinite {
|
||||||
|
unlimited++
|
||||||
|
finite = false
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
total += capacity
|
||||||
|
}
|
||||||
|
return total, finite, unlimited
|
||||||
|
}
|
||||||
@@ -25,7 +25,7 @@ func (s *Store) ListModelCandidates(ctx context.Context, model string, modelType
|
|||||||
COALESCE(p.base_url, ''),
|
COALESCE(p.base_url, ''),
|
||||||
p.auth_type, p.credentials, p.config, p.default_pricing_mode,
|
p.auth_type, p.credentials, p.config, p.default_pricing_mode,
|
||||||
p.default_discount_factor::float8, COALESCE(p.pricing_rule_set_id::text, ''),
|
p.default_discount_factor::float8, COALESCE(p.pricing_rule_set_id::text, ''),
|
||||||
p.retry_policy, p.rate_limit_policy,
|
p.retry_policy, p.rate_limit_policy, COALESCE(b.default_rate_limit_policy, '{}'::jsonb),
|
||||||
COALESCE(p.dynamic_priority, p.priority) AS effective_priority,
|
COALESCE(p.dynamic_priority, p.priority) AS effective_priority,
|
||||||
m.id::text, COALESCE(m.base_model_id::text, ''), COALESCE(b.canonical_model_key, ''),
|
m.id::text, COALESCE(m.base_model_id::text, ''), COALESCE(b.canonical_model_key, ''),
|
||||||
COALESCE(NULLIF(m.provider_model_name, ''), m.model_name), m.model_name, COALESCE(m.model_alias, ''),
|
COALESCE(NULLIF(m.provider_model_name, ''), m.model_name), m.model_name, COALESCE(m.model_alias, ''),
|
||||||
@@ -54,8 +54,11 @@ func (s *Store) ListModelCandidates(ctx context.Context, model string, modelType
|
|||||||
COALESCE(b.base_billing_config, '{}'::jsonb), m.billing_config, m.billing_config_override,
|
COALESCE(b.base_billing_config, '{}'::jsonb), m.billing_config, m.billing_config_override,
|
||||||
m.pricing_mode, COALESCE(m.discount_factor, 0)::float8, COALESCE(m.pricing_rule_set_id::text, ''),
|
m.pricing_mode, COALESCE(m.discount_factor, 0)::float8, COALESCE(m.pricing_rule_set_id::text, ''),
|
||||||
COALESCE(b.pricing_rule_set_id::text, ''),
|
COALESCE(b.pricing_rule_set_id::text, ''),
|
||||||
m.permission_config, m.retry_policy, m.rate_limit_policy, COALESCE(m.runtime_policy_set_id::text, COALESCE(b.runtime_policy_set_id::text, '')),
|
m.permission_config, m.retry_policy, m.rate_limit_policy, m.rate_limit_policy_mode,
|
||||||
|
COALESCE(m.runtime_policy_set_id::text, COALESCE(b.runtime_policy_set_id::text, '')),
|
||||||
|
(m.runtime_policy_set_id IS NOT NULL),
|
||||||
COALESCE(NULLIF(m.runtime_policy_override, '{}'::jsonb), b.runtime_policy_override, '{}'::jsonb),
|
COALESCE(NULLIF(m.runtime_policy_override, '{}'::jsonb), b.runtime_policy_override, '{}'::jsonb),
|
||||||
|
COALESCE(m.runtime_policy_override, '{}'::jsonb),
|
||||||
COALESCE(rp.retry_policy, '{}'::jsonb), COALESCE(rp.rate_limit_policy, '{}'::jsonb),
|
COALESCE(rp.retry_policy, '{}'::jsonb), COALESCE(rp.rate_limit_policy, '{}'::jsonb),
|
||||||
COALESCE(rp.auto_disable_policy, '{}'::jsonb), COALESCE(rp.degrade_policy, '{}'::jsonb),
|
COALESCE(rp.auto_disable_policy, '{}'::jsonb), COALESCE(rp.degrade_policy, '{}'::jsonb),
|
||||||
COALESCE(con.active, 0)::float8,
|
COALESCE(con.active, 0)::float8,
|
||||||
@@ -177,6 +180,7 @@ WHERE p.status = 'enabled'
|
|||||||
var platformConfig []byte
|
var platformConfig []byte
|
||||||
var platformRetryPolicy []byte
|
var platformRetryPolicy []byte
|
||||||
var platformRateLimitPolicy []byte
|
var platformRateLimitPolicy []byte
|
||||||
|
var baseRateLimitPolicy []byte
|
||||||
var capabilities []byte
|
var capabilities []byte
|
||||||
var capabilityOverride []byte
|
var capabilityOverride []byte
|
||||||
var baseBilling []byte
|
var baseBilling []byte
|
||||||
@@ -186,6 +190,7 @@ WHERE p.status = 'enabled'
|
|||||||
var modelRetryPolicy []byte
|
var modelRetryPolicy []byte
|
||||||
var modelRateLimitPolicy []byte
|
var modelRateLimitPolicy []byte
|
||||||
var runtimePolicyOverride []byte
|
var runtimePolicyOverride []byte
|
||||||
|
var rateLimitRuntimeOverride []byte
|
||||||
var runtimeRetryPolicy []byte
|
var runtimeRetryPolicy []byte
|
||||||
var runtimeRateLimitPolicy []byte
|
var runtimeRateLimitPolicy []byte
|
||||||
var autoDisablePolicy []byte
|
var autoDisablePolicy []byte
|
||||||
@@ -222,6 +227,7 @@ WHERE p.status = 'enabled'
|
|||||||
&item.PlatformPricingRuleSetID,
|
&item.PlatformPricingRuleSetID,
|
||||||
&platformRetryPolicy,
|
&platformRetryPolicy,
|
||||||
&platformRateLimitPolicy,
|
&platformRateLimitPolicy,
|
||||||
|
&baseRateLimitPolicy,
|
||||||
&item.PlatformPriority,
|
&item.PlatformPriority,
|
||||||
&item.PlatformModelID,
|
&item.PlatformModelID,
|
||||||
&item.BaseModelID,
|
&item.BaseModelID,
|
||||||
@@ -244,8 +250,11 @@ WHERE p.status = 'enabled'
|
|||||||
&permissionConfig,
|
&permissionConfig,
|
||||||
&modelRetryPolicy,
|
&modelRetryPolicy,
|
||||||
&modelRateLimitPolicy,
|
&modelRateLimitPolicy,
|
||||||
|
&item.ModelRateLimitPolicyMode,
|
||||||
&item.RuntimePolicySetID,
|
&item.RuntimePolicySetID,
|
||||||
|
&item.RuntimePolicyExplicit,
|
||||||
&runtimePolicyOverride,
|
&runtimePolicyOverride,
|
||||||
|
&rateLimitRuntimeOverride,
|
||||||
&runtimeRetryPolicy,
|
&runtimeRetryPolicy,
|
||||||
&runtimeRateLimitPolicy,
|
&runtimeRateLimitPolicy,
|
||||||
&autoDisablePolicy,
|
&autoDisablePolicy,
|
||||||
@@ -274,6 +283,7 @@ WHERE p.status = 'enabled'
|
|||||||
item.PlatformConfig = decodeObject(platformConfig)
|
item.PlatformConfig = decodeObject(platformConfig)
|
||||||
item.PlatformRetryPolicy = decodeObject(platformRetryPolicy)
|
item.PlatformRetryPolicy = decodeObject(platformRetryPolicy)
|
||||||
item.PlatformRateLimitPolicy = decodeObject(platformRateLimitPolicy)
|
item.PlatformRateLimitPolicy = decodeObject(platformRateLimitPolicy)
|
||||||
|
item.BaseRateLimitPolicy = decodeObject(baseRateLimitPolicy)
|
||||||
item.Capabilities = decodeObject(capabilities)
|
item.Capabilities = decodeObject(capabilities)
|
||||||
item.CapabilityOverride = decodeObject(capabilityOverride)
|
item.CapabilityOverride = decodeObject(capabilityOverride)
|
||||||
item.BaseBillingConfig = decodeObject(baseBilling)
|
item.BaseBillingConfig = decodeObject(baseBilling)
|
||||||
@@ -283,6 +293,7 @@ WHERE p.status = 'enabled'
|
|||||||
item.ModelRetryPolicy = decodeObject(modelRetryPolicy)
|
item.ModelRetryPolicy = decodeObject(modelRetryPolicy)
|
||||||
item.ModelRateLimitPolicy = decodeObject(modelRateLimitPolicy)
|
item.ModelRateLimitPolicy = decodeObject(modelRateLimitPolicy)
|
||||||
item.RuntimePolicyOverride = decodeObject(runtimePolicyOverride)
|
item.RuntimePolicyOverride = decodeObject(runtimePolicyOverride)
|
||||||
|
item.RateLimitRuntimeOverride = decodeObject(rateLimitRuntimeOverride)
|
||||||
item.RuntimeRetryPolicy = decodeObject(runtimeRetryPolicy)
|
item.RuntimeRetryPolicy = decodeObject(runtimeRetryPolicy)
|
||||||
item.RuntimeRateLimitPolicy = decodeObject(runtimeRateLimitPolicy)
|
item.RuntimeRateLimitPolicy = decodeObject(runtimeRateLimitPolicy)
|
||||||
item.AutoDisablePolicy = decodeObject(autoDisablePolicy)
|
item.AutoDisablePolicy = decodeObject(autoDisablePolicy)
|
||||||
@@ -303,7 +314,15 @@ WHERE p.status = 'enabled'
|
|||||||
LastObservedUnix: cacheLastObservedUnix,
|
LastObservedUnix: cacheLastObservedUnix,
|
||||||
})
|
})
|
||||||
applyRuntimeCandidateLoad(&item, runtimeCandidateLoadInput{
|
applyRuntimeCandidateLoad(&item, runtimeCandidateLoadInput{
|
||||||
Policy: effectiveModelRateLimitPolicy(item.PlatformRateLimitPolicy, item.RuntimeRateLimitPolicy, item.RuntimePolicySetID, item.RuntimePolicyOverride, item.ModelRateLimitPolicy),
|
Policy: EffectiveRateLimitPolicy(EffectiveRateLimitPolicyInput{
|
||||||
|
BasePolicy: item.BaseRateLimitPolicy,
|
||||||
|
PlatformPolicy: item.PlatformRateLimitPolicy,
|
||||||
|
RuntimePolicy: item.RuntimeRateLimitPolicy,
|
||||||
|
RuntimePolicyExplicit: item.RuntimePolicyExplicit,
|
||||||
|
RuntimePolicyOverride: item.RateLimitRuntimeOverride,
|
||||||
|
ModelPolicy: item.ModelRateLimitPolicy,
|
||||||
|
ModelPolicyMode: item.ModelRateLimitPolicyMode,
|
||||||
|
}),
|
||||||
ConcurrentActive: concurrentActive,
|
ConcurrentActive: concurrentActive,
|
||||||
QueuedWaiting: queuedWaiting,
|
QueuedWaiting: queuedWaiting,
|
||||||
RPMUsed: rpmUsed,
|
RPMUsed: rpmUsed,
|
||||||
|
|||||||
@@ -143,18 +143,13 @@ SELECT EXISTS (
|
|||||||
// soon as the base pricing rule changes and can mask the authoritative rule.
|
// soon as the base pricing rule changes and can mask the authoritative rule.
|
||||||
billingConfig := input.BillingConfig
|
billingConfig := input.BillingConfig
|
||||||
explicitRuntimePolicySetID := strings.TrimSpace(input.RuntimePolicySetID)
|
explicitRuntimePolicySetID := strings.TrimSpace(input.RuntimePolicySetID)
|
||||||
|
rateLimitPolicyMode := NormalizeRateLimitPolicyMode(input.RateLimitPolicyMode, input.RateLimitPolicy != nil)
|
||||||
rateLimitPolicy := input.RateLimitPolicy
|
rateLimitPolicy := input.RateLimitPolicy
|
||||||
if len(rateLimitPolicy) == 0 && explicitRuntimePolicySetID == "" {
|
if rateLimitPolicyMode == RateLimitPolicyModeInherit {
|
||||||
rateLimitPolicy = base.DefaultRateLimitPolicy
|
rateLimitPolicy = nil
|
||||||
}
|
}
|
||||||
runtimePolicySetID := explicitRuntimePolicySetID
|
runtimePolicySetID := explicitRuntimePolicySetID
|
||||||
if runtimePolicySetID == "" {
|
|
||||||
runtimePolicySetID = base.RuntimePolicySetID
|
|
||||||
}
|
|
||||||
runtimePolicyOverride := input.RuntimePolicyOverride
|
runtimePolicyOverride := input.RuntimePolicyOverride
|
||||||
if len(runtimePolicyOverride) == 0 {
|
|
||||||
runtimePolicyOverride = base.RuntimePolicyOverride
|
|
||||||
}
|
|
||||||
|
|
||||||
capabilityOverrideJSON, _ := json.Marshal(emptyObjectIfNil(input.CapabilityOverride))
|
capabilityOverrideJSON, _ := json.Marshal(emptyObjectIfNil(input.CapabilityOverride))
|
||||||
capabilitiesJSON, _ := json.Marshal(emptyObjectIfNil(capabilities))
|
capabilitiesJSON, _ := json.Marshal(emptyObjectIfNil(capabilities))
|
||||||
@@ -189,14 +184,14 @@ SELECT EXISTS (
|
|||||||
INSERT INTO platform_models (
|
INSERT INTO platform_models (
|
||||||
platform_id, base_model_id, model_name, provider_model_name, model_alias, model_type, display_name,
|
platform_id, base_model_id, model_name, provider_model_name, model_alias, model_type, display_name,
|
||||||
capability_override, capabilities, pricing_mode, discount_factor,
|
capability_override, capabilities, pricing_mode, discount_factor,
|
||||||
pricing_rule_set_id, billing_config_override, billing_config, permission_config, retry_policy, rate_limit_policy,
|
pricing_rule_set_id, billing_config_override, billing_config, permission_config, retry_policy, rate_limit_policy, rate_limit_policy_mode,
|
||||||
runtime_policy_set_id, runtime_policy_override, enabled
|
runtime_policy_set_id, runtime_policy_override, enabled
|
||||||
)
|
)
|
||||||
VALUES (
|
VALUES (
|
||||||
$1::uuid, $2::uuid, $3, NULLIF($4, ''), NULLIF($5, ''), $6::jsonb, $7,
|
$1::uuid, $2::uuid, $3, NULLIF($4, ''), NULLIF($5, ''), $6::jsonb, $7,
|
||||||
$8::jsonb, $9::jsonb, $10, $11::numeric,
|
$8::jsonb, $9::jsonb, $10, $11::numeric,
|
||||||
NULLIF($12, '')::uuid, $13::jsonb, $14::jsonb, $15::jsonb, $16::jsonb, $17::jsonb,
|
NULLIF($12, '')::uuid, $13::jsonb, $14::jsonb, $15::jsonb, $16::jsonb, $17::jsonb, $18,
|
||||||
NULLIF($18, '')::uuid, $19::jsonb, true
|
NULLIF($19, '')::uuid, $20::jsonb, true
|
||||||
)
|
)
|
||||||
ON CONFLICT (platform_id, model_name) DO UPDATE
|
ON CONFLICT (platform_id, model_name) DO UPDATE
|
||||||
SET base_model_id = EXCLUDED.base_model_id,
|
SET base_model_id = EXCLUDED.base_model_id,
|
||||||
@@ -213,6 +208,7 @@ SET base_model_id = EXCLUDED.base_model_id,
|
|||||||
permission_config = EXCLUDED.permission_config,
|
permission_config = EXCLUDED.permission_config,
|
||||||
retry_policy = EXCLUDED.retry_policy,
|
retry_policy = EXCLUDED.retry_policy,
|
||||||
rate_limit_policy = EXCLUDED.rate_limit_policy,
|
rate_limit_policy = EXCLUDED.rate_limit_policy,
|
||||||
|
rate_limit_policy_mode = EXCLUDED.rate_limit_policy_mode,
|
||||||
runtime_policy_set_id = EXCLUDED.runtime_policy_set_id,
|
runtime_policy_set_id = EXCLUDED.runtime_policy_set_id,
|
||||||
runtime_policy_override = EXCLUDED.runtime_policy_override,
|
runtime_policy_override = EXCLUDED.runtime_policy_override,
|
||||||
enabled = true,
|
enabled = true,
|
||||||
@@ -221,7 +217,7 @@ RETURNING id::text, platform_id::text, COALESCE(base_model_id::text, ''), model_
|
|||||||
COALESCE(NULLIF(provider_model_name, ''), model_name), COALESCE(model_alias, ''), model_type, display_name, capability_override,
|
COALESCE(NULLIF(provider_model_name, ''), model_name), COALESCE(model_alias, ''), model_type, display_name, capability_override,
|
||||||
capabilities, pricing_mode, COALESCE(discount_factor, 0)::float8,
|
capabilities, pricing_mode, COALESCE(discount_factor, 0)::float8,
|
||||||
COALESCE(pricing_rule_set_id::text, ''), billing_config_override, billing_config,
|
COALESCE(pricing_rule_set_id::text, ''), billing_config_override, billing_config,
|
||||||
permission_config, retry_policy, rate_limit_policy, COALESCE(runtime_policy_set_id::text, ''), runtime_policy_override,
|
permission_config, retry_policy, rate_limit_policy, rate_limit_policy_mode, COALESCE(runtime_policy_set_id::text, ''), runtime_policy_override,
|
||||||
COALESCE(to_char(cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''),
|
COALESCE(to_char(cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''),
|
||||||
enabled, created_at, updated_at`,
|
enabled, created_at, updated_at`,
|
||||||
input.PlatformID,
|
input.PlatformID,
|
||||||
@@ -241,6 +237,7 @@ RETURNING id::text, platform_id::text, COALESCE(base_model_id::text, ''), model_
|
|||||||
string(permissionJSON),
|
string(permissionJSON),
|
||||||
string(retryJSON),
|
string(retryJSON),
|
||||||
string(rateLimitJSON),
|
string(rateLimitJSON),
|
||||||
|
rateLimitPolicyMode,
|
||||||
runtimePolicySetID,
|
runtimePolicySetID,
|
||||||
string(runtimePolicyOverrideJSON),
|
string(runtimePolicyOverrideJSON),
|
||||||
).Scan(
|
).Scan(
|
||||||
@@ -262,6 +259,7 @@ RETURNING id::text, platform_id::text, COALESCE(base_model_id::text, ''), model_
|
|||||||
&permissionBytes,
|
&permissionBytes,
|
||||||
&retryPolicyBytes,
|
&retryPolicyBytes,
|
||||||
&rateLimitPolicyBytes,
|
&rateLimitPolicyBytes,
|
||||||
|
&model.RateLimitPolicyMode,
|
||||||
&model.RuntimePolicySetID,
|
&model.RuntimePolicySetID,
|
||||||
&runtimePolicyOverrideBytes,
|
&runtimePolicyOverrideBytes,
|
||||||
&model.CooldownUntil,
|
&model.CooldownUntil,
|
||||||
|
|||||||
@@ -235,6 +235,8 @@ type PlatformModel struct {
|
|||||||
Capabilities map[string]any `json:"capabilities,omitempty"`
|
Capabilities map[string]any `json:"capabilities,omitempty"`
|
||||||
BaseCapabilities map[string]any `json:"-"`
|
BaseCapabilities map[string]any `json:"-"`
|
||||||
BaseBillingConfig map[string]any `json:"-"`
|
BaseBillingConfig map[string]any `json:"-"`
|
||||||
|
BaseRateLimitPolicy map[string]any `json:"-"`
|
||||||
|
BaseRuntimePolicySetID string `json:"-"`
|
||||||
BasePricingRuleSetID string `json:"-"`
|
BasePricingRuleSetID string `json:"-"`
|
||||||
PlatformPricingRuleSetID string `json:"-"`
|
PlatformPricingRuleSetID string `json:"-"`
|
||||||
PricingMode string `json:"pricingMode"`
|
PricingMode string `json:"pricingMode"`
|
||||||
@@ -245,6 +247,7 @@ type PlatformModel struct {
|
|||||||
PermissionConfig map[string]any `json:"permissionConfig,omitempty"`
|
PermissionConfig map[string]any `json:"permissionConfig,omitempty"`
|
||||||
RetryPolicy map[string]any `json:"retryPolicy,omitempty"`
|
RetryPolicy map[string]any `json:"retryPolicy,omitempty"`
|
||||||
RateLimitPolicy map[string]any `json:"rateLimitPolicy,omitempty"`
|
RateLimitPolicy map[string]any `json:"rateLimitPolicy,omitempty"`
|
||||||
|
RateLimitPolicyMode string `json:"rateLimitPolicyMode" enums:"inherit,override"`
|
||||||
RuntimePolicySetID string `json:"runtimePolicySetId,omitempty"`
|
RuntimePolicySetID string `json:"runtimePolicySetId,omitempty"`
|
||||||
RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride,omitempty"`
|
RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride,omitempty"`
|
||||||
CooldownUntil string `json:"cooldownUntil,omitempty"`
|
CooldownUntil string `json:"cooldownUntil,omitempty"`
|
||||||
@@ -956,7 +959,9 @@ SELECT m.id::text, m.platform_id::text, COALESCE(m.base_model_id::text, ''), p.p
|
|||||||
COALESCE(b.base_billing_config, '{}'::jsonb), COALESCE(b.pricing_rule_set_id::text, ''),
|
COALESCE(b.base_billing_config, '{}'::jsonb), COALESCE(b.pricing_rule_set_id::text, ''),
|
||||||
COALESCE(p.pricing_rule_set_id::text, ''), m.pricing_mode, COALESCE(m.discount_factor, 0)::float8,
|
COALESCE(p.pricing_rule_set_id::text, ''), m.pricing_mode, COALESCE(m.discount_factor, 0)::float8,
|
||||||
COALESCE(m.pricing_rule_set_id::text, ''), m.billing_config_override, m.billing_config,
|
COALESCE(m.pricing_rule_set_id::text, ''), m.billing_config_override, m.billing_config,
|
||||||
m.permission_config, m.retry_policy, m.rate_limit_policy, COALESCE(m.runtime_policy_set_id::text, ''), m.runtime_policy_override,
|
m.permission_config, m.retry_policy, m.rate_limit_policy, m.rate_limit_policy_mode,
|
||||||
|
COALESCE(m.runtime_policy_set_id::text, ''), m.runtime_policy_override,
|
||||||
|
COALESCE(b.default_rate_limit_policy, '{}'::jsonb), COALESCE(b.runtime_policy_set_id::text, ''),
|
||||||
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"'), ''),
|
||||||
m.enabled, m.created_at, m.updated_at
|
m.enabled, m.created_at, m.updated_at
|
||||||
FROM platform_models m
|
FROM platform_models m
|
||||||
@@ -964,6 +969,7 @@ JOIN integration_platforms p ON p.id = m.platform_id
|
|||||||
LEFT JOIN LATERAL (
|
LEFT JOIN LATERAL (
|
||||||
SELECT catalog.invocation_name, catalog.display_name, catalog.model_type,
|
SELECT catalog.invocation_name, catalog.display_name, catalog.model_type,
|
||||||
catalog.capabilities, catalog.base_billing_config, catalog.pricing_rule_set_id,
|
catalog.capabilities, catalog.base_billing_config, catalog.pricing_rule_set_id,
|
||||||
|
catalog.default_rate_limit_policy, catalog.runtime_policy_set_id,
|
||||||
COALESCE((
|
COALESCE((
|
||||||
SELECT jsonb_agg(DISTINCT compatibility_alias.alias ORDER BY compatibility_alias.alias)
|
SELECT jsonb_agg(DISTINCT compatibility_alias.alias ORDER BY compatibility_alias.alias)
|
||||||
FROM model_compatibility_aliases compatibility_alias
|
FROM model_compatibility_aliases compatibility_alias
|
||||||
@@ -1004,6 +1010,7 @@ ORDER BY m.model_type ASC, m.model_name ASC`, args...)
|
|||||||
var retryPolicy []byte
|
var retryPolicy []byte
|
||||||
var rateLimitPolicy []byte
|
var rateLimitPolicy []byte
|
||||||
var runtimePolicyOverride []byte
|
var runtimePolicyOverride []byte
|
||||||
|
var baseRateLimitPolicy []byte
|
||||||
var modelTypeBytes []byte
|
var modelTypeBytes []byte
|
||||||
var legacyAliasesBytes []byte
|
var legacyAliasesBytes []byte
|
||||||
if err := rows.Scan(
|
if err := rows.Scan(
|
||||||
@@ -1032,8 +1039,11 @@ ORDER BY m.model_type ASC, m.model_name ASC`, args...)
|
|||||||
&permissionConfig,
|
&permissionConfig,
|
||||||
&retryPolicy,
|
&retryPolicy,
|
||||||
&rateLimitPolicy,
|
&rateLimitPolicy,
|
||||||
|
&model.RateLimitPolicyMode,
|
||||||
&model.RuntimePolicySetID,
|
&model.RuntimePolicySetID,
|
||||||
&runtimePolicyOverride,
|
&runtimePolicyOverride,
|
||||||
|
&baseRateLimitPolicy,
|
||||||
|
&model.BaseRuntimePolicySetID,
|
||||||
&model.CooldownUntil,
|
&model.CooldownUntil,
|
||||||
&model.Enabled,
|
&model.Enabled,
|
||||||
&model.CreatedAt,
|
&model.CreatedAt,
|
||||||
@@ -1053,6 +1063,7 @@ ORDER BY m.model_type ASC, m.model_name ASC`, args...)
|
|||||||
model.RetryPolicy = decodeObject(retryPolicy)
|
model.RetryPolicy = decodeObject(retryPolicy)
|
||||||
model.RateLimitPolicy = decodeObject(rateLimitPolicy)
|
model.RateLimitPolicy = decodeObject(rateLimitPolicy)
|
||||||
model.RuntimePolicyOverride = decodeObject(runtimePolicyOverride)
|
model.RuntimePolicyOverride = decodeObject(runtimePolicyOverride)
|
||||||
|
model.BaseRateLimitPolicy = decodeObject(baseRateLimitPolicy)
|
||||||
models = append(models, model)
|
models = append(models, model)
|
||||||
}
|
}
|
||||||
return models, rows.Err()
|
return models, rows.Err()
|
||||||
|
|||||||
@@ -0,0 +1,167 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
RateLimitPolicyModeInherit = "inherit"
|
||||||
|
RateLimitPolicyModeOverride = "override"
|
||||||
|
)
|
||||||
|
|
||||||
|
type EffectiveRateLimitPolicyInput struct {
|
||||||
|
BasePolicy map[string]any
|
||||||
|
PlatformPolicy map[string]any
|
||||||
|
RuntimePolicy map[string]any
|
||||||
|
RuntimePolicyExplicit bool
|
||||||
|
RuntimePolicyOverride map[string]any
|
||||||
|
ModelPolicy map[string]any
|
||||||
|
ModelPolicyMode string
|
||||||
|
}
|
||||||
|
|
||||||
|
// EffectiveRateLimitPolicy resolves one authoritative policy. Platform
|
||||||
|
// policies are defaults for every bound model; explicit model/runtime settings
|
||||||
|
// replace the complete policy instead of merging individual metrics.
|
||||||
|
func EffectiveRateLimitPolicy(input EffectiveRateLimitPolicyInput) map[string]any {
|
||||||
|
policy := input.BasePolicy
|
||||||
|
if policySpecified(input.PlatformPolicy) {
|
||||||
|
policy = input.PlatformPolicy
|
||||||
|
}
|
||||||
|
if input.RuntimePolicyExplicit {
|
||||||
|
policy = input.RuntimePolicy
|
||||||
|
}
|
||||||
|
if raw, ok := input.RuntimePolicyOverride["rateLimitPolicy"]; ok {
|
||||||
|
policy, _ = raw.(map[string]any)
|
||||||
|
}
|
||||||
|
mode := NormalizeRateLimitPolicyMode(input.ModelPolicyMode, input.ModelPolicy != nil)
|
||||||
|
if mode == RateLimitPolicyModeOverride {
|
||||||
|
policy = input.ModelPolicy
|
||||||
|
}
|
||||||
|
return NormalizeRateLimitPolicy(policy)
|
||||||
|
}
|
||||||
|
|
||||||
|
func NormalizeRateLimitPolicyMode(mode string, policyProvided bool) string {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(mode)) {
|
||||||
|
case RateLimitPolicyModeOverride:
|
||||||
|
return RateLimitPolicyModeOverride
|
||||||
|
case RateLimitPolicyModeInherit:
|
||||||
|
return RateLimitPolicyModeInherit
|
||||||
|
default:
|
||||||
|
if policyProvided {
|
||||||
|
return RateLimitPolicyModeOverride
|
||||||
|
}
|
||||||
|
return RateLimitPolicyModeInherit
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func RateLimitPolicyMetric(policy map[string]any, metric string) (float64, bool) {
|
||||||
|
rules, _ := NormalizeRateLimitPolicy(policy)["rules"].([]any)
|
||||||
|
for _, rawRule := range rules {
|
||||||
|
rule, _ := rawRule.(map[string]any)
|
||||||
|
if strings.TrimSpace(stringValue(rule["metric"])) != metric {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
value := floatValue(rule["limit"])
|
||||||
|
return value, true
|
||||||
|
}
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeRateLimitPolicy keeps the canonical rules contract while accepting
|
||||||
|
// policies imported from server-main before that contract existed. Runtime
|
||||||
|
// enforcement and worker sizing must agree on these legacy limits; otherwise a
|
||||||
|
// configured max_concurrent_requests silently becomes unlimited.
|
||||||
|
func NormalizeRateLimitPolicy(policy map[string]any) map[string]any {
|
||||||
|
if policy == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
out := clonePolicy(policy)
|
||||||
|
rules, _ := out["rules"].([]any)
|
||||||
|
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 = 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 len(rules) > 0 {
|
||||||
|
out["rules"] = rules
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func ConcurrentPolicyCapacity(policy map[string]any) (int, bool) {
|
||||||
|
limit, ok := RateLimitPolicyMetric(policy, "concurrent")
|
||||||
|
if !ok || limit <= 0 {
|
||||||
|
return 0, false
|
||||||
|
}
|
||||||
|
if limit < 1 {
|
||||||
|
return 1, true
|
||||||
|
}
|
||||||
|
if limit >= float64(math.MaxInt) {
|
||||||
|
return math.MaxInt, true
|
||||||
|
}
|
||||||
|
return int(math.Floor(limit)), true
|
||||||
|
}
|
||||||
|
|
||||||
|
func policySpecified(policy map[string]any) bool {
|
||||||
|
rules, _ := NormalizeRateLimitPolicy(policy)["rules"].([]any)
|
||||||
|
return len(rules) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func lowestPositiveLegacyLimit(scopes []map[string]any, keys ...string) (float64, bool) {
|
||||||
|
limit := 0.0
|
||||||
|
found := false
|
||||||
|
for _, scope := range scopes {
|
||||||
|
for _, key := range keys {
|
||||||
|
value := floatValue(scope[key])
|
||||||
|
if value <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if !found || value < limit {
|
||||||
|
limit = value
|
||||||
|
found = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return limit, found
|
||||||
|
}
|
||||||
|
|
||||||
|
func clonePolicy(policy map[string]any) map[string]any {
|
||||||
|
if policy == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
out := make(map[string]any, len(policy))
|
||||||
|
for key, value := range policy {
|
||||||
|
out[key] = value
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
@@ -0,0 +1,120 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRateLimitPolicyModeMigrationClassifiesExistingRows(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 rate-limit migration PostgreSQL integration tests")
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
schema := fmt.Sprintf("rate_limit_mode_%d", time.Now().UnixNano())
|
||||||
|
adminPool, err := pgxpool.New(ctx, databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect migration test database: %v", err)
|
||||||
|
}
|
||||||
|
defer adminPool.Close()
|
||||||
|
if _, err := adminPool.Exec(ctx, `CREATE SCHEMA `+schema); err != nil {
|
||||||
|
t.Fatalf("create migration test schema: %v", err)
|
||||||
|
}
|
||||||
|
defer adminPool.Exec(context.Background(), `DROP SCHEMA IF EXISTS `+schema+` CASCADE`)
|
||||||
|
|
||||||
|
schemaURL, err := url.Parse(databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parse test database url: %v", err)
|
||||||
|
}
|
||||||
|
query := schemaURL.Query()
|
||||||
|
query.Set("search_path", schema)
|
||||||
|
schemaURL.RawQuery = query.Encode()
|
||||||
|
pool, err := pgxpool.New(ctx, schemaURL.String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect migration test schema: %v", err)
|
||||||
|
}
|
||||||
|
defer pool.Close()
|
||||||
|
if _, err := pool.Exec(ctx, `
|
||||||
|
CREATE TABLE base_model_catalog (
|
||||||
|
id uuid PRIMARY KEY,
|
||||||
|
default_rate_limit_policy jsonb NOT NULL DEFAULT '{}'::jsonb,
|
||||||
|
runtime_policy_set_id uuid,
|
||||||
|
runtime_policy_override jsonb NOT NULL DEFAULT '{}'::jsonb
|
||||||
|
);
|
||||||
|
CREATE TABLE integration_platforms (
|
||||||
|
id uuid PRIMARY KEY,
|
||||||
|
rate_limit_policy jsonb NOT NULL DEFAULT '{}'::jsonb,
|
||||||
|
updated_at timestamptz NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
CREATE TABLE platform_models (
|
||||||
|
id uuid PRIMARY KEY,
|
||||||
|
base_model_id uuid,
|
||||||
|
rate_limit_policy jsonb NOT NULL DEFAULT '{}'::jsonb,
|
||||||
|
runtime_policy_set_id uuid,
|
||||||
|
runtime_policy_override jsonb NOT NULL DEFAULT '{}'::jsonb,
|
||||||
|
updated_at timestamptz NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
INSERT INTO integration_platforms (id, rate_limit_policy)
|
||||||
|
VALUES ('20000000-0000-0000-0000-000000000001', '{"rules":[]}');
|
||||||
|
INSERT INTO base_model_catalog (id, default_rate_limit_policy)
|
||||||
|
VALUES ('00000000-0000-0000-0000-000000000001', '{"rules":[{"metric":"concurrent","limit":4}]}');
|
||||||
|
INSERT INTO platform_models (id, base_model_id, rate_limit_policy, runtime_policy_override) VALUES
|
||||||
|
('10000000-0000-0000-0000-000000000001', '00000000-0000-0000-0000-000000000001', '{"rules":[{"metric":"concurrent","limit":4}]}', '{}'),
|
||||||
|
('10000000-0000-0000-0000-000000000002', '00000000-0000-0000-0000-000000000001', '{"rules":[{"metric":"concurrent","limit":8}]}', '{}'),
|
||||||
|
('10000000-0000-0000-0000-000000000003', '00000000-0000-0000-0000-000000000001', '{"rules":[{"metric":"concurrent","limit":4}]}', '{"rateLimitPolicy":{"rules":[]}}'),
|
||||||
|
('10000000-0000-0000-0000-000000000004', NULL, '{}', '{}');`); err != nil {
|
||||||
|
t.Fatalf("seed pre-migration rows: %v", err)
|
||||||
|
}
|
||||||
|
_, currentFile, _, _ := runtime.Caller(0)
|
||||||
|
migrationPath := filepath.Join(filepath.Dir(currentFile), "..", "..", "migrations", "0080_platform_model_rate_limit_policy_mode.sql")
|
||||||
|
migration, err := os.ReadFile(migrationPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read migration: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := pool.Exec(ctx, string(migration)); err != nil {
|
||||||
|
t.Fatalf("apply migration: %v", err)
|
||||||
|
}
|
||||||
|
rows, err := pool.Query(ctx, `SELECT id::text, rate_limit_policy_mode FROM platform_models ORDER BY id`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read classified rows: %v", err)
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
want := []string{"inherit", "override", "override", "inherit"}
|
||||||
|
index := 0
|
||||||
|
for rows.Next() {
|
||||||
|
var id, mode string
|
||||||
|
if err := rows.Scan(&id, &mode); err != nil {
|
||||||
|
t.Fatalf("scan classified row: %v", err)
|
||||||
|
}
|
||||||
|
if index >= len(want) || mode != want[index] {
|
||||||
|
t.Fatalf("row %s mode=%s, want=%s", id, mode, want[index])
|
||||||
|
}
|
||||||
|
index++
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
t.Fatalf("read classified rows: %v", err)
|
||||||
|
}
|
||||||
|
if index != len(want) {
|
||||||
|
t.Fatalf("classified %d rows, want %d", index, len(want))
|
||||||
|
}
|
||||||
|
var normalizedPlatformPolicy string
|
||||||
|
if err := pool.QueryRow(ctx, `
|
||||||
|
SELECT rate_limit_policy::text
|
||||||
|
FROM integration_platforms
|
||||||
|
WHERE id = '20000000-0000-0000-0000-000000000001'`).Scan(&normalizedPlatformPolicy); err != nil {
|
||||||
|
t.Fatalf("read normalized platform policy: %v", err)
|
||||||
|
}
|
||||||
|
if normalizedPlatformPolicy != "{}" {
|
||||||
|
t.Fatalf("normalized platform policy=%s, want={}", normalizedPlatformPolicy)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestEffectiveRateLimitPolicyPrecedence(t *testing.T) {
|
||||||
|
policy := func(limit float64) map[string]any {
|
||||||
|
return map[string]any{"rules": []any{map[string]any{"metric": "concurrent", "limit": limit}}}
|
||||||
|
}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input EffectiveRateLimitPolicyInput
|
||||||
|
want float64
|
||||||
|
ok bool
|
||||||
|
}{
|
||||||
|
{name: "base", input: EffectiveRateLimitPolicyInput{BasePolicy: policy(2)}, want: 2, ok: true},
|
||||||
|
{name: "platform", input: EffectiveRateLimitPolicyInput{BasePolicy: policy(2), PlatformPolicy: policy(4)}, want: 4, ok: true},
|
||||||
|
{name: "empty platform inherits base", input: EffectiveRateLimitPolicyInput{BasePolicy: policy(2), PlatformPolicy: map[string]any{"rules": []any{}}}, want: 2, ok: true},
|
||||||
|
{name: "runtime", input: EffectiveRateLimitPolicyInput{PlatformPolicy: policy(4), RuntimePolicy: policy(8), RuntimePolicyExplicit: true}, want: 8, ok: true},
|
||||||
|
{name: "runtime override", input: EffectiveRateLimitPolicyInput{PlatformPolicy: policy(4), RuntimePolicyOverride: map[string]any{"rateLimitPolicy": policy(16)}}, want: 16, ok: true},
|
||||||
|
{name: "model override", input: EffectiveRateLimitPolicyInput{PlatformPolicy: policy(4), ModelPolicy: policy(32), ModelPolicyMode: "override"}, want: 32, ok: true},
|
||||||
|
{name: "model explicit unlimited", input: EffectiveRateLimitPolicyInput{PlatformPolicy: policy(4), ModelPolicy: map[string]any{}, ModelPolicyMode: "override"}, ok: false},
|
||||||
|
{name: "model inherit", input: EffectiveRateLimitPolicyInput{PlatformPolicy: policy(4), ModelPolicy: policy(32), ModelPolicyMode: "inherit"}, want: 4, ok: true},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, ok := ConcurrentPolicyCapacity(EffectiveRateLimitPolicy(tt.input))
|
||||||
|
if ok != tt.ok || (ok && got != int(tt.want)) {
|
||||||
|
t.Fatalf("capacity = (%d, %v), want (%d, %v)", got, ok, int(tt.want), tt.ok)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeRateLimitPolicyLegacyShapes(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
policy map[string]any
|
||||||
|
metric string
|
||||||
|
want float64
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "platform concurrent",
|
||||||
|
policy: map[string]any{"platformLimits": map[string]any{"max_concurrent_requests": 5.0}},
|
||||||
|
metric: "concurrent",
|
||||||
|
want: 5,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "model concurrent camel case",
|
||||||
|
policy: map[string]any{"modelLimits": map[string]any{"maxConcurrentRequests": 10.0}},
|
||||||
|
metric: "concurrent",
|
||||||
|
want: 10,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "stricter duplicate wins",
|
||||||
|
policy: map[string]any{
|
||||||
|
"platformLimits": map[string]any{"max_concurrent_requests": 8.0},
|
||||||
|
"modelLimits": map[string]any{"max_concurrent_requests": 3.0},
|
||||||
|
},
|
||||||
|
metric: "concurrent",
|
||||||
|
want: 3,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "requests per minute",
|
||||||
|
policy: map[string]any{"model_limits": map[string]any{"max_request_per_minute": 60.0}},
|
||||||
|
metric: "rpm",
|
||||||
|
want: 60,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, ok := RateLimitPolicyMetric(tt.policy, tt.metric)
|
||||||
|
if !ok || got != tt.want {
|
||||||
|
t.Fatalf("metric %s = (%v, %v), want (%v, true)", tt.metric, got, ok, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConcurrentPolicyCapacityRoundsDown(t *testing.T) {
|
||||||
|
policy := map[string]any{"rules": []any{map[string]any{"metric": "concurrent", "limit": 96.9}}}
|
||||||
|
got, ok := ConcurrentPolicyCapacity(policy)
|
||||||
|
if !ok || got != 96 {
|
||||||
|
t.Fatalf("capacity = (%d, %v), want (96, true)", got, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAsyncWorkerCapacityAggregation(t *testing.T) {
|
||||||
|
policy := func(limit float64) map[string]any {
|
||||||
|
return map[string]any{"rules": []any{map[string]any{"metric": "concurrent", "limit": limit}}}
|
||||||
|
}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
policies []map[string]any
|
||||||
|
hardLimit int
|
||||||
|
wantCapacity int
|
||||||
|
wantDesired int
|
||||||
|
wantCapped bool
|
||||||
|
}{
|
||||||
|
{name: "no enabled models", hardLimit: 2048, wantCapacity: 1, wantDesired: 1},
|
||||||
|
{name: "finite sum", policies: []map[string]any{policy(64), policy(32)}, hardLimit: 2048, wantCapacity: 96, wantDesired: 96},
|
||||||
|
{name: "unlimited model", policies: []map[string]any{policy(64), {}}, hardLimit: 2048, wantCapacity: 2048, wantDesired: 2048},
|
||||||
|
{name: "hard limit cap", policies: []map[string]any{policy(80), policy(80)}, hardLimit: 96, wantCapacity: 96, wantDesired: 160, wantCapped: true},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := asyncWorkerCapacityFromPolicies(tt.policies, tt.hardLimit)
|
||||||
|
if got.Capacity != tt.wantCapacity || got.Desired != tt.wantDesired || got.Capped != tt.wantCapped {
|
||||||
|
t.Fatalf("snapshot=%+v, want capacity=%d desired=%d capped=%v", got, tt.wantCapacity, tt.wantDesired, tt.wantCapped)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAsyncWorkerCapacityRespectsUserGroupCeiling(t *testing.T) {
|
||||||
|
policy := func(limit float64) map[string]any {
|
||||||
|
return map[string]any{"rules": []any{map[string]any{"metric": "concurrent", "limit": limit}}}
|
||||||
|
}
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
models []map[string]any
|
||||||
|
groups []map[string]any
|
||||||
|
hardLimit int
|
||||||
|
wantCapacity int
|
||||||
|
wantDesired int
|
||||||
|
wantCapped bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "group ceiling prevents worker oversubscription",
|
||||||
|
models: []map[string]any{{}},
|
||||||
|
groups: []map[string]any{policy(3), policy(10)},
|
||||||
|
hardLimit: 2048,
|
||||||
|
wantCapacity: 13,
|
||||||
|
wantDesired: 13,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "model ceiling is stricter",
|
||||||
|
models: []map[string]any{policy(5), policy(7)},
|
||||||
|
groups: []map[string]any{policy(300)},
|
||||||
|
hardLimit: 2048,
|
||||||
|
wantCapacity: 12,
|
||||||
|
wantDesired: 12,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "both policy sets unlimited use hard limit",
|
||||||
|
models: []map[string]any{{}},
|
||||||
|
groups: []map[string]any{{}},
|
||||||
|
hardLimit: 256,
|
||||||
|
wantCapacity: 256,
|
||||||
|
wantDesired: 256,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "finite group desired still reports hard cap",
|
||||||
|
models: []map[string]any{{}},
|
||||||
|
groups: []map[string]any{policy(500)},
|
||||||
|
hardLimit: 256,
|
||||||
|
wantCapacity: 256,
|
||||||
|
wantDesired: 500,
|
||||||
|
wantCapped: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := asyncWorkerCapacityFromPolicySets(tt.models, tt.groups, tt.hardLimit)
|
||||||
|
if got.Capacity != tt.wantCapacity || got.Desired != tt.wantDesired || got.Capped != tt.wantCapped {
|
||||||
|
t.Fatalf("snapshot=%+v, want capacity=%d desired=%d capped=%v", got, tt.wantCapacity, tt.wantDesired, tt.wantCapped)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestConcurrencyRetryAfterBoundsPolling(t *testing.T) {
|
||||||
|
if got := concurrencyRetryAfter(time.Time{}); got != 5*time.Second {
|
||||||
|
t.Fatalf("zero expiry retry=%s, want=5s", got)
|
||||||
|
}
|
||||||
|
if got := concurrencyRetryAfter(time.Now().Add(30 * time.Second)); got != 5*time.Second {
|
||||||
|
t.Fatalf("distant expiry retry=%s, want=5s", got)
|
||||||
|
}
|
||||||
|
got := concurrencyRetryAfter(time.Now().Add(2500 * time.Millisecond))
|
||||||
|
if got < 2*time.Second || got > 2500*time.Millisecond {
|
||||||
|
t.Fatalf("near expiry retry=%s, want within [2s,2.5s]", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -143,8 +143,10 @@ func (s *Store) ListModelRateLimitStatuses(ctx context.Context) ([]ModelRateLimi
|
|||||||
p.priority, p.dynamic_priority, COALESCE(p.dynamic_priority, p.priority),
|
p.priority, p.dynamic_priority, COALESCE(p.dynamic_priority, p.priority),
|
||||||
m.model_name, COALESCE(NULLIF(m.provider_model_name, ''), m.model_name), COALESCE(m.model_alias, ''),
|
m.model_name, COALESCE(NULLIF(m.provider_model_name, ''), m.model_name), COALESCE(m.model_alias, ''),
|
||||||
m.model_type, m.display_name, m.enabled,
|
m.model_type, m.display_name, m.enabled,
|
||||||
p.rate_limit_policy, COALESCE(rp.rate_limit_policy, '{}'::jsonb), COALESCE(m.runtime_policy_set_id::text, b.runtime_policy_set_id::text, ''),
|
COALESCE(b.default_rate_limit_policy, '{}'::jsonb), p.rate_limit_policy,
|
||||||
COALESCE(NULLIF(m.runtime_policy_override, '{}'::jsonb), b.runtime_policy_override, '{}'::jsonb), m.rate_limit_policy,
|
COALESCE(rp.rate_limit_policy, '{}'::jsonb),
|
||||||
|
(m.runtime_policy_set_id IS NOT NULL),
|
||||||
|
COALESCE(m.runtime_policy_override, '{}'::jsonb), m.rate_limit_policy, m.rate_limit_policy_mode,
|
||||||
COALESCE(to_char(p.cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''),
|
COALESCE(to_char(p.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(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,
|
||||||
@@ -217,11 +219,13 @@ ORDER BY p.priority ASC, m.model_name ASC`)
|
|||||||
for rows.Next() {
|
for rows.Next() {
|
||||||
var item ModelRateLimitStatus
|
var item ModelRateLimitStatus
|
||||||
var modelTypeBytes []byte
|
var modelTypeBytes []byte
|
||||||
|
var basePolicyBytes []byte
|
||||||
var platformPolicyBytes []byte
|
var platformPolicyBytes []byte
|
||||||
var runtimePolicyBytes []byte
|
var runtimePolicyBytes []byte
|
||||||
var runtimePolicySetID string
|
var runtimePolicyExplicit bool
|
||||||
var runtimeOverrideBytes []byte
|
var runtimeOverrideBytes []byte
|
||||||
var modelPolicyBytes []byte
|
var modelPolicyBytes []byte
|
||||||
|
var modelPolicyMode string
|
||||||
var platformDynamicPriority sql.NullInt64
|
var platformDynamicPriority sql.NullInt64
|
||||||
var platformCooldownUntil string
|
var platformCooldownUntil string
|
||||||
var modelCooldownUntil string
|
var modelCooldownUntil string
|
||||||
@@ -248,11 +252,13 @@ ORDER BY p.priority ASC, m.model_name ASC`)
|
|||||||
&modelTypeBytes,
|
&modelTypeBytes,
|
||||||
&item.DisplayName,
|
&item.DisplayName,
|
||||||
&item.Enabled,
|
&item.Enabled,
|
||||||
|
&basePolicyBytes,
|
||||||
&platformPolicyBytes,
|
&platformPolicyBytes,
|
||||||
&runtimePolicyBytes,
|
&runtimePolicyBytes,
|
||||||
&runtimePolicySetID,
|
&runtimePolicyExplicit,
|
||||||
&runtimeOverrideBytes,
|
&runtimeOverrideBytes,
|
||||||
&modelPolicyBytes,
|
&modelPolicyBytes,
|
||||||
|
&modelPolicyMode,
|
||||||
&platformCooldownUntil,
|
&platformCooldownUntil,
|
||||||
&modelCooldownUntil,
|
&modelCooldownUntil,
|
||||||
&concurrentCurrent,
|
&concurrentCurrent,
|
||||||
@@ -268,13 +274,15 @@ ORDER BY p.priority ASC, m.model_name ASC`)
|
|||||||
}
|
}
|
||||||
item.PlatformDynamicPriority = intPointerFromNull(platformDynamicPriority)
|
item.PlatformDynamicPriority = intPointerFromNull(platformDynamicPriority)
|
||||||
item.ModelType = decodeStringArray(modelTypeBytes)
|
item.ModelType = decodeStringArray(modelTypeBytes)
|
||||||
policy := effectiveModelRateLimitPolicy(
|
policy := EffectiveRateLimitPolicy(EffectiveRateLimitPolicyInput{
|
||||||
decodeObject(platformPolicyBytes),
|
BasePolicy: decodeObject(basePolicyBytes),
|
||||||
decodeObject(runtimePolicyBytes),
|
PlatformPolicy: decodeObject(platformPolicyBytes),
|
||||||
runtimePolicySetID,
|
RuntimePolicy: decodeObject(runtimePolicyBytes),
|
||||||
decodeObject(runtimeOverrideBytes),
|
RuntimePolicyExplicit: runtimePolicyExplicit,
|
||||||
decodeObject(modelPolicyBytes),
|
RuntimePolicyOverride: decodeObject(runtimeOverrideBytes),
|
||||||
)
|
ModelPolicy: decodeObject(modelPolicyBytes),
|
||||||
|
ModelPolicyMode: modelPolicyMode,
|
||||||
|
})
|
||||||
item.PlatformCooldownUntil = platformCooldownUntil
|
item.PlatformCooldownUntil = platformCooldownUntil
|
||||||
item.ModelCooldownUntil = modelCooldownUntil
|
item.ModelCooldownUntil = modelCooldownUntil
|
||||||
item.RateLimitPolicy = policy
|
item.RateLimitPolicy = policy
|
||||||
@@ -491,46 +499,6 @@ func platformPolicyEventFromPayload(id string, taskID string, eventType string,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func effectiveModelRateLimitPolicy(platformPolicy map[string]any, runtimePolicy map[string]any, runtimePolicySetID string, runtimeOverride map[string]any, modelPolicy map[string]any) map[string]any {
|
|
||||||
policy := platformPolicy
|
|
||||||
if strings.TrimSpace(runtimePolicySetID) != "" {
|
|
||||||
policy = runtimePolicy
|
|
||||||
} else if hasRateLimitRules(runtimePolicy) {
|
|
||||||
policy = shallowMergeMap(policy, runtimePolicy)
|
|
||||||
}
|
|
||||||
if _, hasOverride := runtimeOverride["rateLimitPolicy"]; hasOverride {
|
|
||||||
nested, _ := runtimeOverride["rateLimitPolicy"].(map[string]any)
|
|
||||||
if len(nested) == 0 {
|
|
||||||
policy = nil
|
|
||||||
} else {
|
|
||||||
policy = shallowMergeMap(policy, nested)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if hasRateLimitRules(modelPolicy) {
|
|
||||||
policy = shallowMergeMap(policy, modelPolicy)
|
|
||||||
}
|
|
||||||
if hasRateLimitRules(policy) {
|
|
||||||
return policy
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func hasRateLimitRules(policy map[string]any) bool {
|
|
||||||
rules, _ := policy["rules"].([]any)
|
|
||||||
return len(rules) > 0
|
|
||||||
}
|
|
||||||
|
|
||||||
func shallowMergeMap(base map[string]any, override map[string]any) map[string]any {
|
|
||||||
out := map[string]any{}
|
|
||||||
for key, value := range base {
|
|
||||||
out[key] = value
|
|
||||||
}
|
|
||||||
for key, value := range override {
|
|
||||||
out[key] = value
|
|
||||||
}
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
func rateLimitForMetric(policy map[string]any, metric string) float64 {
|
func rateLimitForMetric(policy map[string]any, metric string) float64 {
|
||||||
rules, _ := policy["rules"].([]any)
|
rules, _ := policy["rules"].([]any)
|
||||||
for _, rawRule := range rules {
|
for _, rawRule := range rules {
|
||||||
|
|||||||
@@ -6,23 +6,23 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestEffectiveModelRateLimitPolicyTreatsModelRulesAsAuthoritative(t *testing.T) {
|
func TestEffectiveModelRateLimitPolicyTreatsModelRulesAsAuthoritative(t *testing.T) {
|
||||||
policy := effectiveModelRateLimitPolicy(
|
policy := EffectiveRateLimitPolicy(EffectiveRateLimitPolicyInput{
|
||||||
map[string]any{"rules": []any{
|
PlatformPolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"metric": "rpm", "limit": 500},
|
map[string]any{"metric": "rpm", "limit": 500},
|
||||||
map[string]any{"metric": "tpm_total", "limit": 100000},
|
map[string]any{"metric": "tpm_total", "limit": 100000},
|
||||||
}},
|
}},
|
||||||
map[string]any{"rules": []any{
|
RuntimePolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"metric": "rpm", "limit": 120},
|
map[string]any{"metric": "rpm", "limit": 120},
|
||||||
map[string]any{"metric": "tpm_total", "limit": 240000},
|
map[string]any{"metric": "tpm_total", "limit": 240000},
|
||||||
map[string]any{"metric": "concurrent", "limit": 6},
|
map[string]any{"metric": "concurrent", "limit": 6},
|
||||||
}},
|
}},
|
||||||
"runtime-policy-1",
|
RuntimePolicyExplicit: true,
|
||||||
map[string]any{},
|
ModelPolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"rules": []any{
|
|
||||||
map[string]any{"metric": "rpm", "limit": 30},
|
map[string]any{"metric": "rpm", "limit": 30},
|
||||||
map[string]any{"metric": "concurrent", "limit": 2},
|
map[string]any{"metric": "concurrent", "limit": 2},
|
||||||
}},
|
}},
|
||||||
)
|
ModelPolicyMode: "override",
|
||||||
|
})
|
||||||
|
|
||||||
if got := rateLimitForMetric(policy, "rpm"); got != 30 {
|
if got := rateLimitForMetric(policy, "rpm"); got != 30 {
|
||||||
t.Fatalf("expected model rpm limit to win, got %v", got)
|
t.Fatalf("expected model rpm limit to win, got %v", got)
|
||||||
@@ -36,16 +36,15 @@ func TestEffectiveModelRateLimitPolicyTreatsModelRulesAsAuthoritative(t *testing
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestEffectiveModelRateLimitPolicyTreatsEmptyRuntimePolicyAsUnlimited(t *testing.T) {
|
func TestEffectiveModelRateLimitPolicyTreatsEmptyRuntimePolicyAsUnlimited(t *testing.T) {
|
||||||
policy := effectiveModelRateLimitPolicy(
|
policy := EffectiveRateLimitPolicy(EffectiveRateLimitPolicyInput{
|
||||||
map[string]any{"rules": []any{
|
PlatformPolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"metric": "rpm", "limit": 500},
|
map[string]any{"metric": "rpm", "limit": 500},
|
||||||
map[string]any{"metric": "tpm_total", "limit": 100000},
|
map[string]any{"metric": "tpm_total", "limit": 100000},
|
||||||
}},
|
}},
|
||||||
map[string]any{"rules": []any{}},
|
RuntimePolicy: map[string]any{"rules": []any{}},
|
||||||
"runtime-policy-1",
|
RuntimePolicyExplicit: true,
|
||||||
map[string]any{},
|
ModelPolicyMode: "inherit",
|
||||||
map[string]any{},
|
})
|
||||||
)
|
|
||||||
|
|
||||||
if got := rateLimitForMetric(policy, "rpm"); got != 0 {
|
if got := rateLimitForMetric(policy, "rpm"); got != 0 {
|
||||||
t.Fatalf("expected empty runtime policy rpm to mean unlimited, got %v", got)
|
t.Fatalf("expected empty runtime policy rpm to mean unlimited, got %v", got)
|
||||||
@@ -56,17 +55,16 @@ func TestEffectiveModelRateLimitPolicyTreatsEmptyRuntimePolicyAsUnlimited(t *tes
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestEffectiveModelRateLimitPolicyTreatsNegativeLimitAsUnlimited(t *testing.T) {
|
func TestEffectiveModelRateLimitPolicyTreatsNegativeLimitAsUnlimited(t *testing.T) {
|
||||||
policy := effectiveModelRateLimitPolicy(
|
policy := EffectiveRateLimitPolicy(EffectiveRateLimitPolicyInput{
|
||||||
map[string]any{"rules": []any{
|
PlatformPolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"metric": "rpm", "limit": 500},
|
map[string]any{"metric": "rpm", "limit": 500},
|
||||||
}},
|
}},
|
||||||
map[string]any{"rules": []any{
|
RuntimePolicy: map[string]any{"rules": []any{
|
||||||
map[string]any{"metric": "rpm", "limit": -1},
|
map[string]any{"metric": "rpm", "limit": -1},
|
||||||
}},
|
}},
|
||||||
"runtime-policy-1",
|
RuntimePolicyExplicit: true,
|
||||||
map[string]any{},
|
ModelPolicyMode: "inherit",
|
||||||
map[string]any{},
|
})
|
||||||
)
|
|
||||||
|
|
||||||
if got := rateLimitForMetric(policy, "rpm"); got != -1 {
|
if got := rateLimitForMetric(policy, "rpm"); got != -1 {
|
||||||
t.Fatalf("expected negative runtime rpm marker to be preserved, got %v", got)
|
t.Fatalf("expected negative runtime rpm marker to be preserved, got %v", got)
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sort"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
@@ -17,6 +18,8 @@ type RuntimeRecoveryResult struct {
|
|||||||
RequeuedAsyncTasks int64 `json:"requeuedAsyncTasks"`
|
RequeuedAsyncTasks int64 `json:"requeuedAsyncTasks"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var ErrConcurrencyLeaseLost = errors.New("concurrency lease lost")
|
||||||
|
|
||||||
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 {
|
||||||
@@ -24,6 +27,26 @@ func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID
|
|||||||
}
|
}
|
||||||
defer tx.Rollback(ctx)
|
defer tx.Rollback(ctx)
|
||||||
|
|
||||||
|
lockKeys := make([]string, 0)
|
||||||
|
lockKeySet := make(map[string]struct{})
|
||||||
|
for _, reservation := range reservations {
|
||||||
|
if reservation.Metric != "concurrent" || reservation.Limit <= 0 || reservation.Amount <= 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
key := fmt.Sprintf("%d:%s%d:%s", len(reservation.ScopeType), reservation.ScopeType, len(reservation.ScopeKey), reservation.ScopeKey)
|
||||||
|
if _, exists := lockKeySet[key]; exists {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
lockKeySet[key] = struct{}{}
|
||||||
|
lockKeys = append(lockKeys, key)
|
||||||
|
}
|
||||||
|
sort.Strings(lockKeys)
|
||||||
|
for _, key := range lockKeys {
|
||||||
|
if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 0))`, key); err != nil {
|
||||||
|
return RateLimitResult{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
result := RateLimitResult{}
|
result := RateLimitResult{}
|
||||||
for _, reservation := range reservations {
|
for _, reservation := range reservations {
|
||||||
if reservation.Limit <= 0 || reservation.Amount <= 0 {
|
if reservation.Limit <= 0 || reservation.Amount <= 0 {
|
||||||
@@ -49,11 +72,11 @@ func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID
|
|||||||
reservation.WindowSeconds = 60
|
reservation.WindowSeconds = 60
|
||||||
}
|
}
|
||||||
if reservation.Metric == "concurrent" {
|
if reservation.Metric == "concurrent" {
|
||||||
leaseID, err := reserveConcurrencyLease(ctx, tx, taskID, attemptID, reservation)
|
lease, err := reserveConcurrencyLease(ctx, tx, taskID, attemptID, reservation)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return RateLimitResult{}, err
|
return RateLimitResult{}, err
|
||||||
}
|
}
|
||||||
result.LeaseIDs = append(result.LeaseIDs, leaseID)
|
result.Leases = append(result.Leases, lease)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
normalized, err := reserveCounterWindow(ctx, tx, taskID, attemptID, reservation)
|
normalized, err := reserveCounterWindow(ctx, tx, taskID, attemptID, reservation)
|
||||||
@@ -65,7 +88,7 @@ func (s *Store) ReserveRateLimits(ctx context.Context, taskID string, attemptID
|
|||||||
return result, tx.Commit(ctx)
|
return result, tx.Commit(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
func reserveConcurrencyLease(ctx context.Context, tx pgx.Tx, taskID string, attemptID string, reservation RateLimitReservation) (string, error) {
|
func reserveConcurrencyLease(ctx context.Context, tx pgx.Tx, taskID string, attemptID string, reservation RateLimitReservation) (ConcurrencyLease, error) {
|
||||||
if reservation.LeaseTTLSeconds <= 0 {
|
if reservation.LeaseTTLSeconds <= 0 {
|
||||||
reservation.LeaseTTLSeconds = 120
|
reservation.LeaseTTLSeconds = 120
|
||||||
}
|
}
|
||||||
@@ -83,10 +106,10 @@ WHERE scope_type = $1
|
|||||||
reservation.ScopeKey,
|
reservation.ScopeKey,
|
||||||
reservation.LeaseTTLSeconds,
|
reservation.LeaseTTLSeconds,
|
||||||
).Scan(&active, &nextAvailableAt); err != nil {
|
).Scan(&active, &nextAvailableAt); err != nil {
|
||||||
return "", err
|
return ConcurrencyLease{}, err
|
||||||
}
|
}
|
||||||
if active+reservation.Amount > reservation.Limit {
|
if active+reservation.Amount > reservation.Limit {
|
||||||
return "", &RateLimitExceededError{
|
return ConcurrencyLease{}, &RateLimitExceededError{
|
||||||
ScopeType: reservation.ScopeType,
|
ScopeType: reservation.ScopeType,
|
||||||
ScopeKey: reservation.ScopeKey,
|
ScopeKey: reservation.ScopeKey,
|
||||||
ScopeName: reservation.ScopeName,
|
ScopeName: reservation.ScopeName,
|
||||||
@@ -117,9 +140,9 @@ RETURNING id::text`,
|
|||||||
reservation.Amount,
|
reservation.Amount,
|
||||||
reservation.LeaseTTLSeconds,
|
reservation.LeaseTTLSeconds,
|
||||||
).Scan(&leaseID); err != nil {
|
).Scan(&leaseID); err != nil {
|
||||||
return "", err
|
return ConcurrencyLease{}, err
|
||||||
}
|
}
|
||||||
return leaseID, nil
|
return ConcurrencyLease{ID: leaseID, TTL: time.Duration(reservation.LeaseTTLSeconds) * time.Second}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func reserveCounterWindow(ctx context.Context, tx pgx.Tx, taskID string, attemptID string, reservation RateLimitReservation) (RateLimitReservation, error) {
|
func reserveCounterWindow(ctx context.Context, tx pgx.Tx, taskID string, attemptID string, reservation RateLimitReservation) (RateLimitReservation, error) {
|
||||||
@@ -232,13 +255,16 @@ func retryAfterUntil(when time.Time) time.Duration {
|
|||||||
|
|
||||||
func concurrencyRetryAfter(leaseExpiresAt time.Time) time.Duration {
|
func concurrencyRetryAfter(leaseExpiresAt time.Time) time.Duration {
|
||||||
if leaseExpiresAt.IsZero() {
|
if leaseExpiresAt.IsZero() {
|
||||||
return time.Second
|
return 5 * time.Second
|
||||||
}
|
}
|
||||||
duration := time.Until(leaseExpiresAt)
|
duration := time.Until(leaseExpiresAt)
|
||||||
if duration <= time.Second {
|
if duration <= time.Second {
|
||||||
return time.Second
|
return time.Second
|
||||||
}
|
}
|
||||||
return time.Second
|
if duration > 5*time.Second {
|
||||||
|
return 5 * time.Second
|
||||||
|
}
|
||||||
|
return duration
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) CommitRateLimitReservations(ctx context.Context, reservations []RateLimitReservation, actualByMetric map[string]float64) error {
|
func (s *Store) CommitRateLimitReservations(ctx context.Context, reservations []RateLimitReservation, actualByMetric map[string]float64) error {
|
||||||
@@ -249,26 +275,57 @@ func (s *Store) ReleaseRateLimitReservations(ctx context.Context, reservations [
|
|||||||
return s.finishRateLimitReservations(ctx, reservations, nil, "released", reason)
|
return s.finishRateLimitReservations(ctx, reservations, nil, "released", reason)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) ReleaseConcurrencyLeases(ctx context.Context, leaseIDs []string) error {
|
func (s *Store) ReleaseConcurrencyLeases(ctx context.Context, leases []ConcurrencyLease) error {
|
||||||
|
leaseIDs := concurrencyLeaseIDs(leases)
|
||||||
if len(leaseIDs) == 0 {
|
if len(leaseIDs) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
for _, leaseID := range leaseIDs {
|
_, err := s.pool.Exec(ctx, `
|
||||||
if leaseID == "" {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if _, err := s.pool.Exec(ctx, `
|
|
||||||
UPDATE gateway_concurrency_leases
|
UPDATE gateway_concurrency_leases
|
||||||
SET released_at = now()
|
SET released_at = now()
|
||||||
WHERE id = $1::uuid AND released_at IS NULL`, leaseID); err != nil && !errors.Is(err, ErrRateLimited) {
|
WHERE id = ANY($1::uuid[]) AND released_at IS NULL`, leaseIDs)
|
||||||
return err
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) RenewConcurrencyLeases(ctx context.Context, leases []ConcurrencyLease) error {
|
||||||
|
leaseIDs := make([]string, 0, len(leases))
|
||||||
|
ttlSeconds := make([]int32, 0, len(leases))
|
||||||
|
for _, lease := range leases {
|
||||||
|
if lease.ID == "" {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
ttl := lease.TTL
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = 120 * time.Second
|
||||||
|
}
|
||||||
|
seconds := int32(ttl / time.Second)
|
||||||
|
if seconds < 1 {
|
||||||
|
seconds = 1
|
||||||
|
}
|
||||||
|
leaseIDs = append(leaseIDs, lease.ID)
|
||||||
|
ttlSeconds = append(ttlSeconds, seconds)
|
||||||
|
}
|
||||||
|
if len(leaseIDs) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
tag, err := s.pool.Exec(ctx, `
|
||||||
|
UPDATE gateway_concurrency_leases lease
|
||||||
|
SET expires_at = now() + (renewal.ttl_seconds * interval '1 second')
|
||||||
|
FROM unnest($1::uuid[], $2::int[]) AS renewal(id, ttl_seconds)
|
||||||
|
WHERE lease.id = renewal.id
|
||||||
|
AND lease.released_at IS NULL
|
||||||
|
AND lease.expires_at > now()`, leaseIDs, ttlSeconds)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if tag.RowsAffected() != int64(len(leaseIDs)) {
|
||||||
|
return fmt.Errorf("%w: renewed %d of %d leases", ErrConcurrencyLeaseLost, tag.RowsAffected(), len(leaseIDs))
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) AttachRateLimitResultToAttempt(ctx context.Context, attemptID string, result RateLimitResult) error {
|
func (s *Store) AttachRateLimitResultToAttempt(ctx context.Context, attemptID string, result RateLimitResult) error {
|
||||||
if attemptID == "" || (len(result.Reservations) == 0 && len(result.LeaseIDs) == 0) {
|
if attemptID == "" || (len(result.Reservations) == 0 && len(result.Leases) == 0) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
tx, err := s.pool.Begin(ctx)
|
tx, err := s.pool.Begin(ctx)
|
||||||
@@ -289,7 +346,8 @@ WHERE id = $1::uuid`, reservation.ReservationID, attemptID); err != nil {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for _, leaseID := range result.LeaseIDs {
|
for _, lease := range result.Leases {
|
||||||
|
leaseID := lease.ID
|
||||||
if leaseID == "" {
|
if leaseID == "" {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -303,6 +361,16 @@ WHERE id = $1::uuid`, leaseID, attemptID); err != nil {
|
|||||||
return tx.Commit(ctx)
|
return tx.Commit(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func concurrencyLeaseIDs(leases []ConcurrencyLease) []string {
|
||||||
|
ids := make([]string, 0, len(leases))
|
||||||
|
for _, lease := range leases {
|
||||||
|
if lease.ID != "" {
|
||||||
|
ids = append(ids, lease.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Store) RecoverInterruptedRuntimeState(ctx context.Context) (RuntimeRecoveryResult, error) {
|
func (s *Store) RecoverInterruptedRuntimeState(ctx context.Context) (RuntimeRecoveryResult, error) {
|
||||||
tx, err := s.pool.Begin(ctx)
|
tx, err := s.pool.Begin(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,198 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestConcurrencyLeaseReservationIsAtomicAcrossPools(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 concurrency lease PostgreSQL integration tests")
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 45*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
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()
|
||||||
|
|
||||||
|
scopeKey := "atomic-" + time.Now().UTC().Format("20060102150405.000000000")
|
||||||
|
taskIDs := createLeaseTestTasks(t, ctx, first, 256, scopeKey)
|
||||||
|
defer deleteLeaseTestTasks(t, first, taskIDs)
|
||||||
|
|
||||||
|
var successes atomic.Int64
|
||||||
|
var peak atomic.Int64
|
||||||
|
monitorCtx, stopMonitor := context.WithCancel(ctx)
|
||||||
|
var monitorWG sync.WaitGroup
|
||||||
|
monitorWG.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer monitorWG.Done()
|
||||||
|
ticker := time.NewTicker(2 * time.Millisecond)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-monitorCtx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
var active int64
|
||||||
|
if err := first.Pool().QueryRow(monitorCtx, `
|
||||||
|
SELECT COUNT(*)
|
||||||
|
FROM gateway_concurrency_leases
|
||||||
|
WHERE scope_type = 'platform_model'
|
||||||
|
AND scope_key = $1
|
||||||
|
AND released_at IS NULL
|
||||||
|
AND expires_at > now()`, scopeKey).Scan(&active); err == nil {
|
||||||
|
for active > peak.Load() && !peak.CompareAndSwap(peak.Load(), active) {
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
errs := make(chan error, len(taskIDs))
|
||||||
|
for index, taskID := range taskIDs {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(index int, taskID string) {
|
||||||
|
defer wg.Done()
|
||||||
|
target := first
|
||||||
|
if index%2 == 1 {
|
||||||
|
target = second
|
||||||
|
}
|
||||||
|
_, err := target.ReserveRateLimits(ctx, taskID, "", []RateLimitReservation{{
|
||||||
|
ScopeType: "platform_model",
|
||||||
|
ScopeKey: scopeKey,
|
||||||
|
Metric: "concurrent",
|
||||||
|
Limit: 64,
|
||||||
|
Amount: 1,
|
||||||
|
LeaseTTLSeconds: 30,
|
||||||
|
}})
|
||||||
|
if err == nil {
|
||||||
|
successes.Add(1)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !errors.Is(err, ErrRateLimited) {
|
||||||
|
errs <- err
|
||||||
|
}
|
||||||
|
}(index, taskID)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
stopMonitor()
|
||||||
|
monitorWG.Wait()
|
||||||
|
close(errs)
|
||||||
|
for err := range errs {
|
||||||
|
t.Fatalf("unexpected reservation error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var active int64
|
||||||
|
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()`, scopeKey).Scan(&active); err != nil {
|
||||||
|
t.Fatalf("count active leases: %v", err)
|
||||||
|
}
|
||||||
|
if successes.Load() != 64 || active != 64 {
|
||||||
|
t.Fatalf("successful reservations=%d active leases=%d, want exactly 64", successes.Load(), active)
|
||||||
|
}
|
||||||
|
if peak.Load() > 64 {
|
||||||
|
t.Fatalf("active lease peak=%d, want <=64", peak.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConcurrencyLeaseRenewalExtendsAndReleases(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 concurrency lease PostgreSQL integration tests")
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
db, err := Connect(ctx, databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect store: %v", err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
scopeKey := "renew-" + time.Now().UTC().Format("20060102150405.000000000")
|
||||||
|
taskIDs := createLeaseTestTasks(t, ctx, db, 1, scopeKey)
|
||||||
|
defer deleteLeaseTestTasks(t, db, taskIDs)
|
||||||
|
|
||||||
|
result, err := db.ReserveRateLimits(ctx, taskIDs[0], "", []RateLimitReservation{{
|
||||||
|
ScopeType: "platform_model",
|
||||||
|
ScopeKey: scopeKey,
|
||||||
|
Metric: "concurrent",
|
||||||
|
Limit: 1,
|
||||||
|
Amount: 1,
|
||||||
|
LeaseTTLSeconds: 2,
|
||||||
|
}})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("reserve short lease: %v", err)
|
||||||
|
}
|
||||||
|
time.Sleep(time.Second)
|
||||||
|
if err := db.RenewConcurrencyLeases(ctx, result.Leases); err != nil {
|
||||||
|
t.Fatalf("renew short lease: %v", err)
|
||||||
|
}
|
||||||
|
time.Sleep(1500 * time.Millisecond)
|
||||||
|
var active bool
|
||||||
|
if err := db.Pool().QueryRow(ctx, `
|
||||||
|
SELECT EXISTS (
|
||||||
|
SELECT 1 FROM gateway_concurrency_leases
|
||||||
|
WHERE id = $1::uuid AND released_at IS NULL AND expires_at > now()
|
||||||
|
)`, result.Leases[0].ID).Scan(&active); err != nil {
|
||||||
|
t.Fatalf("read renewed lease: %v", err)
|
||||||
|
}
|
||||||
|
if !active {
|
||||||
|
t.Fatal("renewed lease expired at its original TTL")
|
||||||
|
}
|
||||||
|
if err := db.ReleaseConcurrencyLeases(ctx, result.Leases); err != nil {
|
||||||
|
t.Fatalf("release renewed lease: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func createLeaseTestTasks(t *testing.T, ctx context.Context, db *Store, count int, marker string) []string {
|
||||||
|
t.Helper()
|
||||||
|
rows, err := db.Pool().Query(ctx, `
|
||||||
|
INSERT INTO gateway_tasks (kind, run_mode, user_id, model, model_type, request, status, queue_key)
|
||||||
|
SELECT 'lease-test', 'simulation', $2, 'lease-test', 'text_generate', '{}'::jsonb, 'queued', $2
|
||||||
|
FROM generate_series(1, $1)
|
||||||
|
RETURNING id::text`, count, marker)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("create lease test tasks: %v", err)
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
ids := make([]string, 0, count)
|
||||||
|
for rows.Next() {
|
||||||
|
var id string
|
||||||
|
if err := rows.Scan(&id); err != nil {
|
||||||
|
t.Fatalf("scan lease test task: %v", err)
|
||||||
|
}
|
||||||
|
ids = append(ids, id)
|
||||||
|
}
|
||||||
|
if err := rows.Err(); err != nil {
|
||||||
|
t.Fatalf("create lease test tasks: %v", err)
|
||||||
|
}
|
||||||
|
return ids
|
||||||
|
}
|
||||||
|
|
||||||
|
func deleteLeaseTestTasks(t *testing.T, db *Store, taskIDs []string) {
|
||||||
|
t.Helper()
|
||||||
|
cleanupCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
if _, err := db.Pool().Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = ANY($1::uuid[])`, taskIDs); err != nil {
|
||||||
|
t.Errorf("delete lease test tasks: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -111,6 +111,7 @@ type CreatePlatformModelInput struct {
|
|||||||
PermissionConfig map[string]any `json:"permissionConfig"`
|
PermissionConfig map[string]any `json:"permissionConfig"`
|
||||||
RetryPolicy map[string]any `json:"retryPolicy"`
|
RetryPolicy map[string]any `json:"retryPolicy"`
|
||||||
RateLimitPolicy map[string]any `json:"rateLimitPolicy"`
|
RateLimitPolicy map[string]any `json:"rateLimitPolicy"`
|
||||||
|
RateLimitPolicyMode string `json:"rateLimitPolicyMode" enums:"inherit,override"`
|
||||||
RuntimePolicySetID string `json:"runtimePolicySetId"`
|
RuntimePolicySetID string `json:"runtimePolicySetId"`
|
||||||
RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride"`
|
RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride"`
|
||||||
Enabled bool `json:"enabled"`
|
Enabled bool `json:"enabled"`
|
||||||
@@ -130,6 +131,7 @@ type RuntimeModelCandidate struct {
|
|||||||
DefaultDiscountFactor float64
|
DefaultDiscountFactor float64
|
||||||
PlatformRetryPolicy map[string]any
|
PlatformRetryPolicy map[string]any
|
||||||
PlatformRateLimitPolicy map[string]any
|
PlatformRateLimitPolicy map[string]any
|
||||||
|
BaseRateLimitPolicy map[string]any
|
||||||
PlatformPriority int
|
PlatformPriority int
|
||||||
PlatformModelID string
|
PlatformModelID string
|
||||||
BaseModelID string
|
BaseModelID string
|
||||||
@@ -153,8 +155,11 @@ type RuntimeModelCandidate struct {
|
|||||||
ModelPricingRuleSetID string
|
ModelPricingRuleSetID string
|
||||||
ModelRetryPolicy map[string]any
|
ModelRetryPolicy map[string]any
|
||||||
ModelRateLimitPolicy map[string]any
|
ModelRateLimitPolicy map[string]any
|
||||||
|
ModelRateLimitPolicyMode string
|
||||||
RuntimePolicySetID string
|
RuntimePolicySetID string
|
||||||
|
RuntimePolicyExplicit bool
|
||||||
RuntimePolicyOverride map[string]any
|
RuntimePolicyOverride map[string]any
|
||||||
|
RateLimitRuntimeOverride map[string]any
|
||||||
RuntimeRetryPolicy map[string]any
|
RuntimeRetryPolicy map[string]any
|
||||||
RuntimeRateLimitPolicy map[string]any
|
RuntimeRateLimitPolicy map[string]any
|
||||||
AutoDisablePolicy map[string]any
|
AutoDisablePolicy map[string]any
|
||||||
@@ -219,10 +224,15 @@ type RateLimitReservation struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type RateLimitResult struct {
|
type RateLimitResult struct {
|
||||||
LeaseIDs []string
|
Leases []ConcurrencyLease
|
||||||
Reservations []RateLimitReservation
|
Reservations []RateLimitReservation
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ConcurrencyLease struct {
|
||||||
|
ID string
|
||||||
|
TTL time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
type CreateTaskAttemptInput struct {
|
type CreateTaskAttemptInput struct {
|
||||||
TaskID string
|
TaskID string
|
||||||
AttemptNo int
|
AttemptNo int
|
||||||
|
|||||||
@@ -0,0 +1,58 @@
|
|||||||
|
ALTER TABLE platform_models
|
||||||
|
ADD COLUMN IF NOT EXISTS rate_limit_policy_mode text;
|
||||||
|
|
||||||
|
-- Historical server-main imports wrote {"rules": []} as a placeholder. It
|
||||||
|
-- means "no override", not an explicit unlimited policy.
|
||||||
|
UPDATE integration_platforms
|
||||||
|
SET rate_limit_policy = '{}'::jsonb,
|
||||||
|
updated_at = now()
|
||||||
|
WHERE rate_limit_policy = '{"rules":[]}'::jsonb;
|
||||||
|
|
||||||
|
UPDATE platform_models
|
||||||
|
SET rate_limit_policy = '{}'::jsonb,
|
||||||
|
updated_at = now()
|
||||||
|
WHERE rate_limit_policy = '{"rules":[]}'::jsonb;
|
||||||
|
|
||||||
|
UPDATE platform_models model
|
||||||
|
SET runtime_policy_set_id = NULL,
|
||||||
|
runtime_policy_override = '{}'::jsonb
|
||||||
|
FROM base_model_catalog base
|
||||||
|
WHERE model.base_model_id = base.id
|
||||||
|
AND model.runtime_policy_set_id IS NOT DISTINCT FROM base.runtime_policy_set_id
|
||||||
|
AND model.runtime_policy_override IS NOT DISTINCT FROM base.runtime_policy_override;
|
||||||
|
|
||||||
|
UPDATE platform_models model
|
||||||
|
SET rate_limit_policy_mode = CASE
|
||||||
|
WHEN model.runtime_policy_override ? 'rateLimitPolicy' THEN 'override'
|
||||||
|
WHEN model.rate_limit_policy <> '{}'::jsonb
|
||||||
|
AND (
|
||||||
|
base.id IS NULL
|
||||||
|
OR model.rate_limit_policy IS DISTINCT FROM base.default_rate_limit_policy
|
||||||
|
)
|
||||||
|
THEN 'override'
|
||||||
|
ELSE 'inherit'
|
||||||
|
END
|
||||||
|
FROM (
|
||||||
|
SELECT platform_model.id AS platform_model_id, catalog.id, catalog.default_rate_limit_policy
|
||||||
|
FROM platform_models platform_model
|
||||||
|
LEFT JOIN base_model_catalog catalog ON catalog.id = platform_model.base_model_id
|
||||||
|
) base
|
||||||
|
WHERE model.id = base.platform_model_id
|
||||||
|
AND model.rate_limit_policy_mode IS NULL;
|
||||||
|
|
||||||
|
UPDATE platform_models
|
||||||
|
SET rate_limit_policy_mode = 'inherit'
|
||||||
|
WHERE rate_limit_policy_mode IS NULL;
|
||||||
|
|
||||||
|
ALTER TABLE platform_models
|
||||||
|
ALTER COLUMN rate_limit_policy_mode SET DEFAULT 'inherit';
|
||||||
|
|
||||||
|
ALTER TABLE platform_models
|
||||||
|
ADD CONSTRAINT platform_models_rate_limit_policy_mode_check
|
||||||
|
CHECK (
|
||||||
|
COALESCE(rate_limit_policy_mode, '') IN ('inherit', 'override')
|
||||||
|
)
|
||||||
|
NOT VALID;
|
||||||
|
|
||||||
|
ALTER TABLE platform_models
|
||||||
|
VALIDATE CONSTRAINT platform_models_rate_limit_policy_mode_check;
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
-- The three migrated image identities previously relied on a global
|
||||||
|
-- user-group limit. Give every upstream binding its own canonical concurrency
|
||||||
|
-- lease so raising the service group no longer removes provider protection.
|
||||||
|
WITH limits(invocation_name, policy) AS (
|
||||||
|
VALUES
|
||||||
|
(
|
||||||
|
'gemini-3-pro-image',
|
||||||
|
'{"rules":[{"metric":"concurrent","limit":10,"leaseTtlSeconds":600}]}'::jsonb
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'gemini-3.1-flash-image',
|
||||||
|
'{"rules":[{"metric":"concurrent","limit":10,"leaseTtlSeconds":600}]}'::jsonb
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'gpt-image-2',
|
||||||
|
'{"rules":[{"metric":"concurrent","limit":5,"leaseTtlSeconds":300}]}'::jsonb
|
||||||
|
)
|
||||||
|
),
|
||||||
|
updated_base_models AS (
|
||||||
|
UPDATE base_model_catalog base_model
|
||||||
|
SET default_rate_limit_policy = limits.policy,
|
||||||
|
metadata = COALESCE(base_model.metadata, '{}'::jsonb) || jsonb_build_object(
|
||||||
|
'rateLimitSource', 'image-gateway-migration',
|
||||||
|
'rateLimitUpdatedAt', '2026-07-24'
|
||||||
|
),
|
||||||
|
updated_at = now()
|
||||||
|
FROM limits
|
||||||
|
WHERE base_model.invocation_name = limits.invocation_name
|
||||||
|
RETURNING base_model.id
|
||||||
|
)
|
||||||
|
UPDATE platform_models platform_model
|
||||||
|
SET rate_limit_policy = '{}'::jsonb,
|
||||||
|
rate_limit_policy_mode = 'inherit',
|
||||||
|
updated_at = now()
|
||||||
|
FROM updated_base_models base_model
|
||||||
|
WHERE platform_model.base_model_id = base_model.id
|
||||||
|
AND (
|
||||||
|
platform_model.rate_limit_policy = '{}'::jsonb
|
||||||
|
OR platform_model.rate_limit_policy = '{"rules":[]}'::jsonb
|
||||||
|
OR platform_model.rate_limit_policy ? 'platformLimits'
|
||||||
|
OR platform_model.rate_limit_policy ? 'modelLimits'
|
||||||
|
OR platform_model.rate_limit_policy ? 'platform_limits'
|
||||||
|
OR platform_model.rate_limit_policy ? 'model_limits'
|
||||||
|
);
|
||||||
@@ -1513,6 +1513,7 @@ function mergeExistingPlatformModelInput(input: PlatformModelBindingInput, curre
|
|||||||
discountFactor: (input.discountFactor ?? existing.discountFactor) || undefined,
|
discountFactor: (input.discountFactor ?? existing.discountFactor) || undefined,
|
||||||
pricingRuleSetId: input.pricingRuleSetId ?? existing.pricingRuleSetId,
|
pricingRuleSetId: input.pricingRuleSetId ?? existing.pricingRuleSetId,
|
||||||
rateLimitPolicy: input.rateLimitPolicy ?? existing.rateLimitPolicy,
|
rateLimitPolicy: input.rateLimitPolicy ?? existing.rateLimitPolicy,
|
||||||
|
rateLimitPolicyMode: input.rateLimitPolicyMode ?? existing.rateLimitPolicyMode,
|
||||||
retryPolicy: input.retryPolicy ?? existing.retryPolicy,
|
retryPolicy: input.retryPolicy ?? existing.retryPolicy,
|
||||||
runtimePolicyOverride: input.runtimePolicyOverride ?? (existing.runtimePolicyOverride as Record<string, unknown> | undefined),
|
runtimePolicyOverride: input.runtimePolicyOverride ?? (existing.runtimePolicyOverride as Record<string, unknown> | undefined),
|
||||||
runtimePolicySetId: input.runtimePolicySetId ?? existing.runtimePolicySetId,
|
runtimePolicySetId: input.runtimePolicySetId ?? existing.runtimePolicySetId,
|
||||||
|
|||||||
@@ -409,13 +409,13 @@ function ModelBindingPolicy(props: { form: PlatformWizardForm; onChange: (value:
|
|||||||
</Label>
|
</Label>
|
||||||
</>
|
</>
|
||||||
)}
|
)}
|
||||||
<ToggleField checked={form.modelOverrideRateLimit} label="覆盖模型限流策略" onChange={(checked) => onChange({ ...form, modelOverrideRateLimit: checked })} />
|
<ToggleField checked={form.modelOverrideRateLimit} label="覆盖模型限流策略" onChange={(checked) => onChange({ ...form, modelOverrideRateLimit: checked, modelRateLimitOverrideTouched: true })} />
|
||||||
{form.modelOverrideRateLimit && (
|
{form.modelOverrideRateLimit && (
|
||||||
<>
|
<>
|
||||||
<Label>模型 RPM<Input value={form.modelRpmLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelRpmLimit: event.target.value })} /></Label>
|
<Label>模型 RPM<Input value={form.modelRpmLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelRpmLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
|
||||||
<Label>模型 RPS<Input value={form.modelRpsLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelRpsLimit: event.target.value })} /></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 })} /></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 })} /></Label>
|
<Label>模型并发<Input value={form.modelConcurrencyLimit} placeholder="不填则不限制" inputMode="numeric" onChange={(event) => onChange({ ...form, modelConcurrencyLimit: event.target.value, modelRateLimitOverrideTouched: true })} /></Label>
|
||||||
</>
|
</>
|
||||||
)}
|
)}
|
||||||
</div>
|
</div>
|
||||||
@@ -1004,6 +1004,7 @@ function platformToForm(
|
|||||||
const rateLimitPolicy = platform.rateLimitPolicy ?? {};
|
const rateLimitPolicy = platform.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 ?? {};
|
||||||
return {
|
return {
|
||||||
...createEmptyPlatformForm(platform.provider, defaults),
|
...createEmptyPlatformForm(platform.provider, defaults),
|
||||||
provider: platform.provider,
|
provider: platform.provider,
|
||||||
@@ -1034,6 +1035,12 @@ function platformToForm(
|
|||||||
selectedModelIds: platformModelBaseIds(platform, baseModels, currentModels),
|
selectedModelIds: platformModelBaseIds(platform, baseModels, currentModels),
|
||||||
modelDiscountFactors: platformModelDiscountFactors(platform, baseModels, currentModels),
|
modelDiscountFactors: platformModelDiscountFactors(platform, baseModels, currentModels),
|
||||||
modelNameMappings: platformModelNameMappings(platform, baseModels, currentModels),
|
modelNameMappings: platformModelNameMappings(platform, baseModels, currentModels),
|
||||||
|
modelOverrideRateLimit: currentModels.length > 0 && currentModels.every((model) => model.rateLimitPolicyMode === 'override'),
|
||||||
|
modelRateLimitOverrideTouched: false,
|
||||||
|
modelRpmLimit: readLimit(modelRateLimitPolicy, 'rpm'),
|
||||||
|
modelRpsLimit: readLimit(modelRateLimitPolicy, 'rps'),
|
||||||
|
modelTpmLimit: readLimit(modelRateLimitPolicy, 'tpm_total'),
|
||||||
|
modelConcurrencyLimit: readLimit(modelRateLimitPolicy, 'concurrent'),
|
||||||
selectionMode: 'partial',
|
selectionMode: 'partial',
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,56 @@
|
|||||||
|
import type { BaseModelCatalogItem } from '@easyai-ai-gateway/contracts';
|
||||||
|
import { describe, expect, it } from 'vitest';
|
||||||
|
import { createEmptyPlatformForm, platformModelPayloads } from './platform-form';
|
||||||
|
|
||||||
|
const baseModel = {
|
||||||
|
id: 'model-1',
|
||||||
|
providerKey: 'simulation',
|
||||||
|
canonicalModelKey: 'simulation:model-1',
|
||||||
|
invocationName: 'model-1',
|
||||||
|
providerModelName: 'model-1',
|
||||||
|
modelType: ['video_generate'],
|
||||||
|
modelAlias: 'model-1',
|
||||||
|
displayName: 'Model 1',
|
||||||
|
referenceCount: 0,
|
||||||
|
pricingVersion: 1,
|
||||||
|
status: 'active',
|
||||||
|
createdAt: '2026-07-24T00:00:00Z',
|
||||||
|
updatedAt: '2026-07-24T00:00:00Z',
|
||||||
|
} satisfies BaseModelCatalogItem;
|
||||||
|
|
||||||
|
function selectedForm() {
|
||||||
|
return {
|
||||||
|
...createEmptyPlatformForm('simulation'),
|
||||||
|
selectedModelIds: [baseModel.id],
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
describe('platform model rate limit payload', () => {
|
||||||
|
it('preserves an existing model policy when the edit form was not touched', () => {
|
||||||
|
const [payload] = platformModelPayloads([baseModel], selectedForm());
|
||||||
|
|
||||||
|
expect(payload.rateLimitPolicyMode).toBeUndefined();
|
||||||
|
expect(payload.rateLimitPolicy).toBeUndefined();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('writes inherit only after the override is explicitly disabled', () => {
|
||||||
|
const [payload] = platformModelPayloads([baseModel], {
|
||||||
|
...selectedForm(),
|
||||||
|
modelRateLimitOverrideTouched: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(payload.rateLimitPolicyMode).toBe('inherit');
|
||||||
|
expect(payload.rateLimitPolicy).toBeUndefined();
|
||||||
|
});
|
||||||
|
|
||||||
|
it('writes an explicit unlimited override when the enabled policy has no rules', () => {
|
||||||
|
const [payload] = platformModelPayloads([baseModel], {
|
||||||
|
...selectedForm(),
|
||||||
|
modelOverrideRateLimit: true,
|
||||||
|
modelRateLimitOverrideTouched: true,
|
||||||
|
});
|
||||||
|
|
||||||
|
expect(payload.rateLimitPolicyMode).toBe('override');
|
||||||
|
expect(payload.rateLimitPolicy).toEqual({ rules: [] });
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -41,6 +41,7 @@ export interface PlatformWizardForm {
|
|||||||
modelRetryEnabled: boolean;
|
modelRetryEnabled: boolean;
|
||||||
modelRetryMaxAttempts: string;
|
modelRetryMaxAttempts: string;
|
||||||
modelOverrideRateLimit: boolean;
|
modelOverrideRateLimit: boolean;
|
||||||
|
modelRateLimitOverrideTouched: boolean;
|
||||||
modelRpmLimit: string;
|
modelRpmLimit: string;
|
||||||
modelRpsLimit: string;
|
modelRpsLimit: string;
|
||||||
modelTpmLimit: string;
|
modelTpmLimit: string;
|
||||||
@@ -100,6 +101,7 @@ export function createEmptyPlatformForm(provider = '', defaults?: ProviderConnec
|
|||||||
modelRetryEnabled: true,
|
modelRetryEnabled: true,
|
||||||
modelRetryMaxAttempts: '2',
|
modelRetryMaxAttempts: '2',
|
||||||
modelOverrideRateLimit: false,
|
modelOverrideRateLimit: false,
|
||||||
|
modelRateLimitOverrideTouched: false,
|
||||||
modelRpmLimit: '',
|
modelRpmLimit: '',
|
||||||
modelRpsLimit: '',
|
modelRpsLimit: '',
|
||||||
modelTpmLimit: '',
|
modelTpmLimit: '',
|
||||||
@@ -187,12 +189,15 @@ export function platformModelPayloads(models: BaseModelCatalogItem[], form: Plat
|
|||||||
retryPolicy: form.modelOverrideRetry
|
retryPolicy: form.modelOverrideRetry
|
||||||
? { enabled: form.modelRetryEnabled, maxAttempts: form.modelRetryEnabled ? positiveInt(form.modelRetryMaxAttempts, 2) : 1 }
|
? { enabled: form.modelRetryEnabled, maxAttempts: form.modelRetryEnabled ? positiveInt(form.modelRetryMaxAttempts, 2) : 1 }
|
||||||
: undefined,
|
: undefined,
|
||||||
rateLimitPolicy: form.modelOverrideRateLimit ? rateLimitPolicyPayload({
|
rateLimitPolicy: form.modelRateLimitOverrideTouched && form.modelOverrideRateLimit ? rateLimitPolicyPayload({
|
||||||
rpmLimit: form.modelRpmLimit,
|
rpmLimit: form.modelRpmLimit,
|
||||||
rpsLimit: form.modelRpsLimit,
|
rpsLimit: form.modelRpsLimit,
|
||||||
tpmLimit: form.modelTpmLimit,
|
tpmLimit: form.modelTpmLimit,
|
||||||
concurrencyLimit: form.modelConcurrencyLimit,
|
concurrencyLimit: form.modelConcurrencyLimit,
|
||||||
}) : undefined,
|
}) : undefined,
|
||||||
|
rateLimitPolicyMode: form.modelRateLimitOverrideTouched
|
||||||
|
? (form.modelOverrideRateLimit ? 'override' : 'inherit')
|
||||||
|
: undefined,
|
||||||
runtimePolicyOverride: platformModelRuntimeOverride(form),
|
runtimePolicyOverride: platformModelRuntimeOverride(form),
|
||||||
};
|
};
|
||||||
});
|
});
|
||||||
@@ -304,14 +309,6 @@ function platformModelRuntimeOverride(form: PlatformWizardForm) {
|
|||||||
maxAttempts: form.modelRetryEnabled ? positiveInt(form.modelRetryMaxAttempts, 2) : 1,
|
maxAttempts: form.modelRetryEnabled ? positiveInt(form.modelRetryMaxAttempts, 2) : 1,
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
if (form.modelOverrideRateLimit) {
|
|
||||||
override.rateLimitPolicy = rateLimitPolicyPayload({
|
|
||||||
rpmLimit: form.modelRpmLimit,
|
|
||||||
rpsLimit: form.modelRpsLimit,
|
|
||||||
tpmLimit: form.modelTpmLimit,
|
|
||||||
concurrencyLimit: form.modelConcurrencyLimit,
|
|
||||||
});
|
|
||||||
}
|
|
||||||
return Object.keys(override).length ? override : undefined;
|
return Object.keys(override).length ? override : undefined;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -152,6 +152,7 @@ export interface PlatformModelBindingInput {
|
|||||||
pricingMode?: string;
|
pricingMode?: string;
|
||||||
retryPolicy?: Record<string, unknown>;
|
retryPolicy?: Record<string, unknown>;
|
||||||
rateLimitPolicy?: Record<string, unknown>;
|
rateLimitPolicy?: Record<string, unknown>;
|
||||||
|
rateLimitPolicyMode?: 'inherit' | 'override';
|
||||||
runtimePolicySetId?: string;
|
runtimePolicySetId?: string;
|
||||||
runtimePolicyOverride?: Record<string, unknown>;
|
runtimePolicyOverride?: Record<string, unknown>;
|
||||||
pricingRuleSetId?: string;
|
pricingRuleSetId?: string;
|
||||||
|
|||||||
@@ -723,6 +723,7 @@ export interface PlatformModel {
|
|||||||
permissionConfig?: Record<string, unknown>;
|
permissionConfig?: Record<string, unknown>;
|
||||||
retryPolicy?: Record<string, unknown>;
|
retryPolicy?: Record<string, unknown>;
|
||||||
rateLimitPolicy?: RateLimitPolicy;
|
rateLimitPolicy?: RateLimitPolicy;
|
||||||
|
rateLimitPolicyMode: 'inherit' | 'override';
|
||||||
runtimePolicySetId?: string;
|
runtimePolicySetId?: string;
|
||||||
runtimePolicyOverride?: RuntimePolicyOverride;
|
runtimePolicyOverride?: RuntimePolicyOverride;
|
||||||
cooldownUntil?: string;
|
cooldownUntil?: string;
|
||||||
|
|||||||
Reference in New Issue
Block a user