package routing import ( "errors" "sort" "strconv" "sync" "sync/atomic" "aigw/internal/catalog" "aigw/internal/domain" "aigw/internal/providerhealth" ) 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, 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 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 } sort.SliceStable(routes, func(i, j int) bool { return routes[i].Priority < routes[j].Priority }) result := make([]domain.Route, 0, len(routes)) for start := 0; start < len(routes); { end := start + 1 for end < len(routes) && routes[end].Priority == routes[start].Priority { end++ } result = append(result, r.rotate(modelID, protocol, routes[start:end])...) start = end } 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...) } key := modelID + "\x00" + string(protocol) + "\x00" + strconv.Itoa(routes[0].Priority) counterValue, _ := r.counters.LoadOrStore(key, &atomic.Uint64{}) counter := counterValue.(*atomic.Uint64).Add(1) - 1 totalWeight := 0 for _, route := range routes { totalWeight += route.Weight } position := int(counter % uint64(totalWeight)) selected := 0 for i, route := range routes { if position < route.Weight { selected = i break } position -= route.Weight } result := make([]domain.Route, 0, len(routes)) result = append(result, routes[selected]) for offset := 1; offset < len(routes); offset++ { result = append(result, routes[(selected+offset)%len(routes)]) } return result }