fix(seahorse): preserve active tool-call turn when trimming fresh tail
This commit is contained in:
parent
fe7ded5c13
commit
f0dcba8c5a
3 changed files with 222 additions and 8 deletions
|
|
@ -280,6 +280,74 @@ func TestSeahorseToProviderMessagesWithToolCalls(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSeahorseAssemblePreservesActiveToolTurnAcrossSanitization(t *testing.T) {
|
||||||
|
engine, err := seahorse.NewEngine(seahorse.Config{
|
||||||
|
DBPath: t.TempDir() + "/seahorse.db",
|
||||||
|
}, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewEngine: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
sessionKey := "test:active-tool-turn"
|
||||||
|
_, err = engine.Ingest(ctx, sessionKey, []seahorse.Message{
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: "older context",
|
||||||
|
TokenCount: 20,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Role: "user",
|
||||||
|
Content: "inspect the file",
|
||||||
|
TokenCount: 5,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
TokenCount: 5,
|
||||||
|
Parts: []seahorse.MessagePart{{
|
||||||
|
Type: "tool_use",
|
||||||
|
Name: "read_file",
|
||||||
|
Arguments: `{"path":"/tmp/test.txt"}`,
|
||||||
|
ToolCallID: "tc_1",
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Role: "tool",
|
||||||
|
TokenCount: 200,
|
||||||
|
Parts: []seahorse.MessagePart{{
|
||||||
|
Type: "tool_result",
|
||||||
|
ToolCallID: "tc_1",
|
||||||
|
Text: "very large tool output",
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: "done",
|
||||||
|
TokenCount: 5,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Ingest: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := engine.Assemble(ctx, sessionKey, seahorse.AssembleInput{Budget: 210})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Assemble: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sanitized := sanitizeHistoryForProvider(seahorseToProviderMessages(result))
|
||||||
|
if len(sanitized) != 4 {
|
||||||
|
t.Fatalf("sanitized history len = %d, want 4 protected-turn messages", len(sanitized))
|
||||||
|
}
|
||||||
|
assertRoles(t, sanitized, "user", "assistant", "tool", "assistant")
|
||||||
|
if len(sanitized[1].ToolCalls) != 1 || sanitized[1].ToolCalls[0].ID != "tc_1" {
|
||||||
|
t.Fatalf("assistant tool calls = %+v, want preserved tool call tc_1", sanitized[1].ToolCalls)
|
||||||
|
}
|
||||||
|
if sanitized[2].ToolCallID != "tc_1" {
|
||||||
|
t.Fatalf("tool result id = %q, want tc_1", sanitized[2].ToolCallID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestSeahorseToProviderMessagesToolResult(t *testing.T) {
|
func TestSeahorseToProviderMessagesToolResult(t *testing.T) {
|
||||||
msg := seahorse.Message{
|
msg := seahorse.Message{
|
||||||
Role: "tool",
|
Role: "tool",
|
||||||
|
|
|
||||||
|
|
@ -68,24 +68,33 @@ func (a *Assembler) Assemble(ctx context.Context, convID int64, input AssembleIn
|
||||||
freshTailTokens += r.tokenCount
|
freshTailTokens += r.tokenCount
|
||||||
}
|
}
|
||||||
|
|
||||||
// If the protected tail alone exceeds budget, trim from the oldest end of
|
// If the protected tail alone exceeds budget, trim from the oldest end at
|
||||||
// the tail until the newest items fit within the requested budget.
|
// provider-safe boundaries. The rebuild path later sanitizes leading
|
||||||
|
// assistant(tool_calls)/tool messages, so splitting the active turn here can
|
||||||
|
// silently discard the very context we are trying to protect.
|
||||||
if freshTailTokens > input.Budget {
|
if freshTailTokens > input.Budget {
|
||||||
originalTailCount := len(freshTail)
|
originalTailCount := len(freshTail)
|
||||||
originalFreshTailTokens := freshTailTokens
|
originalFreshTailTokens := freshTailTokens
|
||||||
for freshTailTokens > input.Budget && len(freshTail) > 0 {
|
var preservedActiveTurn bool
|
||||||
freshTailTokens -= freshTail[0].tokenCount
|
freshTail, freshTailTokens, preservedActiveTurn = trimFreshTailToSafeBudget(freshTail, input.Budget)
|
||||||
freshTail = freshTail[1:]
|
logFields := map[string]any{
|
||||||
}
|
|
||||||
logger.InfoCF("seahorse", "assemble: trimmed fresh tail to budget", map[string]any{
|
|
||||||
"budget": input.Budget,
|
"budget": input.Budget,
|
||||||
"fresh_tail_tokens": freshTailTokens,
|
"fresh_tail_tokens": freshTailTokens,
|
||||||
"fresh_tail_count": len(freshTail),
|
"fresh_tail_count": len(freshTail),
|
||||||
"trimmed_fresh_items": originalTailCount - len(freshTail),
|
"trimmed_fresh_items": originalTailCount - len(freshTail),
|
||||||
"original_fresh_tokens": originalFreshTailTokens,
|
"original_fresh_tokens": originalFreshTailTokens,
|
||||||
})
|
"preserved_active_turn": preservedActiveTurn,
|
||||||
|
}
|
||||||
|
if preservedActiveTurn {
|
||||||
|
logger.WarnCF("seahorse", "assemble: preserving active turn over budget", logFields)
|
||||||
|
} else {
|
||||||
|
logger.InfoCF("seahorse", "assemble: trimmed fresh tail to safe boundary", logFields)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
remainingBudget := input.Budget - freshTailTokens
|
remainingBudget := input.Budget - freshTailTokens
|
||||||
|
if remainingBudget < 0 {
|
||||||
|
remainingBudget = 0
|
||||||
|
}
|
||||||
|
|
||||||
var selected []resolvedItem
|
var selected []resolvedItem
|
||||||
evictableTokens := 0
|
evictableTokens := 0
|
||||||
|
|
@ -189,6 +198,81 @@ func (a *Assembler) Assemble(ctx context.Context, convID int64, input AssembleIn
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func trimFreshTailToSafeBudget(tail []resolvedItem, budget int) ([]resolvedItem, int, bool) {
|
||||||
|
tailTokens := resolvedItemsTokenCount(tail)
|
||||||
|
if tailTokens <= budget {
|
||||||
|
return tail, tailTokens, false
|
||||||
|
}
|
||||||
|
|
||||||
|
latestTurnStart := lastUserMessageIndex(tail)
|
||||||
|
if latestTurnStart >= 0 {
|
||||||
|
latestTurnTokens := resolvedItemsTokenCount(tail[latestTurnStart:])
|
||||||
|
if latestTurnTokens > budget {
|
||||||
|
return tail[latestTurnStart:], latestTurnTokens, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
start := 0
|
||||||
|
for tailTokens > budget && start < len(tail) {
|
||||||
|
tailTokens -= tail[start].tokenCount
|
||||||
|
start++
|
||||||
|
}
|
||||||
|
for start < len(tail) && !isProviderSafeHistoryStart(tail[start:]) {
|
||||||
|
tailTokens -= tail[start].tokenCount
|
||||||
|
start++
|
||||||
|
}
|
||||||
|
|
||||||
|
return tail[start:], tailTokens, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolvedItemsTokenCount(items []resolvedItem) int {
|
||||||
|
total := 0
|
||||||
|
for _, item := range items {
|
||||||
|
total += item.tokenCount
|
||||||
|
}
|
||||||
|
return total
|
||||||
|
}
|
||||||
|
|
||||||
|
func lastUserMessageIndex(items []resolvedItem) int {
|
||||||
|
for i := len(items) - 1; i >= 0; i-- {
|
||||||
|
if items[i].itemType != "message" || items[i].message == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if items[i].message.Role == "user" {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
|
func isProviderSafeHistoryStart(items []resolvedItem) bool {
|
||||||
|
for _, item := range items {
|
||||||
|
if item.itemType != "message" || item.message == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if item.message.Role == "tool" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if item.message.Role == "assistant" && messageHasToolUse(item.message) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func messageHasToolUse(msg *Message) bool {
|
||||||
|
if msg == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, part := range msg.Parts {
|
||||||
|
if part.Type == "tool_use" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// resolveItem loads the full message or summary for a context item.
|
// resolveItem loads the full message or summary for a context item.
|
||||||
func (a *Assembler) resolveItem(ctx context.Context, item ContextItem) (resolvedItem, error) {
|
func (a *Assembler) resolveItem(ctx context.Context, item ContextItem) (resolvedItem, error) {
|
||||||
if item.ItemType == "message" {
|
if item.ItemType == "message" {
|
||||||
|
|
|
||||||
|
|
@ -171,6 +171,68 @@ func TestAssemblerBudgetEvictsOldest(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestAssemblerBudgetPreservesLatestToolTurnWhenItExceedsBudget(t *testing.T) {
|
||||||
|
s, convID := setupAssemblerStore(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
oldMsg, _ := s.AddMessage(ctx, convID, "assistant", "older context", 20)
|
||||||
|
userMsg, _ := s.AddMessage(ctx, convID, "user", "inspect the file", 5)
|
||||||
|
assistantToolMsg, _ := s.AddMessageWithParts(ctx, convID, "assistant", []MessagePart{
|
||||||
|
{
|
||||||
|
Type: "tool_use",
|
||||||
|
Name: "read_file",
|
||||||
|
Arguments: `{"path":"/tmp/test.txt"}`,
|
||||||
|
ToolCallID: "tc_1",
|
||||||
|
},
|
||||||
|
}, 5)
|
||||||
|
toolResultMsg, _ := s.AddMessageWithParts(ctx, convID, "tool", []MessagePart{
|
||||||
|
{
|
||||||
|
Type: "tool_result",
|
||||||
|
ToolCallID: "tc_1",
|
||||||
|
Text: "very large tool output",
|
||||||
|
},
|
||||||
|
}, 200)
|
||||||
|
finalAssistantMsg, _ := s.AddMessage(ctx, convID, "assistant", "done", 5)
|
||||||
|
|
||||||
|
s.UpsertContextItems(ctx, convID, []ContextItem{
|
||||||
|
{Ordinal: 100, ItemType: "message", MessageID: oldMsg.ID, TokenCount: 20},
|
||||||
|
{Ordinal: 200, ItemType: "message", MessageID: userMsg.ID, TokenCount: 5},
|
||||||
|
{Ordinal: 300, ItemType: "message", MessageID: assistantToolMsg.ID, TokenCount: 5},
|
||||||
|
{Ordinal: 400, ItemType: "message", MessageID: toolResultMsg.ID, TokenCount: 200},
|
||||||
|
{Ordinal: 500, ItemType: "message", MessageID: finalAssistantMsg.ID, TokenCount: 5},
|
||||||
|
})
|
||||||
|
|
||||||
|
a := &Assembler{store: s, config: Config{}}
|
||||||
|
result, err := a.Assemble(ctx, convID, AssembleInput{Budget: 210})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Assemble: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(result.Messages) != 4 {
|
||||||
|
t.Fatalf("Messages = %d, want 4 protected-turn messages", len(result.Messages))
|
||||||
|
}
|
||||||
|
if result.Messages[0].ID != userMsg.ID {
|
||||||
|
t.Fatalf("first message ID = %d, want current user message %d", result.Messages[0].ID, userMsg.ID)
|
||||||
|
}
|
||||||
|
if result.Messages[1].ID != assistantToolMsg.ID {
|
||||||
|
t.Fatalf("second message ID = %d, want assistant tool-call %d", result.Messages[1].ID, assistantToolMsg.ID)
|
||||||
|
}
|
||||||
|
if result.Messages[2].ID != toolResultMsg.ID {
|
||||||
|
t.Fatalf("third message ID = %d, want tool result %d", result.Messages[2].ID, toolResultMsg.ID)
|
||||||
|
}
|
||||||
|
if result.Messages[3].ID != finalAssistantMsg.ID {
|
||||||
|
t.Fatalf("fourth message ID = %d, want final assistant %d", result.Messages[3].ID, finalAssistantMsg.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
totalTokens := 0
|
||||||
|
for _, msg := range result.Messages {
|
||||||
|
totalTokens += msg.TokenCount
|
||||||
|
}
|
||||||
|
if totalTokens <= 210 {
|
||||||
|
t.Fatalf("assembled tokens = %d, want protected turn to remain over budget", totalTokens)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestAssemblerBudgetFitsAll(t *testing.T) {
|
func TestAssemblerBudgetFitsAll(t *testing.T) {
|
||||||
s, convID := setupAssemblerStore(t)
|
s, convID := setupAssemblerStore(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue