feat(pico): emit real per-turn token usage on finalized message

This commit is contained in:
Andy Lo-A-Foe
2026-06-22 13:15:35 +02:00
parent 052c742fe7
commit 697e94fe8c
3 changed files with 54 additions and 2 deletions
+7
View File
@@ -54,6 +54,7 @@ func (p *Pipeline) tryConfiguredStreamingLLM(
channel: ts.channel,
chatID: ts.chatID,
modelName: exec.llmModelName,
ts: ts,
}
logger.DebugCF("agent", "configured streaming enabled", map[string]any{
@@ -376,6 +377,7 @@ type streamingChunkPublisher struct {
published bool
reasoningPublished bool
err error
ts *turnState
}
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 {
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
if streamer, ok := p.streamer.(bus.ContextUsageStreamer); ok {
err = streamer.FinalizeWithContext(ctx, content, contextUsage)
+13 -1
View File
@@ -105,6 +105,8 @@ type PicoChannel struct {
cancel context.CancelFunc
progress *channels.ToolFeedbackAnimator
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.
@@ -674,8 +676,9 @@ func (s *picoStreamer) sendLocked(ctx context.Context, content string, contextUs
payload[PayloadKeyModelName] = s.modelName
}
setContextUsagePayload(payload, contextUsage)
setTurnUsagePayload(payload, s.turnInputTokens, s.turnOutputTokens)
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
}
} else if content != s.lastContent || contextUsage != nil {
@@ -686,6 +689,7 @@ func (s *picoStreamer) sendLocked(ctx context.Context, content string, contextUs
if 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 {
return err
}
@@ -945,6 +949,14 @@ func (c *PicoChannel) handleMediaDownload(w http.ResponseWriter, r *http.Request
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.
func (c *PicoChannel) broadcastToSession(chatID string, msg PicoMessage) error {
// chatID format: "pico:<sessionID>"
+34 -1
View File
@@ -1,6 +1,9 @@
package pico
import "testing"
import (
"context"
"testing"
)
func TestSetTurnUsagePayload(t *testing.T) {
t.Run("populates usage block when counts present", func(t *testing.T) {
@@ -34,3 +37,33 @@ func TestSetTurnUsagePayload(t *testing.T) {
}
})
}
// 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)
}
}