Merge pull request #3156 from loafoe/feat/session-token-usage

feat(pico): emit per-turn LLM token usage on finalized message
This commit is contained in:
Mauro
2026-07-03 09:14:36 +02:00
committed by GitHub
7 changed files with 217 additions and 14 deletions
+9 -4
View File
@@ -495,11 +495,16 @@ func (p *Pipeline) CallLLM(
} }
} }
// Save finishReason to turnState for SubTurn truncation detection // Save finishReason and usage on the turn state. Use ts directly (the
if innerTS := turnStateFromContext(ctx); innerTS != nil { // authoritative turn state for this call) rather than a context lookup:
innerTS.SetLastFinishReason(exec.response.FinishReason) // the raw ctx passed to CallLLM is not seeded with turnState (only turnCtx
// is), so turnStateFromContext(ctx) returns nil here and silently dropped
// both the finish reason and the per-turn token usage. ts is also exactly
// what the streaming publisher reads via GetLastUsage at finalize.
if ts != nil {
ts.SetLastFinishReason(exec.response.FinishReason)
if exec.response.Usage != nil { if exec.response.Usage != nil {
innerTS.SetLastUsage(exec.response.Usage) ts.SetLastUsage(exec.response.Usage)
} }
} }
+7
View File
@@ -54,6 +54,7 @@ func (p *Pipeline) tryConfiguredStreamingLLM(
channel: ts.channel, channel: ts.channel,
chatID: ts.chatID, chatID: ts.chatID,
modelName: exec.llmModelName, modelName: exec.llmModelName,
ts: ts,
} }
logger.DebugCF("agent", "configured streaming enabled", map[string]any{ logger.DebugCF("agent", "configured streaming enabled", map[string]any{
@@ -376,6 +377,7 @@ type streamingChunkPublisher struct {
published bool published bool
reasoningPublished bool reasoningPublished bool
err error err error
ts *turnState
} }
func (p *streamingChunkPublisher) Update(ctx context.Context, accumulated string) { func (p *streamingChunkPublisher) Update(ctx context.Context, accumulated string) {
@@ -445,6 +447,11 @@ func (p *streamingChunkPublisher) Finalize(ctx context.Context, content string,
if setter, ok := p.streamer.(interface{ SetModelName(modelName string) }); ok { if setter, ok := p.streamer.(interface{ SetModelName(modelName string) }); ok {
setter.SetModelName(p.modelName) setter.SetModelName(p.modelName)
} }
if usage := p.ts.GetLastUsage(); usage != nil {
if setter, ok := p.streamer.(interface{ SetTurnUsage(in, out int) }); ok {
setter.SetTurnUsage(usage.PromptTokens, usage.CompletionTokens)
}
}
var err error var err error
if streamer, ok := p.streamer.(bus.ContextUsageStreamer); ok { if streamer, ok := p.streamer.(bus.ContextUsageStreamer); ok {
err = streamer.FinalizeWithContext(ctx, content, contextUsage) err = streamer.FinalizeWithContext(ctx, content, contextUsage)
+38 -9
View File
@@ -699,18 +699,34 @@ func setStreamerModelName(streamer any, modelName string) {
setter.SetModelName(modelName) setter.SetModelName(modelName)
} }
type turnUsageStreamer interface {
SetTurnUsage(inputTokens, outputTokens int)
}
// setStreamerTurnUsage forwards real per-turn token usage to a streamer that
// supports it, transparently unwrapping the manager's streamer wrappers.
func setStreamerTurnUsage(streamer any, inputTokens, outputTokens int) {
setter, ok := streamer.(turnUsageStreamer)
if !ok {
return
}
setter.SetTurnUsage(inputTokens, outputTokens)
}
// splitMarkerStreamer turns accumulated streaming text containing // splitMarkerStreamer turns accumulated streaming text containing
// MessageSplitMarker into separate channel stream messages. // MessageSplitMarker into separate channel stream messages.
type splitMarkerStreamer struct { type splitMarkerStreamer struct {
mu sync.Mutex mu sync.Mutex
current bus.Streamer current bus.Streamer
reasoning bus.ReasoningStreamer reasoning bus.ReasoningStreamer
begin func(context.Context) (bus.Streamer, error) begin func(context.Context) (bus.Streamer, error)
completedParts int completedParts int
finalized bool finalized bool
onFinalize func(context.Context, string) onFinalize func(context.Context, string)
clearMarker func() clearMarker func()
modelName string modelName string
turnInputTokens int
turnOutputTokens int
} }
func (s *splitMarkerStreamer) Update(ctx context.Context, content string) error { func (s *splitMarkerStreamer) Update(ctx context.Context, content string) error {
@@ -761,6 +777,14 @@ func (s *splitMarkerStreamer) SetModelName(modelName string) {
setStreamerModelName(s.reasoning, s.modelName) setStreamerModelName(s.reasoning, s.modelName)
} }
func (s *splitMarkerStreamer) SetTurnUsage(inputTokens, outputTokens int) {
s.mu.Lock()
defer s.mu.Unlock()
s.turnInputTokens = inputTokens
s.turnOutputTokens = outputTokens
setStreamerTurnUsage(s.current, s.turnInputTokens, s.turnOutputTokens)
}
func (s *splitMarkerStreamer) Cancel(ctx context.Context) { func (s *splitMarkerStreamer) Cancel(ctx context.Context) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
@@ -840,6 +864,7 @@ func (s *splitMarkerStreamer) ensureCurrentLocked(ctx context.Context) error {
} }
s.current = streamer s.current = streamer
setStreamerModelName(s.current, s.modelName) setStreamerModelName(s.current, s.modelName)
setStreamerTurnUsage(s.current, s.turnInputTokens, s.turnOutputTokens)
return nil return nil
} }
@@ -928,6 +953,10 @@ func (s *finalizeHookStreamer) SetModelName(modelName string) {
setStreamerModelName(s.Streamer, strings.TrimSpace(modelName)) setStreamerModelName(s.Streamer, strings.TrimSpace(modelName))
} }
func (s *finalizeHookStreamer) SetTurnUsage(inputTokens, outputTokens int) {
setStreamerTurnUsage(s.Streamer, inputTokens, outputTokens)
}
func (s *finalizeHookStreamer) runFinalizeHook(ctx context.Context, content string) { func (s *finalizeHookStreamer) runFinalizeHook(ctx context.Context, content string) {
if s.onFinalize != nil { if s.onFinalize != nil {
s.onFinalize(ctx, content) s.onFinalize(ctx, content)
+53
View File
@@ -3383,3 +3383,56 @@ func TestManager_SendPlaceholder(t *testing.T) {
t.Error("expected SendPlaceholder to fail for unknown channel") t.Error("expected SendPlaceholder to fail for unknown channel")
} }
} }
// turnUsageTrackingStreamer is a mockStreamer that records SetTurnUsage calls,
// used to verify the manager's streamer wrappers forward per-turn token usage
// to the inner streamer (regression: the wrappers previously dropped it because
// SetTurnUsage is not part of the bus.Streamer interface).
type turnUsageTrackingStreamer struct {
mockStreamer
inputTokens int
outputTokens int
usageCalls int
}
func (m *turnUsageTrackingStreamer) SetTurnUsage(inputTokens, outputTokens int) {
m.usageCalls++
m.inputTokens = inputTokens
m.outputTokens = outputTokens
}
func TestFinalizeHookStreamerForwardsTurnUsage(t *testing.T) {
inner := &turnUsageTrackingStreamer{}
wrapper := &finalizeHookStreamer{Streamer: inner}
setter, ok := any(wrapper).(turnUsageStreamer)
if !ok {
t.Fatal("finalizeHookStreamer does not satisfy turnUsageStreamer")
}
setter.SetTurnUsage(1234, 567)
if inner.usageCalls != 1 {
t.Fatalf("inner SetTurnUsage calls = %d, want 1", inner.usageCalls)
}
if inner.inputTokens != 1234 || inner.outputTokens != 567 {
t.Errorf("inner usage = (%d, %d), want (1234, 567)", inner.inputTokens, inner.outputTokens)
}
}
func TestSplitMarkerStreamerForwardsTurnUsage(t *testing.T) {
inner := &turnUsageTrackingStreamer{}
wrapper := &splitMarkerStreamer{current: inner}
setter, ok := any(wrapper).(turnUsageStreamer)
if !ok {
t.Fatal("splitMarkerStreamer does not satisfy turnUsageStreamer")
}
setter.SetTurnUsage(1234, 567)
if inner.usageCalls != 1 {
t.Fatalf("inner SetTurnUsage calls = %d, want 1", inner.usageCalls)
}
if inner.inputTokens != 1234 || inner.outputTokens != 567 {
t.Errorf("inner usage = (%d, %d), want (1234, 567)", inner.inputTokens, inner.outputTokens)
}
}
+40 -1
View File
@@ -105,6 +105,8 @@ type PicoChannel struct {
cancel context.CancelFunc cancel context.CancelFunc
progress *channels.ToolFeedbackAnimator progress *channels.ToolFeedbackAnimator
deleteMessageFn func(context.Context, string, string) error deleteMessageFn func(context.Context, string, string) error
// broadcastFn lets tests intercept outbound broadcasts. nil → broadcastToSession.
broadcastFn func(chatID string, msg PicoMessage) error
} }
// NewPicoChannel creates a new Pico Protocol channel. // NewPicoChannel creates a new Pico Protocol channel.
@@ -531,6 +533,8 @@ type picoStreamer struct {
channel *PicoChannel channel *PicoChannel
chatID string chatID string
modelName string modelName string
turnInputTokens int
turnOutputTokens int
messageID string messageID string
reasoningID string reasoningID string
throttleInterval time.Duration throttleInterval time.Duration
@@ -553,6 +557,17 @@ func (s *picoStreamer) SetModelName(modelName string) {
s.modelName = strings.TrimSpace(modelName) s.modelName = strings.TrimSpace(modelName)
} }
// SetTurnUsage records the real per-turn LLM token usage to emit on finalize.
func (s *picoStreamer) SetTurnUsage(inputTokens, outputTokens int) {
if s == nil {
return
}
s.mu.Lock()
defer s.mu.Unlock()
s.turnInputTokens = inputTokens
s.turnOutputTokens = outputTokens
}
func (s *picoStreamer) Update(ctx context.Context, content string) error { func (s *picoStreamer) Update(ctx context.Context, content string) error {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
@@ -661,8 +676,9 @@ func (s *picoStreamer) sendLocked(ctx context.Context, content string, contextUs
payload[PayloadKeyModelName] = s.modelName payload[PayloadKeyModelName] = s.modelName
} }
setContextUsagePayload(payload, contextUsage) setContextUsagePayload(payload, contextUsage)
setTurnUsagePayload(payload, s.turnInputTokens, s.turnOutputTokens)
outMsg := newMessage(TypeMessageCreate, payload) outMsg := newMessage(TypeMessageCreate, payload)
if err := s.channel.broadcastToSession(s.chatID, outMsg); err != nil { if err := s.channel.broadcast(s.chatID, outMsg); err != nil {
return err return err
} }
} else if content != s.lastContent || contextUsage != nil { } else if content != s.lastContent || contextUsage != nil {
@@ -673,6 +689,7 @@ func (s *picoStreamer) sendLocked(ctx context.Context, content string, contextUs
if s.modelName != "" { if s.modelName != "" {
payload[PayloadKeyModelName] = s.modelName payload[PayloadKeyModelName] = s.modelName
} }
setTurnUsagePayload(payload, s.turnInputTokens, s.turnOutputTokens)
if err := s.channel.editMessagePayload(ctx, s.chatID, s.messageID, payload, contextUsage); err != nil { if err := s.channel.editMessagePayload(ctx, s.chatID, s.messageID, payload, contextUsage); err != nil {
return err return err
} }
@@ -932,6 +949,14 @@ func (c *PicoChannel) handleMediaDownload(w http.ResponseWriter, r *http.Request
http.ServeContent(w, r, filename, info.ModTime(), file) http.ServeContent(w, r, filename, info.ModTime(), file)
} }
// broadcast routes through broadcastFn when set (tests), else broadcastToSession.
func (c *PicoChannel) broadcast(chatID string, msg PicoMessage) error {
if c.broadcastFn != nil {
return c.broadcastFn(chatID, msg)
}
return c.broadcastToSession(chatID, msg)
}
// broadcastToSession sends a message to all connections with a matching session. // broadcastToSession sends a message to all connections with a matching session.
func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error { func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error {
// chatID format: "pico:<sessionID>" // chatID format: "pico:<sessionID>"
@@ -1403,6 +1428,20 @@ func setContextUsagePayload(payload map[string]any, u *bus.ContextUsage) {
} }
} }
// setTurnUsagePayload attaches real per-turn LLM token usage to the payload.
// Input and output are kept separate (billed at different rates); total is a
// convenience sum. Omitted entirely when both counts are zero.
func setTurnUsagePayload(payload map[string]any, inputTokens, outputTokens int) {
if inputTokens <= 0 && outputTokens <= 0 {
return
}
payload[PayloadKeyUsage] = map[string]any{
"input_tokens": inputTokens,
"output_tokens": outputTokens,
"total_tokens": inputTokens + outputTokens,
}
}
func picoToolCallsPayload(msg bus.OutboundMessage) ([]utils.VisibleToolCall, bool) { func picoToolCallsPayload(msg bus.OutboundMessage) ([]utils.VisibleToolCall, bool) {
raw := strings.TrimSpace(msg.Context.Raw[PayloadKeyToolCalls]) raw := strings.TrimSpace(msg.Context.Raw[PayloadKeyToolCalls])
if raw == "" { if raw == "" {
+69
View File
@@ -0,0 +1,69 @@
package pico
import (
"context"
"testing"
)
func TestSetTurnUsagePayload(t *testing.T) {
t.Run("populates usage block when counts present", func(t *testing.T) {
payload := map[string]any{PayloadKeyContent: "hi"}
setTurnUsagePayload(payload, 1234, 567)
raw, ok := payload[PayloadKeyUsage]
if !ok {
t.Fatalf("expected %q key in payload", PayloadKeyUsage)
}
usage, ok := raw.(map[string]any)
if !ok {
t.Fatalf("usage block is not a map: %T", raw)
}
if usage["input_tokens"] != 1234 {
t.Errorf("input_tokens = %v, want 1234", usage["input_tokens"])
}
if usage["output_tokens"] != 567 {
t.Errorf("output_tokens = %v, want 567", usage["output_tokens"])
}
if usage["total_tokens"] != 1801 {
t.Errorf("total_tokens = %v, want 1801", usage["total_tokens"])
}
})
t.Run("omits usage block when both counts zero", func(t *testing.T) {
payload := map[string]any{PayloadKeyContent: "hi"}
setTurnUsagePayload(payload, 0, 0)
if _, ok := payload[PayloadKeyUsage]; ok {
t.Errorf("expected no %q key when counts are zero", PayloadKeyUsage)
}
})
}
// newCaptureStreamer returns a streamer whose broadcasts are captured into the
// returned map pointer, so tests need no live websocket.
func newCaptureStreamer() (*picoStreamer, *map[string]any) {
var last map[string]any
ch := &PicoChannel{}
ch.broadcastFn = func(chatID string, msg PicoMessage) error {
last = msg.Payload
return nil
}
s := &picoStreamer{channel: ch, chatID: "c1"}
return s, &last
}
func TestStreamerEmitsUsageOnFinalize(t *testing.T) {
s, last := newCaptureStreamer()
s.SetTurnUsage(100, 40)
// sendLocked with empty messageID takes the create branch, which attaches
// usage from the streamer's stored counts.
s.mu.Lock()
err := s.sendLocked(context.Background(), "answer", nil)
s.mu.Unlock()
if err != nil {
t.Fatalf("sendLocked: %v", err)
}
if _, ok := (*last)[PayloadKeyUsage]; !ok {
t.Fatalf("expected usage in payload, got %+v", *last)
}
}
+1
View File
@@ -28,6 +28,7 @@ const (
PayloadKeyPlaceholder = "placeholder" PayloadKeyPlaceholder = "placeholder"
PayloadKeyToolCalls = "tool_calls" PayloadKeyToolCalls = "tool_calls"
PayloadKeyModelName = "model_name" PayloadKeyModelName = "model_name"
PayloadKeyUsage = "usage"
MessageKindThought = "thought" MessageKindThought = "thought"
MessageKindToolCalls = "tool_calls" MessageKindToolCalls = "tool_calls"