mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 03:13:24 +00:00
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:
1 parent
b459d40531
commit
e95851e690
2 files changed
+67
-1
No files matched your search
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user