summaryrefslogtreecommitdiff
path: root/internal/provider
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/provider
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 '')
-rw-r--r--internal/provider/forwarder.go56
-rw-r--r--internal/provider/forwarder_test.go21
-rw-r--r--internal/providerhealth/tracker.go200
-rw-r--r--internal/providerhealth/tracker_test.go50
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)
+ }
+}