From 3f702084d20b3c3a3ea916f3110e99b22bda60b3 Mon Sep 17 00:00:00 2001 From: Chia Date: Thu, 6 Aug 2026 15:58:57 +1200 Subject: feat: complete commercial developer workflows 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. --- internal/httpapi/api_test.go | 161 +++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 157 insertions(+), 4 deletions(-) (limited to 'internal/httpapi/api_test.go') diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go index cbb3967..906e136 100644 --- a/internal/httpapi/api_test.go +++ b/internal/httpapi/api_test.go @@ -74,10 +74,14 @@ func TestInferenceBrowserOriginCORS(t *testing.T) { type fakeBillingMeter struct { authorizeErr error + authorized chan billing.Authorization settled chan domain.UsageEvent } -func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) error { +func (m *fakeBillingMeter) Authorize(_ context.Context, authorization billing.Authorization) error { + if m.authorized != nil { + m.authorized <- authorization + } return m.authorizeErr } @@ -132,7 +136,7 @@ func TestOpenAIProxyRewritesModelAndEmitsUsage(t *testing.T) { select { case event := <-sink.events: - if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || !event.Success { + if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || event.TTFTMS < 1 || !event.Success { t.Fatalf("unexpected usage event: %+v", event) } case <-time.After(time.Second): @@ -178,6 +182,53 @@ func TestProxyFailsOverBeforeWritingResponse(t *testing.T) { } } +func TestGatewayLearnsTTFTAndPrefersFasterProvider(t *testing.T) { + var fastCalls atomic.Int64 + fast := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fastCalls.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer fast.Close() + var slowCalls atomic.Int64 + slow := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + slowCalls.Add(1) + time.Sleep(25 * time.Millisecond) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer slow.Close() + + gateway, sink := newTestGateway(t, + []config.ProviderConfig{ + {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: fast.URL + "/v1", APIKey: "one"}, + {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: slow.URL + "/v1", APIKey: "two"}, + }, + []config.RouteConfig{ + {Provider: "fast", UpstreamModel: "model", Weight: 1}, + {Provider: "slow", UpstreamModel: "model", Weight: 1}, + }, + ) + defer gateway.Close() + + for range 20 { + response := postOpenAI(t, gateway.URL, false) + _, _ = io.Copy(io.Discard, response.Body) + _ = response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("unexpected gateway status: %d", response.StatusCode) + } + select { + case <-sink.events: + case <-time.After(time.Second): + t.Fatal("usage event was not emitted") + } + } + if fastCalls.Load() != 16 || slowCalls.Load() != 4 { + t.Fatalf("TTFT feedback was not applied with bounded exploration: fast=%d slow=%d", fastCalls.Load(), slowCalls.Load()) + } +} + func TestProviderSelectorPinsRouteAndKeepsCanonicalUsageModel(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -465,11 +516,106 @@ func TestOpenAIResponsesProxyRewritesModelAndEmitsUsage(t *testing.T) { t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) } event := <-sink.events - if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || !event.UsageReported { + if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || event.TTFTMS < 1 || !event.UsageReported { t.Fatalf("unexpected Responses usage event: %+v", event) } } +func TestOpenAIEmbeddingsProxyRewritesModelAndMetersInputTokens(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/embeddings" { + t.Errorf("unexpected path: %s", r.URL.Path) + } + if r.Header.Get("Authorization") != "Bearer upstream-secret" { + t.Errorf("unexpected upstream authorization: %q", r.Header.Get("Authorization")) + } + var request map[string]any + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + if request["model"] != "embedding-upstream" || request["input"] != "hello vector" { + t.Errorf("unexpected Embeddings request: %+v", request) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"object":"list","data":[{"object":"embedding","index":0,"embedding":[0.1,0.2]}],"model":"embedding-upstream","usage":{"prompt_tokens":8,"total_tokens":8}}`) + })) + defer upstream.Close() + + meter := &fakeBillingMeter{authorized: make(chan billing.Authorization, 1), settled: make(chan domain.UsageEvent, 1)} + gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ + ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret", + }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, meter) + defer gateway.Close() + + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello vector"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + payload, _ := io.ReadAll(response.Body) + t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) + } + authorization := <-meter.authorized + if authorization.Protocol != domain.ProtocolOpenAIEmbeddings { + t.Fatalf("authorization protocol = %q", authorization.Protocol) + } + event := <-meter.settled + if event.Protocol != domain.ProtocolOpenAIEmbeddings || event.UpstreamModel != "embedding-upstream" || + event.Usage.InputTokens != 8 || event.Usage.OutputTokens != 0 || event.Usage.TotalTokens != 8 || !event.UsageReported { + t.Fatalf("unexpected Embeddings usage event: %+v", event) + } +} + +func TestOpenAIEmbeddingsRequiresDeclaredCapability(t *testing.T) { + providerConfig := config.ProviderConfig{ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "secret"} + modelCatalog := catalog.New(config.Config{Providers: []config.ProviderConfig{providerConfig}, Models: []config.ModelConfig{{ + ID: "public/model", Routes: []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, + }}}) + authenticator, err := auth.NewStatic(`[{"key":"client-secret","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`, false) + if err != nil { + t.Fatal(err) + } + api := New(Options{Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog), Metrics: &telemetry.Metrics{}, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), MaxBodyBytes: 1 << 20}) + request := httptest.NewRequest(http.MethodPost, "/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response := httptest.NewRecorder() + api.Handler().ServeHTTP(response, request) + if response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), "unsupported_capability") { + t.Fatalf("unexpected response %d: %s", response.Code, response.Body.String()) + } +} + +func TestOpenAIEmbeddingsRejectsNonJSONSuccessBeforeResponseStarts(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = io.WriteString(w, "provider console") + })) + defer upstream.Close() + meter := &fakeBillingMeter{settled: make(chan domain.UsageEvent, 1)} + gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ + ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL, APIKey: "secret", + }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}}, meter) + defer gateway.Close() + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, _ := io.ReadAll(response.Body) + if response.StatusCode != http.StatusBadGateway || !strings.Contains(string(body), "invalid_provider_response") || strings.Contains(string(body), "provider console") { + t.Fatalf("unexpected response %d: %s", response.StatusCode, body) + } + event := <-meter.settled + if event.Success || event.StatusCode != http.StatusBadGateway || event.ErrorType != "invalid_provider_response" || event.UsageReported { + t.Fatalf("unexpected invalid provider usage event: %+v", event) + } +} + func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -542,9 +688,16 @@ func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig, if err != nil { t.Fatal(err) } + modelConfig := config.ModelConfig{ID: "public/model", OwnedBy: "test", Routes: routes} + for _, providerConfig := range providers { + if providerConfig.WireAPI == "embeddings" { + modelConfig.Capabilities = []string{"embeddings"} + break + } + } cfg := config.Config{ Providers: providers, - Models: []config.ModelConfig{{ID: "public/model", OwnedBy: "test", Routes: routes}}, + Models: []config.ModelConfig{modelConfig}, UpstreamHTTP: config.UpstreamHTTPConfig{ MaxIdleConnections: 100, MaxIdleConnectionsPerHost: 20, IdleConnectionTimeoutSecs: 10, ResponseHeaderTimeoutSecs: 2, -- cgit v1.2.3