diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/routing/router_test.go | 45 |
1 files changed, 45 insertions, 0 deletions
diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go index 2ad4685..f421184 100644 --- a/internal/routing/router_test.go +++ b/internal/routing/router_test.go @@ -96,15 +96,56 @@ func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { } } +func TestPlanPrefersLowerTTFTAndStillExplores(t *testing.T) { + health := providerhealth.New(providerhealth.Options{}) + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: "https://fast.test", APIKey: "one"}, + {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: "https://slow.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "fast", UpstreamModel: "model", Weight: 1}, + {Provider: "slow", UpstreamModel: "model", Weight: 1}, + }}}, + } + for _, providerID := range []string{"fast", "slow"} { + key := providerhealth.RouteKey{ModelID: "public/model", ProviderID: providerID, WireAPI: "chat_completions"} + for range 5 { + health.Observe(key, providerhealth.Observation{StatusCode: 200, Latency: 20 * time.Millisecond}) + } + latency := 100 * time.Millisecond + if providerID == "slow" { + latency = 500 * time.Millisecond + } + for range 3 { + health.ObserveTTFT(key, latency) + } + } + router := New(catalog.New(cfg), health) + counts := map[string]int{} + for range 200 { + plan, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil { + t.Fatal(err) + } + counts[plan[0].Provider.ID]++ + } + if counts["fast"] != 190 || counts["slow"] != 10 { + t.Fatalf("adaptive selection must prefer fast route while preserving 5%% exploration: %+v", counts) + } +} + func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) { cfg := config.Config{ Providers: []config.ProviderConfig{ {ID: "chat", Protocol: domain.ProtocolOpenAI, WireAPI: "chat_completions", BaseURL: "https://chat.test", APIKey: "one"}, {ID: "responses", Protocol: domain.ProtocolOpenAI, WireAPI: "responses", BaseURL: "https://responses.test", APIKey: "two"}, + {ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://embeddings.test", APIKey: "three"}, }, Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ {Provider: "chat", UpstreamModel: "chat-model", Weight: 1}, {Provider: "responses", UpstreamModel: "responses-model", Weight: 1}, + {Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}, }}}, } router := New(catalog.New(cfg)) @@ -116,6 +157,10 @@ func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) { if err != nil || len(responses) != 1 || responses[0].Provider.ID != "responses" { t.Fatalf("unexpected Responses plan: %+v err=%v", responses, err) } + embeddings, err := router.Plan("public/model", domain.ProtocolOpenAIEmbeddings) + if err != nil || len(embeddings) != 1 || embeddings[0].Provider.ID != "embeddings" { + t.Fatalf("unexpected Embeddings plan: %+v err=%v", embeddings, err) + } } func TestPlanProviderPinsWithoutFallbackToOtherProviders(t *testing.T) { |
