summaryrefslogtreecommitdiff
path: root/internal/routing
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 09:29:41 +1200
committerChia <Chia@93.nz>2026-08-06 09:32:46 +1200
commit41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch)
treec730526150e55e39b822d5197e4a20318ecaa449 /internal/routing
parenteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (diff)
feat: complete commercial control plane, billing, auth, and model catalog
- add PostgreSQL control-plane persistence with Redis-degraded hot reload - implement prepaid balance, usage ledger, Stripe top-up and reconciliation - add registration, email verification, password reset, invitations and RBAC - support TOTP, Passkey MFA, device sessions, quotas and rate limits - add tenant billing profiles, audit logs and operational readiness checks - build authenticated admin console, Quickstart, Playground and usage analytics - add public model catalog with pricing, filtering and cost estimation - support OpenAI Responses providers and provider health failover - validate real upstream usage reporting and balance settlement
Diffstat (limited to 'internal/routing')
-rw-r--r--internal/routing/router.go55
-rw-r--r--internal/routing/router_test.go94
2 files changed, 145 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...)
diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go
index 62dc656..2ad4685 100644
--- a/internal/routing/router_test.go
+++ b/internal/routing/router_test.go
@@ -1,11 +1,14 @@
package routing
import (
+ "errors"
"testing"
+ "time"
"aigw/internal/catalog"
"aigw/internal/config"
"aigw/internal/domain"
+ "aigw/internal/providerhealth"
)
func TestPlanHonorsPriorityAndProtocol(t *testing.T) {
@@ -34,6 +37,40 @@ func TestPlanHonorsPriorityAndProtocol(t *testing.T) {
}
}
+func TestPlanSkipsOpenCircuitAndRecoversAfterCooldown(t *testing.T) {
+ now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC)
+ health := providerhealth.New(providerhealth.Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }})
+ cfg := config.Config{
+ Providers: []config.ProviderConfig{
+ {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"},
+ {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: "https://fallback.test", APIKey: "two"},
+ },
+ Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{
+ {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1},
+ {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1},
+ }}},
+ }
+ for range 3 {
+ health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "primary", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true})
+ }
+ router := New(catalog.New(cfg), health)
+ plan, err := router.Plan("public/model", domain.ProtocolOpenAI)
+ if err != nil || len(plan) != 1 || plan[0].Provider.ID != "fallback" {
+ t.Fatalf("open primary was not skipped: plan=%+v err=%v", plan, err)
+ }
+ for range 3 {
+ health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "fallback", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true})
+ }
+ if _, err := router.Plan("public/model", domain.ProtocolOpenAI); !errors.Is(err, ErrNoHealthyRoute) {
+ t.Fatalf("Plan() error = %v, want ErrNoHealthyRoute", err)
+ }
+ now = now.Add(31 * time.Second)
+ plan, err = router.Plan("public/model", domain.ProtocolOpenAI)
+ if err != nil || len(plan) != 2 || plan[0].Provider.ID != "primary" {
+ t.Fatalf("routes did not recover after cooldown: plan=%+v err=%v", plan, err)
+ }
+}
+
func TestPlanUsesWeightsForPrimarySelection(t *testing.T) {
cfg := config.Config{
Providers: []config.ProviderConfig{
@@ -58,3 +95,60 @@ func TestPlanUsesWeightsForPrimarySelection(t *testing.T) {
t.Fatalf("unexpected weighted distribution: %+v", counts)
}
}
+
+func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) {
+ cfg := config.Config{
+ Providers: []config.ProviderConfig{
+ {ID: "chat", Protocol: domain.ProtocolOpenAI, WireAPI: "chat_completions", BaseURL: "https://chat.test", APIKey: "one"},
+ {ID: "responses", Protocol: domain.ProtocolOpenAI, WireAPI: "responses", BaseURL: "https://responses.test", APIKey: "two"},
+ },
+ Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{
+ {Provider: "chat", UpstreamModel: "chat-model", Weight: 1},
+ {Provider: "responses", UpstreamModel: "responses-model", Weight: 1},
+ }}},
+ }
+ router := New(catalog.New(cfg))
+ chat, err := router.Plan("public/model", domain.ProtocolOpenAI)
+ if err != nil || len(chat) != 1 || chat[0].Provider.ID != "chat" {
+ t.Fatalf("unexpected Chat plan: %+v err=%v", chat, err)
+ }
+ responses, err := router.Plan("public/model", domain.ProtocolOpenAIResponses)
+ if err != nil || len(responses) != 1 || responses[0].Provider.ID != "responses" {
+ t.Fatalf("unexpected Responses plan: %+v err=%v", responses, err)
+ }
+}
+
+func TestPlanProviderPinsWithoutFallbackToOtherProviders(t *testing.T) {
+ cfg := config.Config{
+ Providers: []config.ProviderConfig{
+ {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"},
+ {ID: "backup", Protocol: domain.ProtocolOpenAI, BaseURL: "https://backup.test", APIKey: "two"},
+ },
+ Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{
+ {Provider: "primary", UpstreamModel: "primary-model", Priority: 0, Weight: 1},
+ {Provider: "backup", UpstreamModel: "backup-model", Priority: 10, Weight: 1},
+ }}},
+ }
+ router := New(catalog.New(cfg))
+ plan, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "backup")
+ if err != nil || len(plan) != 1 || plan[0].Provider.ID != "backup" {
+ t.Fatalf("unexpected pinned plan: %+v err=%v", plan, err)
+ }
+ if _, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "missing"); !errors.Is(err, ErrProviderNotFound) {
+ t.Fatalf("missing provider error = %v, want ErrProviderNotFound", err)
+ }
+}
+
+func TestPlanProviderHonorsCircuitBreaker(t *testing.T) {
+ now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC)
+ health := providerhealth.New(providerhealth.Options{FailureThreshold: 1, OpenDuration: time.Minute, Now: func() time.Time { return now }})
+ cfg := config.Config{
+ Providers: []config.ProviderConfig{{ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://primary.test", APIKey: "one"}},
+ Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}}},
+ }
+ health.Observe(providerhealth.RouteKey{ModelID: "public/model", ProviderID: "primary", WireAPI: "chat_completions"}, providerhealth.Observation{StatusCode: 503, Failed: true})
+ router := New(catalog.New(cfg), health)
+ if _, err := router.PlanProvider("public/model", domain.ProtocolOpenAI, "primary"); !errors.Is(err, ErrNoHealthyRoute) {
+ t.Fatalf("open pinned provider error = %v, want ErrNoHealthyRoute", err)
+ }
+}