fix session recovery for stale tool-call tails
This commit is contained in:
parent
15a70ac45c
commit
2455cf3c0f
4 changed files with 112 additions and 2 deletions
|
|
@ -38,7 +38,7 @@ func (b *JSONLBackend) GetHistory(key string) []providers.Message {
|
||||||
log.Printf("session: get history: %v", err)
|
log.Printf("session: get history: %v", err)
|
||||||
return []providers.Message{}
|
return []providers.Message{}
|
||||||
}
|
}
|
||||||
return msgs
|
return sanitizeRecoveredHistory(msgs)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *JSONLBackend) GetSummary(key string) string {
|
func (b *JSONLBackend) GetSummary(key string) string {
|
||||||
|
|
|
||||||
|
|
@ -97,7 +97,7 @@ func (sm *SessionManager) GetHistory(key string) []providers.Message {
|
||||||
|
|
||||||
history := make([]providers.Message, len(session.Messages))
|
history := make([]providers.Message, len(session.Messages))
|
||||||
copy(history, session.Messages)
|
copy(history, session.Messages)
|
||||||
return history
|
return sanitizeRecoveredHistory(history)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sm *SessionManager) GetSummary(key string) string {
|
func (sm *SessionManager) GetSummary(key string) string {
|
||||||
|
|
|
||||||
52
pkg/session/recovery.go
Normal file
52
pkg/session/recovery.go
Normal file
|
|
@ -0,0 +1,52 @@
|
||||||
|
package session
|
||||||
|
|
||||||
|
import "github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
|
||||||
|
// sanitizeRecoveredHistory drops any dangling tail that starts with an
|
||||||
|
// assistant tool-call message whose tool results were never fully persisted.
|
||||||
|
//
|
||||||
|
// This prevents a restarted agent from restoring an unfinished runtime state
|
||||||
|
// (assistant tool calls plus later steering/user messages) as if it were valid
|
||||||
|
// history. We keep completed tool-call sequences intact and only trim the
|
||||||
|
// incomplete suffix.
|
||||||
|
func sanitizeRecoveredHistory(history []providers.Message) []providers.Message {
|
||||||
|
for i := 0; i < len(history); i++ {
|
||||||
|
msg := history[i]
|
||||||
|
if msg.Role != "assistant" || len(msg.ToolCalls) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
expected := make(map[string]bool, len(msg.ToolCalls))
|
||||||
|
for _, tc := range msg.ToolCalls {
|
||||||
|
expected[tc.ID] = false
|
||||||
|
}
|
||||||
|
|
||||||
|
j := i + 1
|
||||||
|
for ; j < len(history); j++ {
|
||||||
|
next := history[j]
|
||||||
|
if next.ToolCallID == "" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if _, ok := expected[next.ToolCallID]; ok {
|
||||||
|
expected[next.ToolCallID] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
complete := true
|
||||||
|
for _, found := range expected {
|
||||||
|
if !found {
|
||||||
|
complete = false
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !complete {
|
||||||
|
return append([]providers.Message(nil), history[:i]...)
|
||||||
|
}
|
||||||
|
|
||||||
|
if j > i+1 {
|
||||||
|
i = j - 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return append([]providers.Message(nil), history...)
|
||||||
|
}
|
||||||
58
pkg/session/recovery_test.go
Normal file
58
pkg/session/recovery_test.go
Normal file
|
|
@ -0,0 +1,58 @@
|
||||||
|
package session
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSanitizeRecoveredHistory_DropsDanglingToolCallTail(t *testing.T) {
|
||||||
|
history := []providers.Message{
|
||||||
|
{Role: "user", Content: "hello"},
|
||||||
|
{Role: "assistant", Content: "working", ToolCalls: []providers.ToolCall{{ID: "call_1"}}},
|
||||||
|
{Role: "user", Content: "?"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := sanitizeRecoveredHistory(history)
|
||||||
|
if len(got) != 1 {
|
||||||
|
t.Fatalf("len(got) = %d, want 1", len(got))
|
||||||
|
}
|
||||||
|
if got[0].Content != "hello" {
|
||||||
|
t.Fatalf("got[0].Content = %q, want hello", got[0].Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSanitizeRecoveredHistory_KeepsCompletedToolCallSequence(t *testing.T) {
|
||||||
|
history := []providers.Message{
|
||||||
|
{Role: "user", Content: "hello"},
|
||||||
|
{Role: "assistant", Content: "working", ToolCalls: []providers.ToolCall{{ID: "call_1"}, {ID: "call_2"}}},
|
||||||
|
{Role: "tool", ToolCallID: "call_1", Content: "done 1"},
|
||||||
|
{Role: "tool", ToolCallID: "call_2", Content: "done 2"},
|
||||||
|
{Role: "assistant", Content: "all set"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := sanitizeRecoveredHistory(history)
|
||||||
|
if len(got) != len(history) {
|
||||||
|
t.Fatalf("len(got) = %d, want %d", len(got), len(history))
|
||||||
|
}
|
||||||
|
if got[len(got)-1].Content != "all set" {
|
||||||
|
t.Fatalf("last content = %q, want all set", got[len(got)-1].Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSanitizeRecoveredHistory_DropsPartialToolResultsAndFollowingMessages(t *testing.T) {
|
||||||
|
history := []providers.Message{
|
||||||
|
{Role: "user", Content: "hello"},
|
||||||
|
{Role: "assistant", Content: "working", ToolCalls: []providers.ToolCall{{ID: "call_1"}, {ID: "call_2"}}},
|
||||||
|
{Role: "tool", ToolCallID: "call_1", Content: "done 1"},
|
||||||
|
{Role: "user", Content: "still there?"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got := sanitizeRecoveredHistory(history)
|
||||||
|
if len(got) != 1 {
|
||||||
|
t.Fatalf("len(got) = %d, want 1", len(got))
|
||||||
|
}
|
||||||
|
if got[0].Content != "hello" {
|
||||||
|
t.Fatalf("got[0].Content = %q, want hello", got[0].Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue