diff options
| author | Chia <Chia@93.nz> | 2026-08-04 19:58:52 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-04 20:43:23 +1200 |
| commit | 5b651488b081b65fda8a323f228e139adb79a35d (patch) | |
| tree | 08baf40efb8fe103b32721cd991ff712323e3173 /internal/httpapi/api_test.go | |
Build AI gateway control plane and admin UI
Diffstat (limited to '')
| -rw-r--r-- | internal/httpapi/api_test.go | 224 |
1 files changed, 224 insertions, 0 deletions
diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go new file mode 100644 index 0000000..66a344b --- /dev/null +++ b/internal/httpapi/api_test.go @@ -0,0 +1,224 @@ +package httpapi + +import ( + "bufio" + "bytes" + "encoding/json" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + "aigw/internal/auth" + "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 +} + +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 newTestGateway(t *testing.T, providers []config.ProviderConfig, routes []config.RouteConfig) (*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, + 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 +} |
