package httpapi import ( "bufio" "bytes" "context" "encoding/json" "io" "log/slog" "net/http" "net/http/httptest" "strings" "sync/atomic" "testing" "time" "aigw/internal/auth" "aigw/internal/billing" "aigw/internal/catalog" "aigw/internal/config" "aigw/internal/domain" "aigw/internal/provider" "aigw/internal/routing" "aigw/internal/telemetry" ) type captureUsageSink struct { events chan domain.UsageEvent } type fakeBillingMeter struct { authorizeErr error settled chan domain.UsageEvent } func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) error { return m.authorizeErr } func (m *fakeBillingMeter) Settle(_ context.Context, event domain.UsageEvent) error { m.settled <- event return nil } func (s *captureUsageSink) Publish(event domain.UsageEvent) { s.events <- event } func TestOpenAIProxyRewritesModelAndEmitsUsage(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/v1/chat/completions" { t.Errorf("unexpected path: %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer upstream-secret" { t.Errorf("upstream authorization leaked or missing: %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"] != "upstream-model" { t.Errorf("model was not rewritten: %+v", request) } w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"id":"chat-1","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`) })) defer upstream.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "upstream-model", Weight: 1}}) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model","messages":[{"role":"user","content":"hello"}]}`)) 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) } if response.Header.Get("X-AIGW-Request-ID") == "" { t.Fatal("missing gateway request id") } select { case event := <-sink.events: if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || !event.Success { t.Fatalf("unexpected usage event: %+v", event) } case <-time.After(time.Second): t.Fatal("usage event was not emitted") } } func TestProxyFailsOverBeforeWritingResponse(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { primaryCalls.Add(1) w.WriteHeader(http.StatusServiceUnavailable) })) defer primary.Close() var fallbackCalls atomic.Int64 fallback := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fallbackCalls.Add(1) w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) })) defer fallback.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{ {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: primary.URL + "/v1", APIKey: "one"}, {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: fallback.URL + "/v1", APIKey: "two"}, }, []config.RouteConfig{ {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1}, {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1}, }, ) defer gateway.Close() response := postOpenAI(t, gateway.URL, false) defer response.Body.Close() if response.StatusCode != http.StatusOK || primaryCalls.Load() != 1 || fallbackCalls.Load() != 1 { t.Fatalf("failover did not complete: status=%d primary=%d fallback=%d", response.StatusCode, primaryCalls.Load(), fallbackCalls.Load()) } event := <-sink.events if event.Attempts != 2 || event.ProviderID != "fallback" { t.Fatalf("unexpected failover event: %+v", event) } } func TestSSEIsFlushedBeforeUpstreamCompletes(t *testing.T) { release := make(chan struct{}) upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "text/event-stream") _, _ = io.WriteString(w, "data: {\"id\":\"chunk-1\",\"choices\":[]}\n\n") w.(http.Flusher).Flush() <-release _, _ = io.WriteString(w, "data: [DONE]\n\n") })) defer upstream.Close() gateway, _ := newTestGateway(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}) defer gateway.Close() response := postOpenAI(t, gateway.URL, true) defer response.Body.Close() reader := bufio.NewReader(response.Body) line, err := reader.ReadString('\n') if err != nil { close(release) t.Fatal(err) } if !strings.Contains(line, "chunk-1") { close(release) t.Fatalf("unexpected first SSE line: %q", line) } close(release) } func TestAnthropicHeadersAndPath(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/v1/messages" || r.Header.Get("x-api-key") != "anthropic-upstream" || r.Header.Get("anthropic-version") != "2023-06-01" { t.Errorf("unexpected anthropic request: path=%s key=%q version=%q", r.URL.Path, r.Header.Get("x-api-key"), r.Header.Get("anthropic-version")) } _, _ = io.WriteString(w, `{"type":"message","content":[],"usage":{"input_tokens":4,"output_tokens":6}}`) })) defer upstream.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{{ ID: "anthropic", Protocol: domain.ProtocolAnthropic, BaseURL: upstream.URL + "/v1", APIKey: "anthropic-upstream", }}, []config.RouteConfig{{Provider: "anthropic", UpstreamModel: "claude-upstream", Weight: 1}}) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/anthropic/v1/messages", strings.NewReader(`{"model":"public/model","max_tokens":10,"messages":[{"role":"user","content":"hello"}]}`)) request.Header.Set("x-api-key", "client-secret") request.Header.Set("anthropic-version", "2023-06-01") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() if response.StatusCode != http.StatusOK { t.Fatalf("unexpected status: %d", response.StatusCode) } event := <-sink.events if event.Usage.InputTokens != 4 || event.Usage.OutputTokens != 6 { t.Fatalf("unexpected anthropic usage: %+v", event.Usage) } } func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { calls.Add(1) w.WriteHeader(http.StatusOK) })) defer upstream.Close() meter := &fakeBillingMeter{authorizeErr: billing.ErrInsufficientBalance, settled: make(chan domain.UsageEvent, 1)} gateway, sink := newTestGatewayWithBilling(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}, meter) defer gateway.Close() response := postOpenAI(t, gateway.URL, false) defer response.Body.Close() if response.StatusCode != http.StatusPaymentRequired { t.Fatalf("status = %d, want 402", response.StatusCode) } if calls.Load() != 0 { t.Fatalf("upstream calls = %d, want 0", calls.Load()) } select { case <-meter.settled: t.Fatal("rejected request was settled") default: } select { case event := <-sink.events: if event.StatusCode != http.StatusPaymentRequired || event.ErrorType != "insufficient_balance" { t.Fatalf("unexpected rejection usage: %+v", event) } case <-time.After(time.Second): t.Fatal("rejected request did not emit usage") } } func TestSuccessfulRequestIsSettled(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":2,"completion_tokens":3,"total_tokens":5}}`) })) defer upstream.Close() meter := &fakeBillingMeter{settled: make(chan domain.UsageEvent, 1)} gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}, meter) defer gateway.Close() response := postOpenAI(t, gateway.URL, false) defer response.Body.Close() if response.StatusCode != http.StatusOK { t.Fatalf("status = %d, want 200", response.StatusCode) } select { case event := <-meter.settled: if event.Usage.TotalTokens != 5 { t.Fatalf("settled usage = %+v", event.Usage) } case <-time.After(time.Second): t.Fatal("request was not settled") } } func newTestGateway(t *testing.T, providers []config.ProviderConfig, routes []config.RouteConfig) (*httptest.Server, *captureUsageSink) { return newTestGatewayWithBilling(t, providers, routes, nil) } func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig, routes []config.RouteConfig, meter billing.Meter) (*httptest.Server, *captureUsageSink) { t.Helper() 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) } cfg := config.Config{ Providers: providers, Models: []config.ModelConfig{{ID: "public/model", OwnedBy: "test", Routes: routes}}, UpstreamHTTP: config.UpstreamHTTPConfig{ MaxIdleConnections: 100, MaxIdleConnectionsPerHost: 20, IdleConnectionTimeoutSecs: 10, ResponseHeaderTimeoutSecs: 2, }, } metrics := &telemetry.Metrics{} modelCatalog := catalog.New(cfg) sink := &captureUsageSink{events: make(chan domain.UsageEvent, 10)} api := New(Options{ Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog), Forwarder: provider.New(cfg.UpstreamHTTP, metrics), UsageSink: sink, BillingMeter: meter, Metrics: metrics, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), MaxBodyBytes: 1 << 20, }) return httptest.NewServer(api.Handler()), sink } func postOpenAI(t *testing.T, gatewayURL string, stream bool) *http.Response { t.Helper() payload := []byte(`{"model":"public/model","messages":[{"role":"user","content":"hello"}],"stream":false}`) if stream { payload = []byte(`{"model":"public/model","messages":[{"role":"user","content":"hello"}],"stream":true}`) } request, _ := http.NewRequest(http.MethodPost, gatewayURL+"/v1/chat/completions", bytes.NewReader(payload)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } return response }