feat: improve model rate limit tracking
This commit is contained in:
@@ -3,6 +3,7 @@ package store
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
)
|
||||
@@ -38,6 +39,7 @@ WHERE p.status = 'enabled'
|
||||
AND m.enabled = true
|
||||
AND m.model_type @> jsonb_build_array($2)
|
||||
AND (p.cooldown_until IS NULL OR p.cooldown_until <= now())
|
||||
AND (m.cooldown_until IS NULL OR m.cooldown_until <= now())
|
||||
AND (
|
||||
(COALESCE(m.model_alias, '') <> '' AND m.model_alias = $1)
|
||||
OR (
|
||||
@@ -151,6 +153,11 @@ ORDER BY effective_priority ASC,
|
||||
return nil, err
|
||||
}
|
||||
if len(items) == 0 {
|
||||
if unavailableErr, err := s.modelCandidateCooldownError(ctx, model, modelType); err != nil {
|
||||
return nil, err
|
||||
} else if unavailableErr != nil {
|
||||
return nil, unavailableErr
|
||||
}
|
||||
return nil, ErrNoModelCandidate
|
||||
}
|
||||
items, err = s.filterCandidatesByAccessRules(ctx, user, items)
|
||||
@@ -162,3 +169,89 @@ ORDER BY effective_priority ASC,
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func (s *Store) modelCandidateCooldownError(ctx context.Context, model string, modelType string) (error, error) {
|
||||
rows, err := s.pool.Query(ctx, `
|
||||
SELECT p.name,
|
||||
COALESCE(NULLIF(m.display_name, ''), NULLIF(m.model_alias, ''), m.model_name),
|
||||
COALESCE(to_char(p.cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''),
|
||||
GREATEST(COALESCE(EXTRACT(EPOCH FROM p.cooldown_until - now()), 0), 0)::float8,
|
||||
COALESCE(to_char(m.cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''),
|
||||
GREATEST(COALESCE(EXTRACT(EPOCH FROM m.cooldown_until - now()), 0), 0)::float8
|
||||
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
|
||||
WHERE p.status = 'enabled'
|
||||
AND p.deleted_at IS NULL
|
||||
AND m.enabled = true
|
||||
AND m.model_type @> jsonb_build_array($2)
|
||||
AND (
|
||||
(COALESCE(m.model_alias, '') <> '' AND m.model_alias = $1)
|
||||
OR (
|
||||
COALESCE(m.model_alias, '') = ''
|
||||
AND (
|
||||
m.model_name = $1
|
||||
OR b.canonical_model_key = $1
|
||||
OR b.provider_model_name = $1
|
||||
)
|
||||
)
|
||||
)
|
||||
ORDER BY GREATEST(COALESCE(p.cooldown_until, to_timestamp(0)), COALESCE(m.cooldown_until, to_timestamp(0))) DESC,
|
||||
p.priority ASC,
|
||||
m.created_at ASC`, model, modelType)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
var platformName string
|
||||
var displayName string
|
||||
var platformCooldownUntil string
|
||||
var platformRemainingSeconds float64
|
||||
var modelCooldownUntil string
|
||||
var modelRemainingSeconds float64
|
||||
if err := rows.Scan(
|
||||
&platformName,
|
||||
&displayName,
|
||||
&platformCooldownUntil,
|
||||
&platformRemainingSeconds,
|
||||
&modelCooldownUntil,
|
||||
&modelRemainingSeconds,
|
||||
); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if modelRemainingSeconds > 0 {
|
||||
return &ModelCandidateUnavailableError{
|
||||
Code: "model_cooling_down",
|
||||
Message: cooldownErrorMessage("模型", displayName, modelRemainingSeconds, modelCooldownUntil),
|
||||
}, nil
|
||||
}
|
||||
if platformRemainingSeconds > 0 {
|
||||
return &ModelCandidateUnavailableError{
|
||||
Code: "platform_cooling_down",
|
||||
Message: cooldownErrorMessage("平台", platformName, platformRemainingSeconds, platformCooldownUntil),
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func cooldownErrorMessage(scope string, name string, remainingSeconds float64, cooldownUntil string) string {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
name = "候选"
|
||||
}
|
||||
remainingMinutes := remainingSeconds / 60
|
||||
if remainingMinutes < 0.1 {
|
||||
remainingMinutes = 0.1
|
||||
}
|
||||
message := fmt.Sprintf("%s %s 冷却中,剩余 %.1f 分钟", scope, name, remainingMinutes)
|
||||
if strings.TrimSpace(cooldownUntil) != "" {
|
||||
message += ",预计恢复时间 " + cooldownUntil
|
||||
}
|
||||
return message
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user