feat(pico): emit real per-turn token usage on finalized message
This commit is contained in:
parent
052c742fe7
commit
697e94fe8c
3 changed files with 54 additions and 2 deletions
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
@ -674,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 {
|
||||||
|
|
@ -686,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
|
||||||
}
|
}
|
||||||
|
|
@ -945,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>"
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,9 @@
|
||||||
package pico
|
package pico
|
||||||
|
|
||||||
import "testing"
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
func TestSetTurnUsagePayload(t *testing.T) {
|
func TestSetTurnUsagePayload(t *testing.T) {
|
||||||
t.Run("populates usage block when counts present", func(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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue