package convert import ( "encoding/json" "testing" ) // ---- ExtractUsageJSON: 非流式各协议 ---- func TestExtractUsageJSONChat(t *testing.T) { body := []byte(`{ "id": "chatcmpl-1", "object": "chat.completion", "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}}], "usage": { "prompt_tokens": 11, "completion_tokens": 7, "total_tokens": 18, "prompt_tokens_details": {"cached_tokens": 4} } }`) u, ok := ExtractUsageJSON(body, ProtoChat) if !ok { t.Fatalf("expected ok=true") } if u.InputTokens != 11 || u.OutputTokens != 7 { t.Fatalf("chat usage = %+v, want input=11 output=7", u) } if u.CacheReadTokens != 4 { t.Fatalf("chat cacheRead = %d, want 4", u.CacheReadTokens) } } func TestExtractUsageJSONMessages(t *testing.T) { body := []byte(`{ "id": "msg_1", "type": "message", "role": "assistant", "content": [{"type": "text", "text": "hi"}], "usage": { "input_tokens": 15, "output_tokens": 8, "cache_read_input_tokens": 3, "cache_creation_input_tokens": 2 } }`) u, ok := ExtractUsageJSON(body, ProtoMessages) if !ok { t.Fatalf("expected ok=true") } if u.InputTokens != 15 || u.OutputTokens != 8 || u.CacheReadTokens != 3 || u.CacheCreationTokens != 2 { t.Fatalf("messages usage = %+v", u) } } func TestExtractUsageJSONResponses(t *testing.T) { body := []byte(`{ "id": "resp_1", "object": "response", "output": [], "usage": { "input_tokens": 13, "output_tokens": 9, "input_tokens_details": {"cached_tokens": 5} } }`) u, ok := ExtractUsageJSON(body, ProtoResponses) if !ok { t.Fatalf("expected ok=true") } if u.InputTokens != 13 || u.OutputTokens != 9 || u.CacheReadTokens != 5 { t.Fatalf("responses usage = %+v", u) } } func TestExtractUsageJSONInvalidAndMissing(t *testing.T) { if _, ok := ExtractUsageJSON([]byte("not json"), ProtoChat); ok { t.Fatalf("invalid json should not report ok") } if _, ok := ExtractUsageJSON([]byte(`{"id": "x"}`), ProtoChat); ok { t.Fatalf("missing usage should not report ok") } // 空对象 usage:全 0 视为无效 if _, ok := ExtractUsageJSON([]byte(`{"usage": {}}`), ProtoChat); ok { t.Fatalf("empty usage should not report ok") } } // ---- StreamUsageAccum: 流式各协议 ---- func feedLines(t *testing.T, proto string, lines ...string) TokenUsage { t.Helper() acc := NewStreamUsageAccum() for _, ln := range lines { acc.Feed([]byte(ln), proto) } return acc.Usage() } func TestStreamUsageChatFinalChunk(t *testing.T) { // 前面的 chunk 不带 usage;最后一个 chunk 带完整 usage u := feedLines(t, ProtoChat, `{"id":"c1","object":"chat.completion.chunk","choices":[{"delta":{"content":"he"}}]}`, `{"id":"c1","object":"chat.completion.chunk","choices":[{"delta":{"content":"llo"}}]}`, `{"id":"c1","object":"chat.completion.chunk","choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"prompt_tokens_details":{"cached_tokens":4}}}`, ) if u.InputTokens != 11 || u.OutputTokens != 7 || u.CacheReadTokens != 4 { t.Fatalf("chat stream usage = %+v", u) } } func TestStreamUsageMessagesStartAndDelta(t *testing.T) { // message_start 带 input/cache,message_delta 带 output;逐字段取 max 合并 u := feedLines(t, ProtoMessages, `{"type":"message_start","message":{"id":"msg_1","usage":{"input_tokens":15,"cache_read_input_tokens":3,"cache_creation_input_tokens":2}}}`, `{"type":"content_block_delta","delta":{"type":"text_delta","text":"hi"}}`, `{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":8}}`, ) if u.InputTokens != 15 || u.OutputTokens != 8 || u.CacheReadTokens != 3 || u.CacheCreationTokens != 2 { t.Fatalf("messages stream usage = %+v", u) } } func TestStreamUsageResponsesCompleted(t *testing.T) { // response.completed 事件的用量嵌在 response.usage 下 u := feedLines(t, ProtoResponses, `{"type":"response.output_text.delta","delta":"hi"}`, `{"type":"response.completed","response":{"id":"resp_1","usage":{"input_tokens":13,"output_tokens":9,"input_tokens_details":{"cached_tokens":5}}}}`, ) if u.InputTokens != 13 || u.OutputTokens != 9 || u.CacheReadTokens != 5 { t.Fatalf("responses stream usage = %+v", u) } } func TestStreamUsageIgnoresNonDataPayloads(t *testing.T) { // [DONE]、垃圾行、空对象都不应产生用量 u := feedLines(t, ProtoChat, `[DONE]`, `{`, ``, `{"choices":[]}`) if u.has() { t.Fatalf("expected zero usage, got %+v", u) } } func TestStreamUsageFeedKeepsMaxAcrossEvents(t *testing.T) { // 同一字段在多个事件出现时取较大值(防乱序/重复) u := feedLines(t, ProtoMessages, `{"type":"message_start","message":{"usage":{"input_tokens":15}}}`, `{"type":"message_delta","usage":{"output_tokens":5}}`, `{"type":"message_delta","usage":{"output_tokens":8}}`, ) if u.InputTokens != 15 || u.OutputTokens != 8 { t.Fatalf("max-merge usage = %+v", u) } } // ---- usage JSON 结构合法性(防止手写 struct 漂移)---- func TestUsageJSONRoundTrip(t *testing.T) { u := TokenUsage{InputTokens: 10, OutputTokens: 5, CacheReadTokens: 2, CacheCreationTokens: 1} b, err := json.Marshal(u) if err != nil { t.Fatalf("marshal: %v", err) } var back TokenUsage if err := json.Unmarshal(b, &back); err != nil { t.Fatalf("unmarshal: %v", err) } if back != u { t.Fatalf("round trip = %+v, want %+v", back, u) } }