diff options
| author | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
| commit | 3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch) | |
| tree | 517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/routing | |
| parent | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (diff) | |
feat: complete commercial developer workflowspublish-commercial-control-plane
Add tenant-safe usage observability, prepaid billing controls, API key lifecycle management, Embeddings metering, configurable billing alerts, and resilient provider health propagation. Harden Stripe failure handling, migrations, readiness, and the authenticated control-plane UI with end-to-end verification evidence.
Diffstat (limited to 'internal/routing')
| -rw-r--r-- | internal/routing/router.go | 123 | ||||
| -rw-r--r-- | internal/routing/router_test.go | 45 |
2 files changed, 155 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 +} diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go index 2ad4685..f421184 100644 --- a/internal/routing/router_test.go +++ b/internal/routing/router_test.go @@ -96,15 +96,56 @@ func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { } } +func TestPlanPrefersLowerTTFTAndStillExplores(t *testing.T) { + health := providerhealth.New(providerhealth.Options{}) + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: "https://fast.test", APIKey: "one"}, + {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: "https://slow.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "fast", UpstreamModel: "model", Weight: 1}, + {Provider: "slow", UpstreamModel: "model", Weight: 1}, + }}}, + } + for _, providerID := range []string{"fast", "slow"} { + key := providerhealth.RouteKey{ModelID: "public/model", ProviderID: providerID, WireAPI: "chat_completions"} + for range 5 { + health.Observe(key, providerhealth.Observation{StatusCode: 200, Latency: 20 * time.Millisecond}) + } + latency := 100 * time.Millisecond + if providerID == "slow" { + latency = 500 * time.Millisecond + } + for range 3 { + health.ObserveTTFT(key, latency) + } + } + router := New(catalog.New(cfg), health) + counts := map[string]int{} + for range 200 { + plan, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil { + t.Fatal(err) + } + counts[plan[0].Provider.ID]++ + } + if counts["fast"] != 190 || counts["slow"] != 10 { + t.Fatalf("adaptive selection must prefer fast route while preserving 5%% exploration: %+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"}, + {ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://embeddings.test", APIKey: "three"}, }, Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ {Provider: "chat", UpstreamModel: "chat-model", Weight: 1}, {Provider: "responses", UpstreamModel: "responses-model", Weight: 1}, + {Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}, }}}, } router := New(catalog.New(cfg)) @@ -116,6 +157,10 @@ func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) { if err != nil || len(responses) != 1 || responses[0].Provider.ID != "responses" { t.Fatalf("unexpected Responses plan: %+v err=%v", responses, err) } + embeddings, err := router.Plan("public/model", domain.ProtocolOpenAIEmbeddings) + if err != nil || len(embeddings) != 1 || embeddings[0].Provider.ID != "embeddings" { + t.Fatalf("unexpected Embeddings plan: %+v err=%v", embeddings, err) + } } func TestPlanProviderPinsWithoutFallbackToOtherProviders(t *testing.T) { |
