diff options
| author | Chia <Chia@93.nz> | 2026-08-06 09:29:41 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 09:32:46 +1200 |
| commit | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch) | |
| tree | c730526150e55e39b822d5197e4a20318ecaa449 /internal/provider | |
| parent | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (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 '')
| -rw-r--r-- | internal/provider/forwarder.go | 56 | ||||
| -rw-r--r-- | internal/provider/forwarder_test.go | 21 | ||||
| -rw-r--r-- | internal/providerhealth/tracker.go | 200 | ||||
| -rw-r--r-- | internal/providerhealth/tracker_test.go | 50 |
4 files changed, 316 insertions, 11 deletions
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index 0d40bb1..d7d850d 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -14,6 +14,7 @@ import ( "aigw/internal/config" "aigw/internal/domain" + "aigw/internal/providerhealth" "aigw/internal/telemetry" ) @@ -26,9 +27,10 @@ type Result struct { type Forwarder struct { client *http.Client metrics *telemetry.Metrics + health *providerhealth.Tracker } -func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder { +func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics, trackers ...*providerhealth.Tracker) *Forwarder { transport := &http.Transport{ Proxy: http.ProxyFromEnvironment, DialContext: (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext, @@ -40,31 +42,41 @@ func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder { ResponseHeaderTimeout: time.Duration(cfg.ResponseHeaderTimeoutSecs) * time.Second, ExpectContinueTimeout: time.Second, } - return &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} + forwarder := &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} + if len(trackers) > 0 { + forwarder.health = trackers[0] + } + return forwarder } -func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID string, originalBody []byte, sourceHeaders http.Header, routes []domain.Route) (Result, error) { +func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID, modelID string, originalBody []byte, sourceHeaders http.Header, routes []domain.Route) (Result, error) { var lastErr error for i, route := range routes { if err := ctx.Err(); err != nil { return Result{Attempts: i}, err } - body, err := rewriteRequest(originalBody, route.UpstreamModel, protocol) + body, err := rewriteRequestWithWireAPI(originalBody, route.UpstreamModel, protocol, route.Provider.EffectiveWireAPI()) if err != nil { return Result{Attempts: i}, err } - request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider.BaseURL, protocol), bytes.NewReader(body)) + request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider, protocol), bytes.NewReader(body)) if err != nil { return Result{Attempts: i}, fmt.Errorf("build upstream request: %w", err) } setHeaders(request.Header, sourceHeaders, route.Provider, protocol, requestID) f.metrics.UpstreamAttempt() + attemptStarted := time.Now() response, err := f.client.Do(request) if err != nil { + if ctx.Err() != nil { + return Result{Attempts: i + 1}, ctx.Err() + } + f.observe(modelID, route, 0, time.Since(attemptStarted), true) lastErr = err continue } attempts := i + 1 + f.observe(modelID, route, response.StatusCode, time.Since(attemptStarted), retryableStatus(response.StatusCode)) if retryableStatus(response.StatusCode) && attempts < len(routes) { _, _ = io.CopyN(io.Discard, response.Body, 8<<10) _ = response.Body.Close() @@ -79,18 +91,34 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque return Result{Attempts: len(routes)}, lastErr } +func (f *Forwarder) observe(modelID string, route domain.Route, statusCode int, latency time.Duration, failed bool) { + if f.health == nil { + return + } + f.health.Observe(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}, + providerhealth.Observation{StatusCode: statusCode, Latency: latency, Failed: failed}) +} + func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { - return rewriteRequest(body, upstreamModel, "") + return rewriteRequest(body, upstreamModel, domain.ProtocolOpenAI) } func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) ([]byte, error) { + wireAPI := "chat_completions" + if protocol == domain.ProtocolAnthropic { + wireAPI = "messages" + } + return rewriteRequestWithWireAPI(body, upstreamModel, protocol, wireAPI) +} + +func rewriteRequestWithWireAPI(body []byte, upstreamModel string, protocol domain.Protocol, wireAPI string) ([]byte, error) { var object map[string]json.RawMessage if err := json.Unmarshal(body, &object); err != nil { return nil, fmt.Errorf("decode request body: %w", err) } encoded, _ := json.Marshal(upstreamModel) object["model"] = encoded - if protocol == domain.ProtocolOpenAI { + if protocol == domain.ProtocolOpenAI && wireAPI == "chat_completions" { var stream bool _ = json.Unmarshal(object["stream"], &stream) if stream { @@ -110,12 +138,18 @@ func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) return result, nil } -func endpointURL(baseURL string, protocol domain.Protocol) string { - baseURL = strings.TrimRight(baseURL, "/") - if protocol == domain.ProtocolAnthropic { +func endpointURL(provider domain.Provider, _ domain.Protocol) string { + baseURL := strings.TrimRight(provider.BaseURL, "/") + switch provider.EffectiveWireAPI() { + case "responses": + return baseURL + "/responses" + case "messages": return baseURL + "/messages" + case "chat_completions": + fallthrough + default: + return baseURL + "/chat/completions" } - return baseURL + "/chat/completions" } func setHeaders(target, source http.Header, provider domain.Provider, protocol domain.Protocol, requestID string) { diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go index 2e9b4a1..e9751ce 100644 --- a/internal/provider/forwarder_test.go +++ b/internal/provider/forwarder_test.go @@ -37,3 +37,24 @@ func TestRewriteRequestDoesNotAddStreamOptionsToAnthropic(t *testing.T) { t.Fatalf("unexpected OpenAI stream options in Anthropic request: %s", result) } } + +func TestResponsesWireAPIUsesResponsesEndpointWithoutChatStreamOptions(t *testing.T) { + result, err := rewriteRequestWithWireAPI([]byte(`{"model":"public/model","input":"hello","stream":true}`), "gpt-upstream", domain.ProtocolOpenAIResponses, "responses") + if err != nil { + t.Fatal(err) + } + var body map[string]json.RawMessage + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if string(body["model"]) != `"gpt-upstream"` { + t.Fatalf("model was not rewritten: %s", result) + } + if _, exists := body["stream_options"]; exists { + t.Fatalf("Responses request contains Chat Completions stream options: %s", result) + } + provider := domain.Provider{BaseURL: "https://example.test", Protocol: domain.ProtocolOpenAI, WireAPI: "responses"} + if got := endpointURL(provider, domain.ProtocolOpenAIResponses); got != "https://example.test/responses" { + t.Fatalf("endpoint URL = %q", got) + } +} diff --git a/internal/providerhealth/tracker.go b/internal/providerhealth/tracker.go new file mode 100644 index 0000000..a5af2b2 --- /dev/null +++ b/internal/providerhealth/tracker.go @@ -0,0 +1,200 @@ +package providerhealth + +import ( + "sort" + "sync" + "time" +) + +const recentWindow = 100 + +type RouteKey struct { + ModelID string + ProviderID string + WireAPI string +} + +type Observation struct { + StatusCode int + Latency time.Duration + Failed bool + ObservedAt time.Time +} + +type Status struct { + ModelID string `json:"model_id"` + ProviderID string `json:"provider_id"` + WireAPI string `json:"wire_api"` + State string `json:"state"` + Attempts uint64 `json:"attempts"` + RecentSamples int `json:"recent_samples"` + AvailabilityPercent float64 `json:"availability_percent"` + HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"` + ConsecutiveFailures uint64 `json:"consecutive_failures"` + LastStatusCode int `json:"last_status_code,omitempty"` + LastObservedAt *time.Time `json:"last_observed_at,omitempty"` + LastHealthyAt *time.Time `json:"last_healthy_at,omitempty"` + CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"` +} + +type Options struct { + FailureThreshold uint64 + OpenDuration time.Duration + Now func() time.Time +} + +type Tracker struct { + states sync.Map + failureThreshold uint64 + openDuration time.Duration + now func() time.Time +} + +type routeState struct { + mu sync.RWMutex + attempts uint64 + consecutiveFailures uint64 + lastStatusCode int + lastObservedAt time.Time + lastHealthyAt time.Time + openUntil time.Time + headerLatencyEWMA float64 + recent [recentWindow]bool + recentCount int + recentPosition int + recentHealthy int +} + +func New(options Options) *Tracker { + if options.FailureThreshold == 0 { + options.FailureThreshold = 3 + } + if options.OpenDuration <= 0 { + options.OpenDuration = 30 * time.Second + } + if options.Now == nil { + options.Now = time.Now + } + return &Tracker{failureThreshold: options.FailureThreshold, openDuration: options.OpenDuration, now: options.Now} +} + +func (t *Tracker) Observe(key RouteKey, observation Observation) { + if t == nil || key.ProviderID == "" { + return + } + if observation.ObservedAt.IsZero() { + observation.ObservedAt = t.now() + } + value, _ := t.states.LoadOrStore(key, &routeState{}) + state := value.(*routeState) + state.mu.Lock() + defer state.mu.Unlock() + + state.attempts++ + state.lastStatusCode = observation.StatusCode + state.lastObservedAt = observation.ObservedAt + if observation.Latency > 0 { + latency := float64(observation.Latency.Milliseconds()) + if latency < 1 { + latency = 1 + } + if state.headerLatencyEWMA == 0 { + state.headerLatencyEWMA = latency + } else { + state.headerLatencyEWMA = state.headerLatencyEWMA*0.8 + latency*0.2 + } + } + state.addRecent(!observation.Failed) + if observation.Failed { + state.consecutiveFailures++ + if state.consecutiveFailures >= t.failureThreshold { + state.openUntil = observation.ObservedAt.Add(t.openDuration) + } + return + } + state.consecutiveFailures = 0 + state.openUntil = time.Time{} + state.lastHealthyAt = observation.ObservedAt +} + +func (s *routeState) addRecent(healthy bool) { + if s.recentCount == recentWindow { + if s.recent[s.recentPosition] { + s.recentHealthy-- + } + } else { + s.recentCount++ + } + s.recent[s.recentPosition] = healthy + if healthy { + s.recentHealthy++ + } + s.recentPosition = (s.recentPosition + 1) % recentWindow +} + +func (t *Tracker) CircuitOpen(key RouteKey) bool { + if t == nil { + return false + } + value, ok := t.states.Load(key) + if !ok { + return false + } + state := value.(*routeState) + state.mu.RLock() + defer state.mu.RUnlock() + return state.openUntil.After(t.now()) +} + +func (t *Tracker) Snapshot() []Status { + if t == nil { + return []Status{} + } + now := t.now() + result := make([]Status, 0) + t.states.Range(func(rawKey, rawState any) bool { + key := rawKey.(RouteKey) + state := rawState.(*routeState) + state.mu.RLock() + item := statusFromState(key, state, now) + state.mu.RUnlock() + result = append(result, item) + return true + }) + sort.Slice(result, func(i, j int) bool { + if result[i].ModelID != result[j].ModelID { + return result[i].ModelID < result[j].ModelID + } + return result[i].ProviderID < result[j].ProviderID + }) + return result +} + +func statusFromState(key RouteKey, state *routeState, now time.Time) Status { + item := Status{ModelID: key.ModelID, ProviderID: key.ProviderID, WireAPI: key.WireAPI, + Attempts: state.attempts, RecentSamples: state.recentCount, HeaderLatencyEWMA: int64(state.headerLatencyEWMA + 0.5), + ConsecutiveFailures: state.consecutiveFailures, LastStatusCode: state.lastStatusCode} + if state.recentCount > 0 { + item.AvailabilityPercent = float64(state.recentHealthy) / float64(state.recentCount) * 100 + } + if !state.lastObservedAt.IsZero() { + value := state.lastObservedAt + item.LastObservedAt = &value + } + if !state.lastHealthyAt.IsZero() { + value := state.lastHealthyAt + item.LastHealthyAt = &value + } + if state.openUntil.After(now) { + value := state.openUntil + item.CircuitOpenUntil = &value + item.State = "open" + } else if state.attempts == 0 { + item.State = "unknown" + } else if state.consecutiveFailures > 0 || (state.recentCount >= 5 && item.AvailabilityPercent < 95) { + item.State = "degraded" + } else { + item.State = "healthy" + } + return item +} diff --git a/internal/providerhealth/tracker_test.go b/internal/providerhealth/tracker_test.go new file mode 100644 index 0000000..73bd6d5 --- /dev/null +++ b/internal/providerhealth/tracker_test.go @@ -0,0 +1,50 @@ +package providerhealth + +import ( + "sync" + "testing" + "time" +) + +func TestTrackerOpensAndRecoversCircuit(t *testing.T) { + now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) + tracker := New(Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }}) + key := RouteKey{ModelID: "openai/test", ProviderID: "primary", WireAPI: "responses"} + for _, status := range []int{503, 429, 504} { + tracker.Observe(key, Observation{StatusCode: status, Latency: 100 * time.Millisecond, Failed: true}) + } + if !tracker.CircuitOpen(key) { + t.Fatal("circuit did not open after consecutive retryable failures") + } + snapshot := tracker.Snapshot() + if len(snapshot) != 1 || snapshot[0].State != "open" || snapshot[0].AvailabilityPercent != 0 || snapshot[0].CircuitOpenUntil == nil { + t.Fatalf("unexpected open snapshot: %+v", snapshot) + } + now = now.Add(31 * time.Second) + if tracker.CircuitOpen(key) { + t.Fatal("circuit did not permit a recovery attempt after cooldown") + } + tracker.Observe(key, Observation{StatusCode: 200, Latency: 50 * time.Millisecond}) + snapshot = tracker.Snapshot() + if snapshot[0].State != "healthy" || snapshot[0].ConsecutiveFailures != 0 || snapshot[0].LastHealthyAt == nil || snapshot[0].HeaderLatencyEWMA < 50 { + t.Fatalf("unexpected recovered snapshot: %+v", snapshot[0]) + } +} + +func TestTrackerConcurrentObservations(t *testing.T) { + tracker := New(Options{FailureThreshold: 1000}) + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "chat_completions"} + var group sync.WaitGroup + for index := range 200 { + group.Add(1) + go func(failed bool) { + defer group.Done() + tracker.Observe(key, Observation{StatusCode: 200, Latency: time.Millisecond, Failed: failed}) + }(index%2 == 0) + } + group.Wait() + snapshot := tracker.Snapshot() + if len(snapshot) != 1 || snapshot[0].Attempts != 200 || snapshot[0].RecentSamples != recentWindow || snapshot[0].AvailabilityPercent < 0 || snapshot[0].AvailabilityPercent > 100 { + t.Fatalf("unexpected concurrent snapshot: %+v", snapshot) + } +} |
