diff options
Diffstat (limited to 'internal/routing')
| -rw-r--r-- | internal/routing/router.go | 55 | ||||
| -rw-r--r-- | internal/routing/router_test.go | 94 |
2 files changed, 145 insertions, 4 deletions
diff --git a/internal/routing/router.go b/internal/routing/router.go index 53e5261..37de633 100644 --- a/internal/routing/router.go +++ b/internal/routing/router.go @@ -9,31 +9,65 @@ import ( "aigw/internal/catalog" "aigw/internal/domain" + "aigw/internal/providerhealth" ) -var ErrNoRoute = errors.New("no compatible upstream route") +var ( + ErrNoRoute = errors.New("no compatible upstream route") + ErrNoHealthyRoute = errors.New("all compatible upstream routes have open circuits") + ErrProviderNotFound = errors.New("requested provider is not configured for this model and protocol") +) type Router struct { catalog *catalog.Catalog + health *providerhealth.Tracker counters sync.Map } -func New(catalog *catalog.Catalog) *Router { - return &Router{catalog: catalog} +func New(catalog *catalog.Catalog, trackers ...*providerhealth.Tracker) *Router { + router := &Router{catalog: catalog} + if len(trackers) > 0 { + router.health = trackers[0] + } + return router } func (r *Router) Plan(modelID string, protocol domain.Protocol) ([]domain.Route, error) { + return r.plan(modelID, protocol, "") +} + +func (r *Router) PlanProvider(modelID string, protocol domain.Protocol, providerSlug string) ([]domain.Route, error) { + return r.plan(modelID, protocol, providerSlug) +} + +func (r *Router) plan(modelID string, protocol domain.Protocol, providerSlug string) ([]domain.Route, error) { model, err := r.catalog.Model(modelID) if err != nil { return nil, err } routes := make([]domain.Route, 0, len(model.Routes)) + compatible := 0 + matched := 0 for _, route := range model.Routes { - if route.Provider.Protocol == protocol { + if protocolCompatible(route.Provider, protocol) { + compatible++ + if providerSlug != "" && route.Provider.EffectiveSlug() != providerSlug { + continue + } + matched++ + if r.health != nil && r.health.CircuitOpen(providerhealth.RouteKey{ModelID: model.ID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}) { + continue + } routes = append(routes, route) } } if len(routes) == 0 { + if providerSlug != "" && matched == 0 { + return nil, ErrProviderNotFound + } + if compatible > 0 { + return nil, ErrNoHealthyRoute + } return nil, ErrNoRoute } @@ -50,6 +84,19 @@ func (r *Router) Plan(modelID string, protocol domain.Protocol) ([]domain.Route, return result, nil } +func protocolCompatible(provider domain.Provider, requestProtocol domain.Protocol) bool { + switch requestProtocol { + case domain.ProtocolOpenAI: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions" + case domain.ProtocolOpenAIResponses: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses" + case domain.ProtocolAnthropic: + return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages" + default: + return false + } +} + func (r *Router) rotate(modelID string, protocol domain.Protocol, routes []domain.Route) []domain.Route { if len(routes) < 2 { return append([]domain.Route(nil), routes...) diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go index 62dc656..2ad4685 100644 --- a/internal/routing/router_test.go +++ b/internal/routing/router_test.go @@ -1,11 +1,14 @@ package routing import ( + "errors" "testing" + "time" "aigw/internal/catalog" "aigw/internal/config" "aigw/internal/domain" + "aigw/internal/providerhealth" ) func TestPlanHonorsPriorityAndProtocol(t *testing.T) { @@ -34,6 +37,40 @@ func TestPlanHonorsPriorityAndProtocol(t *testing.T) { } } +func TestPlanSkipsOpenCircuitAndRecoversAfterCooldown(t *testing.T) { + now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) + health := providerhealth.New(providerhealth.Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }}) + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}, + {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: "https://fallback.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1}, + {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1}, + }}}, + } + for range 3 { + health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "primary", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true}) + } + router := New(catalog.New(cfg), health) + plan, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil || len(plan) != 1 || plan[0].Provider.ID != "fallback" { + t.Fatalf("open primary was not skipped: plan=%+v err=%v", plan, err) + } + for range 3 { + health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "fallback", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true}) + } + if _, err := router.Plan("public/model", domain.ProtocolOpenAI); !errors.Is(err, ErrNoHealthyRoute) { + t.Fatalf("Plan() error = %v, want ErrNoHealthyRoute", err) + } + now = now.Add(31 * time.Second) + plan, err = router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil || len(plan) != 2 || plan[0].Provider.ID != "primary" { + t.Fatalf("routes did not recover after cooldown: plan=%+v err=%v", plan, err) + } +} + func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { cfg := config.Config{ Providers: []config.ProviderConfig{ @@ -58,3 +95,60 @@ func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { t.Fatalf("unexpected weighted distribution: %+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"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "chat", UpstreamModel: "chat-model", Weight: 1}, + {Provider: "responses", UpstreamModel: "responses-model", Weight: 1}, + }}}, + } + router := New(catalog.New(cfg)) + chat, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil || len(chat) != 1 || chat[0].Provider.ID != "chat" { + t.Fatalf("unexpected Chat plan: %+v err=%v", chat, err) + } + responses, err := router.Plan("public/model", domain.ProtocolOpenAIResponses) + if err != nil || len(responses) != 1 || responses[0].Provider.ID != "responses" { + t.Fatalf("unexpected Responses plan: %+v err=%v", responses, err) + } +} + +func TestPlanProviderPinsWithoutFallbackToOtherProviders(t *testing.T) { + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}, + {ID: "backup", Protocol: domain.ProtocolOpenAI, BaseURL: "https://backup.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "primary", UpstreamModel: "primary-model", Priority: 0, Weight: 1}, + {Provider: "backup", UpstreamModel: "backup-model", Priority: 10, Weight: 1}, + }}}, + } + router := New(catalog.New(cfg)) + plan, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "backup") + if err != nil || len(plan) != 1 || plan[0].Provider.ID != "backup" { + t.Fatalf("unexpected pinned plan: %+v err=%v", plan, err) + } + if _, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "missing"); !errors.Is(err, ErrProviderNotFound) { + t.Fatalf("missing provider error = %v, want ErrProviderNotFound", err) + } +} + +func TestPlanProviderHonorsCircuitBreaker(t *testing.T) { + now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) + health := providerhealth.New(providerhealth.Options{FailureThreshold: 1, OpenDuration: time.Minute, Now: func() time.Time { return now }}) + cfg := config.Config{ + Providers: []config.ProviderConfig{{ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}}, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}}}, + } + health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "primary", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true}) + router := New(catalog.New(cfg), health) + if _, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "primary"); !errors.Is(err, ErrNoHealthyRoute) { + t.Fatalf("open pinned provider error = %v, want ErrNoHealthyRoute", err) + } +} |
