feat(api): add load-aware client fallback
This commit is contained in:
@@ -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,
|
||||
°radePolicy,
|
||||
&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)
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user