package gateway import ( "context" "encoding/json" "errors" "io" "net/http" "net/http/httptest" "strings" "testing" ) // sse stands up an endpoint that replays the given event lines. func sse(t *testing.T, events ...string) *OpenAIGateway { t.Helper() srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") for _, e := range events { _, _ = io.WriteString(w, e+"\n") } })) t.Cleanup(srv.Close) return NewOpenAI(Config{ APIKey: "test-key", BaseURL: srv.URL, Balanced: Routing{Model: "m-balanced", Effort: EffortHigh}, }) } func TestStreamDeliversTextAsItArrives(t *testing.T) { gw := sse(t, `data: {"model":"m-1","choices":[{"delta":{"content":"Three "}}]}`, `data: {"choices":[{"delta":{"content":"are "}}]}`, `data: {"choices":[{"delta":{"content":"free."},"finish_reason":"stop"}]}`, `data: {"choices":[],"usage":{"prompt_tokens":40,"completion_tokens":4}}`, `data: [DONE]`, ) var deltas []string resp, err := gw.Stream(context.Background(), ask("who is free?"), func(d string) { deltas = append(deltas, d) }) if err != nil { t.Fatalf("Stream: %v", err) } if strings.Join(deltas, "") != "Three are free." { t.Errorf("deltas joined to %q", strings.Join(deltas, "")) } if len(deltas) != 3 { t.Errorf("got %d deltas, want 3 — text must arrive in fragments, not in one lump", len(deltas)) } if resp.Text != "Three are free." { t.Errorf("Text = %q", resp.Text) } // Usage arrives in a trailing chunk with no choices. Missing it would mean // a streamed run cost nothing on the ledger, and I3 cannot enforce a budget // it cannot measure. if resp.Usage.Total() != 44 { t.Errorf("Usage.Total() = %d, want 44 — the trailing usage chunk was dropped", resp.Usage.Total()) } if resp.StopReason != "end_turn" { t.Errorf("StopReason = %q", resp.StopReason) } } // THE ONE THAT IS EASY TO GET WRONG. // // Providers interleave the fragments of parallel tool calls, so arrival order // is not call order. Appending fragments as they land splices one call's // arguments onto another's — producing two calls that are each valid JSON and // both wrong, which is the worst possible failure: the tools run, with the // wrong inputs, and nothing errors. func TestStreamAccumulatesInterleavedToolCallsByIndex(t *testing.T) { gw := sse(t, `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_a","function":{"name":"find_workers","arguments":"{\"day\""}}]}}]}`, `data: {"choices":[{"delta":{"tool_calls":[{"index":1,"id":"call_b","function":{"name":"open_shifts","arguments":"{\"week\""}}]}}]}`, `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":":\"friday\"}"}}]}}]}`, `data: {"choices":[{"delta":{"tool_calls":[{"index":1,"function":{"arguments":":\"next\"}"}}]}}]}`, `data: {"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`, `data: [DONE]`, ) resp, err := gw.Stream(context.Background(), ask("cover friday"), nil) if err != nil { t.Fatalf("Stream: %v", err) } if len(resp.ToolCalls) != 2 { t.Fatalf("got %d tool calls, want 2: %+v", len(resp.ToolCalls), resp.ToolCalls) } want := []struct{ id, name, day string }{ {"call_a", "find_workers", "friday"}, {"call_b", "open_shifts", "next"}, } for i, w := range want { got := resp.ToolCalls[i] if got.ID != w.id || got.Name != w.name { t.Errorf("call %d = {%s %s}, want {%s %s}", i, got.ID, got.Name, w.id, w.name) } // Each must be valid JSON on its own. A spliced pair usually is too, // which is exactly why the value is asserted and not just the parse. var args map[string]string if err := json.Unmarshal(got.Input, &args); err != nil { t.Fatalf("call %d input %q is not valid JSON: %v", i, got.Input, err) } if len(args) != 1 { t.Errorf("call %d carried %d args, want 1 — fragments from another call were spliced in: %v", i, len(args), args) } for _, v := range args { if v != w.day { t.Errorf("call %d arg = %q, want %q", i, v, w.day) } } } if resp.StopReason != "tool_use" { t.Errorf("StopReason = %q, want tool_use", resp.StopReason) } } // A tool call is buffered until the stream closes: a half-built argument object // is not a smaller version of the finished one, and dispatching on it would run // a tool with arguments the model had not finished choosing. func TestStreamNeverEmitsPartialToolArguments(t *testing.T) { gw := sse(t, `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"c","function":{"name":"t","arguments":"{\"a\":"}}]}}]}`, `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"1}"}}]}}]}`, `data: {"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`, `data: [DONE]`, ) var streamed strings.Builder resp, err := gw.Stream(context.Background(), ask("go"), func(d string) { streamed.WriteString(d) }) if err != nil { t.Fatalf("Stream: %v", err) } if streamed.String() != "" { t.Errorf("tool-call JSON reached the reader as text: %q", streamed.String()) } if string(resp.ToolCalls[0].Input) != `{"a":1}` { t.Errorf("Input = %q, want the assembled object", resp.ToolCalls[0].Input) } } // Keep-alives, comment lines and provider-specific events are not failures. A // stream that died on one would fail against providers that are working fine. func TestStreamIgnoresNoiseEvents(t *testing.T) { gw := sse(t, `: keep-alive`, ``, `event: ping`, `data: {"not":"a chunk"`, `data:{"choices":[{"delta":{"content":"ok"},"finish_reason":"stop"}]}`, `data: [DONE]`, ) resp, err := gw.Stream(context.Background(), ask("hi"), nil) if err != nil { t.Fatalf("Stream: %v", err) } if resp.Text != "ok" { t.Errorf("Text = %q, want ok", resp.Text) } } // A streamed refusal must come back as the same structured outcome the // non-streaming path produces, and must not be shown to the reader as though // it were the answer. func TestStreamRefusalIsNotShownToTheReader(t *testing.T) { gw := sse(t, `data: {"choices":[{"delta":{"refusal":"I cannot help with that."},"finish_reason":"stop"}]}`, `data: [DONE]`, ) var streamed strings.Builder _, err := gw.Stream(context.Background(), ask("do something disallowed"), func(d string) { streamed.WriteString(d) }) var gwErr *Error if !errors.As(err, &gwErr) || gwErr.Code != CodeRefused { t.Fatalf("err = %v, want a %s", err, CodeRefused) } if streamed.String() != "" { t.Errorf("a refusal was streamed to the reader as an answer: %q", streamed.String()) } } // StreamComplete has to reach the streaming path for a gateway that has one. // The fallback exists for gateways that do not, and silently taking it here // would turn every streamed answer into one lump with no error to trace it to. func TestStreamCompleteUsesTheStreamingPath(t *testing.T) { gw := sse(t, `data: {"choices":[{"delta":{"content":"a"}}]}`, `data: {"choices":[{"delta":{"content":"b"},"finish_reason":"stop"}]}`, `data: [DONE]`, ) var deltas int resp, err := StreamComplete(context.Background(), gw, ask("hi"), func(string) { deltas++ }) if err != nil { t.Fatalf("StreamComplete: %v", err) } if deltas != 2 { t.Errorf("got %d deltas, want 2 — the non-streaming fallback was taken", deltas) } if resp.Text != "ab" { t.Errorf("Text = %q", resp.Text) } } // The streamed shape of TestToolCallProviderMetadataIsRoundTripped: the // metadata arrives on one fragment, and later fragments that carry only // argument text must not erase it. func TestStreamKeepsToolCallProviderMetadata(t *testing.T) { const sig = `{"google":{"thought_signature":"El4KXAFpFH0T4CM3"}}` acc, err := accumulateSSE(strings.NewReader(strings.Join([]string{ `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_a","function":{"name":"open_positions","arguments":""},"extra_content":` + sig + `}]}}]}`, `data: {"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{}"}}]}}]}`, `data: {"choices":[{"delta":{},"finish_reason":"tool_calls"}]}`, `data: [DONE]`, }, "\n\n")), func(string) {}) if err != nil { t.Fatalf("accumulateSSE: %v", err) } msg := acc.message() if len(msg.ToolCalls) != 1 { t.Fatalf("got %d tool calls, want 1", len(msg.ToolCalls)) } if string(msg.ToolCalls[0].ExtraContent) != sig { t.Errorf("extra_content after streaming = %s, want %s", msg.ToolCalls[0].ExtraContent, sig) } if msg.ToolCalls[0].Function.Arguments != "{}" { t.Errorf("arguments = %q, want {}", msg.ToolCalls[0].Function.Arguments) } }