fix(seahorse): enforce budget on fresh tail and rebuild paths
This commit is contained in:
parent
941bac2332
commit
1502636bf0
6 changed files with 216 additions and 27 deletions
|
|
@ -115,3 +115,55 @@ func isOverContextBudget(
|
||||||
|
|
||||||
return total > contextWindow
|
return total > contextWindow
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// trimHistoryToFitContextWindow rebuilds the prompt from progressively newer
|
||||||
|
// history slices until it fits within the context window. Oldest complete turns
|
||||||
|
// are dropped first so tool-call sequences remain intact.
|
||||||
|
func trimHistoryToFitContextWindow(
|
||||||
|
history []providers.Message,
|
||||||
|
build func([]providers.Message) []providers.Message,
|
||||||
|
contextWindow int,
|
||||||
|
toolDefs []providers.ToolDefinition,
|
||||||
|
maxTokens int,
|
||||||
|
) ([]providers.Message, []providers.Message, bool) {
|
||||||
|
messages := build(history)
|
||||||
|
if !isOverContextBudget(contextWindow, messages, toolDefs, maxTokens) {
|
||||||
|
return history, messages, true
|
||||||
|
}
|
||||||
|
|
||||||
|
trimmedHistory := append([]providers.Message(nil), history...)
|
||||||
|
for len(trimmedHistory) > 0 {
|
||||||
|
dropUntil := nextHistoryTrimStart(trimmedHistory)
|
||||||
|
if dropUntil <= 0 || dropUntil >= len(trimmedHistory) {
|
||||||
|
trimmedHistory = nil
|
||||||
|
} else {
|
||||||
|
trimmedHistory = append([]providers.Message(nil), trimmedHistory[dropUntil:]...)
|
||||||
|
}
|
||||||
|
|
||||||
|
messages = build(trimmedHistory)
|
||||||
|
if !isOverContextBudget(contextWindow, messages, toolDefs, maxTokens) {
|
||||||
|
return trimmedHistory, messages, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, messages, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func nextHistoryTrimStart(history []providers.Message) int {
|
||||||
|
if len(history) == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
turns := parseTurnBoundaries(history)
|
||||||
|
if len(turns) >= 2 {
|
||||||
|
return turns[1]
|
||||||
|
}
|
||||||
|
if len(turns) == 1 {
|
||||||
|
if turns[0] > 0 {
|
||||||
|
return turns[0]
|
||||||
|
}
|
||||||
|
return len(history)
|
||||||
|
}
|
||||||
|
|
||||||
|
return len(history)
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -844,3 +844,64 @@ func TestIsOverContextBudget_RealisticSession(t *testing.T) {
|
||||||
t.Error("realistic session should exceed 500 context window")
|
t.Error("realistic session should exceed 500 context window")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestTrimHistoryToFitContextWindow_DropsOldestTurns(t *testing.T) {
|
||||||
|
history := []providers.Message{
|
||||||
|
msgUser(strings.Repeat("u1 ", 120)),
|
||||||
|
msgAssistant(strings.Repeat("a1 ", 120)),
|
||||||
|
msgUser(strings.Repeat("u2 ", 120)),
|
||||||
|
msgAssistant(strings.Repeat("a2 ", 120)),
|
||||||
|
msgUser(strings.Repeat("u3 ", 120)),
|
||||||
|
msgAssistant(strings.Repeat("a3 ", 120)),
|
||||||
|
}
|
||||||
|
|
||||||
|
build := func(history []providers.Message) []providers.Message {
|
||||||
|
return append([]providers.Message(nil), history...)
|
||||||
|
}
|
||||||
|
|
||||||
|
trimmedHistory, messages, fit := trimHistoryToFitContextWindow(
|
||||||
|
history,
|
||||||
|
build,
|
||||||
|
700,
|
||||||
|
nil,
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
if !fit {
|
||||||
|
t.Fatal("expected trimmed history to fit context window")
|
||||||
|
}
|
||||||
|
if len(trimmedHistory) != 4 {
|
||||||
|
t.Fatalf("trimmed history len = %d, want 4", len(trimmedHistory))
|
||||||
|
}
|
||||||
|
if trimmedHistory[0].Content != history[2].Content {
|
||||||
|
t.Fatalf("first kept message = %q, want second turn start", trimmedHistory[0].Content)
|
||||||
|
}
|
||||||
|
if isOverContextBudget(700, messages, nil, 0) {
|
||||||
|
t.Fatal("trimmed messages should be within budget")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTrimHistoryToFitContextWindow_ClearsSingleOversizedTurn(t *testing.T) {
|
||||||
|
history := []providers.Message{
|
||||||
|
msgUser(strings.Repeat("oversized ", 200)),
|
||||||
|
msgAssistant(strings.Repeat("oversized ", 200)),
|
||||||
|
}
|
||||||
|
|
||||||
|
trimmedHistory, messages, fit := trimHistoryToFitContextWindow(
|
||||||
|
history,
|
||||||
|
func(history []providers.Message) []providers.Message {
|
||||||
|
return append([]providers.Message(nil), history...)
|
||||||
|
},
|
||||||
|
200,
|
||||||
|
nil,
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
if !fit {
|
||||||
|
t.Fatal("expected empty history rebuild to fit context window")
|
||||||
|
}
|
||||||
|
if len(trimmedHistory) != 0 {
|
||||||
|
t.Fatalf("trimmed history len = %d, want 0", len(trimmedHistory))
|
||||||
|
}
|
||||||
|
if len(messages) != 0 {
|
||||||
|
t.Fatalf("messages len = %d, want 0", len(messages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -369,14 +369,50 @@ func (p *Pipeline) CallLLM(
|
||||||
contextualSkills = ts.agent.ContextBuilder.ResolveActiveSkillsForContext(ts.activeSkills)
|
contextualSkills = ts.agent.ContextBuilder.ResolveActiveSkillsForContext(ts.activeSkills)
|
||||||
}
|
}
|
||||||
ts.recordSkillContextSnapshot(skillContextTriggerContextRetryRebuild, contextualSkills)
|
ts.recordSkillContextSnapshot(skillContextTriggerContextRetryRebuild, contextualSkills)
|
||||||
rebuildPromptReq := promptBuildRequestForTurn(ts, exec.history, exec.summary, "", nil)
|
buildMessages := func(trimmedHistory []providers.Message) []providers.Message {
|
||||||
rebuildPromptReq.ActiveSkills = append([]string(nil), contextualSkills...)
|
rebuildPromptReq := promptBuildRequestForTurn(ts, trimmedHistory, exec.summary, "", nil)
|
||||||
exec.messages = ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq)
|
rebuildPromptReq.ActiveSkills = append([]string(nil), contextualSkills...)
|
||||||
exec.callMessages = exec.messages
|
return ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq)
|
||||||
|
}
|
||||||
|
originalHistoryCount := len(exec.history)
|
||||||
|
var fit bool
|
||||||
|
exec.history, exec.callMessages, fit = trimHistoryToFitContextWindow(
|
||||||
|
exec.history,
|
||||||
|
func(trimmedHistory []providers.Message) []providers.Message {
|
||||||
|
rebuilt := buildMessages(trimmedHistory)
|
||||||
|
if exec.gracefulTerminal {
|
||||||
|
return append(append([]providers.Message(nil), rebuilt...), ts.interruptHintMessage())
|
||||||
|
}
|
||||||
|
return rebuilt
|
||||||
|
},
|
||||||
|
ts.agent.ContextWindow,
|
||||||
|
exec.providerToolDefs,
|
||||||
|
ts.agent.MaxTokens,
|
||||||
|
)
|
||||||
|
exec.messages = buildMessages(exec.history)
|
||||||
if exec.gracefulTerminal {
|
if exec.gracefulTerminal {
|
||||||
msgs := append([]providers.Message(nil), exec.messages...)
|
msgs := append([]providers.Message(nil), exec.messages...)
|
||||||
exec.callMessages = append(msgs, ts.interruptHintMessage())
|
exec.callMessages = append(msgs, ts.interruptHintMessage())
|
||||||
}
|
}
|
||||||
|
if dropped := originalHistoryCount - len(exec.history); dropped > 0 {
|
||||||
|
logger.WarnCF("agent", "Trimmed rebuilt history after context retry compaction", map[string]any{
|
||||||
|
"session_key": ts.sessionKey,
|
||||||
|
"retry": retry,
|
||||||
|
"dropped_msgs": dropped,
|
||||||
|
"remaining_msgs": len(exec.history),
|
||||||
|
"context_window": ts.agent.ContextWindow,
|
||||||
|
"max_tokens": ts.agent.MaxTokens,
|
||||||
|
"still_overlimit": !fit,
|
||||||
|
})
|
||||||
|
} else if !fit {
|
||||||
|
logger.WarnCF("agent", "Context still exceeds budget after retry compaction rebuild", map[string]any{
|
||||||
|
"session_key": ts.sessionKey,
|
||||||
|
"retry": retry,
|
||||||
|
"history_msgs": len(exec.history),
|
||||||
|
"context_window": ts.agent.ContextWindow,
|
||||||
|
"max_tokens": ts.agent.MaxTokens,
|
||||||
|
})
|
||||||
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
|
|
|
||||||
|
|
@ -66,10 +66,38 @@ func (p *Pipeline) SetupTurn(ctx context.Context, ts *turnState) (*turnExecution
|
||||||
history = resp.History
|
history = resp.History
|
||||||
summary = resp.Summary
|
summary = resp.Summary
|
||||||
}
|
}
|
||||||
rebuildPromptReq := promptBuildRequestForTurn(ts, history, summary, ts.userMessage, ts.media)
|
originalHistoryCount := len(history)
|
||||||
rebuildPromptReq.ActiveSkills = append([]string(nil), contextualSkills...)
|
var fit bool
|
||||||
messages = ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq)
|
history, messages, fit = trimHistoryToFitContextWindow(
|
||||||
messages = resolveMediaRefs(messages, p.MediaStore, maxMediaSize)
|
history,
|
||||||
|
func(trimmedHistory []providers.Message) []providers.Message {
|
||||||
|
rebuildPromptReq := promptBuildRequestForTurn(ts, trimmedHistory, summary, ts.userMessage, ts.media)
|
||||||
|
rebuildPromptReq.ActiveSkills = append([]string(nil), contextualSkills...)
|
||||||
|
rebuilt := ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq)
|
||||||
|
return resolveMediaRefs(rebuilt, p.MediaStore, maxMediaSize)
|
||||||
|
},
|
||||||
|
ts.agent.ContextWindow,
|
||||||
|
toolDefs,
|
||||||
|
ts.agent.MaxTokens,
|
||||||
|
)
|
||||||
|
if dropped := originalHistoryCount - len(history); dropped > 0 {
|
||||||
|
logger.WarnCF("agent", "Trimmed rebuilt history after proactive compaction", map[string]any{
|
||||||
|
"session_key": ts.sessionKey,
|
||||||
|
"dropped_msgs": dropped,
|
||||||
|
"remaining_msgs": len(history),
|
||||||
|
"context_window": ts.agent.ContextWindow,
|
||||||
|
"max_tokens": ts.agent.MaxTokens,
|
||||||
|
"still_overlimit": !fit,
|
||||||
|
})
|
||||||
|
} else if !fit {
|
||||||
|
logger.WarnCF("agent", "Context still exceeds budget "+
|
||||||
|
"after proactive compaction rebuild", map[string]any{
|
||||||
|
"session_key": ts.sessionKey,
|
||||||
|
"history_msgs": len(history),
|
||||||
|
"context_window": ts.agent.ContextWindow,
|
||||||
|
"max_tokens": ts.agent.MaxTokens,
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -68,19 +68,24 @@ func (a *Assembler) Assemble(ctx context.Context, convID int64, input AssembleIn
|
||||||
freshTailTokens += r.tokenCount
|
freshTailTokens += r.tokenCount
|
||||||
}
|
}
|
||||||
|
|
||||||
// Budget-aware selection of evictable items
|
// If the protected tail alone exceeds budget, trim from the oldest end of
|
||||||
remainingBudget := input.Budget - freshTailTokens
|
// the tail until the newest items fit within the requested budget.
|
||||||
if remainingBudget < 0 {
|
if freshTailTokens > input.Budget {
|
||||||
// Fresh tail alone exceeds budget - we keep it anyway (design decision)
|
originalTailCount := len(freshTail)
|
||||||
// Log for debugging retry/overflow issues
|
originalFreshTailTokens := freshTailTokens
|
||||||
logger.InfoCF("seahorse", "assemble: fresh tail exceeds budget", map[string]any{
|
for freshTailTokens > input.Budget && len(freshTail) > 0 {
|
||||||
"budget": input.Budget,
|
freshTailTokens -= freshTail[0].tokenCount
|
||||||
"fresh_tail_tokens": freshTailTokens,
|
freshTail = freshTail[1:]
|
||||||
"fresh_tail_count": len(freshTail),
|
}
|
||||||
"over_budget_by": freshTailTokens - input.Budget,
|
logger.InfoCF("seahorse", "assemble: trimmed fresh tail to budget", 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,
|
||||||
})
|
})
|
||||||
remainingBudget = 0
|
|
||||||
}
|
}
|
||||||
|
remainingBudget := input.Budget - freshTailTokens
|
||||||
|
|
||||||
var selected []resolvedItem
|
var selected []resolvedItem
|
||||||
evictableTokens := 0
|
evictableTokens := 0
|
||||||
|
|
|
||||||
|
|
@ -145,22 +145,29 @@ func TestAssemblerBudgetEvictsOldest(t *testing.T) {
|
||||||
s.UpsertContextItems(ctx, convID, items)
|
s.UpsertContextItems(ctx, convID, items)
|
||||||
|
|
||||||
// Budget of 200 tokens with FreshTailCount=32
|
// Budget of 200 tokens with FreshTailCount=32
|
||||||
// Fresh tail = last 32 messages (320 tokens, over budget, but always included)
|
// Fresh tail = last 32 messages (320 tokens, over budget)
|
||||||
// Evictable = first 8 messages (80 tokens)
|
// Evictable = first 8 messages (80 tokens)
|
||||||
// Budget after tail: max(0, 200-320) = 0 → no evictable items included
|
// The oldest messages from the fresh tail should be dropped so only the
|
||||||
|
// newest 20 messages remain within the 200-token budget.
|
||||||
a := &Assembler{store: s, config: Config{}}
|
a := &Assembler{store: s, config: Config{}}
|
||||||
result, err := a.Assemble(ctx, convID, AssembleInput{Budget: 200})
|
result, err := a.Assemble(ctx, convID, AssembleInput{Budget: 200})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Assemble: %v", err)
|
t.Fatalf("Assemble: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Should only include the 32-item fresh tail
|
if len(result.Messages) != 20 {
|
||||||
if len(result.Messages) != 32 {
|
t.Errorf("Messages = %d, want 20", len(result.Messages))
|
||||||
t.Errorf("Messages = %d, want 32 (fresh tail)", len(result.Messages))
|
|
||||||
}
|
}
|
||||||
// Should be the LAST 32 messages
|
if result.Messages[0].ID != msgs[20].ID {
|
||||||
if result.Messages[0].ID != msgs[8].ID {
|
t.Errorf("first message ID = %d, want %d (msgs[20])", result.Messages[0].ID, msgs[20].ID)
|
||||||
t.Errorf("first message ID = %d, want %d (msgs[8])", result.Messages[0].ID, msgs[8].ID)
|
}
|
||||||
|
|
||||||
|
totalTokens := 0
|
||||||
|
for _, msg := range result.Messages {
|
||||||
|
totalTokens += msg.TokenCount
|
||||||
|
}
|
||||||
|
if totalTokens > 200 {
|
||||||
|
t.Errorf("assembled tokens = %d, want <= 200", totalTokens)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue