fix(pipeline): aggregate cache tokens across agent turns (#1419)

ThinkStage folded each turn's usage into state.Think.TotalUsage but summed
only prompt/completion/total/thinking tokens, dropping CacheReadTokens,
CacheCreationTokens and the PromptTokensIncludeCachedSegments flag. The
aggregated RunResult.Usage is the source for webhook usage responses and
usage_events analytics, so async multi-turn agent runs reported zero cache
tokens even when traces showed heavy prompt-cache use (cache is read
per-response for spans, but never summed into the aggregate).

Sum cache read/creation across turns and OR the include-cached flag (a
provider-level property, consistent across a run). Restores correct cache
accounting for webhook usage and usage_events without any envelope change.
This commit is contained in:
thotam authored and GitHub committed 2026-07-10 20:25:06 +07:00
1 parent b459d40531
commit e95851e690
2 files changed
+67 -1

No files matched your search

+9 -1
View File
@@ -81,12 +81,20 @@ func (s *ThinkStage) Execute(ctx context.Context, state *RunState) error {
return fmt.Errorf("llm call: %w", err)
}
// 5. Accumulate usage (including ThinkingTokens for reasoning models)
// 5. Accumulate usage across turns: base tokens, ThinkingTokens for reasoning
// models, and cache tokens. Cache read/creation must be summed too — otherwise
// the aggregated RunResult.Usage (and thus webhook usage + usage_events) reports
// zero cache even when each turn hit the prompt cache heavily.
if resp.Usage != nil {
state.Think.TotalUsage.PromptTokens += resp.Usage.PromptTokens
state.Think.TotalUsage.CompletionTokens += resp.Usage.CompletionTokens
state.Think.TotalUsage.TotalTokens += resp.Usage.TotalTokens
state.Think.TotalUsage.ThinkingTokens += resp.Usage.ThinkingTokens
state.Think.TotalUsage.CacheReadTokens += resp.Usage.CacheReadTokens
state.Think.TotalUsage.CacheCreationTokens += resp.Usage.CacheCreationTokens
if resp.Usage.PromptTokensIncludeCachedSegments {
state.Think.TotalUsage.PromptTokensIncludeCachedSegments = true
}
}
if isEmptyLengthResponse(resp) {
@@ -0,0 +1,58 @@
package pipeline
import (
"context"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
)
// TestThinkStage_AccumulatesCacheTokens guards the pipeline aggregation of cache
// tokens across turns. Regression: async agent runs reported 0 cache tokens in
// webhook usage / usage_events despite traces showing heavy cache use, because
// ThinkStage summed only prompt/completion/total/thinking and dropped the cache
// fields when folding each turn's usage into state.Think.TotalUsage.
func TestThinkStage_AccumulatesCacheTokens(t *testing.T) {
t.Parallel()
deps := &PipelineDeps{
Config: PipelineConfig{MaxIterations: 10, MaxTokens: 1000},
CallLLM: func(_ context.Context, _ *RunState, _ providers.ChatRequest) (*providers.ChatResponse, error) {
return &providers.ChatResponse{
FinishReason: "stop",
Content: "done",
Usage: &providers.Usage{
PromptTokens: 100,
CompletionTokens: 20,
TotalTokens: 120,
CacheReadTokens: 80,
CacheCreationTokens: 10,
PromptTokensIncludeCachedSegments: true,
},
}, nil
},
}
stage := NewThinkStage(deps)
state := defaultState()
// Two iterations to prove accumulation (+=), not assignment.
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() 1: %v", err)
}
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() 2: %v", err)
}
u := state.Think.TotalUsage
if u.CacheReadTokens != 160 {
t.Errorf("CacheReadTokens = %d, want 160", u.CacheReadTokens)
}
if u.CacheCreationTokens != 20 {
t.Errorf("CacheCreationTokens = %d, want 20", u.CacheCreationTokens)
}
if !u.PromptTokensIncludeCachedSegments {
t.Error("PromptTokensIncludeCachedSegments = false, want true")
}
if u.PromptTokens != 200 || u.CompletionTokens != 40 {
t.Errorf("base tokens = %d/%d, want 200/40", u.PromptTokens, u.CompletionTokens)
}
}