mirror of
https://github.com/sipeed/picoclaw.git
synced 2026-07-28 01:27:58 +00:00
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:
@@ -495,11 +495,16 @@ func (p *Pipeline) CallLLM(
|
||||
}
|
||||
}
|
||||
|
||||
// Save finishReason to turnState for SubTurn truncation detection
|
||||
if innerTS := turnStateFromContext(ctx); innerTS != nil {
|
||||
innerTS.SetLastFinishReason(exec.response.FinishReason)
|
||||
// Save finishReason and usage on the turn state. Use ts directly (the
|
||||
// authoritative turn state for this call) rather than a context lookup:
|
||||
// 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 {
|
||||
innerTS.SetLastUsage(exec.response.Usage)
|
||||
ts.SetLastUsage(exec.response.Usage)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -699,6 +699,20 @@ func setStreamerModelName(streamer any, modelName string) {
|
||||
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
|
||||
// MessageSplitMarker into separate channel stream messages.
|
||||
type splitMarkerStreamer struct {
|
||||
@@ -711,6 +725,8 @@ type splitMarkerStreamer struct {
|
||||
onFinalize func(context.Context, string)
|
||||
clearMarker func()
|
||||
modelName string
|
||||
turnInputTokens int
|
||||
turnOutputTokens int
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
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) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
@@ -840,6 +864,7 @@ func (s *splitMarkerStreamer) ensureCurrentLocked(ctx context.Context) error {
|
||||
}
|
||||
s.current = streamer
|
||||
setStreamerModelName(s.current, s.modelName)
|
||||
setStreamerTurnUsage(s.current, s.turnInputTokens, s.turnOutputTokens)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -928,6 +953,10 @@ func (s *finalizeHookStreamer) SetModelName(modelName string) {
|
||||
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) {
|
||||
if s.onFinalize != nil {
|
||||
s.onFinalize(ctx, content)
|
||||
|
||||
@@ -3383,3 +3383,56 @@ func TestManager_SendPlaceholder(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
@@ -531,6 +533,8 @@ type picoStreamer struct {
|
||||
channel *PicoChannel
|
||||
chatID string
|
||||
modelName string
|
||||
turnInputTokens int
|
||||
turnOutputTokens int
|
||||
messageID string
|
||||
reasoningID string
|
||||
throttleInterval time.Duration
|
||||
@@ -553,6 +557,17 @@ func (s *picoStreamer) SetModelName(modelName string) {
|
||||
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 {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
@@ -661,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 {
|
||||
@@ -673,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
|
||||
}
|
||||
@@ -932,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>"
|
||||
@@ -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) {
|
||||
raw := strings.TrimSpace(msg.Context.Raw[PayloadKeyToolCalls])
|
||||
if raw == "" {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -28,6 +28,7 @@ const (
|
||||
PayloadKeyPlaceholder = "placeholder"
|
||||
PayloadKeyToolCalls = "tool_calls"
|
||||
PayloadKeyModelName = "model_name"
|
||||
PayloadKeyUsage = "usage"
|
||||
|
||||
MessageKindThought = "thought"
|
||||
MessageKindToolCalls = "tool_calls"
|
||||
|
||||
Reference in New Issue
Block a user