feat(api): add load-aware client fallback

This commit is contained in:
2026-05-12 16:59:51 +08:00
parent c5cede2359
commit c2696e7bbe
7 changed files with 558 additions and 5 deletions
+192 -5
View File
@@ -3,6 +3,7 @@ package store
import (
"context"
"fmt"
"sort"
"strings"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
@@ -23,10 +24,18 @@ SELECT p.id::text, p.platform_key, p.name, p.provider,
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, ''),
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, '')),
COALESCE(NULLIF(m.runtime_policy_override, '{}'::jsonb), b.runtime_policy_override, '{}'::jsonb),
COALESCE(rp.retry_policy, '{}'::jsonb), COALESCE(rp.rate_limit_policy, '{}'::jsonb),
COALESCE(rp.auto_disable_policy, '{}'::jsonb), COALESCE(rp.degrade_policy, '{}'::jsonb)
m.permission_config, m.retry_policy, m.rate_limit_policy, COALESCE(m.runtime_policy_set_id::text, COALESCE(b.runtime_policy_set_id::text, '')),
COALESCE(NULLIF(m.runtime_policy_override, '{}'::jsonb), b.runtime_policy_override, '{}'::jsonb),
COALESCE(rp.retry_policy, '{}'::jsonb), COALESCE(rp.rate_limit_policy, '{}'::jsonb),
COALESCE(rp.auto_disable_policy, '{}'::jsonb), COALESCE(rp.degrade_policy, '{}'::jsonb),
COALESCE(con.active, 0)::float8,
COALESCE(queued.waiting, 0)::float8,
COALESCE(rpm.used_value, 0)::float8, COALESCE(rpm.reserved_value, 0)::float8,
COALESCE(tpm.used_value, 0)::float8, COALESCE(tpm.reserved_value, 0)::float8,
COALESCE(s.running_count, 0)::float8,
COALESCE(s.waiting_count, 0)::float8,
COALESCE(s.limiter_ratio, 0)::float8,
COALESCE(EXTRACT(EPOCH FROM s.last_assigned_at), 0)::float8
FROM platform_models m
JOIN integration_platforms p ON p.id = m.platform_id
LEFT JOIN model_catalog_providers cp ON cp.provider_key = p.provider OR cp.provider_code = p.provider
@@ -34,6 +43,57 @@ 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)
LEFT JOIN runtime_client_states s
ON s.client_id = p.platform_key || ':' || $2::text || ':' || COALESCE(NULLIF(m.provider_model_name, ''), m.model_name)
LEFT JOIN (
SELECT scope_key, SUM(lease_value) AS active
FROM gateway_concurrency_leases
WHERE scope_type = 'platform_model'
AND released_at IS NULL
AND expires_at > now()
GROUP BY scope_key
) con ON con.scope_key = m.id::text
LEFT JOIN (
SELECT queued_sources.platform_model_id, COUNT(DISTINCT queued_sources.task_id) AS waiting
FROM (
SELECT t.id::text AS task_id, qm.id::text AS platform_model_id
FROM gateway_tasks t
JOIN integration_platforms qp ON TRUE
JOIN platform_models qm ON qm.platform_id = qp.id
WHERE t.async_mode = true
AND t.status = 'queued'
AND NULLIF(t.model_type, '') IS NOT NULL
AND qm.model_type @> jsonb_build_array(t.model_type)
AND t.queue_key = qp.platform_key || ':' || t.model_type || ':' || COALESCE(NULLIF(qm.provider_model_name, ''), qm.model_name)
AND NOT EXISTS (SELECT 1 FROM gateway_task_attempts existing_attempt WHERE existing_attempt.task_id = t.id)
UNION ALL
SELECT latest.task_id::text AS task_id, latest.platform_model_id
FROM (
SELECT DISTINCT ON (a.task_id) a.task_id, a.platform_model_id::text AS platform_model_id
FROM gateway_tasks t
JOIN gateway_task_attempts a ON a.task_id = t.id
WHERE t.async_mode = true
AND t.status = 'queued'
AND a.platform_model_id IS NOT NULL
ORDER BY a.task_id, a.attempt_no DESC, a.started_at DESC
) latest
) queued_sources
GROUP BY queued_sources.platform_model_id
) queued ON queued.platform_model_id = m.id::text
LEFT JOIN (
SELECT DISTINCT ON (scope_key) scope_key, used_value, reserved_value
FROM gateway_rate_limit_counters
WHERE scope_type = 'platform_model'
AND metric = 'rpm'
AND reset_at > now()
ORDER BY scope_key, window_start DESC
) rpm ON rpm.scope_key = m.id::text
LEFT JOIN (
SELECT scope_key, SUM(used_value) AS used_value, SUM(reserved_value) AS reserved_value
FROM gateway_rate_limit_counters
WHERE scope_type = 'platform_model'
AND metric LIKE 'tpm%'
AND reset_at > now()
GROUP BY scope_key
) tpm ON tpm.scope_key = m.id::text
WHERE p.status = 'enabled'
AND p.deleted_at IS NULL
AND m.enabled = true
@@ -52,7 +112,6 @@ WHERE p.status = 'enabled'
)
)
ORDER BY effective_priority ASC,
COALESCE(s.limiter_ratio, 0) ASC,
COALESCE(s.running_count, 0) ASC,
COALESCE(s.waiting_count, 0) ASC,
COALESCE(s.last_assigned_at, to_timestamp(0)) ASC,
@@ -82,6 +141,16 @@ ORDER BY effective_priority ASC,
var runtimeRateLimitPolicy []byte
var autoDisablePolicy []byte
var degradePolicy []byte
var concurrentActive float64
var queuedWaiting float64
var rpmUsed float64
var rpmReserved float64
var tpmUsed float64
var tpmReserved float64
var stateRunningCount float64
var stateWaitingCount float64
var stateLimiterRatio float64
var lastAssignedUnix float64
if err := rows.Scan(
&item.PlatformID,
&item.PlatformKey,
@@ -124,6 +193,16 @@ ORDER BY effective_priority ASC,
&runtimeRateLimitPolicy,
&autoDisablePolicy,
&degradePolicy,
&concurrentActive,
&queuedWaiting,
&rpmUsed,
&rpmReserved,
&tpmUsed,
&tpmReserved,
&stateRunningCount,
&stateWaitingCount,
&stateLimiterRatio,
&lastAssignedUnix,
); err != nil {
return nil, err
}
@@ -147,6 +226,21 @@ ORDER BY effective_priority ASC,
upstreamModelName := firstNonEmpty(item.ProviderModelName, item.ModelName)
item.ClientID = fmt.Sprintf("%s:%s:%s", item.PlatformKey, item.ModelType, upstreamModelName)
item.QueueKey = item.ClientID
item.RunningCount = stateRunningCount
item.WaitingCount = maxFloat(queuedWaiting, stateWaitingCount)
item.LastAssignedUnix = lastAssignedUnix
applyRuntimeCandidateLoad(&item, runtimeCandidateLoadInput{
Policy: effectiveModelRateLimitPolicy(item.PlatformRateLimitPolicy, item.RuntimeRateLimitPolicy, item.RuntimePolicyOverride, item.ModelRateLimitPolicy),
ConcurrentActive: concurrentActive,
QueuedWaiting: queuedWaiting,
RPMUsed: rpmUsed,
RPMReserved: rpmReserved,
TPMUsed: tpmUsed,
TPMReserved: tpmReserved,
StateRunningCount: stateRunningCount,
StateWaitingCount: stateWaitingCount,
StateLimiterRatio: stateLimiterRatio,
})
items = append(items, item)
}
if err := rows.Err(); err != nil {
@@ -167,9 +261,102 @@ ORDER BY effective_priority ASC,
if len(items) == 0 {
return nil, ErrNoModelCandidate
}
sortRuntimeModelCandidates(items)
return items, nil
}
type runtimeCandidateLoadInput struct {
Policy map[string]any
ConcurrentActive float64
QueuedWaiting float64
RPMUsed float64
RPMReserved float64
TPMUsed float64
TPMReserved float64
StateRunningCount float64
StateWaitingCount float64
StateLimiterRatio float64
}
func applyRuntimeCandidateLoad(candidate *RuntimeModelCandidate, input runtimeCandidateLoadInput) {
rpmLimit := rateLimitForMetric(input.Policy, "rpm")
tpmLimitValue := tpmLimit(input.Policy)
concurrentLimit := rateLimitForMetric(input.Policy, "concurrent")
rpmCurrent := input.RPMUsed + input.RPMReserved
tpmCurrent := input.TPMUsed + input.TPMReserved
concurrentCurrent := input.ConcurrentActive + input.QueuedWaiting
metrics := RuntimeCandidateLoadMetrics{
RPMCurrent: rpmCurrent,
RPMLimit: rpmLimit,
RPMRatio: ratioIfLimited(rpmCurrent, rpmLimit),
TPMCurrent: tpmCurrent,
TPMLimit: tpmLimitValue,
TPMRatio: ratioIfLimited(tpmCurrent, tpmLimitValue),
ConcurrentCurrent: concurrentCurrent,
ConcurrentLimit: concurrentLimit,
ConcurrentRatio: ratioIfLimited(concurrentCurrent, concurrentLimit),
QueuedCount: input.QueuedWaiting,
StateRunningCount: input.StateRunningCount,
StateWaitingCount: input.StateWaitingCount,
StateLimiterRatio: input.StateLimiterRatio,
}
candidate.LoadMetrics = metrics
candidate.LoadLimited = rpmLimit > 0 || tpmLimitValue > 0 || concurrentLimit > 0
candidate.LoadRatio = maxFloat(metrics.RPMRatio, metrics.TPMRatio, metrics.ConcurrentRatio)
}
func ratioIfLimited(current float64, limit float64) float64 {
if limit <= 0 {
return 0
}
return current / limit
}
func sortRuntimeModelCandidates(items []RuntimeModelCandidate) {
hasFull := false
hasNonFull := false
for index := range items {
items[index].LoadAvoided = false
if runtimeCandidateFull(items[index]) {
hasFull = true
} else {
hasNonFull = true
}
}
if hasFull && hasNonFull {
for index := range items {
items[index].LoadAvoided = runtimeCandidateFull(items[index])
}
}
sort.SliceStable(items, func(i, j int) bool {
aFull := runtimeCandidateFull(items[i])
bFull := runtimeCandidateFull(items[j])
if aFull != bFull {
return !aFull
}
if items[i].PlatformPriority != items[j].PlatformPriority {
return items[i].PlatformPriority < items[j].PlatformPriority
}
if items[i].LoadRatio != items[j].LoadRatio {
return items[i].LoadRatio < items[j].LoadRatio
}
if items[i].RunningCount != items[j].RunningCount {
return items[i].RunningCount < items[j].RunningCount
}
if items[i].WaitingCount != items[j].WaitingCount {
return items[i].WaitingCount < items[j].WaitingCount
}
if items[i].LastAssignedUnix != items[j].LastAssignedUnix {
return items[i].LastAssignedUnix < items[j].LastAssignedUnix
}
return false
})
}
func runtimeCandidateFull(candidate RuntimeModelCandidate) bool {
return candidate.LoadLimited && candidate.LoadRatio >= 1
}
func (s *Store) modelCandidateCooldownError(ctx context.Context, model string, modelType string) (error, error) {
rows, err := s.pool.Query(ctx, `
SELECT p.name,
@@ -0,0 +1,61 @@
package store
import "testing"
func TestRuntimeCandidateLoadUsesMaxLimitedMetric(t *testing.T) {
candidate := RuntimeModelCandidate{}
applyRuntimeCandidateLoad(&candidate, runtimeCandidateLoadInput{
Policy: map[string]any{"rules": []any{
map[string]any{"metric": "rpm", "limit": 100},
map[string]any{"metric": "tpm_total", "limit": 1000},
map[string]any{"metric": "concurrent", "limit": 10},
}},
RPMUsed: 40,
RPMReserved: 10,
TPMUsed: 900,
ConcurrentActive: 3,
QueuedWaiting: 2,
})
if !candidate.LoadLimited {
t.Fatal("expected load to be limited when rate limit rules exist")
}
if candidate.LoadMetrics.RPMRatio != 0.5 {
t.Fatalf("expected rpm ratio 0.5, got %v", candidate.LoadMetrics.RPMRatio)
}
if candidate.LoadMetrics.TPMRatio != 0.9 {
t.Fatalf("expected tpm ratio 0.9, got %v", candidate.LoadMetrics.TPMRatio)
}
if candidate.LoadMetrics.ConcurrentRatio != 0.5 {
t.Fatalf("expected concurrent ratio 0.5, got %v", candidate.LoadMetrics.ConcurrentRatio)
}
if candidate.LoadRatio != 0.9 {
t.Fatalf("expected max load ratio 0.9, got %v", candidate.LoadRatio)
}
}
func TestRuntimeCandidateSortingAvoidsFullCandidatesButKeepsFallback(t *testing.T) {
candidates := []RuntimeModelCandidate{
{
PlatformID: "high-priority-full",
PlatformPriority: 1,
LoadLimited: true,
LoadRatio: 1.2,
},
{
PlatformID: "lower-priority-available",
PlatformPriority: 50,
LoadLimited: true,
LoadRatio: 0.2,
},
}
sortRuntimeModelCandidates(candidates)
if candidates[0].PlatformID != "lower-priority-available" {
t.Fatalf("expected non-full candidate to be tried first, got %+v", candidates)
}
if candidates[1].PlatformID != "high-priority-full" || !candidates[1].LoadAvoided {
t.Fatalf("expected full high-priority candidate to remain as avoided fallback, got %+v", candidates)
}
}
+23
View File
@@ -137,6 +137,29 @@ type RuntimeModelCandidate struct {
DegradePolicy map[string]any
ClientID string
QueueKey string
LoadRatio float64
LoadLimited bool
LoadAvoided bool
LoadMetrics RuntimeCandidateLoadMetrics
RunningCount float64
WaitingCount float64
LastAssignedUnix float64
}
type RuntimeCandidateLoadMetrics struct {
RPMCurrent float64
RPMLimit float64
RPMRatio float64
TPMCurrent float64
TPMLimit float64
TPMRatio float64
ConcurrentCurrent float64
ConcurrentLimit float64
ConcurrentRatio float64
QueuedCount float64
StateRunningCount float64
StateWaitingCount float64
StateLimiterRatio float64
}
type RateLimitReservation struct {