diff --git a/internal/runtime/claude_progress.go b/internal/runtime/claude_progress.go index 72142d01bf..c0192d59c1 100644 --- a/internal/runtime/claude_progress.go +++ b/internal/runtime/claude_progress.go @@ -118,6 +118,7 @@ func parseClaudeStream(r io.Reader, onEvent func(AgentEvent)) error { // per-message token tracking for throttled TokensEvent totalInput int totalOutput int + msgReasoning int // per-message thinking tokens (reset on message_start) totalCacheRead int totalCacheWrite int lastEmittedTotal int @@ -128,6 +129,7 @@ func parseClaudeStream(r io.Reader, onEvent func(AgentEvent)) error { cumulativeCacheRead int cumulativeCacheWrite int seenResult bool + totalReasoning int // accumulated thinking tokens across all messages (for ResultEvent) ) // Emit a final cumulative TokensEvent when the stream ends without @@ -275,6 +277,7 @@ func parseClaudeStream(r io.Reader, onEvent func(AgentEvent)) error { totalInput = msg.Message.Usage.InputTokens totalOutput = 0 + msgReasoning = 0 totalCacheRead = msg.Message.Usage.CacheReadInputTokens totalCacheWrite = msg.Message.Usage.CacheCreationInputTokens } @@ -282,20 +285,26 @@ func parseClaudeStream(r io.Reader, onEvent func(AgentEvent)) error { case "message_delta": var md struct { Usage struct { - OutputTokens int `json:"output_tokens"` + OutputTokens int `json:"output_tokens"` + OutputTokensDetails struct { + ThinkingTokens int `json:"thinking_tokens"` + } `json:"output_tokens_details"` } `json:"usage"` } if err := json.Unmarshal(wrapper.Event, &md); err == nil && md.Usage.OutputTokens > 0 { totalOutput = md.Usage.OutputTokens + msgReasoning = md.Usage.OutputTokensDetails.ThinkingTokens + totalReasoning += msgReasoning total := cumulativeInput + totalInput + cumulativeOutput + totalOutput + - cumulativeCacheRead + totalCacheRead + cumulativeCacheWrite + totalCacheWrite + msgReasoning + cumulativeCacheRead + totalCacheRead + cumulativeCacheWrite + totalCacheWrite if total-lastEmittedTotal >= tokenThreshold { lastEmittedTotal = total onEvent(TokensEvent{ - InputTokens: cumulativeInput + totalInput, - OutputTokens: cumulativeOutput + totalOutput, - CacheRead: cumulativeCacheRead + totalCacheRead, - CacheWrite: cumulativeCacheWrite + totalCacheWrite, + InputTokens: cumulativeInput + totalInput, + OutputTokens: cumulativeOutput + totalOutput, + ReasoningTokens: msgReasoning, + CacheRead: cumulativeCacheRead + totalCacheRead, + CacheWrite: cumulativeCacheWrite + totalCacheWrite, }) } } @@ -315,6 +324,7 @@ func parseClaudeStream(r io.Reader, onEvent func(AgentEvent)) error { Subtype: re.Subtype, InputTokens: re.Usage.InputTokens, OutputTokens: re.Usage.OutputTokens, + ReasoningTokens: totalReasoning, CacheCreationInputTokens: re.Usage.CacheCreationInputTokens, CacheReadInputTokens: re.Usage.CacheReadInputTokens, }) diff --git a/internal/runtime/claude_progress_test.go b/internal/runtime/claude_progress_test.go index 3aed0969ae..0bd7b1a7e6 100644 --- a/internal/runtime/claude_progress_test.go +++ b/internal/runtime/claude_progress_test.go @@ -942,6 +942,111 @@ func TestParseClaudeStreamTokensEvent(t *testing.T) { } } +func TestParseClaudeStreamTokensEventWithReasoningTokens(t *testing.T) { + lines := []string{ + `{"type":"stream_event","event":{"type":"message_start","message":{"usage":{"input_tokens":4000,"cache_read_input_tokens":500,"cache_creation_input_tokens":200}}}}`, + `{"type":"stream_event","event":{"type":"message_delta","usage":{"output_tokens":1000,"output_tokens_details":{"thinking_tokens":300}}}}`, + } + events := collectEvents(t, strings.Join(lines, "\n")) + + var tokens []TokensEvent + for _, e := range events { + if te, ok := e.(TokensEvent); ok { + tokens = append(tokens, te) + } + } + // Total = 4000 + 1000 + 300 + 500 + 200 = 6000, crosses 5k threshold + if len(tokens) != 1 { + t.Fatalf("expected 1 tokens event, got %d", len(tokens)) + } + if tokens[0].ReasoningTokens != 300 { + t.Errorf("expected 300 reasoning tokens, got %d", tokens[0].ReasoningTokens) + } + if tokens[0].OutputTokens != 1000 { + t.Errorf("expected 1000 output tokens, got %d", tokens[0].OutputTokens) + } +} + +func TestParseClaudeStreamResultEventAccumulatesReasoningTokens(t *testing.T) { + lines := []string{ + // First message turn with 200 thinking tokens. + `{"type":"stream_event","event":{"type":"message_start","message":{"usage":{"input_tokens":4000,"cache_read_input_tokens":500,"cache_creation_input_tokens":200}}}}`, + `{"type":"stream_event","event":{"type":"message_delta","usage":{"output_tokens":1000,"output_tokens_details":{"thinking_tokens":200}}}}`, + // Second message turn with 150 thinking tokens. + `{"type":"stream_event","event":{"type":"message_start","message":{"usage":{"input_tokens":6000,"cache_read_input_tokens":500,"cache_creation_input_tokens":200}}}}`, + `{"type":"stream_event","event":{"type":"message_delta","usage":{"output_tokens":800,"output_tokens_details":{"thinking_tokens":150}}}}`, + // Result event. + `{"type":"result","num_turns":2,"total_cost_usd":0.50,"usage":{"input_tokens":10000,"output_tokens":1800,"cache_creation_input_tokens":400,"cache_read_input_tokens":1000}}`, + } + events := collectEvents(t, strings.Join(lines, "\n")) + + var results []ResultEvent + for _, e := range events { + if re, ok := e.(ResultEvent); ok { + results = append(results, re) + } + } + if len(results) != 1 { + t.Fatalf("expected 1 result event, got %d", len(results)) + } + // Accumulated: 200 + 150 = 350. + if results[0].ReasoningTokens != 350 { + t.Errorf("expected 350 accumulated reasoning tokens, got %d", results[0].ReasoningTokens) + } +} + +func TestParseClaudeStreamNoThinkingTokensBackwardCompat(t *testing.T) { + lines := []string{ + `{"type":"stream_event","event":{"type":"message_start","message":{"usage":{"input_tokens":4000,"cache_read_input_tokens":500,"cache_creation_input_tokens":200}}}}`, + `{"type":"stream_event","event":{"type":"message_delta","usage":{"output_tokens":1000}}}`, + `{"type":"result","num_turns":1,"total_cost_usd":0.10,"usage":{"input_tokens":4000,"output_tokens":1000,"cache_creation_input_tokens":200,"cache_read_input_tokens":500}}`, + } + events := collectEvents(t, strings.Join(lines, "\n")) + + var tokens []TokensEvent + var results []ResultEvent + for _, e := range events { + switch ev := e.(type) { + case TokensEvent: + tokens = append(tokens, ev) + case ResultEvent: + results = append(results, ev) + } + } + // TokensEvent: reasoning should be 0 when no thinking tokens present. + if len(tokens) == 1 && tokens[0].ReasoningTokens != 0 { + t.Errorf("expected 0 reasoning tokens when absent, got %d", tokens[0].ReasoningTokens) + } + // ResultEvent: reasoning should be 0. + if len(results) != 1 { + t.Fatalf("expected 1 result event, got %d", len(results)) + } + if results[0].ReasoningTokens != 0 { + t.Errorf("expected 0 reasoning tokens in result when absent, got %d", results[0].ReasoningTokens) + } +} + +func TestProgressParserCapturesReasoningTokensInMetrics(t *testing.T) { + lines := []string{ + `{"type":"stream_event","event":{"type":"message_start","message":{"usage":{"input_tokens":4000,"cache_read_input_tokens":500,"cache_creation_input_tokens":200}}}}`, + `{"type":"stream_event","event":{"type":"message_delta","usage":{"output_tokens":1000,"output_tokens_details":{"thinking_tokens":250}}}}`, + `{"type":"result","num_turns":1,"total_cost_usd":0.10,"usage":{"input_tokens":4000,"output_tokens":1000,"cache_creation_input_tokens":200,"cache_read_input_tokens":500}}`, + } + + input := strings.NewReader(strings.Join(lines, "\n")) + var buf bytes.Buffer + printer := ui.New(&buf) + metrics := &RunMetrics{} + + if err := progressParser(input, printer, metrics); err != nil { + t.Fatalf("progressParser returned error: %v", err) + } + + if metrics.ReasoningTokens != 250 { + t.Errorf("expected 250 reasoning tokens in metrics, got %d", metrics.ReasoningTokens) + } +} + func TestParseClaudeStreamTokensEventThrottled(t *testing.T) { lines := []string{ `{"type":"stream_event","event":{"type":"message_start","message":{"usage":{"input_tokens":4000}}}}`,