From 41e322c53d7b4b796eb377d0df9c29ecd10ba431 Mon Sep 17 00:00:00 2001 From: Chia Date: Thu, 6 Aug 2026 09:29:41 +1200 Subject: 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 --- internal/provider/forwarder.go | 56 +++++++++++++++++++++++++++++-------- internal/provider/forwarder_test.go | 21 ++++++++++++++ 2 files changed, 66 insertions(+), 11 deletions(-) (limited to 'internal/provider') 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) + } +} -- cgit v1.2.3