1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
|
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
}
|