summaryrefslogtreecommitdiff
path: root/internal/routing/router.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/routing/router.go55
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...)