fix: make TruncateHistory pair-aware to prevent orphaned tool messages
When the cut point falls inside a tool_use/tool_result group, snap backward to include the full group.
This commit is contained in:
parent
3664ab3914
commit
3d1d480975
2 changed files with 64 additions and 1 deletions
|
|
@ -141,7 +141,17 @@ func (sm *SessionManager) TruncateHistory(key string, keepLast int) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
session.Messages = session.Messages[len(session.Messages)-keepLast:]
|
// Start with the naive cut point
|
||||||
|
cutIdx := len(session.Messages) - keepLast
|
||||||
|
|
||||||
|
// Snap the cut point backward: if the message at cutIdx is a "tool"
|
||||||
|
// (tool_result), walk backward to include the preceding assistant
|
||||||
|
// message that owns the tool_call group.
|
||||||
|
for cutIdx > 0 && session.Messages[cutIdx].Role == "tool" {
|
||||||
|
cutIdx--
|
||||||
|
}
|
||||||
|
|
||||||
|
session.Messages = session.Messages[cutIdx:]
|
||||||
session.Updated = time.Now()
|
session.Updated = time.Now()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@ import (
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestSanitizeFilename(t *testing.T) {
|
func TestSanitizeFilename(t *testing.T) {
|
||||||
|
|
@ -72,3 +74,54 @@ func TestSave_RejectsPathTraversal(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestTruncateHistory_PreservesToolPairs(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
sm := NewSessionManager(tmpDir)
|
||||||
|
|
||||||
|
key := "test-truncate"
|
||||||
|
sm.GetOrCreate(key)
|
||||||
|
|
||||||
|
// Build: user, assistant+2tools, tool1, tool2, user, assistant = 6 messages
|
||||||
|
sm.AddFullMessage(key, providers.Message{Role: "user", Content: "q1"})
|
||||||
|
sm.AddFullMessage(key, providers.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: "checking",
|
||||||
|
ToolCalls: []providers.ToolCall{
|
||||||
|
{ID: "c1", Name: "exec"},
|
||||||
|
{ID: "c2", Name: "web"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
sm.AddFullMessage(key, providers.Message{Role: "tool", Content: "r1", ToolCallID: "c1"})
|
||||||
|
sm.AddFullMessage(key, providers.Message{Role: "tool", Content: "r2", ToolCallID: "c2"})
|
||||||
|
sm.AddFullMessage(key, providers.Message{Role: "user", Content: "q2"})
|
||||||
|
sm.AddFullMessage(key, providers.Message{Role: "assistant", Content: "done"})
|
||||||
|
|
||||||
|
// keepLast=4 naively keeps: [tool2, user, assistant_done] or similar
|
||||||
|
// which orphans tool messages. Should snap to include/exclude full group.
|
||||||
|
sm.TruncateHistory(key, 4)
|
||||||
|
history := sm.GetHistory(key)
|
||||||
|
|
||||||
|
// Verify no orphaned tool messages
|
||||||
|
toolCallIDs := map[string]bool{}
|
||||||
|
toolResultIDs := map[string]bool{}
|
||||||
|
for _, m := range history {
|
||||||
|
for _, tc := range m.ToolCalls {
|
||||||
|
toolCallIDs[tc.ID] = true
|
||||||
|
}
|
||||||
|
if m.Role == "tool" && m.ToolCallID != "" {
|
||||||
|
toolResultIDs[m.ToolCallID] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for id := range toolResultIDs {
|
||||||
|
if !toolCallIDs[id] {
|
||||||
|
t.Errorf("orphaned tool_result %q after truncation", id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for id := range toolCallIDs {
|
||||||
|
if !toolResultIDs[id] {
|
||||||
|
t.Errorf("orphaned tool_call %q after truncation", id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue