fix(pico): clean up tool feedback progress state
This commit is contained in:
parent
7d03023987
commit
a36d36a368
6 changed files with 252 additions and 0 deletions
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue