summaryrefslogtreecommitdiff
path: root/internal/httpapi/api_test.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 15:58:57 +1200
committerChia <Chia@93.nz>2026-08-06 15:58:57 +1200
commit3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch)
tree517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/httpapi/api_test.go
parent41e322c53d7b4b796eb377d0df9c29ecd10ba431 (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/httpapi/api_test.go')
-rw-r--r--internal/httpapi/api_test.go161
1 files changed, 157 insertions, 4 deletions
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, "<!doctype html><title>provider console</title>")
+ }))
+ 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,