package gateway import ( "context" "encoding/json" "errors" "io" "net/http" "net/http/httptest" "strings" "testing" ) // serve stands up a fake OpenAI-compatible endpoint and returns a gateway // pointed at it, plus a pointer to the last request body it received. func serve(t *testing.T, handler func(w http.ResponseWriter, body *oaiRequest)) (*OpenAIGateway, *oaiRequest) { t.Helper() var captured oaiRequest srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { raw, _ := io.ReadAll(r.Body) if err := json.Unmarshal(raw, &captured); err != nil { t.Errorf("request body was not valid JSON: %v", err) } handler(w, &captured) })) t.Cleanup(srv.Close) gw := NewOpenAI(Config{ Provider: ProviderOpenAI, APIKey: "test-key", BaseURL: srv.URL, Fast: Routing{Model: "m-fast", Effort: EffortLow}, Balanced: Routing{Model: "m-balanced", Effort: EffortHigh}, Deep: Routing{Model: "m-deep", Effort: EffortXhigh}, MaxOutputTokens: 4096, }) return gw, &captured } func ask(text string) Request { return Request{Tier: TierBalanced, Messages: []Message{{Role: RoleUser, Text: text}}} } // THE REGRESSION THIS FILE EXISTS FOR. // // OpenAI reports prompt_tokens INCLUSIVE of the cached prefix; Anthropic // reports input tokens EXCLUSIVE of it. Usage.Total() adds all four fields, so // copying both numbers across verbatim bills the cached prefix twice — and it // does it worst on long conversations, which is exactly where I3's budget // matters most. A wrong total here is invisible: the run still answers, it just // terminates BudgetExceeded earlier than it should. func TestUsageDoesNotDoubleCountCachedTokens(t *testing.T) { usage := oaiUsage{PromptTokens: 1000, CompletionTokens: 200} usage.PromptTokensDetails.CachedTokens = 800 got := usage.normalise() if got.InputTokens != 200 { t.Errorf("InputTokens = %d, want 200 (1000 prompt less 800 cached)", got.InputTokens) } if got.CacheReadTokens != 800 { t.Errorf("CacheReadTokens = %d, want 800", got.CacheReadTokens) } if got.Total() != 1200 { t.Errorf("Total() = %d, want 1200 — the wire billed 1000 prompt + 200 output, "+ "and anything higher is the cached prefix counted twice", got.Total()) } } // A provider reporting more cached tokens than prompt tokens is wrong, but the // failure must not hand the run free budget: a negative charge would reduce the // total, which is the one direction a bug must never go. func TestUsageClampsImpossibleCacheReport(t *testing.T) { usage := oaiUsage{PromptTokens: 100, CompletionTokens: 10} usage.PromptTokensDetails.CachedTokens = 500 got := usage.normalise() if got.InputTokens < 0 { t.Fatalf("InputTokens = %d, want no negative charge", got.InputTokens) } if got.Total() < got.OutputTokens { t.Errorf("Total() = %d is below OutputTokens = %d", got.Total(), got.OutputTokens) } } // Tool results are blocks inside one user turn on the Anthropic wire and // standalone role:"tool" messages here. Getting the split wrong detaches a // result from the call it answers, which most providers reject outright and // some silently mis-attribute. func TestEncodeMessagesSplitsToolResults(t *testing.T) { msgs := []Message{ {Role: RoleUser, Text: "who is free friday?"}, {Role: RoleAssistant, ToolCalls: []ToolCall{ {ID: "call_1", Name: "find_workers", Input: json.RawMessage(`{"day":"friday"}`)}, {ID: "call_2", Name: "open_shifts", Input: json.RawMessage(`{}`)}, }}, {Role: RoleUser, ToolResults: []ToolResult{ {CallID: "call_1", Content: `{"workers":3}`}, {CallID: "call_2", Content: `{"shifts":1}`}, }}, } got := encodeOpenAIMessages("you are a scheduler", msgs) wantRoles := []string{"system", "user", "assistant", "tool", "tool"} if len(got) != len(wantRoles) { t.Fatalf("got %d messages, want %d: %+v", len(got), len(wantRoles), got) } for i, want := range wantRoles { if got[i].Role != want { t.Errorf("messages[%d].Role = %q, want %q", i, got[i].Role, want) } } if got[0].Content != "you are a scheduler" { t.Errorf("system message = %q", got[0].Content) } if len(got[2].ToolCalls) != 2 { t.Fatalf("assistant turn carried %d tool calls, want 2", len(got[2].ToolCalls)) } // The call id is the model's own handle. A result carrying a different one // is a result attached to the wrong question. if got[3].ToolCallID != "call_1" || got[4].ToolCallID != "call_2" { t.Errorf("tool results correlated to %q and %q, want call_1 and call_2", got[3].ToolCallID, got[4].ToolCallID) } } // A turn that is only tool results carries no text, and dropping it would strip // every answer the tools produced. func TestEncodeMessagesKeepsResultOnlyTurn(t *testing.T) { got := encodeOpenAIMessages("", []Message{ {Role: RoleUser, Text: "hi"}, {Role: RoleUser, ToolResults: []ToolResult{{CallID: "c1", Content: "{}"}}}, }) if len(got) != 2 || got[1].Role != "tool" { t.Fatalf("result-only turn was not encoded: %+v", got) } } func TestCompleteDecodesTextAndUsage(t *testing.T) { gw, captured := serve(t, func(w http.ResponseWriter, _ *oaiRequest) { _, _ = io.WriteString(w, `{ "model":"m-balanced-0625", "choices":[{"message":{"role":"assistant","content":"Three are free."}, "finish_reason":"stop"}], "usage":{"prompt_tokens":120,"completion_tokens":8} }`) }) resp, err := gw.Complete(context.Background(), ask("who is free?")) if err != nil { t.Fatalf("Complete: %v", err) } if resp.Text != "Three are free." { t.Errorf("Text = %q", resp.Text) } // The id ACTUALLY used, not the tier that was asked for — a change of // routing has to be visible in the trajectory rather than inferred. if resp.Model != "m-balanced-0625" { t.Errorf("Model = %q, want the id the provider reported", resp.Model) } if resp.StopReason != "end_turn" { t.Errorf("StopReason = %q, want end_turn", resp.StopReason) } if resp.Usage.Total() != 128 { t.Errorf("Usage.Total() = %d, want 128", resp.Usage.Total()) } if captured.Model != "m-balanced" { t.Errorf("requested model = %q, want the balanced tier's", captured.Model) } } func TestCompleteDecodesToolCalls(t *testing.T) { gw, _ := serve(t, func(w http.ResponseWriter, _ *oaiRequest) { _, _ = io.WriteString(w, `{ "choices":[{"message":{"role":"assistant","tool_calls":[ {"id":"call_x","type":"function", "function":{"name":"find_workers","arguments":"{\"day\":\"friday\"}"}}]}, "finish_reason":"tool_calls"}], "usage":{"prompt_tokens":10,"completion_tokens":5} }`) }) resp, err := gw.Complete(context.Background(), ask("who is free?")) if err != nil { t.Fatalf("Complete: %v", err) } if len(resp.ToolCalls) != 1 { t.Fatalf("got %d tool calls, want 1", len(resp.ToolCalls)) } call := resp.ToolCalls[0] if call.ID != "call_x" || call.Name != "find_workers" { t.Errorf("call = %+v", call) } // The loop branches on len(ToolCalls), but the trajectory records the stop // reason, and it has to read the same as the Anthropic path's. if resp.StopReason != "tool_use" { t.Errorf("StopReason = %q, want tool_use", resp.StopReason) } var args map[string]string if err := json.Unmarshal(call.Input, &args); err != nil { t.Fatalf("tool input was not valid JSON: %v", err) } if args["day"] != "friday" { t.Errorf("args = %v", args) } } // An argumentless call arrives as "" on this wire, which is not valid JSON. The // handler's decoder would reject it for a reason that has nothing to do with // the request. func TestEmptyToolArgumentsBecomeEmptyObject(t *testing.T) { gw, _ := serve(t, func(w http.ResponseWriter, _ *oaiRequest) { _, _ = io.WriteString(w, `{"choices":[{"message":{"tool_calls":[ {"id":"c1","function":{"name":"workspace_summary","arguments":""}}]}, "finish_reason":"tool_calls"}]}`) }) resp, err := gw.Complete(context.Background(), ask("summarise")) if err != nil { t.Fatalf("Complete: %v", err) } if string(resp.ToolCalls[0].Input) != "{}" { t.Errorf("Input = %q, want {}", resp.ToolCalls[0].Input) } } // A refusal is a successful HTTP response and one of the six terminations. It // is still billed: a refusal that cost nothing on the ledger is one the loop // would happily repeat. func TestRefusalIsStructuredAndStillBilled(t *testing.T) { gw, _ := serve(t, func(w http.ResponseWriter, _ *oaiRequest) { _, _ = io.WriteString(w, `{"choices":[{"message":{"role":"assistant", "refusal":"I cannot help with that."},"finish_reason":"stop"}], "usage":{"prompt_tokens":50,"completion_tokens":6}}`) }) resp, err := gw.Complete(context.Background(), ask("do something disallowed")) var gwErr *Error if !errors.As(err, &gwErr) || gwErr.Code != CodeRefused { t.Fatalf("err = %v, want a %s", err, CodeRefused) } if gwErr.Retryable() { t.Error("a refusal must not be retryable — re-sending it burns the budget on one turn") } if resp == nil { t.Fatal("a refusal must still carry its usage") } if resp.Usage.Total() != 56 { t.Errorf("Usage.Total() = %d, want 56", resp.Usage.Total()) } } func TestErrorsMapToRetryability(t *testing.T) { cases := []struct { status int wantCode string retryable bool }{ {400, CodeInvalidRequest, false}, // A model id that does not exist on this endpoint is a configuration // mistake and will fail identically next time. {404, CodeInvalidRequest, false}, {401, CodeUnauthorized, false}, {429, CodeRateLimited, true}, {500, CodeUpstream, true}, {503, CodeUpstream, true}, } for _, c := range cases { err := translateOpenAI(c.status, []byte(`{"error":{"message":"upstream detail"}}`)) var gwErr *Error if !errors.As(err, &gwErr) { t.Fatalf("http %d: not a gateway error", c.status) } if gwErr.Code != c.wantCode { t.Errorf("http %d: code = %s, want %s", c.status, gwErr.Code, c.wantCode) } if gwErr.Retryable() != c.retryable { t.Errorf("http %d: Retryable() = %v, want %v", c.status, gwErr.Retryable(), c.retryable) } // The upstream reason has to survive: the trajectory records only the // message, and "the model call failed" costs an hour to diagnose. if !strings.Contains(gwErr.Message, "upstream detail") { t.Errorf("http %d: message %q dropped the upstream detail", c.status, gwErr.Message) } } } // Most non-reasoning models reject the whole request rather than ignoring an // unknown key, so the field must be absent unless a deployment opted in. func TestReasoningEffortIsOptIn(t *testing.T) { gw, captured := serve(t, func(w http.ResponseWriter, _ *oaiRequest) { _, _ = io.WriteString(w, `{"choices":[{"message":{"content":"ok"},"finish_reason":"stop"}]}`) }) if _, err := gw.Complete(context.Background(), ask("hi")); err != nil { t.Fatalf("Complete: %v", err) } if captured.ReasoningEffort != "" { t.Errorf("reasoning_effort = %q, want it omitted by default", captured.ReasoningEffort) } gw.cfg.SendReasoningEffort = true if _, err := gw.Complete(context.Background(), Request{ Tier: TierDeep, Messages: []Message{{Role: RoleUser, Text: "hi"}}, }); err != nil { t.Fatalf("Complete: %v", err) } // Ordering preserved, not spelling: their scale runs minimal/low/medium/ // high, so the platform's xhigh is their high. if captured.ReasoningEffort != "high" { t.Errorf("deep tier sent reasoning_effort = %q, want high", captured.ReasoningEffort) } } // A local model needs no credential. Requiring one would make the zero-cost // development path impossible to configure. func TestLocalEndpointNeedsNoCredential(t *testing.T) { local := NewOpenAI(Config{BaseURL: "http://localhost:11434/v1"}) if local.needsCredential() { t.Error("a localhost endpoint must not require a key") } hosted := NewOpenAI(Config{BaseURL: "https://api.groq.com/openai/v1"}) if !hosted.needsCredential() { t.Error("a hosted endpoint must require a key") } if _, err := NewOpenAI(Config{BaseURL: "https://api.groq.com/openai/v1"}). Complete(context.Background(), ask("hi")); err == nil { t.Error("a hosted call without a key must fail as NotConfigured") } } func TestBaseURLDefaultsAndTrimsSlash(t *testing.T) { if got := NewOpenAI(Config{}).endpoint(); got != DefaultOpenAIBaseURL+"/chat/completions" { t.Errorf("endpoint = %q", got) } if got := NewOpenAI(Config{BaseURL: "https://x.test/v1/"}).endpoint(); got != "https://x.test/v1/chat/completions" { t.Errorf("endpoint = %q, want the trailing slash collapsed", got) } }