feat(provider): support per-model cooldown scopes

This commit is contained in:
Alix-007 2026-03-17 21:26:58 +08:00
parent e70928cc6f
commit e55208c450
6 changed files with 222 additions and 6 deletions

View file

@ -167,8 +167,9 @@ func NewAgentInstance(
summarizeTokenPercent = 75
}
// Resolve fallback candidates
// Resolve fallback candidates.
candidates := resolveModelCandidates(cfg, defaults.Provider, model, fallbacks)
candidates = applyCooldownKeys(cfg, candidates)
// Model routing setup: pre-resolve light model candidates at creation time
// to avoid repeated model_list lookups on every incoming message.
@ -192,7 +193,7 @@ func NewAgentInstance(
LightModel: rc.LightModel,
Threshold: rc.Threshold,
})
lightCandidates = resolved
lightCandidates = applyCooldownKeys(cfg, resolved)
lightProvider = lp
}
}
@ -228,6 +229,64 @@ func NewAgentInstance(
}
}
func applyCooldownKeys(cfg *config.Config, candidates []providers.FallbackCandidate) []providers.FallbackCandidate {
if len(candidates) == 0 {
return candidates
}
resolved := make([]providers.FallbackCandidate, len(candidates))
copy(resolved, candidates)
for i := range resolved {
resolved[i].CooldownKey = resolveCooldownKey(cfg, resolved[i])
}
return resolved
}
func resolveCooldownKey(cfg *config.Config, candidate providers.FallbackCandidate) string {
if candidate.Provider == "" {
return ""
}
if cfg != nil && candidateUsesPerModelCooldown(cfg, candidate) {
return providers.ModelKey(candidate.Provider, candidate.Model)
}
return candidate.Provider
}
func candidateUsesPerModelCooldown(cfg *config.Config, candidate providers.FallbackCandidate) bool {
if cfg == nil {
return false
}
candidateKey := providers.ModelKey(candidate.Provider, candidate.Model)
for i := range cfg.ModelList {
ref := providers.ParseModelRef(cfg.ModelList[i].Model, "openai")
if ref == nil {
continue
}
if providers.ModelKey(ref.Provider, ref.Model) != candidateKey {
continue
}
if normalizeCooldownStrategy(cfg.ModelList[i].CooldownStrategy) == "model" {
return true
}
}
return false
}
func normalizeCooldownStrategy(strategy string) string {
switch strings.ReplaceAll(strings.ToLower(strings.TrimSpace(strategy)), "_", "-") {
case "model", "per-model":
return "model"
default:
return "provider"
}
}
// resolveAgentWorkspace determines the workspace directory for an agent.
func resolveAgentWorkspace(agentCfg *config.AgentConfig, defaults *config.AgentDefaults) string {
if agentCfg != nil && strings.TrimSpace(agentCfg.Workspace) != "" {

View file

@ -165,6 +165,52 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
}
}
func TestNewAgentInstance_ResolvePerModelCooldownKeys(t *testing.T) {
tmpDir := t.TempDir()
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "router-a",
ModelFallbacks: []string{"router-b", "shared-provider"},
},
},
ModelList: []config.ModelConfig{
{
ModelName: "router-a",
Model: "litellm/openai/gpt-4o-mini",
CooldownStrategy: "per-model",
},
{
ModelName: "router-b",
Model: "litellm/openai/gpt-4o",
CooldownStrategy: "model",
},
{
ModelName: "shared-provider",
Model: "litellm/openai/gpt-4.1",
},
},
}
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, &mockProvider{})
if len(agent.Candidates) != 3 {
t.Fatalf("len(Candidates) = %d, want 3", len(agent.Candidates))
}
if got := agent.Candidates[0].CooldownKey; got != "litellm/openai/gpt-4o-mini" {
t.Fatalf("candidate[0] cooldown key = %q, want %q", got, "litellm/openai/gpt-4o-mini")
}
if got := agent.Candidates[1].CooldownKey; got != "litellm/openai/gpt-4o" {
t.Fatalf("candidate[1] cooldown key = %q, want %q", got, "litellm/openai/gpt-4o")
}
if got := agent.Candidates[2].CooldownKey; got != "litellm" {
t.Fatalf("candidate[2] cooldown key = %q, want %q", got, "litellm")
}
}
func TestNewAgentInstance_AllowsMediaTempDirForReadListAndExec(t *testing.T) {
workspace := t.TempDir()
mediaDir := media.TempDir()

View file

@ -671,6 +671,10 @@ type ModelConfig struct {
MaxTokensField string `json:"max_tokens_field,omitempty"` // Field name for max tokens (e.g., "max_completion_tokens")
RequestTimeout int `json:"request_timeout,omitempty"`
ThinkingLevel string `json:"thinking_level,omitempty"` // Extended thinking: off|low|medium|high|xhigh|adaptive
// CooldownStrategy controls the scope of cooldown tracking for this model.
// - "provider" (default): cooldown is shared across all models in the provider
// - "model" or "per-model": cooldown is isolated to this specific model
CooldownStrategy string `json:"cooldown_strategy,omitempty"`
ExtraBody map[string]any `json:"extra_body,omitempty"` // Additional fields to inject into request body
APIKeys SecureStrings `json:"api_keys,omitzero" yaml:"api_keys,omitempty"` // API authentication keys (multiple keys for failover)
@ -701,9 +705,30 @@ func (c *ModelConfig) Validate() error {
if c.Model == "" {
return fmt.Errorf("model is required")
}
if !isValidCooldownStrategy(c.CooldownStrategy) {
return fmt.Errorf("cooldown_strategy must be one of: provider, model, per-model (or per_model)")
}
return nil
}
// NormalizeCooldownStrategy canonicalizes cooldown strategy aliases.
// It returns "provider" for the default/shared scope, "model" for per-model
// scope, and "" for invalid values.
func NormalizeCooldownStrategy(strategy string) string {
switch strings.ReplaceAll(strings.ToLower(strings.TrimSpace(strategy)), "_", "-") {
case "", "provider":
return "provider"
case "model", "per-model":
return "model"
default:
return ""
}
}
func isValidCooldownStrategy(strategy string) bool {
return NormalizeCooldownStrategy(strategy) != ""
}
func (c *ModelConfig) SetAPIKey(value string) {
if len(c.APIKeys) > 0 {
c.APIKeys[0].Set(value)

View file

@ -194,6 +194,15 @@ func TestModelConfig_Validate(t *testing.T) {
},
wantErr: false,
},
{
name: "valid per-model cooldown strategy",
config: ModelConfig{
ModelName: "router-model",
Model: "litellm/openai/gpt-4o",
CooldownStrategy: "per-model",
},
wantErr: false,
},
{
name: "missing model_name",
config: ModelConfig{
@ -213,6 +222,15 @@ func TestModelConfig_Validate(t *testing.T) {
config: ModelConfig{},
wantErr: true,
},
{
name: "invalid cooldown strategy",
config: ModelConfig{
ModelName: "test",
Model: "openai/gpt-4o",
CooldownStrategy: "backend",
},
wantErr: true,
},
}
for _, tt := range tests {

View file

@ -16,6 +16,7 @@ type FallbackChain struct {
type FallbackCandidate struct {
Provider string
Model string
CooldownKey string
}
// FallbackResult contains the successful response and metadata about all attempts.
@ -41,6 +42,13 @@ func NewFallbackChain(cooldown *CooldownTracker) *FallbackChain {
return &FallbackChain{cooldown: cooldown}
}
func (c FallbackCandidate) cooldownKey() string {
if strings.TrimSpace(c.CooldownKey) != "" {
return c.CooldownKey
}
return ModelKey(c.Provider, c.Model)
}
// ResolveCandidates parses model config into a deduplicated candidate list.
func ResolveCandidates(cfg ModelConfig, defaultProvider string) []FallbackCandidate {
return ResolveCandidatesWithLookup(cfg, defaultProvider, nil)
@ -112,6 +120,8 @@ func (fc *FallbackChain) Execute(
}
for i, candidate := range candidates {
cooldownKey := candidate.cooldownKey()
// Check context before each attempt.
if ctx.Err() == context.Canceled {
return nil, context.Canceled
@ -119,7 +129,6 @@ func (fc *FallbackChain) Execute(
// Check cooldown (per provider/model, not just provider).
// This allows multi-key failover where different keys use different model names.
cooldownKey := ModelKey(candidate.Provider, candidate.Model)
if !fc.cooldown.IsAvailable(cooldownKey) {
remaining := fc.cooldown.CooldownRemaining(cooldownKey)
result.Attempts = append(result.Attempts, FallbackAttempt{

View file

@ -11,6 +11,14 @@ func makeCandidate(provider, model string) FallbackCandidate {
return FallbackCandidate{Provider: provider, Model: model}
}
func makePerModelCandidate(provider, model string) FallbackCandidate {
return FallbackCandidate{
Provider: provider,
Model: model,
CooldownKey: ModelKey(provider, model),
}
}
func successRun(content string) func(ctx context.Context, provider, model string) (*LLMResponse, error) {
return func(ctx context.Context, provider, model string) (*LLMResponse, error) {
return &LLMResponse{Content: content, FinishReason: "stop"}, nil
@ -191,6 +199,57 @@ func TestFallback_CooldownSkip(t *testing.T) {
}
}
func TestFallback_PerModelCooldownDoesNotSkipSiblingModel(t *testing.T) {
now := time.Now()
ct, _ := newTestTracker(now)
fc := NewFallbackChain(ct)
candidates := []FallbackCandidate{
makePerModelCandidate("litellm", "openai/gpt-4o-mini"),
makePerModelCandidate("litellm", "openai/gpt-4o"),
}
ct.MarkFailure(candidates[0].CooldownKey, FailoverRateLimit)
called := []string{}
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
called = append(called, provider+"/"+model)
return &LLMResponse{Content: "ok", FinishReason: "stop"}, nil
}
result, err := fc.Execute(context.Background(), candidates, run)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result.Model != "openai/gpt-4o" {
t.Fatalf("result model = %q, want %q", result.Model, "openai/gpt-4o")
}
if len(called) != 1 || called[0] != "litellm/openai/gpt-4o" {
t.Fatalf("called = %v, want only second model", called)
}
}
func TestFallback_DefaultProviderCooldownSkipsSiblingModel(t *testing.T) {
now := time.Now()
ct, _ := newTestTracker(now)
fc := NewFallbackChain(ct)
candidates := []FallbackCandidate{
makeCandidate("litellm", "openai/gpt-4o-mini"),
makeCandidate("litellm", "openai/gpt-4o"),
}
ct.MarkFailure("litellm", FailoverRateLimit)
_, err := fc.Execute(context.Background(), candidates, func(ctx context.Context, provider, model string) (*LLMResponse, error) {
t.Fatal("run should not be called when provider cooldown is shared")
return nil, nil
})
if err == nil {
t.Fatal("expected error when all same-provider candidates are skipped")
}
}
func TestFallback_AllInCooldown(t *testing.T) {
ct := NewCooldownTracker()
fc := NewFallbackChain(ct)