summaryrefslogtreecommitdiff
path: root/internal/routing/router.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/routing/router.go')
-rw-r--r--internal/routing/router.go123
1 files changed, 110 insertions, 13 deletions
diff --git a/internal/routing/router.go b/internal/routing/router.go
index 37de633..4465c09 100644
--- a/internal/routing/router.go
+++ b/internal/routing/router.go
@@ -24,6 +24,13 @@ type Router struct {
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 {
@@ -90,6 +97,8 @@ func protocolCompatible(provider domain.Provider, requestProtocol domain.Protoco
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:
@@ -104,25 +113,113 @@ func (r *Router) rotate(modelID string, protocol domain.Protocol, routes []domai
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
+ indices := make([]int, len(routes))
+ for index := range routes {
+ indices[index] = index
}
- position := int(counter % uint64(totalWeight))
- selected := 0
- for i, route := range routes {
- if position < route.Weight {
- selected = i
- break
+ 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
}
- position -= route.Weight
}
+ selectedPosition := weightedPosition(routes, primaryPool, counter)
+ selected := primaryPool[selectedPosition]
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)])
+ 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
+}