package proxy import ( "encoding/json" "io" "testing" ) func TestScanUsageChat(t *testing.T) { line := []byte(`data: {"id":"x","choices":[],"usage":{"prompt_tokens":12,"completion_tokens":9,"total_tokens":21}}`) raw := scanUsage(line) if raw == nil { t.Fatal("chat usage not detected") } var us usageShape if err := json.Unmarshal(raw, &us); err != nil { t.Fatal(err) } if us.PromptTokens != 12 || us.CompletionTokens != 9 { t.Fatalf("usage mismatch: %+v", us) } } func TestScanUsageResponsesNested(t *testing.T) { line := []byte(`data: {"response":{"id":"r","status":"completed","usage":{"input_tokens":15,"output_tokens":11,"total_tokens":26}},"type":"response.completed"}`) raw := scanUsage(line) if raw == nil { t.Fatal("responses nested usage not detected") } var us usageShape if err := json.Unmarshal(raw, &us); err != nil { t.Fatal(err) } if us.InputTokens != 15 || us.OutputTokens != 11 { t.Fatalf("usage mismatch: %+v", us) } } func TestScanUsageIgnoresNonData(t *testing.T) { if scanUsage([]byte("event: response.completed")) != nil { t.Fatal("event line should be ignored") } if scanUsage([]byte("data: [DONE]")) != nil { t.Fatal("[DONE] should be ignored") } if scanUsage([]byte("data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}")) != nil { t.Fatal("content chunk without usage should be ignored") } } func TestExtractUsageFromFullBody(t *testing.T) { body := []byte(`{"id":"x","choices":[{"message":{"content":"hi"}}],"usage":{"prompt_tokens":1,"completion_tokens":2}}`) raw := extractUsage(body) if raw == nil { t.Fatal("usage not extracted from full body") } var us usageShape _ = json.Unmarshal(raw, &us) if us.PromptTokens != 1 || us.CompletionTokens != 2 { t.Fatalf("usage mismatch: %+v", us) } } func TestSSEScannerLines(t *testing.T) { // 模拟分块写入的 SSE 流 data := "data: {\"a\":1}\n\ndata: {\"usage\":{\"input_tokens\":3}}\n\n" parts := [][]byte{[]byte(data[:10]), []byte(data[10:20]), []byte(data[20:])} reader := newChunkReader(parts) s := newSSEScanner(reader) var lines [][]byte for { line, err := s.Next() if line != nil { lines = append(lines, line) } if err != nil { break } } if len(lines) != 4 { t.Fatalf("expected 4 lines, got %d", len(lines)) } // 合并后应能还原原始数据 joined := "" for _, l := range lines { joined += string(l) } if joined != string(data) { t.Fatalf("stream corrupted:\n got: %q\nwant: %q", joined, data) } } type chunkReader struct { parts [][]byte idx int } func newChunkReader(parts [][]byte) *chunkReader { return &chunkReader{parts: parts} } func (r *chunkReader) Read(p []byte) (int, error) { if r.idx >= len(r.parts) { return 0, io.EOF } n := copy(p, r.parts[r.idx]) r.idx++ return n, nil }