diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/routing/router.go | 55 |
1 files changed, 51 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...) |
