package proxy_test import ( "bufio" "encoding/json" "fmt" "io" "net" "net/http" "net/http/httptest" "net/url" "strings" "sync" "testing" "time" "git.misaka.ren/M1saka/token_thief/config" "git.misaka.ren/M1saka/token_thief/logger" "git.misaka.ren/M1saka/token_thief/proxy" ) type captureSubmitter struct { mu sync.Mutex entries []*logger.LogEntry } func (c *captureSubmitter) Submit(e *logger.LogEntry) { c.mu.Lock() defer c.mu.Unlock() c.entries = append(c.entries, e) } func (c *captureSubmitter) Len() int { c.mu.Lock() defer c.mu.Unlock() return len(c.entries) } func (c *captureSubmitter) Entry(i int) *logger.LogEntry { c.mu.Lock() defer c.mu.Unlock() return c.entries[i] } // fakeOpenAIStreamUpstream 模拟一个 OpenAI 兼容的 SSE 上游: // 分 5 次往响应里 write 一行 SSE 数据,每次都 Flush。 func fakeOpenAIStreamUpstream() *httptest.Server { chunks := []string{ `data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"你"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"好"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-1","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}` + "\n\n", "data: [DONE]\n\n", } return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.WriteHeader(http.StatusOK) flusher := w.(http.Flusher) for _, ch := range chunks { _, _ = io.WriteString(w, ch) flusher.Flush() time.Sleep(5 * time.Millisecond) } })) } func fakeAnthropicStreamUpstream() *httptest.Server { chunks := []string{ `event: message_start` + "\n" + `data: {"type":"message_start","message":{"id":"msg-1","type":"message","role":"assistant","content":[],"model":"claude-3-5-sonnet","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}` + "\n\n", `event: content_block_start` + "\n" + `data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}` + "\n\n", `event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"你"}}` + "\n\n", `event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"好"}}` + "\n\n", `event: message_delta` + "\n" + `data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":3}}` + "\n\n", `event: message_stop` + "\n" + `data: {"type":"message_stop"}` + "\n\n", } return newSSEUpstream(chunks) } func fakeGeminiStreamUpstream() *httptest.Server { chunks := []string{ `data: {"candidates":[{"content":{"parts":[{"text":"你"}],"role":"model"},"index":0}]}` + "\n\n", `data: {"candidates":[{"content":{"parts":[{"text":"好"}],"role":"model"},"finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":2,"totalTokenCount":4}}` + "\n\n", } return newSSEUpstream(chunks) } func fakeUnknownStreamUpstream() *httptest.Server { return newSSEUpstream([]string{ `event: custom` + "\n" + `data: not-json` + "\n\n", }) } func fakeOpenAIStreamUpstreamThatStaysOpen(release <-chan struct{}) *httptest.Server { chunks := []string{ `data: {"id":"chatcmpl-hang","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-hang","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-hang","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}` + "\n\n", "data: [DONE]\n\n", } return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) flusher := w.(http.Flusher) for _, ch := range chunks { _, _ = io.WriteString(w, ch) flusher.Flush() } <-release })) } func newSSEUpstream(chunks []string) *httptest.Server { return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.WriteHeader(http.StatusOK) flusher := w.(http.Flusher) for _, ch := range chunks { _, _ = io.WriteString(w, ch) flusher.Flush() time.Sleep(5 * time.Millisecond) } })) } func TestSSEChunkAssembly(t *testing.T) { upstream := fakeOpenAIStreamUpstream() defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &captureSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 1024*1024) proxySrv := httptest.NewServer(h) defer proxySrv.Close() // 客户端走原始 TCP,逐字节读 + 打印,验证流式实时到达 pu, _ := url.Parse(proxySrv.URL) conn, err := net.DialTimeout("tcp", pu.Host, 3*time.Second) if err != nil { t.Fatal(err) } defer conn.Close() fmt.Fprintf(conn, "POST /v1/chat/completions HTTP/1.1\r\nHost: %s\r\nContent-Length: 0\r\n\r\n", pu.Host) _ = conn.SetReadDeadline(time.Now().Add(3 * time.Second)) br := bufio.NewReader(conn) resp, err := http.ReadResponse(br, nil) if err != nil { t.Fatal(err) } clientBody, _ := io.ReadAll(resp.Body) // 等待 proxy.ServeHTTP 返回并把 entry 提交 deadline := time.Now().Add(2 * time.Second) for sub.Len() == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if sub.Len() == 0 { t.Fatal("no log entry captured") } e := sub.Entry(0) t.Logf("\n========== 客户端收到的字节 ==========\n%s", clientBody) t.Logf("\n========== 数据库 response_body 字段(按字节原样存储)==========\n%s", e.ResponseBody) t.Logf("\n========== 元数据 ==========") t.Logf("is_stream = %v", e.IsStream) t.Logf("status_code = %d", e.StatusCode) t.Logf("len(body) = %d bytes", len(e.ResponseBody)) t.Logf("response_truncated = %v", e.ResponseTruncated) if !strings.Contains(string(clientBody), `"content":"你"`) || !strings.Contains(string(clientBody), `"content":"好"`) || !strings.Contains(string(clientBody), "[DONE]") { t.Errorf("client response body 缺少预期 chunk 内容") } var captured struct { Choices []struct { Message struct { Role string `json:"role"` Content string `json:"content"` } `json:"message"` FinishReason string `json:"finish_reason"` } `json:"choices"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("stream response_body should be assembled JSON: %v; body=%q", err, e.ResponseBody) } if len(captured.Choices) != 1 { t.Fatalf("assembled JSON choices length=%d, want 1", len(captured.Choices)) } if captured.Choices[0].Message.Role != "assistant" { t.Errorf("assembled role=%q, want assistant", captured.Choices[0].Message.Role) } if captured.Choices[0].Message.Content != "你好" { t.Errorf("assembled content=%q, want 你好", captured.Choices[0].Message.Content) } if captured.Choices[0].FinishReason != "stop" { t.Errorf("assembled finish_reason=%q, want stop", captured.Choices[0].FinishReason) } if !e.IsStream { t.Errorf("is_stream 应为 true") } } func TestOpenAIStreamPreservesReasoningAndMetadata(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"think "},"finish_reason":null,"native_finish_reason":null}]}` + "\n\n", `data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{"reasoning_content":"hard"},"finish_reason":null,"native_finish_reason":null}]}` + "\n\n", `data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{"content":"final"},"finish_reason":null,"native_finish_reason":null}]}` + "\n\n", `data: {"id":"resp-1","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini-2026-03-17","choices":[{"index":0,"delta":{},"finish_reason":"stop","native_finish_reason":"stop"}]}` + "\n\n", "data: [DONE]\n\n", }) e := requestStreamEntry(t, upstream, "/v1/chat/completions") var captured struct { ID string `json:"id"` Object string `json:"object"` Created int64 `json:"created"` Model string `json:"model"` Choices []struct { Index int `json:"index"` Message struct { Role string `json:"role"` Content string `json:"content"` ReasoningContent string `json:"reasoning_content"` } `json:"message"` FinishReason string `json:"finish_reason"` NativeFinishReason string `json:"native_finish_reason"` } `json:"choices"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("openai stream response_body should be JSON: %v; body=%q", err, e.ResponseBody) } if captured.ID != "resp-1" || captured.Object != "chat.completion" || captured.Created != 1779335544 || captured.Model != "gpt-5.4-mini-2026-03-17" { t.Fatalf("unexpected metadata: %+v", captured) } if len(captured.Choices) != 1 { t.Fatalf("choices length=%d, want 1", len(captured.Choices)) } choice := captured.Choices[0] if choice.Message.Role != "assistant" || choice.Message.Content != "final" || choice.Message.ReasoningContent != "think hard" { t.Fatalf("unexpected message: %+v", choice.Message) } if choice.FinishReason != "stop" || choice.NativeFinishReason != "stop" { t.Fatalf("unexpected finish reasons: %+v", choice) } } func TestOpenAIStreamDoneDoesNotSubmitBeforeUpstreamCloses(t *testing.T) { release := make(chan struct{}) upstream := fakeOpenAIStreamUpstreamThatStaysOpen(release) u, _ := url.Parse(upstream.URL) sub := &captureSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 1024*1024) proxySrv := httptest.NewServer(h) released := false defer func() { if !released { close(release) } proxySrv.Close() upstream.Close() }() clientDone := make(chan error, 1) go func() { resp, err := http.Post(proxySrv.URL+"/v1/chat/completions", "application/json", strings.NewReader(`{"stream":true}`)) if err != nil { clientDone <- err return } _, _ = io.ReadAll(resp.Body) _ = resp.Body.Close() clientDone <- nil }() time.Sleep(100 * time.Millisecond) if sub.Len() != 0 { t.Fatal("terminal SSE event must not submit before ReverseProxy returns") } close(release) released = true select { case err := <-clientDone: if err != nil { t.Fatal(err) } if sub.Len() != 1 { t.Fatalf("stream should be submitted once, got %d entries", sub.Len()) } case <-time.After(2 * time.Second): t.Fatal("client did not finish after upstream closed") } } func TestOpenAIStreamPreservesUsageChunk(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"id":"resp-usage","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","system_fingerprint":"fp_123","choices":[{"index":0,"delta":{"role":"assistant","content":"ok"},"finish_reason":null}]}` + "\n\n", `data: {"id":"resp-usage","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","system_fingerprint":"fp_123","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":null}` + "\n\n", `data: {"id":"resp-usage","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","system_fingerprint":"fp_123","choices":[],"usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7}}` + "\n\n", "data: [DONE]\n\n", }) e := requestStreamEntry(t, upstream, "/v1/chat/completions") var captured struct { ID string `json:"id"` SystemFingerprint string `json:"system_fingerprint"` Choices []struct { Message struct { Content string `json:"content"` } `json:"message"` } `json:"choices"` Usage struct { PromptTokens int `json:"prompt_tokens"` CompletionTokens int `json:"completion_tokens"` TotalTokens int `json:"total_tokens"` } `json:"usage"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("openai stream with usage should be assembled JSON: %v; body=%q", err, e.ResponseBody) } if captured.ID != "resp-usage" || captured.SystemFingerprint != "fp_123" { t.Fatalf("metadata not preserved: %+v", captured) } if len(captured.Choices) != 1 || captured.Choices[0].Message.Content != "ok" { t.Fatalf("choices not assembled: %+v", captured.Choices) } if captured.Usage.PromptTokens != 5 || captured.Usage.CompletionTokens != 2 || captured.Usage.TotalTokens != 7 { t.Fatalf("usage not preserved: %+v", captured.Usage) } } func TestOpenAIStreamAssemblesToolCalls(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"id":"resp-tools","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"lookup","arguments":"{\"q\":"}}]},"finish_reason":null}]}` + "\n\n", `data: {"id":"resp-tools","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"weather\"}"}}]},"finish_reason":null}]}` + "\n\n", `data: {"id":"resp-tools","object":"chat.completion.chunk","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}` + "\n\n", "data: [DONE]\n\n", }) e := requestStreamEntry(t, upstream, "/v1/chat/completions") var captured struct { Choices []struct { Message struct { Role string `json:"role"` Content string `json:"content"` ToolCalls []struct { ID string `json:"id"` Type string `json:"type"` Function struct { Name string `json:"name"` Arguments string `json:"arguments"` } `json:"function"` } `json:"tool_calls"` } `json:"message"` FinishReason string `json:"finish_reason"` } `json:"choices"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("openai tool stream should be assembled JSON: %v; body=%q", err, e.ResponseBody) } if len(captured.Choices) != 1 || captured.Choices[0].FinishReason != "tool_calls" { t.Fatalf("unexpected choices: %+v", captured.Choices) } message := captured.Choices[0].Message if message.Role != "assistant" || message.Content != "" { t.Fatalf("unexpected message basics: %+v", message) } if len(message.ToolCalls) != 1 { t.Fatalf("tool_calls length=%d, want 1; body=%s", len(message.ToolCalls), e.ResponseBody) } tool := message.ToolCalls[0] if tool.ID != "call_1" || tool.Type != "function" || tool.Function.Name != "lookup" || tool.Function.Arguments != `{"q":"weather"}` { t.Fatalf("unexpected tool call: %+v", tool) } } func TestOpenAIResponsesStreamAssemblesCompletedResponse(t *testing.T) { upstream := newSSEUpstream([]string{ `event: response.created` + "\n" + `data: {"type":"response.created","response":{"id":"resp-1","object":"response","status":"in_progress","model":"gpt-5.4-mini","output":[]}}` + "\n\n", `event: response.output_text.delta` + "\n" + `data: {"type":"response.output_text.delta","item_id":"msg-1","output_index":0,"content_index":0,"delta":"hello"}` + "\n\n", `event: response.completed` + "\n" + `data: {"type":"response.completed","response":{"id":"resp-1","object":"response","status":"completed","model":"gpt-5.4-mini","output":[{"id":"msg-1","type":"message","role":"assistant","content":[{"type":"output_text","text":"hello"}]}],"usage":{"input_tokens":3,"output_tokens":1,"total_tokens":4}}}` + "\n\n", }) e := requestStreamEntry(t, upstream, "/v1/responses") var captured struct { ID string `json:"id"` Object string `json:"object"` Status string `json:"status"` Model string `json:"model"` Output []struct { Type string `json:"type"` Role string `json:"role"` Content []struct { Type string `json:"type"` Text string `json:"text"` } `json:"content"` } `json:"output"` Usage struct { InputTokens int `json:"input_tokens"` OutputTokens int `json:"output_tokens"` TotalTokens int `json:"total_tokens"` } `json:"usage"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("responses stream should be completed response JSON: %v; body=%q", err, e.ResponseBody) } if captured.ID != "resp-1" || captured.Object != "response" || captured.Status != "completed" || captured.Model != "gpt-5.4-mini" { t.Fatalf("unexpected response metadata: %+v", captured) } if len(captured.Output) != 1 || len(captured.Output[0].Content) != 1 || captured.Output[0].Content[0].Text != "hello" { t.Fatalf("unexpected response output: %+v", captured.Output) } if captured.Usage.TotalTokens != 4 { t.Fatalf("usage not preserved: %+v", captured.Usage) } } func TestOpenAICompletionsStreamAssemblesText(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"id":"cmpl-1","object":"text_completion","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"text":"hello","finish_reason":null}]}` + "\n\n", `data: {"id":"cmpl-1","object":"text_completion","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"text":" world","finish_reason":null}]}` + "\n\n", `data: {"id":"cmpl-1","object":"text_completion","created":1779335544,"model":"gpt-5.4-mini","choices":[{"index":0,"text":"","finish_reason":"stop"}],"usage":{"prompt_tokens":2,"completion_tokens":2,"total_tokens":4}}` + "\n\n", "data: [DONE]\n\n", }) e := requestStreamEntry(t, upstream, "/v1/completions") var captured struct { ID string `json:"id"` Object string `json:"object"` Created int64 `json:"created"` Model string `json:"model"` Choices []struct { Index int `json:"index"` Text string `json:"text"` FinishReason string `json:"finish_reason"` } `json:"choices"` Usage struct { TotalTokens int `json:"total_tokens"` } `json:"usage"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("completions stream should be assembled JSON: %v; body=%q", err, e.ResponseBody) } if captured.ID != "cmpl-1" || captured.Object != "text_completion" || captured.Model != "gpt-5.4-mini" || captured.Created != 1779335544 { t.Fatalf("unexpected metadata: %+v", captured) } if len(captured.Choices) != 1 || captured.Choices[0].Text != "hello world" || captured.Choices[0].FinishReason != "stop" { t.Fatalf("unexpected choices: %+v", captured.Choices) } if captured.Usage.TotalTokens != 4 { t.Fatalf("usage not preserved: %+v", captured.Usage) } } func TestAnthropicSSEAssemblesNativeMessageJSON(t *testing.T) { e := requestStreamEntry(t, fakeAnthropicStreamUpstream(), "/v1/messages") var captured struct { ID string `json:"id"` Type string `json:"type"` Role string `json:"role"` Content []struct { Type string `json:"type"` Text string `json:"text"` } `json:"content"` StopReason string `json:"stop_reason"` Usage struct { InputTokens int `json:"input_tokens"` OutputTokens int `json:"output_tokens"` } `json:"usage"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("anthropic stream response_body should be native JSON: %v; body=%q", err, e.ResponseBody) } if captured.ID != "msg-1" || captured.Type != "message" || captured.Role != "assistant" { t.Fatalf("unexpected anthropic message metadata: %+v", captured) } if len(captured.Content) != 1 || captured.Content[0].Type != "text" || captured.Content[0].Text != "你好" { t.Fatalf("unexpected anthropic content: %+v", captured.Content) } if captured.StopReason != "end_turn" { t.Errorf("stop_reason=%q, want end_turn", captured.StopReason) } if captured.Usage.InputTokens != 10 || captured.Usage.OutputTokens != 3 { t.Errorf("usage=%+v, want input=10 output=3", captured.Usage) } } func TestAnthropicSSEAssemblesToolUseContent(t *testing.T) { upstream := newSSEUpstream([]string{ `event: message_start` + "\n" + `data: {"type":"message_start","message":{"id":"msg-tool","type":"message","role":"assistant","content":[],"model":"claude-3-5-sonnet","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}` + "\n\n", `event: content_block_start` + "\n" + `data: {"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"toolu_1","name":"lookup","input":{}}}` + "\n\n", `event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"q\":"}}` + "\n\n", `event: content_block_delta` + "\n" + `data: {"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"\"weather\"}"}}` + "\n\n", `event: message_delta` + "\n" + `data: {"type":"message_delta","delta":{"stop_reason":"tool_use","stop_sequence":null},"usage":{"output_tokens":8}}` + "\n\n", `event: message_stop` + "\n" + `data: {"type":"message_stop"}` + "\n\n", }) e := requestStreamEntry(t, upstream, "/v1/messages") var captured struct { Content []struct { Type string `json:"type"` ID string `json:"id"` Name string `json:"name"` Input map[string]any `json:"input"` } `json:"content"` StopReason string `json:"stop_reason"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("anthropic tool stream should be JSON: %v; body=%q", err, e.ResponseBody) } if len(captured.Content) != 1 { t.Fatalf("content length=%d, want 1; body=%s", len(captured.Content), e.ResponseBody) } tool := captured.Content[0] if tool.Type != "tool_use" || tool.ID != "toolu_1" || tool.Name != "lookup" || tool.Input["q"] != "weather" { t.Fatalf("unexpected tool content: %+v", tool) } if captured.StopReason != "tool_use" { t.Fatalf("stop_reason=%q, want tool_use", captured.StopReason) } } func TestGeminiSSEAssemblesNativeGenerateContentJSON(t *testing.T) { e := requestStreamEntry(t, fakeGeminiStreamUpstream(), "/v1beta/models/gemini-1.5-pro:generateContent") var captured struct { Candidates []struct { Content struct { Role string `json:"role"` Parts []struct { Text string `json:"text"` } `json:"parts"` } `json:"content"` FinishReason string `json:"finishReason"` Index int `json:"index"` } `json:"candidates"` UsageMetadata struct { PromptTokenCount int `json:"promptTokenCount"` CandidatesTokenCount int `json:"candidatesTokenCount"` TotalTokenCount int `json:"totalTokenCount"` } `json:"usageMetadata"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("gemini stream response_body should be native JSON: %v; body=%q", err, e.ResponseBody) } if len(captured.Candidates) != 1 { t.Fatalf("candidates length=%d, want 1", len(captured.Candidates)) } candidate := captured.Candidates[0] if candidate.Content.Role != "model" || len(candidate.Content.Parts) != 1 || candidate.Content.Parts[0].Text != "你好" { t.Fatalf("unexpected gemini content: %+v", candidate.Content) } if candidate.FinishReason != "STOP" { t.Errorf("finishReason=%q, want STOP", candidate.FinishReason) } if captured.UsageMetadata.TotalTokenCount != 4 { t.Errorf("usageMetadata=%+v, want totalTokenCount=4", captured.UsageMetadata) } } func TestGeminiStreamPreservesSafetyRatingsAndUsageMetadata(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"candidates":[{"content":{"parts":[{"text":"你"}],"role":"model"},"finishReason":null,"index":0,"safetyRatings":[{"category":"HARM_CATEGORY_HARASSMENT","probability":"NEGLIGIBLE"}]}],"usageMetadata":{"promptTokenCount":8,"toolUsePromptTokenCount":0,"candidatesTokenCount":0,"totalTokenCount":8,"thoughtsTokenCount":10}}` + "\n\n", `data: {"candidates":[{"content":{"parts":[{"text":"好"}],"role":"model"},"finishReason":"STOP","index":0,"safetyRatings":[{"category":"HARM_CATEGORY_HARASSMENT","probability":"NEGLIGIBLE"}]}],"usageMetadata":{"promptTokenCount":8,"toolUsePromptTokenCount":0,"candidatesTokenCount":2,"totalTokenCount":20,"thoughtsTokenCount":10}}` + "\n\n", }) e := requestStreamEntry(t, upstream, "/v1beta/models/gemini-1.5-pro:streamGenerateContent") var captured struct { Candidates []struct { Content struct { Parts []struct { Text string `json:"text"` } `json:"parts"` } `json:"content"` SafetyRatings []struct { Category string `json:"category"` Probability string `json:"probability"` } `json:"safetyRatings"` } `json:"candidates"` UsageMetadata struct { PromptTokenCount int `json:"promptTokenCount"` ToolUsePromptTokenCount int `json:"toolUsePromptTokenCount"` CandidatesTokenCount int `json:"candidatesTokenCount"` TotalTokenCount int `json:"totalTokenCount"` ThoughtsTokenCount int `json:"thoughtsTokenCount"` } `json:"usageMetadata"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("gemini stream response_body should be JSON: %v; body=%q", err, e.ResponseBody) } if len(captured.Candidates) != 1 || len(captured.Candidates[0].SafetyRatings) != 1 { t.Fatalf("expected safetyRatings to be preserved, got %+v", captured.Candidates) } if captured.Candidates[0].SafetyRatings[0].Category != "HARM_CATEGORY_HARASSMENT" { t.Fatalf("unexpected safetyRatings: %+v", captured.Candidates[0].SafetyRatings) } if captured.UsageMetadata.ToolUsePromptTokenCount != 0 || captured.UsageMetadata.ThoughtsTokenCount != 10 || captured.UsageMetadata.TotalTokenCount != 20 { t.Fatalf("usageMetadata fields not preserved: %+v", captured.UsageMetadata) } } func TestGeminiStreamPreservesFunctionCallParts(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"candidates":[{"content":{"parts":[{"functionCall":{"name":"lookup","args":{"q":"weather"}}}],"role":"model"},"finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":2,"candidatesTokenCount":2,"totalTokenCount":4}}` + "\n\n", }) e := requestStreamEntry(t, upstream, "/v1beta/models/gemini-1.5-pro:streamGenerateContent") var captured struct { Candidates []struct { Content struct { Parts []struct { FunctionCall struct { Name string `json:"name"` Args map[string]any `json:"args"` } `json:"functionCall"` } `json:"parts"` } `json:"content"` } `json:"candidates"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("gemini functionCall stream should be JSON: %v; body=%q", err, e.ResponseBody) } call := captured.Candidates[0].Content.Parts[0].FunctionCall if call.Name != "lookup" || call.Args["q"] != "weather" { t.Fatalf("functionCall not preserved: %+v; body=%s", call, e.ResponseBody) } } func TestGeminiStreamPreservesMultiplePartsByIndex(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"candidates":[{"content":{"parts":[{"text":"hello "},{"functionCall":{"name":"lookup","args":{"q":"weather"}}},{"text":"world"}],"role":"model"},"index":0}]}` + "\n\n", `data: {"candidates":[{"content":{"parts":[{"text":"again"},{"functionCall":{"name":"lookup","args":{"q":"weather"}}},{"text":"!"}],"role":"model"},"finishReason":"STOP","index":0}]}` + "\n\n", }) e := requestStreamEntry(t, upstream, "/v1beta/models/gemini-1.5-pro:streamGenerateContent") var captured struct { Candidates []struct { Content struct { Parts []map[string]any `json:"parts"` } `json:"content"` } `json:"candidates"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("unmarshal assembled Gemini response: %v; body=%s", err, e.ResponseBody) } parts := captured.Candidates[0].Content.Parts if len(parts) != 3 || parts[0]["text"] != "hello again" || parts[2]["text"] != "world!" { t.Fatalf("multiple Gemini parts not preserved: %+v", parts) } if _, ok := parts[1]["functionCall"]; !ok { t.Fatalf("middle functionCall part missing: %+v", parts) } } func TestGeminiStreamAppendsDifferentPartKindsAcrossChunks(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"candidates":[{"content":{"parts":[{"text":"answer"}],"role":"model"},"index":0}]}` + "\n\n", `data: {"candidates":[{"content":{"parts":[{"functionCall":{"name":"lookup","args":{"q":"weather"}}}],"role":"model"},"finishReason":"STOP","index":0}]}` + "\n\n", }) e := requestStreamEntry(t, upstream, "/v1beta/models/gemini-1.5-pro:streamGenerateContent") var captured struct { Candidates []struct { Content struct { Parts []map[string]any `json:"parts"` } `json:"content"` } `json:"candidates"` } if err := json.Unmarshal(e.ResponseBody, &captured); err != nil { t.Fatalf("unmarshal assembled Gemini response: %v; body=%s", err, e.ResponseBody) } parts := captured.Candidates[0].Content.Parts if len(parts) != 2 || parts[0]["text"] != "answer" { t.Fatalf("different Gemini part kinds were merged: %+v", parts) } if _, ok := parts[1]["functionCall"]; !ok { t.Fatalf("functionCall part missing: %+v", parts) } } func TestUnknownSSEKeepsRawBody(t *testing.T) { e := requestStreamEntry(t, fakeUnknownStreamUpstream(), "/v1/chat/completions") if string(e.ResponseBody) != "event: custom\ndata: not-json\n\n" { t.Fatalf("unknown stream should keep raw body, got %q", e.ResponseBody) } } func TestTruncatedSSEKeepsCapturedRawBody(t *testing.T) { upstream := fakeOpenAIStreamUpstream() defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &captureSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 520) proxySrv := httptest.NewServer(h) defer proxySrv.Close() resp, err := http.Post(proxySrv.URL+"/v1/chat/completions", "application/json", nil) if err != nil { t.Fatal(err) } _, _ = io.ReadAll(resp.Body) _ = resp.Body.Close() deadline := time.Now().Add(2 * time.Second) for sub.Len() == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if sub.Len() == 0 { t.Fatal("no log entry captured") } e := sub.Entry(0) if !e.ResponseTruncated { t.Fatal("response should be marked truncated") } if !strings.HasPrefix(string(e.ResponseBody), "data: ") { t.Fatalf("truncated stream should keep captured raw body, got %q", e.ResponseBody) } if json.Valid(e.ResponseBody) { t.Fatalf("truncated stream should not be assembled as JSON, got %q", e.ResponseBody) } } func TestMultimodalStreamRequestBodyIsCaptured(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":null}]}` + "\n\n", `data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}` + "\n\n", "data: [DONE]\n\n", }) defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &captureSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 1024*1024) proxySrv := httptest.NewServer(h) defer proxySrv.Close() reqBody := `{"model":"gpt-5.4-mini","messages":[{"role":"user","content":[{"type":"text","text":"describe"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA"}}]}],"stream":true}` resp, err := http.Post(proxySrv.URL+"/v1/chat/completions", "application/json", strings.NewReader(reqBody)) if err != nil { t.Fatal(err) } _, _ = io.ReadAll(resp.Body) _ = resp.Body.Close() deadline := time.Now().Add(2 * time.Second) for sub.Len() == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if sub.Len() == 0 { t.Fatal("no log entry captured") } e := sub.Entry(0) if e.RequestTruncated { t.Fatal("multimodal request should not be truncated") } requestText := string(e.RequestBody) if !strings.Contains(requestText, `"image_url"`) || !strings.Contains(requestText, `data:image/png;base64,AAAA`) { t.Fatalf("request_body should contain image input, got %q", requestText) } if !e.IsStream { t.Fatal("response should still be marked stream") } } func TestChunkedMultimodalRequestBodyIsCaptured(t *testing.T) { upstream := newSSEUpstream([]string{ `data: {"id":"chatcmpl-image","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"ok"},"finish_reason":"stop"}]}` + "\n\n", "data: [DONE]\n\n", }) defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &captureSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 1024*1024) proxySrv := httptest.NewServer(h) defer proxySrv.Close() reqBody := `{"model":"gpt-5.4-mini","messages":[{"role":"user","content":[{"type":"text","text":"describe"},{"type":"image_url","image_url":{"url":"data:image/png;base64,AAAA"}}]}],"stream":true}` req, err := http.NewRequest(http.MethodPost, proxySrv.URL+"/v1/chat/completions", strings.NewReader(reqBody)) if err != nil { t.Fatal(err) } req.ContentLength = -1 req.Header.Set("Content-Type", "application/json") req.TransferEncoding = []string{"chunked"} resp, err := http.DefaultClient.Do(req) if err != nil { t.Fatal(err) } _, _ = io.ReadAll(resp.Body) _ = resp.Body.Close() deadline := time.Now().Add(2 * time.Second) for sub.Len() == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if sub.Len() == 0 { t.Fatal("no log entry captured") } e := sub.Entry(0) requestText := string(e.RequestBody) if !strings.Contains(requestText, `"image_url"`) || !strings.Contains(requestText, `data:image/png;base64,AAAA`) { t.Fatalf("chunked request_body should contain image input, got %q", requestText) } } func requestStreamEntry(t *testing.T, upstream *httptest.Server, path string) *logger.LogEntry { t.Helper() defer upstream.Close() u, _ := url.Parse(upstream.URL) sub := &captureSubmitter{} filter, err := config.NewFilter(config.FilterDisabled, nil) if err != nil { t.Fatal(err) } h := proxy.New(u, filter, sub, 1024*1024) proxySrv := httptest.NewServer(h) defer proxySrv.Close() resp, err := http.Post(proxySrv.URL+path, "application/json", nil) if err != nil { t.Fatal(err) } _, _ = io.ReadAll(resp.Body) _ = resp.Body.Close() deadline := time.Now().Add(2 * time.Second) for sub.Len() == 0 && time.Now().Before(deadline) { time.Sleep(10 * time.Millisecond) } if sub.Len() == 0 { t.Fatal("no log entry captured") } return sub.Entry(0) }