summaryrefslogtreecommitdiff
path: root/internal/httpapi/api_test.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/httpapi/api_test.go224
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
+}