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 } const ( adaptiveMinimumAvailabilitySamples = 5 adaptiveMinimumTTFTSamples = 3 adaptiveExplorationInterval = 20 adaptivePreferenceThreshold = 0.90 ) 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.ProtocolOpenAIEmbeddings: return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "embeddings" 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 indices := make([]int, len(routes)) for index := range routes { indices[index] = index } primaryPool := indices if r.health != nil && counter%adaptiveExplorationInterval != adaptiveExplorationInterval-1 { if preferred := r.preferredRoutes(modelID, routes); len(preferred) > 0 && len(preferred) < len(routes) { indices = append(preferred, difference(indices, preferred)...) primaryPool = preferred } } selectedPosition := weightedPosition(routes, primaryPool, counter) selected := primaryPool[selectedPosition] result := make([]domain.Route, 0, len(routes)) result = append(result, routes[selected]) for _, index := range indices { if index != selected { result = append(result, routes[index]) } } return result } func (r *Router) preferredRoutes(modelID string, routes []domain.Route) []int { type measuredRoute struct { qualified bool status providerhealth.Status } measured := make([]measuredRoute, len(routes)) fastestTTFT := int64(0) for index, route := range routes { status, exists := r.health.StatusFor(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}) if !exists { continue } qualified := status.RecentSamples >= adaptiveMinimumAvailabilitySamples || status.TTFTSamples >= adaptiveMinimumTTFTSamples measured[index] = measuredRoute{qualified: qualified, status: status} if status.TTFTSamples >= adaptiveMinimumTTFTSamples && status.TTFTEWMA > 0 && (fastestTTFT == 0 || status.TTFTEWMA < fastestTTFT) { fastestTTFT = status.TTFTEWMA } } scores := make([]float64, len(routes)) best := 0.0 hasQualified := false for index, item := range measured { score := 1.0 if item.qualified { hasQualified = true if item.status.RecentSamples >= adaptiveMinimumAvailabilitySamples { availability := item.status.AvailabilityPercent / 100 if availability < 0.05 { availability = 0.05 } score *= availability } if fastestTTFT > 0 && item.status.TTFTSamples >= adaptiveMinimumTTFTSamples && item.status.TTFTEWMA > 0 { latencyFactor := float64(fastestTTFT) / float64(item.status.TTFTEWMA) if latencyFactor < 0.10 { latencyFactor = 0.10 } score *= latencyFactor } } scores[index] = score if score > best { best = score } } if !hasQualified { return nil } result := make([]int, 0, len(routes)) for index, score := range scores { if score >= best*adaptivePreferenceThreshold { result = append(result, index) } } return result } func difference(all, selected []int) []int { included := make(map[int]struct{}, len(selected)) for _, index := range selected { included[index] = struct{}{} } result := make([]int, 0, len(all)-len(selected)) for _, index := range all { if _, exists := included[index]; !exists { result = append(result, index) } } return result } func weightedPosition(routes []domain.Route, indices []int, counter uint64) int { totalWeight := 0 for _, index := range indices { totalWeight += routes[index].Weight } position := int(counter % uint64(totalWeight)) for positionIndex, routeIndex := range indices { if position < routes[routeIndex].Weight { return positionIndex } position -= routes[routeIndex].Weight } return 0 }