package proxy import ( "encoding/json" "io" "strings" "testing" ) func TestSSEScannerSplitsLines(t *testing.T) { input := "event: message\ndata: {\"a\":1}\n\n" + "data: {\"b\":2}\r\n\r\n" + "data: [DONE]\n\n" s := newSSEScanner(strings.NewReader(input)) var lines []string for { line, err := s.Next() if line != nil { lines = append(lines, string(line)) } if err == io.EOF { break } if err != nil { t.Fatalf("Next: %v", err) } } want := []string{ "event: message\n", "data: {\"a\":1}\n", "\n", "data: {\"b\":2}\r\n", "\r\n", "data: [DONE]\n", "\n", } if len(lines) != len(want) { t.Fatalf("line count = %d, want %d (lines: %q)", len(lines), len(want), lines) } for i := range want { if lines[i] != want[i] { t.Fatalf("line[%d] = %q, want %q", i, lines[i], want[i]) } } } func TestScanUsageChatStream(t *testing.T) { chunk := `data: {"id":"x","choices":[],"usage":{"prompt_tokens":12,"completion_tokens":9,"total_tokens":21}}` raw := scanUsage([]byte(chunk + "\n")) if raw == nil { t.Fatal("expected usage extracted") } var us usageShape if err := json.Unmarshal(raw, &us); err != nil { t.Fatalf("unmarshal: %v", err) } if us.PromptTokens != 12 || us.CompletionTokens != 9 { t.Fatalf("usage mismatch: %+v", us) } } func TestScanUsageResponsesCompleted(t *testing.T) { line := `data: {"type":"response.completed","response":{"id":"r1","status":"completed","usage":{"input_tokens":15,"output_tokens":11}}}` raw := scanUsage([]byte(line + "\n")) if raw == nil { t.Fatal("expected usage extracted from response.completed") } var us usageShape _ = json.Unmarshal(raw, &us) if us.InputTokens != 15 || us.OutputTokens != 11 { t.Fatalf("usage mismatch: %+v", us) } } func TestScanUsageIgnoresNonUsage(t *testing.T) { if raw := scanUsage([]byte(`data: {"type":"response.output_text.delta","delta":"hi"}`)); raw != nil { t.Fatalf("expected nil for non-usage line, got %s", raw) } if raw := scanUsage([]byte(`data: [DONE]`)); raw != nil { t.Fatal("expected nil for [DONE]") } } func TestExtractUsageChatBody(t *testing.T) { body := `{"id":"x","choices":[{"message":{"role":"assistant","content":"hi"}}],"usage":{"prompt_tokens":1,"completion_tokens":2,"total_tokens":3}}` raw := extractUsage([]byte(body)) if raw == nil { t.Fatal("expected usage") } if !strings.Contains(string(raw), `"prompt_tokens":1`) { t.Fatalf("unexpected usage: %s", raw) } } func TestExtractUsageResponsesNested(t *testing.T) { // responses 顶层只有 response 对象,usage 嵌套其中 body := `{"id":"r1","object":"response","status":"completed","response":{"usage":{"input_tokens":7,"output_tokens":8}}}` raw := extractUsage([]byte(body)) if raw == nil { t.Fatal("expected nested usage") } var us usageShape _ = json.Unmarshal(raw, &us) if us.InputTokens != 7 || us.OutputTokens != 8 { t.Fatalf("usage mismatch: %+v", us) } } func TestUsageSinkMergesFields(t *testing.T) { // message_start 给 input,message_delta 给 output,合并后两者都在 s := &usageSink{} s.push(json.RawMessage(`{"input_tokens":14,"output_tokens":0}`)) s.push(json.RawMessage(`{"output_tokens":10}`)) got := s.Shape() if got.InputTokens != 14 || got.OutputTokens != 10 { t.Fatalf("merge mismatch: %+v", got) } // chat 末块同时携带两字段 s2 := &usageSink{} s2.push(json.RawMessage(`{"prompt_tokens":12,"completion_tokens":9}`)) g := s2.Shape() if g.PromptTokens != 12 || g.CompletionTokens != 9 { t.Fatalf("chat usage mismatch: %+v", g) } }