fix(pico): clean up tool feedback progress state

This commit is contained in:
lxowalle 2026-04-22 16:23:28 +08:00
parent 7d03023987
commit a36d36a368
6 changed files with 252 additions and 0 deletions

View file

@ -90,6 +90,7 @@ type PicoChannel struct {
ctx context.Context
cancel context.CancelFunc
progress *channels.ToolFeedbackAnimator
deleteMessageFn func(context.Context, string, string) error
}
// NewPicoChannel creates a new Pico Protocol channel.
@ -131,6 +132,7 @@ func NewPicoChannel(
sessionConnections: make(map[string]map[string]*picoConn),
}
ch.progress = channels.NewToolFeedbackAnimator(ch.EditMessage)
ch.deleteMessageFn = ch.DeleteMessage
return ch, nil
}
@ -332,6 +334,14 @@ func (c *PicoChannel) EditMessage(ctx context.Context, chatID string, messageID
return c.broadcastToSession(chatID, outMsg)
}
// DeleteMessage implements channels.MessageDeleter.
func (c *PicoChannel) DeleteMessage(ctx context.Context, chatID string, messageID string) error {
outMsg := newMessage(TypeMessageDelete, map[string]any{
"message_id": messageID,
})
return c.broadcastToSession(chatID, outMsg)
}
func (c *PicoChannel) currentToolFeedbackMessage(chatID string) (string, bool) {
if c.progress == nil {
return "", false
@ -373,6 +383,11 @@ func (c *PicoChannel) dismissTrackedToolFeedbackMessage(ctx context.Context, cha
return
}
c.ClearToolFeedbackMessage(chatID)
deleteFn := c.deleteMessageFn
if deleteFn == nil {
deleteFn = c.DeleteMessage
}
_ = deleteFn(ctx, chatID, messageID)
}
func (c *PicoChannel) finalizeTrackedToolFeedbackMessage(
@ -442,6 +457,7 @@ func (c *PicoChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessag
if !c.IsRunning() {
return nil, channels.ErrNotRunning
}
trackedMsgID, hasTrackedMsg := c.currentToolFeedbackMessage(msg.ChatID)
store := c.GetMediaStore()
if store == nil {
@ -517,6 +533,9 @@ func (c *PicoChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessag
if err := c.broadcastToSession(msg.ChatID, outMsg); err != nil {
return nil, err
}
if hasTrackedMsg {
c.dismissTrackedToolFeedbackMessage(ctx, msg.ChatID, trackedMsgID)
}
return []string{msgID}, nil
}

View file

@ -4,12 +4,16 @@ import (
"context"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
"github.com/gorilla/websocket"
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels"
@ -60,6 +64,32 @@ func TestFinalizeTrackedToolFeedbackMessage_StopsTrackingBeforeEdit(t *testing.T
}
}
func TestDismissTrackedToolFeedbackMessage_DeletesProgressMessage(t *testing.T) {
ch := &PicoChannel{
progress: channels.NewToolFeedbackAnimator(nil),
}
ch.RecordToolFeedbackMessage("pico:chat-1", "msg-1", "🔧 `read_file`")
var deleted struct {
chatID string
messageID string
}
ch.deleteMessageFn = func(_ context.Context, chatID string, messageID string) error {
deleted.chatID = chatID
deleted.messageID = messageID
return nil
}
ch.DismissToolFeedbackMessage(context.Background(), "pico:chat-1")
if deleted.chatID != "pico:chat-1" || deleted.messageID != "msg-1" {
t.Fatalf("unexpected delete target: %+v", deleted)
}
if _, ok := ch.currentToolFeedbackMessage("pico:chat-1"); ok {
t.Fatal("expected tracked tool feedback to be cleared after dismissal")
}
}
func TestCreateAndAddConnection_RespectsMaxConnectionsConcurrently(t *testing.T) {
ch := newTestPicoChannel(t)
@ -197,6 +227,75 @@ func TestSendMedia_ResolvesMediaBeforeDelivery(t *testing.T) {
}
}
func TestSendMedia_DismissesTrackedToolFeedbackMessage(t *testing.T) {
ch := newTestPicoChannel(t)
store := media.NewFileMediaStore()
ch.SetMediaStore(store)
if err := ch.Start(context.Background()); err != nil {
t.Fatalf("Start() error = %v", err)
}
defer ch.Stop(context.Background())
clientConn, received, cleanup := newTestPicoWebSocket(t)
defer cleanup()
ch.addConnForTest(&picoConn{id: "conn-1", conn: clientConn, sessionID: "sess-1"})
localPath := filepath.Join(t.TempDir(), "report.txt")
if err := os.WriteFile(localPath, []byte("attachment body"), 0o600); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
ref, err := store.Store(localPath, media.MediaMeta{
Filename: "report.txt",
ContentType: "text/plain",
}, "test-scope")
if err != nil {
t.Fatalf("Store() error = %v", err)
}
ch.RecordToolFeedbackMessage("pico:sess-1", "msg-progress", "🔧 `read_file`")
var deleted struct {
chatID string
messageID string
}
ch.deleteMessageFn = func(_ context.Context, chatID string, messageID string) error {
deleted.chatID = chatID
deleted.messageID = messageID
return nil
}
_, err = ch.SendMedia(context.Background(), bus.OutboundMediaMessage{
ChatID: "pico:sess-1",
Parts: []bus.MediaPart{{
Ref: ref,
Type: "file",
Filename: "report.txt",
ContentType: "text/plain",
}},
})
if err != nil {
t.Fatalf("SendMedia() error = %v", err)
}
select {
case msg := <-received:
if msg.Type != TypeMessageCreate {
t.Fatalf("message type = %q, want %q", msg.Type, TypeMessageCreate)
}
case <-time.After(time.Second):
t.Fatal("expected media message to be delivered")
}
if deleted.chatID != "pico:sess-1" || deleted.messageID != "msg-progress" {
t.Fatalf("unexpected delete target: %+v", deleted)
}
if _, ok := ch.currentToolFeedbackMessage("pico:sess-1"); ok {
t.Fatal("expected tracked tool feedback to be cleared after media delivery")
}
}
func TestPicoDownloadURLForRef(t *testing.T) {
got, err := picoDownloadURLForRef("media://attachment-1")
if err != nil {
@ -268,3 +367,39 @@ func (c *PicoChannel) addConnForTest(pc *picoConn) {
}
bySession[pc.id] = pc
}
func newTestPicoWebSocket(t *testing.T) (*websocket.Conn, <-chan PicoMessage, func()) {
t.Helper()
received := make(chan PicoMessage, 4)
upgrader := websocket.Upgrader{}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
t.Errorf("Upgrade() error = %v", err)
return
}
defer conn.Close()
for {
var msg PicoMessage
if err := conn.ReadJSON(&msg); err != nil {
return
}
received <- msg
}
}))
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
clientConn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
if err != nil {
server.Close()
t.Fatalf("Dial() error = %v", err)
}
cleanup := func() {
clientConn.Close()
server.Close()
}
return clientConn, received, cleanup
}

View file

@ -12,6 +12,7 @@ const (
// TypeMessageCreate is sent from server to client.
TypeMessageCreate = "message.create"
TypeMessageUpdate = "message.update"
TypeMessageDelete = "message.delete"
TypeMediaCreate = "media.create"
TypeTypingStart = "typing.start"
TypeTypingStop = "typing.stop"

View file

@ -515,6 +515,7 @@ func visibleSessionMessages(messages []providers.Message, toolFeedbackMaxArgsLen
// must remain visible in restored session history.
if len(msg.ToolCalls) > 0 &&
len(msg.Media) == 0 &&
len(attachments) == 0 &&
assistantToolCallContentDuplicated(msg.Content, toolSummaryMessages, visibleToolMessages) {
continue
}

View file

@ -897,6 +897,90 @@ func TestHandleGetSession_PreservesMediaWhenAssistantToolCallContentDuplicatesSu
}
}
func TestHandleGetSession_PreservesAttachmentsWhenAssistantToolCallContentDuplicatesSummary(t *testing.T) {
configPath, cleanup := setupOAuthTestEnv(t)
defer cleanup()
dir := sessionsTestDir(t, configPath)
store, err := memory.NewJSONLStore(dir)
if err != nil {
t.Fatalf("NewJSONLStore() error = %v", err)
}
sessionKey := picoSessionPrefix + "detail-tool-summary-duplicate-content-with-attachments"
for _, msg := range []providers.Message{
{Role: "user", Content: "check report"},
{
Role: "assistant",
Content: "Reviewing the generated report.",
Attachments: []providers.Attachment{{
Type: "file",
URL: "https://example.com/report.txt",
Filename: "report.txt",
ContentType: "text/plain",
}},
ToolCalls: []providers.ToolCall{
{
ID: "call_1",
Type: "function",
Function: &providers.FunctionCall{
Name: "read_file",
Arguments: `{"path":"report.txt"}`,
},
ExtraContent: &providers.ExtraContent{
ToolFeedbackExplanation: "Reviewing the generated report.",
},
},
},
},
} {
if err := store.AddFullMessage(nil, sessionKey, msg); err != nil {
t.Fatalf("AddFullMessage() error = %v", err)
}
}
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
rec := httptest.NewRecorder()
req := httptest.NewRequest(
http.MethodGet,
"/api/sessions/detail-tool-summary-duplicate-content-with-attachments",
nil,
)
mux.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusOK, rec.Body.String())
}
var resp struct {
Messages []sessionChatMessage `json:"messages"`
}
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
t.Fatalf("Unmarshal() error = %v", err)
}
if len(resp.Messages) != 3 {
t.Fatalf("len(resp.Messages) = %d, want 3", len(resp.Messages))
}
if !strings.Contains(resp.Messages[1].Content, "`read_file`") {
t.Fatalf("tool summary message = %#v, want read_file summary", resp.Messages[1])
}
if resp.Messages[2].Role != "assistant" {
t.Fatalf("assistant message role = %q, want assistant", resp.Messages[2].Role)
}
if resp.Messages[2].Content != "Reviewing the generated report." {
t.Fatalf("assistant content = %q, want preserved duplicated content", resp.Messages[2].Content)
}
if len(resp.Messages[2].Attachments) != 1 {
t.Fatalf("len(assistant.Attachments) = %d, want 1", len(resp.Messages[2].Attachments))
}
if resp.Messages[2].Attachments[0].URL != "https://example.com/report.txt" {
t.Fatalf("attachment url = %q, want report URL", resp.Messages[2].Attachments[0].URL)
}
}
func TestHandleGetSession_UsesConfiguredToolFeedbackMaxArgsLength(t *testing.T) {
configPath, cleanup := setupOAuthTestEnv(t)
defer cleanup()

View file

@ -157,6 +157,18 @@ export function handlePicoMessage(
break
}
case "message.delete": {
const messageId = payload.message_id as string
if (!messageId) {
break
}
updateChatStore((prev) => ({
messages: prev.messages.filter((msg) => msg.id !== messageId),
}))
break
}
case "typing.start":
updateChatStore({ isTyping: true })
break