This commit is contained in:
afjcjsbx 2026-03-30 19:48:06 +02:00
parent 51d61ac970
commit 3bf9e8b484
2 changed files with 59 additions and 45 deletions

View file

@ -293,42 +293,73 @@ func TestFallback_SuccessResetsCooldown(t *testing.T) {
}
}
func TestFallback_LocalRateLimitSkipsToHealthyFallback(t *testing.T) {
func assertLocalRateLimitSkipsToHealthyFallback(
t *testing.T,
primaryKey string,
fallbackKey string,
fallbackProvider string,
fallbackModel string,
execute func(context.Context, *FallbackChain, []FallbackCandidate,
func(context.Context, string, string) (*LLMResponse, error),
) (*FallbackResult, error),
responseContent string,
) {
t.Helper()
ct := NewCooldownTracker()
rl := NewRateLimiterRegistry()
rl.Register("model_name:primary", 1)
if err := rl.Wait(context.Background(), "model_name:primary"); err != nil {
rl.Register(primaryKey, 1)
if err := rl.Wait(context.Background(), primaryKey); err != nil {
t.Fatalf("failed to pre-drain primary limiter: %v", err)
}
fc := NewFallbackChain(ct, rl)
candidates := []FallbackCandidate{
{Provider: "openai", Model: "gpt-4o", IdentityKey: "model_name:primary"},
{Provider: "anthropic", Model: "claude", IdentityKey: "model_name:fallback"},
{Provider: "openai", Model: "gpt-4o", IdentityKey: primaryKey},
{Provider: fallbackProvider, Model: fallbackModel, IdentityKey: fallbackKey},
}
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
if provider != "anthropic" || model != "claude" {
if provider != fallbackProvider || model != fallbackModel {
t.Fatalf("expected fallback candidate to run, got %s/%s", provider, model)
}
return &LLMResponse{Content: "fallback ok", FinishReason: "stop"}, nil
return &LLMResponse{Content: responseContent, FinishReason: "stop"}, nil
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond)
defer cancel()
result, err := fc.Execute(ctx, candidates, run)
result, err := execute(ctx, fc, candidates, run)
if err != nil {
t.Fatalf("expected fallback success, got error: %v", err)
}
if result.Provider != "anthropic" || result.Model != "claude" {
t.Fatalf("result = %s/%s, want anthropic/claude", result.Provider, result.Model)
if result.Provider != fallbackProvider || result.Model != fallbackModel {
t.Fatalf("result = %s/%s, want %s/%s", result.Provider, result.Model, fallbackProvider, fallbackModel)
}
if len(result.Attempts) != 1 || !result.Attempts[0].Skipped {
t.Fatalf("expected one skipped primary attempt, got %+v", result.Attempts)
}
}
func TestFallback_LocalRateLimitSkipsToHealthyFallback(t *testing.T) {
assertLocalRateLimitSkipsToHealthyFallback(
t,
"model_name:primary",
"model_name:fallback",
"anthropic",
"claude",
func(
ctx context.Context,
fc *FallbackChain,
candidates []FallbackCandidate,
run func(context.Context, string, string) (*LLMResponse, error),
) (*FallbackResult, error) {
return fc.Execute(ctx, candidates, run)
},
"fallback ok",
)
}
// --- Image Fallback Tests ---
func TestImageFallback_Success(t *testing.T) {
@ -421,39 +452,22 @@ func TestImageFallback_RetryOnOtherErrors(t *testing.T) {
}
func TestImageFallback_LocalRateLimitSkipsToHealthyFallback(t *testing.T) {
ct := NewCooldownTracker()
rl := NewRateLimiterRegistry()
rl.Register("model_name:primary-image", 1)
if err := rl.Wait(context.Background(), "model_name:primary-image"); err != nil {
t.Fatalf("failed to pre-drain primary image limiter: %v", err)
}
fc := NewFallbackChain(ct, rl)
candidates := []FallbackCandidate{
{Provider: "openai", Model: "gpt-4o", IdentityKey: "model_name:primary-image"},
{Provider: "anthropic", Model: "claude-sonnet", IdentityKey: "model_name:fallback-image"},
}
run := func(ctx context.Context, provider, model string) (*LLMResponse, error) {
if provider != "anthropic" || model != "claude-sonnet" {
t.Fatalf("expected image fallback candidate to run, got %s/%s", provider, model)
}
return &LLMResponse{Content: "image fallback ok", FinishReason: "stop"}, nil
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond)
defer cancel()
result, err := fc.ExecuteImage(ctx, candidates, run)
if err != nil {
t.Fatalf("expected image fallback success, got error: %v", err)
}
if result.Provider != "anthropic" || result.Model != "claude-sonnet" {
t.Fatalf("result = %s/%s, want anthropic/claude-sonnet", result.Provider, result.Model)
}
if len(result.Attempts) != 1 || !result.Attempts[0].Skipped {
t.Fatalf("expected one skipped primary attempt, got %+v", result.Attempts)
}
assertLocalRateLimitSkipsToHealthyFallback(
t,
"model_name:primary-image",
"model_name:fallback-image",
"anthropic",
"claude-sonnet",
func(
ctx context.Context,
fc *FallbackChain,
candidates []FallbackCandidate,
run func(context.Context, string, string) (*LLMResponse, error),
) (*FallbackResult, error) {
return fc.ExecuteImage(ctx, candidates, run)
},
"image fallback ok",
)
}
func TestImageFallback_NoCandidates(t *testing.T) {

View file

@ -38,8 +38,8 @@ func newRateLimiter(rpm int) *RateLimiter {
}
}
// Wait blocks until a token is available or ctx is cancelled.
// Returns ctx.Err() if cancelled while waiting.
// Wait blocks until a token is available or ctx is canceled.
// Returns ctx.Err() if canceled while waiting.
func (rl *RateLimiter) Wait(ctx context.Context) error {
for {
rl.mu.Lock()