diff --git a/pkg/agent/context_seahorse_test.go b/pkg/agent/context_seahorse_test.go index e405ef94..05c83183 100644 --- a/pkg/agent/context_seahorse_test.go +++ b/pkg/agent/context_seahorse_test.go @@ -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) { msg := seahorse.Message{ Role: "tool", diff --git a/pkg/seahorse/short_assembler.go b/pkg/seahorse/short_assembler.go index 0bfd66a6..5533512a 100644 --- a/pkg/seahorse/short_assembler.go +++ b/pkg/seahorse/short_assembler.go @@ -68,24 +68,33 @@ func (a *Assembler) Assemble(ctx context.Context, convID int64, input AssembleIn freshTailTokens += r.tokenCount } - // If the protected tail alone exceeds budget, trim from the oldest end of - // the tail until the newest items fit within the requested budget. + // If the protected tail alone exceeds budget, trim from the oldest end at + // 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 { originalTailCount := len(freshTail) originalFreshTailTokens := freshTailTokens - for freshTailTokens > input.Budget && len(freshTail) > 0 { - freshTailTokens -= freshTail[0].tokenCount - freshTail = freshTail[1:] - } - logger.InfoCF("seahorse", "assemble: trimmed fresh tail to budget", map[string]any{ + var preservedActiveTurn bool + freshTail, freshTailTokens, preservedActiveTurn = trimFreshTailToSafeBudget(freshTail, input.Budget) + logFields := map[string]any{ "budget": input.Budget, "fresh_tail_tokens": freshTailTokens, "fresh_tail_count": len(freshTail), "trimmed_fresh_items": originalTailCount - len(freshTail), "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 + if remainingBudget < 0 { + remainingBudget = 0 + } var selected []resolvedItem evictableTokens := 0 @@ -189,6 +198,81 @@ func (a *Assembler) Assemble(ctx context.Context, convID int64, input AssembleIn }, 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. func (a *Assembler) resolveItem(ctx context.Context, item ContextItem) (resolvedItem, error) { if item.ItemType == "message" { diff --git a/pkg/seahorse/short_assembler_test.go b/pkg/seahorse/short_assembler_test.go index 81918afa..6472bc36 100644 --- a/pkg/seahorse/short_assembler_test.go +++ b/pkg/seahorse/short_assembler_test.go @@ -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) { s, convID := setupAssemblerStore(t) ctx := context.Background()