mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 12:18:59 +00:00
fix: persist mid-loop compaction to stop re-compaction loop and revive episodic
The v3 pipeline compacts session history mid-loop (prune_stage + final-request guard) but only mutates the run's message buffer, never the session store. Each turn reloads full history and re-compacts from scratch: message_tokens climb 129k->156k across turns while every turn compacts back down to ~60k. The lossy compaction differs per run, degrading the agent. The same missing persistence stalls episodic memory: the cumulative compaction count never advances, so the episodic worker's idempotency key (sessionKey:count) is pinned and every cycle after the first is skipped. Observed on live traffic: 8 run.completed since deploy, 0 new episodic. Fixes, all reusing existing machinery (no new store methods, no migrations): - Bug A: emitSessionCompleted reads cumulative GetCompactionCount (matching the legacy v2 path) instead of the per-run counter that resets to 0. - Bug B/anti-loop: finalize passes state.Prune.MidLoopCompacted into maybeSummarize; under pressure it lowers the trigger to a unit-aligned threshold (compactionInputCap - overhead, same MaxRequestShare the guard uses) so the compaction is PERSISTED via the existing TruncateHistory + IncrementCompaction path. Defensive floor prevents over-compaction on pathological config; tool-result-only bloat still skips (history-only). - Bug C: SourceID embeds the count (sessionKey:count) so the eventbus dedup key advances per compaction cycle instead of swallowing rapid same-session turns within the 5m TTL. Tests: episodic compaction, maybe_summarize pressure, request budget. go build (PG + sqliteonly), go vet, go test -race all green.
This commit is contained in:
1 parent
9957a94868
commit
5ca433b8ac
69 files changed
+4110
-371
No files matched your search
@@ -213,6 +213,7 @@ func setupSubagents(providerReg *providers.Registry, cfg *config.Config, msgBus
|
||||
|
||||
manager := tools.NewSubagentManager(provider, providerReg, agentCfg.Model, msgBus, toolsFactory, subCfg)
|
||||
manager.SetUsageCapService(usageCapSvc)
|
||||
manager.SetAgentBudget(agentCfg.ContextWindow, agentCfg.MaxTokens)
|
||||
return manager
|
||||
}
|
||||
|
||||
|
||||
@@ -61,6 +61,12 @@ func processNormalMessage(
|
||||
})
|
||||
return
|
||||
}
|
||||
// Team Work and intent gates run before Loop.injectContext. Propagate the
|
||||
// resolved agent budget here so every classifier sees the same authority.
|
||||
ctx = agent.WithAgentBudget(ctx, agentLoop)
|
||||
if uid := agentLoop.UUID(); uid != uuid.Nil {
|
||||
ctx = store.WithAgentID(ctx, uid)
|
||||
}
|
||||
|
||||
// Build session key based on scope config (matching TS buildAgentPeerSessionKey).
|
||||
peerKind := msg.PeerKind
|
||||
|
||||
+227
-28
@@ -2,6 +2,7 @@ package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
@@ -39,6 +40,12 @@ Conversation to summarize:
|
||||
|
||||
const defaultCompactionTimeout = 120 * time.Second
|
||||
|
||||
const (
|
||||
maxCompactionChunks = 16
|
||||
maxCompactionMergeLevels = 3
|
||||
defaultCompactionShare = 0.85
|
||||
)
|
||||
|
||||
func (l *Loop) compactionTimeout() time.Duration {
|
||||
if l.compactionCfg != nil && l.compactionCfg.TimeoutSeconds > 0 {
|
||||
return time.Duration(l.compactionCfg.TimeoutSeconds) * time.Second
|
||||
@@ -81,43 +88,30 @@ func (l *Loop) compactMessagesInPlace(ctx context.Context, messages []providers.
|
||||
return nil
|
||||
}
|
||||
|
||||
// Build summary input (same pattern as maybeSummarize in loop_history.go).
|
||||
toSummarize := messages[:splitIdx]
|
||||
var sb strings.Builder
|
||||
for _, m := range toSummarize {
|
||||
switch m.Role {
|
||||
case "user":
|
||||
fmt.Fprintf(&sb, "user: %s\n", m.Content)
|
||||
case "assistant":
|
||||
fmt.Fprintf(&sb, "assistant: %s\n", SanitizeAssistantContent(m.Content))
|
||||
}
|
||||
}
|
||||
|
||||
timeout := l.compactionTimeout()
|
||||
sctx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
inTokens := l.estimateSummaryInputTokens(toSummarize)
|
||||
slog.Info("compact_budget",
|
||||
"path", "mid-loop",
|
||||
"agent", l.id,
|
||||
"in_tokens", inTokens,
|
||||
"out_tokens", dynamicSummaryMax(inTokens),
|
||||
"timeout_seconds", int(timeout/time.Second),
|
||||
)
|
||||
chatReq := providers.ChatRequest{
|
||||
Messages: []providers.Message{{
|
||||
Role: "user",
|
||||
Content: compactionSummaryPrompt + sb.String(),
|
||||
}},
|
||||
Model: l.model,
|
||||
Options: map[string]any{"max_tokens": dynamicSummaryMax(inTokens), "temperature": 0.3},
|
||||
inputCap := l.compactionInputCap()
|
||||
if inputCap <= 0 {
|
||||
slog.Warn("mid_loop_compaction_failed", "agent", l.id, "error", "context_window_unresolved")
|
||||
return nil
|
||||
}
|
||||
resp, err := l.callInternalLLMWithUsage(sctx, chatReq, "mid-loop-compaction")
|
||||
units := buildCompactionUnits(toSummarize)
|
||||
summaryContent, chunkCount, err := l.summarizeCompactionUnits(sctx, units, inputCap, 1)
|
||||
if err != nil {
|
||||
slog.Warn("mid_loop_compaction_failed", "agent", l.id, "timeout_seconds", int(timeout/time.Second), "error", err)
|
||||
return nil
|
||||
}
|
||||
slog.Info("compact_budget",
|
||||
"path", "mid-loop",
|
||||
"agent", l.id,
|
||||
"in_tokens", l.estimateSummaryInputTokens(toSummarize),
|
||||
"input_cap_tokens", inputCap,
|
||||
"chunks", chunkCount,
|
||||
"timeout_seconds", int(timeout/time.Second),
|
||||
)
|
||||
|
||||
// Collect MediaRefs from compacted messages (keep up to 30 most recent).
|
||||
const maxPreservedMediaRefs = 30
|
||||
@@ -133,7 +127,7 @@ func (l *Loop) compactMessagesInPlace(ctx context.Context, messages []providers.
|
||||
|
||||
summary := providers.Message{
|
||||
Role: "user",
|
||||
Content: "[Summary of earlier conversation]\n" + SanitizeAssistantContent(resp.Content),
|
||||
Content: "[Summary of earlier conversation]\n" + summaryContent,
|
||||
MediaRefs: preservedRefs,
|
||||
}
|
||||
result := make([]providers.Message, 0, 1+keepCount)
|
||||
@@ -149,6 +143,211 @@ func (l *Loop) compactMessagesInPlace(ctx context.Context, messages []providers.
|
||||
return result
|
||||
}
|
||||
|
||||
func (l *Loop) compactionInputCap() int {
|
||||
contextWindow := l.resolveEffectiveContextWindow()
|
||||
if contextWindow <= 0 {
|
||||
return 0
|
||||
}
|
||||
maxTokens := l.effectiveMaxTokens()
|
||||
hardInputCap := contextWindow - maxTokens
|
||||
share := defaultCompactionShare
|
||||
if l.compactionCfg != nil && l.compactionCfg.MaxRequestShare > 0 && l.compactionCfg.MaxRequestShare <= 1 {
|
||||
share = l.compactionCfg.MaxRequestShare
|
||||
}
|
||||
softTarget := int(float64(contextWindow)*share) - maxTokens
|
||||
return min(hardInputCap, softTarget)
|
||||
}
|
||||
|
||||
func buildCompactionUnits(messages []providers.Message) []string {
|
||||
units := make([]string, 0, len(messages))
|
||||
for i := 0; i < len(messages); {
|
||||
end := i + 1
|
||||
if messages[i].Role == "assistant" && len(messages[i].ToolCalls) > 0 {
|
||||
for end < len(messages) && messages[end].Role == "tool" {
|
||||
end++
|
||||
}
|
||||
}
|
||||
if text := renderCompactionMessages(messages[i:end]); text != "" {
|
||||
units = append(units, text)
|
||||
}
|
||||
i = end
|
||||
}
|
||||
return units
|
||||
}
|
||||
|
||||
func renderCompactionMessages(messages []providers.Message) string {
|
||||
var sb strings.Builder
|
||||
for _, m := range messages {
|
||||
switch m.Role {
|
||||
case "user":
|
||||
fmt.Fprintf(&sb, "user: %s\n", m.Content)
|
||||
case "assistant":
|
||||
if content := SanitizeAssistantContent(m.Content); content != "" {
|
||||
fmt.Fprintf(&sb, "assistant: %s\n", content)
|
||||
}
|
||||
// Tool calls carry the assistant's intent (which tool, what args);
|
||||
// dropping them loses the "why" behind each tool result below.
|
||||
for _, tc := range m.ToolCalls {
|
||||
if args, err := json.Marshal(tc.Arguments); err == nil && len(tc.Arguments) > 0 {
|
||||
fmt.Fprintf(&sb, "assistant tool call %s(%s)\n", tc.Name, string(args))
|
||||
} else {
|
||||
fmt.Fprintf(&sb, "assistant tool call %s()\n", tc.Name)
|
||||
}
|
||||
}
|
||||
case "tool":
|
||||
// Tool results hold the technical payload the summary must retain
|
||||
// (search hits, file contents, API responses). buildCompactionUnits
|
||||
// groups these with their assistant tool_call; the previous renderer
|
||||
// silently dropped them, erasing the data before summarization.
|
||||
if m.Content != "" {
|
||||
fmt.Fprintf(&sb, "tool result: %s\n", m.Content)
|
||||
}
|
||||
}
|
||||
}
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
func (l *Loop) summarizeCompactionUnits(ctx context.Context, units []string, inputCap, level int) (string, int, error) {
|
||||
if len(units) == 0 {
|
||||
return "", 0, fmt.Errorf("no compactable conversation content")
|
||||
}
|
||||
chunks, err := l.packCompactionChunks(units, inputCap)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
if len(chunks) > maxCompactionChunks {
|
||||
return "", 0, fmt.Errorf("compaction chunk limit exceeded: chunks=%d limit=%d", len(chunks), maxCompactionChunks)
|
||||
}
|
||||
|
||||
summaries := make([]string, 0, len(chunks))
|
||||
for i, chunk := range chunks {
|
||||
inputTokens := l.estimateCompactionRequestTokens(chunk)
|
||||
if inputTokens > inputCap {
|
||||
return "", 0, fmt.Errorf("compaction chunk exceeds input cap: chunk=%d input=%d cap=%d", i, inputTokens, inputCap)
|
||||
}
|
||||
outputTokens := dynamicSummaryMax(inputTokens)
|
||||
slog.Debug("compact_chunk_budget",
|
||||
"path", "mid-loop",
|
||||
"agent", l.id,
|
||||
"level", level,
|
||||
"chunk", i+1,
|
||||
"chunks", len(chunks),
|
||||
"in_tokens", inputTokens,
|
||||
"out_tokens", outputTokens,
|
||||
"input_cap_tokens", inputCap,
|
||||
)
|
||||
resp, callErr := l.callInternalLLMWithUsage(ctx, providers.ChatRequest{
|
||||
Messages: []providers.Message{{Role: "user", Content: compactionSummaryPrompt + chunk}},
|
||||
Model: l.model,
|
||||
Options: map[string]any{"max_tokens": outputTokens, "temperature": 0.3},
|
||||
}, "mid-loop-compaction")
|
||||
if callErr != nil {
|
||||
return "", 0, callErr
|
||||
}
|
||||
summary := SanitizeAssistantContent(resp.Content)
|
||||
if strings.TrimSpace(summary) == "" {
|
||||
return "", 0, fmt.Errorf("compaction returned empty summary")
|
||||
}
|
||||
summaries = append(summaries, summary)
|
||||
}
|
||||
if len(summaries) == 1 {
|
||||
return summaries[0], len(chunks), nil
|
||||
}
|
||||
if level >= maxCompactionMergeLevels {
|
||||
return "", 0, fmt.Errorf("compaction merge level exceeded: level=%d limit=%d", level, maxCompactionMergeLevels)
|
||||
}
|
||||
|
||||
mergeUnits := make([]string, len(summaries))
|
||||
for i, summary := range summaries {
|
||||
mergeUnits[i] = fmt.Sprintf("partial summary %d: %s\n", i+1, summary)
|
||||
}
|
||||
merged, mergeChunks, mergeErr := l.summarizeCompactionUnits(ctx, mergeUnits, inputCap, level+1)
|
||||
return merged, len(chunks) + mergeChunks, mergeErr
|
||||
}
|
||||
|
||||
func (l *Loop) packCompactionChunks(units []string, inputCap int) ([]string, error) {
|
||||
var chunks []string
|
||||
var current strings.Builder
|
||||
flush := func() {
|
||||
if current.Len() == 0 {
|
||||
return
|
||||
}
|
||||
chunks = append(chunks, current.String())
|
||||
current.Reset()
|
||||
}
|
||||
|
||||
for _, unit := range units {
|
||||
parts, err := l.splitCompactionUnit(unit, inputCap)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, part := range parts {
|
||||
candidate := current.String() + part
|
||||
if current.Len() > 0 && l.estimateCompactionRequestTokens(candidate) > inputCap {
|
||||
flush()
|
||||
candidate = part
|
||||
}
|
||||
if l.estimateCompactionRequestTokens(candidate) > inputCap {
|
||||
return nil, fmt.Errorf("atomic compaction unit exceeds input cap")
|
||||
}
|
||||
current.WriteString(part)
|
||||
if len(chunks) >= maxCompactionChunks {
|
||||
return nil, fmt.Errorf("compaction chunk limit exceeded: limit=%d", maxCompactionChunks)
|
||||
}
|
||||
}
|
||||
}
|
||||
flush()
|
||||
return chunks, nil
|
||||
}
|
||||
|
||||
func (l *Loop) splitCompactionUnit(unit string, inputCap int) ([]string, error) {
|
||||
if l.estimateCompactionRequestTokens(unit) <= inputCap {
|
||||
return []string{unit}, nil
|
||||
}
|
||||
words := strings.Fields(unit)
|
||||
if len(words) == 0 {
|
||||
return nil, fmt.Errorf("empty oversized compaction unit")
|
||||
}
|
||||
|
||||
var parts []string
|
||||
var current strings.Builder
|
||||
for _, word := range words {
|
||||
candidate := word
|
||||
if current.Len() > 0 {
|
||||
candidate = current.String() + " " + word
|
||||
}
|
||||
if l.estimateCompactionRequestTokens(candidate) <= inputCap {
|
||||
if current.Len() > 0 {
|
||||
current.WriteByte(' ')
|
||||
}
|
||||
current.WriteString(word)
|
||||
continue
|
||||
}
|
||||
if current.Len() == 0 {
|
||||
return nil, fmt.Errorf("atomic compaction token exceeds input cap")
|
||||
}
|
||||
parts = append(parts, current.String()+"\n")
|
||||
current.Reset()
|
||||
current.WriteString("[continued] ")
|
||||
current.WriteString(word)
|
||||
if l.estimateCompactionRequestTokens(current.String()) > inputCap {
|
||||
return nil, fmt.Errorf("atomic compaction token exceeds input cap")
|
||||
}
|
||||
}
|
||||
if current.Len() > 0 {
|
||||
parts = append(parts, current.String()+"\n")
|
||||
}
|
||||
return parts, nil
|
||||
}
|
||||
|
||||
func (l *Loop) estimateCompactionRequestTokens(content string) int {
|
||||
message := providers.Message{Role: "user", Content: compactionSummaryPrompt + content}
|
||||
if l.tokenCounter != nil {
|
||||
return l.tokenCounter.CountMessages(l.model, []providers.Message{message})
|
||||
}
|
||||
return len([]rune(message.Content))/3 + 4
|
||||
}
|
||||
|
||||
// dynamicSummaryMax returns the output-token budget for a compaction or
|
||||
// summarization call, scaled to input size. Formula: in/25 (~4% compression),
|
||||
// clamped to [1024, 8192]. Floor keeps short summaries coherent; cap prevents
|
||||
|
||||
@@ -47,9 +47,10 @@ func TestLoopCompact_Integration_DynamicMaxTokens_VietnameseFixture(t *testing.T
|
||||
cap := &capturingProvider{response: "Tóm tắt cuộc trò chuyện: Đã thảo luận về nhiều chủ đề."}
|
||||
|
||||
loop := &Loop{
|
||||
provider: cap,
|
||||
model: "claude-3-5-sonnet",
|
||||
tokenCounter: tokencount.NewFallbackCounter(),
|
||||
provider: cap,
|
||||
model: "claude-3-5-sonnet",
|
||||
tokenCounter: tokencount.NewFallbackCounter(),
|
||||
contextWindow: 200_000,
|
||||
}
|
||||
|
||||
// 600 messages × ~500 runes each ≈ 300k runes ÷ 3 ≈ 100k tokens total.
|
||||
@@ -97,3 +98,105 @@ func TestLoopCompact_Integration_DynamicMaxTokens_VietnameseFixture(t *testing.T
|
||||
t.Logf("observed: msgs=%d splitIdx=%d estimatedIn=%d max_tokens=%d dynamicSummaryMax=%d",
|
||||
len(msgs), splitIdx, estimatedIn, maxTokens, dynamicSummaryMax(estimatedIn))
|
||||
}
|
||||
|
||||
func TestLoopCompact_ChunksEverySummaryRequestUnderAgentCap(t *testing.T) {
|
||||
cap := &capturingProvider{response: "bounded summary"}
|
||||
loop := &Loop{
|
||||
provider: cap,
|
||||
model: "claude-3-5-sonnet",
|
||||
tokenCounter: tokencount.NewFallbackCounter(),
|
||||
contextWindow: 20_000,
|
||||
}
|
||||
msgs := buildVietnameseMsgs(20, 4_000)
|
||||
|
||||
result := loop.compactMessagesInPlace(context.Background(), msgs)
|
||||
if result == nil {
|
||||
t.Fatal("compactMessagesInPlace returned nil")
|
||||
}
|
||||
if len(cap.captured) < 3 {
|
||||
t.Fatalf("provider calls = %d, want multiple chunks plus merge", len(cap.captured))
|
||||
}
|
||||
inputCap := loop.compactionInputCap()
|
||||
for i, req := range cap.captured {
|
||||
got := loop.tokenCounter.CountMessages(loop.model, req.Messages)
|
||||
if got > inputCap {
|
||||
t.Fatalf("request %d input = %d, exceeds cap %d", i, got, inputCap)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoopCompact_InputCapReservesFixedMaxTokens locks the agent-only budget:
|
||||
// compactionInputCap() reserves the CALLING agent's fixed effectiveMaxTokens()
|
||||
// from the window, bounded by the request share. The dynamic-output term
|
||||
// (dynamicSummaryMax) governs only the per-chunk output budget — it is NOT an
|
||||
// authority over the input cap, so it does not appear here.
|
||||
func TestLoopCompact_InputCapReservesFixedMaxTokens(t *testing.T) {
|
||||
const (
|
||||
window = 20_000
|
||||
maxTokens = 8_192
|
||||
)
|
||||
loop := &Loop{
|
||||
provider: &capturingProvider{},
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: window,
|
||||
maxTokens: maxTokens,
|
||||
}
|
||||
inputCap := loop.compactionInputCap()
|
||||
|
||||
hardCap := window - maxTokens // 11808
|
||||
softTarget := int(float64(window)*defaultCompactionShare) - maxTokens // 8808
|
||||
want := min(hardCap, softTarget)
|
||||
if inputCap != want {
|
||||
t.Fatalf("compactionInputCap() = %d, want %d (min(window-maxTokens, share*window-maxTokens))", inputCap, want)
|
||||
}
|
||||
// Invariant: input cap plus the fixed reserve never exceeds the window.
|
||||
if inputCap+maxTokens > window {
|
||||
t.Fatalf("inputCap %d + reserve %d exceeds window %d", inputCap, maxTokens, window)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoopCompact_RejectsMoreThanSixteenChunksBeforeProvider(t *testing.T) {
|
||||
cap := &capturingProvider{response: "unused"}
|
||||
loop := &Loop{
|
||||
provider: cap,
|
||||
model: "claude-3-5-sonnet",
|
||||
tokenCounter: tokencount.NewFallbackCounter(),
|
||||
contextWindow: 200_000,
|
||||
}
|
||||
unit := strings.Repeat("bounded-word ", 1_500)
|
||||
unitCap := loop.estimateCompactionRequestTokens(unit)
|
||||
units := make([]string, maxCompactionChunks+1)
|
||||
for i := range units {
|
||||
units[i] = unit
|
||||
}
|
||||
|
||||
if _, _, err := loop.summarizeCompactionUnits(context.Background(), units, unitCap, 1); err == nil {
|
||||
t.Fatal("expected compaction chunk limit error")
|
||||
}
|
||||
if len(cap.captured) != 0 {
|
||||
t.Fatalf("provider calls = %d, want 0 before chunk-limit rejection", len(cap.captured))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCompactionUnits_KeepsToolCycleTogether(t *testing.T) {
|
||||
units := buildCompactionUnits([]providers.Message{
|
||||
{Role: "assistant", Content: "calling tool", ToolCalls: []providers.ToolCall{{ID: "call-1", Name: "read_file", Arguments: map[string]any{"path": "a.go"}}}},
|
||||
{Role: "tool", Content: "file body here", ToolCallID: "call-1"},
|
||||
{Role: "user", Content: "continue"},
|
||||
})
|
||||
if len(units) != 2 {
|
||||
t.Fatalf("units = %d, want 2 (tool cycle + following user)", len(units))
|
||||
}
|
||||
if !strings.Contains(units[0], "calling tool") || !strings.Contains(units[1], "continue") {
|
||||
t.Fatalf("unexpected units: %#v", units)
|
||||
}
|
||||
// Regression: the tool RESULT payload and the tool CALL must survive into the
|
||||
// compaction unit. Previously renderCompactionMessages dropped role=tool and
|
||||
// assistant tool_calls, erasing technical data before summarization.
|
||||
if !strings.Contains(units[0], "file body here") {
|
||||
t.Fatalf("tool result dropped from compaction unit: %#v", units[0])
|
||||
}
|
||||
if !strings.Contains(units[0], "read_file") {
|
||||
t.Fatalf("tool call dropped from compaction unit: %#v", units[0])
|
||||
}
|
||||
}
|
||||
@@ -31,8 +31,9 @@ func TestCompactMessagesInPlace_MaxTokensDynamic(t *testing.T) {
|
||||
cap := &capturingProvider{response: "Summary of conversation."}
|
||||
|
||||
loop := &Loop{
|
||||
provider: cap,
|
||||
model: "claude-3-5-sonnet",
|
||||
provider: cap,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: 200_000,
|
||||
// tokenCounter nil → estimateSummaryInputTokens uses rune/3 fallback
|
||||
}
|
||||
|
||||
|
||||
@@ -25,8 +25,9 @@ func TestCompactMessagesInPlace_UsesDefaultTimeout(t *testing.T) {
|
||||
capturingProvider: capturingProvider{response: "Summary of conversation."},
|
||||
}
|
||||
loop := &Loop{
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: 200_000,
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
@@ -43,8 +44,9 @@ func TestCompactMessagesInPlace_UsesConfiguredTimeout(t *testing.T) {
|
||||
capturingProvider: capturingProvider{response: "Summary of conversation."},
|
||||
}
|
||||
loop := &Loop{
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: 200_000,
|
||||
compactionCfg: &config.CompactionConfig{
|
||||
TimeoutSeconds: 45,
|
||||
},
|
||||
@@ -64,8 +66,9 @@ func TestCompactMessagesInPlace_NonPositiveTimeoutFallsBackToDefault(t *testing.
|
||||
capturingProvider: capturingProvider{response: "Summary of conversation."},
|
||||
}
|
||||
loop := &Loop{
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: 200_000,
|
||||
compactionCfg: &config.CompactionConfig{
|
||||
TimeoutSeconds: -1,
|
||||
},
|
||||
|
||||
@@ -38,6 +38,9 @@ func (l *Loop) injectContext(ctx context.Context, req *RunRequest) (contextSetup
|
||||
if l.tenantID != uuid.Nil {
|
||||
ctx = store.WithTenantID(ctx, l.tenantID)
|
||||
}
|
||||
// Propagate the configured agent budget to every nested model call.
|
||||
ctx = store.WithAgentContextWindow(ctx, l.contextWindow)
|
||||
ctx = store.WithAgentMaxTokens(ctx, l.effectiveMaxTokens())
|
||||
// Inject user ID into context for per-user scoping (memory, context files, etc.)
|
||||
if req.UserID != "" {
|
||||
ctx = store.WithUserID(ctx, req.UserID)
|
||||
|
||||
@@ -0,0 +1,208 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/eventbus"
|
||||
)
|
||||
|
||||
// recordingBus captures every published DomainEvent for assertion.
|
||||
// Distinct from the eventbus package's own test bus — this one does no
|
||||
// dispatch, it only records what the emit callback published.
|
||||
type recordingBus struct {
|
||||
mu sync.Mutex
|
||||
published []eventbus.DomainEvent
|
||||
}
|
||||
|
||||
func (r *recordingBus) Publish(event eventbus.DomainEvent) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.published = append(r.published, event)
|
||||
}
|
||||
func (r *recordingBus) Subscribe(_ eventbus.EventType, _ eventbus.DomainEventHandler) func() {
|
||||
return func() {}
|
||||
}
|
||||
func (r *recordingBus) Start(_ context.Context) {}
|
||||
func (r *recordingBus) Drain(_ time.Duration) error { return nil }
|
||||
func (r *recordingBus) events() []eventbus.DomainEvent {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
out := make([]eventbus.DomainEvent, len(r.published))
|
||||
copy(out, r.published)
|
||||
return out
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// V1-A: emit reads the CUMULATIVE session compaction count, not the per-run 0.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestEmitSessionCompleted_UsesCumulativeCount(t *testing.T) {
|
||||
bus := &recordingBus{}
|
||||
sessions := &nopSessionStore{
|
||||
compactionCount: 5, // cumulative session count
|
||||
summary: "prior-cycle summary",
|
||||
}
|
||||
loop := &Loop{
|
||||
id: "test-agent",
|
||||
agentUUID: uuid.New(),
|
||||
tenantID: uuid.New(),
|
||||
domainBus: bus,
|
||||
sessions: sessions,
|
||||
}
|
||||
|
||||
loop.emitSessionCompleted(context.Background(), "sess-1", "user-1", 42, 9000)
|
||||
|
||||
evs := bus.events()
|
||||
if len(evs) != 1 {
|
||||
t.Fatalf("published %d events, want 1", len(evs))
|
||||
}
|
||||
ev := evs[0]
|
||||
// Bug A: payload.CompactionCount (the field the worker uses for idempotency)
|
||||
// must be the cumulative count (5), not the per-run 0.
|
||||
pl, ok := ev.Payload.(*eventbus.SessionCompletedPayload)
|
||||
if !ok {
|
||||
t.Fatalf("payload type = %T, want *SessionCompletedPayload", ev.Payload)
|
||||
}
|
||||
if pl.CompactionCount != 5 {
|
||||
t.Errorf("payload.CompactionCount = %d, want 5 (cumulative)", pl.CompactionCount)
|
||||
}
|
||||
// Bug C: SourceID embeds the cumulative count.
|
||||
if ev.SourceID != "sess-1:5" {
|
||||
t.Errorf("SourceID = %q, want %q", ev.SourceID, "sess-1:5")
|
||||
}
|
||||
// Summary from the previous cycle is attached when count > 0.
|
||||
if pl.Summary != "prior-cycle summary" {
|
||||
t.Errorf("payload.Summary = %q, want prior-cycle summary", pl.Summary)
|
||||
}
|
||||
if pl.MessageCount != 42 || pl.TokensUsed != 9000 {
|
||||
t.Errorf("msgCount/tokens = %d/%d, want 42/9000", pl.MessageCount, pl.TokensUsed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmitSessionCompleted_ZeroCountOmitsSummary(t *testing.T) {
|
||||
bus := &recordingBus{}
|
||||
// count == 0 → no prior cycle → summary must NOT be fetched/attached
|
||||
// (worker would otherwise get a stale/empty summary and skip the LLM path).
|
||||
sessions := &nopSessionStore{
|
||||
compactionCount: 0,
|
||||
summary: "should-not-be-attached",
|
||||
}
|
||||
loop := &Loop{
|
||||
id: "test-agent",
|
||||
agentUUID: uuid.New(),
|
||||
tenantID: uuid.New(),
|
||||
domainBus: bus,
|
||||
sessions: sessions,
|
||||
}
|
||||
|
||||
loop.emitSessionCompleted(context.Background(), "sess-1", "user-1", 1, 100)
|
||||
|
||||
evs := bus.events()
|
||||
if len(evs) != 1 {
|
||||
t.Fatalf("published %d events, want 1", len(evs))
|
||||
}
|
||||
pl := evs[0].Payload.(*eventbus.SessionCompletedPayload)
|
||||
if pl.CompactionCount != 0 {
|
||||
t.Errorf("payload.CompactionCount = %d, want 0", pl.CompactionCount)
|
||||
}
|
||||
if pl.Summary != "" {
|
||||
t.Errorf("payload.Summary = %q, want empty (count==0)", pl.Summary)
|
||||
}
|
||||
if evs[0].SourceID != "sess-1:0" {
|
||||
t.Errorf("SourceID = %q, want sess-1:0", evs[0].SourceID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmitSessionCompleted_NoBusIsNoop(t *testing.T) {
|
||||
loop := &Loop{
|
||||
id: "test-agent",
|
||||
agentUUID: uuid.New(),
|
||||
tenantID: uuid.New(),
|
||||
domainBus: nil, // no bus
|
||||
sessions: &nopSessionStore{compactionCount: 3},
|
||||
}
|
||||
// Must not panic when domainBus is nil.
|
||||
loop.emitSessionCompleted(context.Background(), "sess-1", "user-1", 1, 100)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// V1-C: eventbus dedup — distinct counts both pass, same count deduped.
|
||||
// Uses the REAL bus + real dedupSet (not the recording double) so we exercise
|
||||
// the actual dedup key = Type + ":" + SourceID.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestEventbusDedup_DistinctCompactionCountsPass(t *testing.T) {
|
||||
bus := eventbus.NewDomainEventBus(eventbus.Config{
|
||||
QueueSize: 100,
|
||||
WorkerCount: 2,
|
||||
RetryAttempts: 1,
|
||||
RetryDelay: time.Millisecond,
|
||||
DedupTTL: time.Minute,
|
||||
})
|
||||
bus.Start(context.Background())
|
||||
defer func() { _ = bus.Drain(time.Second) }()
|
||||
|
||||
var mu sync.Mutex
|
||||
var seen []string
|
||||
bus.Subscribe(eventbus.EventSessionCompleted, func(_ context.Context, e eventbus.DomainEvent) error {
|
||||
mu.Lock()
|
||||
seen = append(seen, e.SourceID)
|
||||
mu.Unlock()
|
||||
return nil
|
||||
})
|
||||
|
||||
// Same session, three DIFFERENT compaction cycles → all three must pass dedup.
|
||||
bus.Publish(eventbus.DomainEvent{Type: eventbus.EventSessionCompleted, SourceID: "sess-1:1"})
|
||||
bus.Publish(eventbus.DomainEvent{Type: eventbus.EventSessionCompleted, SourceID: "sess-1:2"})
|
||||
bus.Publish(eventbus.DomainEvent{Type: eventbus.EventSessionCompleted, SourceID: "sess-1:3"})
|
||||
// Rapid duplicate of cycle 2 within TTL → must be deduped.
|
||||
bus.Publish(eventbus.DomainEvent{Type: eventbus.EventSessionCompleted, SourceID: "sess-1:2"})
|
||||
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
|
||||
mu.Lock()
|
||||
got := len(seen)
|
||||
mu.Unlock()
|
||||
if got != 3 {
|
||||
t.Errorf("handler saw %d events, want 3 (cycles 1,2,3; duplicate of 2 deduped)", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEventbusDedup_SameCompactionCountDeduped(t *testing.T) {
|
||||
bus := eventbus.NewDomainEventBus(eventbus.Config{
|
||||
QueueSize: 100,
|
||||
WorkerCount: 2,
|
||||
RetryAttempts: 1,
|
||||
RetryDelay: time.Millisecond,
|
||||
DedupTTL: time.Minute,
|
||||
})
|
||||
bus.Start(context.Background())
|
||||
defer func() { _ = bus.Drain(time.Second) }()
|
||||
|
||||
var count int
|
||||
var mu sync.Mutex
|
||||
bus.Subscribe(eventbus.EventSessionCompleted, func(_ context.Context, _ eventbus.DomainEvent) error {
|
||||
mu.Lock()
|
||||
count++
|
||||
mu.Unlock()
|
||||
return nil
|
||||
})
|
||||
|
||||
// Same SourceID (same session, same cycle) published 3× → only 1 passes.
|
||||
for range 3 {
|
||||
bus.Publish(eventbus.DomainEvent{Type: eventbus.EventSessionCompleted, SourceID: "sess-1:7"})
|
||||
}
|
||||
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
|
||||
mu.Lock()
|
||||
got := count
|
||||
mu.Unlock()
|
||||
if got != 1 {
|
||||
t.Errorf("handler called %d times, want 1 (same cycle deduped)", got)
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -186,10 +187,9 @@ func (l *Loop) finalizeRun(
|
||||
l.sessions.UpdateMetadata(ctx, req.SessionKey, l.model, l.provider.Name(), req.Channel)
|
||||
l.sessions.AccumulateTokens(ctx, req.SessionKey, int64(rs.totalUsage.PromptTokens), int64(rs.totalUsage.CompletionTokens))
|
||||
|
||||
// Calibrate token estimation: store actual prompt tokens + message count.
|
||||
if rs.totalUsage.PromptTokens > 0 {
|
||||
msgCount := len(history) + rs.checkpointFlushedMsgs + len(rs.pendingMsgs)
|
||||
l.sessions.SetLastPromptTokens(ctx, req.SessionKey, rs.totalUsage.PromptTokens, msgCount)
|
||||
// Calibrate token estimation using the last LLM request, not the total run.
|
||||
if rs.lastUsage.PromptTokens > 0 && rs.lastUsageMsgCount > 0 {
|
||||
l.sessions.SetLastPromptTokens(ctx, req.SessionKey, rs.lastUsage.PromptTokens, rs.lastUsageMsgCount)
|
||||
}
|
||||
|
||||
l.sessions.Save(ctx, req.SessionKey)
|
||||
@@ -203,31 +203,43 @@ func (l *Loop) finalizeRun(
|
||||
}
|
||||
|
||||
// 9. Maybe summarize
|
||||
l.maybeSummarize(ctx, req.SessionKey)
|
||||
// Legacy v2 path has no mid-loop pressure signal (that lives in v3 RunState.Prune),
|
||||
// so pass false — this preserves the original baseline-threshold behavior here.
|
||||
l.maybeSummarize(ctx, req.SessionKey, false)
|
||||
|
||||
// V3: emit session.completed for consolidation pipeline (episodic → semantic → dreaming)
|
||||
if l.domainBus != nil {
|
||||
// Bug C parity: include count in SourceID so the eventbus dedup key advances
|
||||
// per compaction cycle (matches the v3 adapter emit path). This path is
|
||||
// currently disabled (FinalizeStage owns finalization) but kept in sync.
|
||||
finalizeCount := l.sessions.GetCompactionCount(ctx, req.SessionKey)
|
||||
l.domainBus.Publish(eventbus.DomainEvent{
|
||||
Type: eventbus.EventSessionCompleted,
|
||||
TenantID: l.tenantID.String(),
|
||||
AgentID: l.agentUUID.String(),
|
||||
UserID: req.UserID,
|
||||
SourceID: req.SessionKey,
|
||||
SourceID: fmt.Sprintf("%s:%d", req.SessionKey, finalizeCount),
|
||||
Payload: &eventbus.SessionCompletedPayload{
|
||||
SessionKey: req.SessionKey,
|
||||
MessageCount: len(history) + len(rs.pendingMsgs),
|
||||
TokensUsed: rs.totalUsage.PromptTokens + rs.totalUsage.CompletionTokens,
|
||||
CompactionCount: l.sessions.GetCompactionCount(ctx, req.SessionKey),
|
||||
CompactionCount: finalizeCount,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
var lastUsage *providers.Usage
|
||||
if rs.lastUsage.PromptTokens > 0 || rs.lastUsage.CompletionTokens > 0 || rs.lastUsage.TotalTokens > 0 {
|
||||
lastUsage = &rs.lastUsage
|
||||
}
|
||||
|
||||
return &RunResult{
|
||||
Content: rs.finalContent,
|
||||
Thinking: rs.finalThinking,
|
||||
RunID: req.RunID,
|
||||
Iterations: rs.iteration,
|
||||
Usage: &rs.totalUsage,
|
||||
LastUsage: lastUsage,
|
||||
Media: rs.mediaResults,
|
||||
Deliverables: rs.deliverables,
|
||||
BlockReplies: rs.blockReplies,
|
||||
|
||||
@@ -213,16 +213,30 @@ func hasPendingToolResultAhead(msgs []providers.Message, start int, idQueue map[
|
||||
return false
|
||||
}
|
||||
|
||||
func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
|
||||
// maybeSummarize truncates+summarizes session history when it grows large.
|
||||
//
|
||||
// midLoopCompacted signals that the final-request guard already had to compact
|
||||
// this session mid-loop. Because mid-loop compaction only mutates the run's
|
||||
// message buffer (never the session store), it is thrown away every turn —
|
||||
// causing an unbounded re-compaction loop AND stalling the cumulative compaction
|
||||
// count episodic depends on. When set, we lower the trigger threshold to a
|
||||
// unit-aligned "pressure threshold" so the compaction gets PERSISTED here.
|
||||
func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string, midLoopCompacted bool) {
|
||||
history := l.sessions.GetHistory(ctx, sessionKey)
|
||||
|
||||
// Use calibrated token estimation, adjusted for overhead.
|
||||
// lastPromptTokens includes everything (system prompt, tools, context files, history).
|
||||
// We subtract estimated overhead so the threshold comparison is history-only.
|
||||
lastPT, lastMC := l.sessions.GetLastPromptTokens(ctx, sessionKey)
|
||||
overheadEstimate := l.estimateOverhead(history, lastPT, lastMC)
|
||||
adjustedLastPT := max(lastPT-overheadEstimate, 0)
|
||||
tokenEstimate := EstimateTokensWithCalibration(history, adjustedLastPT, lastMC)
|
||||
calibrationPT, calibrationMC := lastPT, lastMC
|
||||
calibrationInvalid := l.contextWindow > 0 && lastPT > l.contextWindow
|
||||
if calibrationInvalid {
|
||||
calibrationPT = 0
|
||||
calibrationMC = 0
|
||||
}
|
||||
overheadEstimate := l.estimateOverhead(history, calibrationPT, calibrationMC)
|
||||
adjustedLastPT := max(calibrationPT-overheadEstimate, 0)
|
||||
tokenEstimate := EstimateTokensWithCalibration(history, adjustedLastPT, calibrationMC)
|
||||
|
||||
// Resolve compaction threshold from config: token-only (no message count guard).
|
||||
// Industry standard — Claude Code, Anthropic API, LangChain all use token-based thresholds.
|
||||
@@ -231,9 +245,28 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
|
||||
historyShare = l.compactionCfg.MaxHistoryShare
|
||||
}
|
||||
|
||||
// Baseline (floor) threshold: history-only, MaxHistoryShare. Used when there was
|
||||
// no mid-loop pressure this run — keeps normal sessions from over-compacting.
|
||||
threshold := int(float64(l.contextWindow) * historyShare)
|
||||
if tokenEstimate <= threshold {
|
||||
l.logCompactionDecision(sessionKey, "skip", "under_threshold", tokenEstimate, threshold, historyShare, lastPT, lastMC, adjustedLastPT, overheadEstimate)
|
||||
effectiveThreshold := threshold
|
||||
if midLoopCompacted {
|
||||
// Pressure threshold: align the unit + constants with the final-request guard by
|
||||
// reusing compactionInputCap() (min(cw-maxTokens, cw*MaxRequestShare-maxTokens)),
|
||||
// then subtract the fixed overhead estimate to compare against history-only
|
||||
// tokenEstimate. The guard reads MaxRequestShare while this path historically read
|
||||
// MaxHistoryShare — a DIFFERENT config field — so we must NOT rebuild the formula
|
||||
// from historyShare or the two would diverge under custom config.
|
||||
pressureThreshold := max(l.compactionInputCap()-overheadEstimate, 0)
|
||||
// Defensive floor: a pathological config (maxTokens >= cw*MaxRequestShare) could
|
||||
// drive compactionInputCap() toward 0, making pressureThreshold≈0 and forcing a
|
||||
// summarize on every mid-loop turn. Never drop below half the baseline threshold.
|
||||
if minFloor := threshold / 2; pressureThreshold < minFloor {
|
||||
pressureThreshold = minFloor
|
||||
}
|
||||
effectiveThreshold = pressureThreshold
|
||||
}
|
||||
if tokenEstimate <= effectiveThreshold {
|
||||
l.logCompactionDecision(sessionKey, "skip", "under_threshold", tokenEstimate, effectiveThreshold, historyShare, lastPT, lastMC, adjustedLastPT, overheadEstimate, calibrationInvalid)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -243,12 +276,12 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
|
||||
muI, _ := l.summarizeMu.LoadOrStore(sessionKey, &sync.Mutex{})
|
||||
sessionMu := muI.(*sync.Mutex)
|
||||
if !sessionMu.TryLock() {
|
||||
l.logCompactionDecision(sessionKey, "skip", "already_in_progress", tokenEstimate, threshold, historyShare, lastPT, lastMC, adjustedLastPT, overheadEstimate)
|
||||
l.logCompactionDecision(sessionKey, "skip", "already_in_progress", tokenEstimate, effectiveThreshold, historyShare, lastPT, lastMC, adjustedLastPT, overheadEstimate, calibrationInvalid)
|
||||
slog.Debug("summarization already in progress, skipping", "session", sessionKey)
|
||||
return
|
||||
}
|
||||
|
||||
l.logCompactionDecision(sessionKey, "trigger", "", tokenEstimate, threshold, historyShare, lastPT, lastMC, adjustedLastPT, overheadEstimate)
|
||||
l.logCompactionDecision(sessionKey, "trigger", "", tokenEstimate, effectiveThreshold, historyShare, lastPT, lastMC, adjustedLastPT, overheadEstimate, calibrationInvalid)
|
||||
|
||||
// Memory flush runs synchronously INSIDE the guard
|
||||
// (so concurrent runs don't both trigger flush for the same compaction cycle).
|
||||
@@ -282,42 +315,49 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
|
||||
summary := l.sessions.GetSummary(sctx, sessionKey)
|
||||
toSummarize := history[:len(history)-keepLast]
|
||||
|
||||
var sb strings.Builder
|
||||
// Resolve the same per-agent request budget the mid-loop compaction path
|
||||
// uses so no single summarization request can exceed the window. The old
|
||||
// path concatenated the entire history into one prompt, which could blow
|
||||
// the hard ceiling on large sessions; chunk + map-reduce keeps every
|
||||
// request within cap.
|
||||
inputCap := l.compactionInputCap()
|
||||
if inputCap <= 0 {
|
||||
slog.Warn("summarization failed", "session", sessionKey, "error", "context_window_unresolved")
|
||||
return
|
||||
}
|
||||
|
||||
var mediaKinds []string
|
||||
for _, m := range toSummarize {
|
||||
if m.Role == "user" {
|
||||
sb.WriteString(fmt.Sprintf("user: %s\n", m.Content))
|
||||
} else if m.Role == "assistant" {
|
||||
sb.WriteString(fmt.Sprintf("assistant: %s\n", SanitizeAssistantContent(m.Content)))
|
||||
}
|
||||
for _, ref := range m.MediaRefs {
|
||||
mediaKinds = append(mediaKinds, ref.Kind)
|
||||
}
|
||||
}
|
||||
|
||||
var prompt strings.Builder
|
||||
prompt.WriteString(compactionSummaryPrompt)
|
||||
// Build cap-respecting units, prepending any existing summary and a media
|
||||
// note as their own leading units so the packer can split them if needed.
|
||||
var leading []string
|
||||
if summary != "" {
|
||||
leading = append(leading, "Existing context: "+summary+"\n\n")
|
||||
}
|
||||
if len(mediaKinds) > 0 {
|
||||
// Deduplicate and count media types for a compact note.
|
||||
counts := make(map[string]int)
|
||||
for _, k := range mediaKinds {
|
||||
counts[k]++
|
||||
}
|
||||
prompt.WriteString("Note: user shared media files (")
|
||||
var note strings.Builder
|
||||
note.WriteString("Note: user shared media files (")
|
||||
first := true
|
||||
for k, n := range counts {
|
||||
if !first {
|
||||
prompt.WriteString(", ")
|
||||
note.WriteString(", ")
|
||||
}
|
||||
prompt.WriteString(fmt.Sprintf("%d %s(s)", n, k))
|
||||
note.WriteString(fmt.Sprintf("%d %s(s)", n, k))
|
||||
first = false
|
||||
}
|
||||
prompt.WriteString(") which are no longer in context. Mention briefly if relevant.\n\n")
|
||||
note.WriteString(") which are no longer in context. Mention briefly if relevant.\n\n")
|
||||
leading = append(leading, note.String())
|
||||
}
|
||||
if summary != "" {
|
||||
prompt.WriteString("Existing context: " + summary + "\n\n")
|
||||
}
|
||||
prompt.WriteString(sb.String())
|
||||
units := append(leading, buildCompactionUnits(toSummarize)...)
|
||||
|
||||
inTokens := l.estimateSummaryInputTokens(toSummarize)
|
||||
slog.Info("compact_budget",
|
||||
@@ -327,22 +367,19 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
|
||||
"in_tokens", inTokens,
|
||||
"out_tokens", dynamicSummaryMax(inTokens),
|
||||
"context_window", l.contextWindow,
|
||||
"input_cap_tokens", inputCap,
|
||||
"threshold", threshold,
|
||||
"token_estimate", tokenEstimate,
|
||||
"max_history_share", historyShare,
|
||||
"reserve_tokens_floor", l.resolveReserveTokens(),
|
||||
"timeout_seconds", int(timeout/time.Second),
|
||||
)
|
||||
chatReq := providers.ChatRequest{
|
||||
Messages: []providers.Message{{Role: "user", Content: prompt.String()}},
|
||||
Model: l.model,
|
||||
Options: map[string]any{"max_tokens": dynamicSummaryMax(inTokens), "temperature": 0.3},
|
||||
}
|
||||
resp, err := l.callInternalLLMWithUsage(sctx, chatReq, "session-summarization")
|
||||
summaryContent, _, err := l.summarizeCompactionUnits(sctx, units, inputCap, 1)
|
||||
if err != nil {
|
||||
slog.Warn("summarization failed", "session", sessionKey, "error", err)
|
||||
return
|
||||
}
|
||||
resp := &providers.ChatResponse{Content: summaryContent}
|
||||
|
||||
// Collect MediaRefs from messages about to be truncated (keep up to 30 most recent).
|
||||
const maxPreservedMediaRefs = 30
|
||||
@@ -381,7 +418,7 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
|
||||
}()
|
||||
}
|
||||
|
||||
func (l *Loop) logCompactionDecision(sessionKey, decision, skipReason string, tokenEstimate, threshold int, historyShare float64, lastPromptTokens, lastMessageCount, adjustedLastPromptTokens, overheadEstimate int) {
|
||||
func (l *Loop) logCompactionDecision(sessionKey, decision, skipReason string, tokenEstimate, threshold int, historyShare float64, lastPromptTokens, lastMessageCount, adjustedLastPromptTokens, overheadEstimate int, calibrationInvalid bool) {
|
||||
args := []any{
|
||||
"path", "post-turn",
|
||||
"agent", l.id,
|
||||
@@ -396,6 +433,10 @@ func (l *Loop) logCompactionDecision(sessionKey, decision, skipReason string, to
|
||||
"last_message_count", lastMessageCount,
|
||||
"adjusted_last_prompt_tokens", adjustedLastPromptTokens,
|
||||
"overhead_estimate", overheadEstimate,
|
||||
"calibration_invalid", calibrationInvalid,
|
||||
}
|
||||
if calibrationInvalid {
|
||||
args = append(args, "invalid_last_prompt_tokens", lastPromptTokens)
|
||||
}
|
||||
if skipReason != "" {
|
||||
args = append(args, "skip_reason", skipReason)
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -21,6 +22,21 @@ type nopSessionStore struct {
|
||||
history []providers.Message
|
||||
lastPromptTokens int
|
||||
lastMsgCount int
|
||||
inputTokens int64
|
||||
outputTokens int64
|
||||
setLastTokens int
|
||||
setLastMsgCount int
|
||||
|
||||
// Configurable/recording fields for compaction-count and truncation tests.
|
||||
// Guarded by mu because maybeSummarize mutates them from a background goroutine.
|
||||
mu sync.Mutex
|
||||
summary string // returned by GetSummary
|
||||
compactionCount int // returned by GetCompactionCount
|
||||
incrementCalls int // counts IncrementCompaction invocations
|
||||
truncateCalls int // counts TruncateHistory invocations
|
||||
truncateKeepLast int // last keepLast passed to TruncateHistory
|
||||
setSummaryCalls int // counts SetSummary invocations
|
||||
lastSetSummary string // last summary passed to SetSummary
|
||||
}
|
||||
|
||||
// SessionCoreStore methods
|
||||
@@ -30,24 +46,61 @@ func (n *nopSessionStore) GetOrCreate(_ context.Context, _ string) *store.Sessio
|
||||
func (n *nopSessionStore) Get(_ context.Context, _ string) *store.SessionData { return nil }
|
||||
func (n *nopSessionStore) AddMessage(_ context.Context, _ string, _ providers.Message) {}
|
||||
func (n *nopSessionStore) GetHistory(_ context.Context, _ string) []providers.Message {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
return n.history
|
||||
}
|
||||
func (n *nopSessionStore) GetSummary(_ context.Context, _ string) string { return "" }
|
||||
func (n *nopSessionStore) SetSummary(_ context.Context, _, _ string) {}
|
||||
func (n *nopSessionStore) GetSummary(_ context.Context, _ string) string {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
return n.summary
|
||||
}
|
||||
func (n *nopSessionStore) SetSummary(_ context.Context, _, s string) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.setSummaryCalls++
|
||||
n.lastSetSummary = s
|
||||
}
|
||||
func (n *nopSessionStore) GetLabel(_ context.Context, _ string) string { return "" }
|
||||
func (n *nopSessionStore) SetLabel(_ context.Context, _, _ string) {}
|
||||
func (n *nopSessionStore) SetAgentInfo(_ context.Context, _ string, _ uuid.UUID, _ string) {}
|
||||
func (n *nopSessionStore) TruncateHistory(_ context.Context, _ string, _ int) {}
|
||||
func (n *nopSessionStore) SetHistory(_ context.Context, _ string, _ []providers.Message) {}
|
||||
func (n *nopSessionStore) TruncateHistory(_ context.Context, _ string, keepLast int) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.truncateCalls++
|
||||
n.truncateKeepLast = keepLast
|
||||
// Mirror real-store semantics: keep only the last keepLast messages so
|
||||
// anti-loop tests observe an actually-shrunken history on the next turn.
|
||||
if keepLast >= 0 && keepLast < len(n.history) {
|
||||
n.history = n.history[len(n.history)-keepLast:]
|
||||
}
|
||||
}
|
||||
func (n *nopSessionStore) SetHistory(_ context.Context, _ string, msgs []providers.Message) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.history = msgs
|
||||
}
|
||||
func (n *nopSessionStore) Reset(_ context.Context, _ string) {}
|
||||
func (n *nopSessionStore) Delete(_ context.Context, _ string) error { return nil }
|
||||
func (n *nopSessionStore) Save(_ context.Context, _ string) error { return nil }
|
||||
|
||||
// SessionMetadataStore methods
|
||||
func (n *nopSessionStore) UpdateMetadata(_ context.Context, _, _, _, _ string) {}
|
||||
func (n *nopSessionStore) AccumulateTokens(_ context.Context, _ string, _, _ int64) {}
|
||||
func (n *nopSessionStore) IncrementCompaction(_ context.Context, _ string) {}
|
||||
func (n *nopSessionStore) GetCompactionCount(_ context.Context, _ string) int { return 0 }
|
||||
func (n *nopSessionStore) UpdateMetadata(_ context.Context, _, _, _, _ string) {}
|
||||
func (n *nopSessionStore) AccumulateTokens(_ context.Context, _ string, input, output int64) {
|
||||
n.inputTokens += input
|
||||
n.outputTokens += output
|
||||
}
|
||||
func (n *nopSessionStore) IncrementCompaction(_ context.Context, _ string) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
n.incrementCalls++
|
||||
n.compactionCount++
|
||||
}
|
||||
func (n *nopSessionStore) GetCompactionCount(_ context.Context, _ string) int {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
return n.compactionCount
|
||||
}
|
||||
func (n *nopSessionStore) GetMemoryFlushCompactionCount(_ context.Context, _ string) int { return 0 }
|
||||
func (n *nopSessionStore) SetMemoryFlushDone(_ context.Context, _ string) {}
|
||||
func (n *nopSessionStore) GetSessionMetadata(_ context.Context, _ string) map[string]string {
|
||||
@@ -57,7 +110,12 @@ func (n *nopSessionStore) SetSessionMetadata(_ context.Context, _ string, _ map[
|
||||
func (n *nopSessionStore) SetSpawnInfo(_ context.Context, _, _ string, _ int) {}
|
||||
func (n *nopSessionStore) SetContextWindow(_ context.Context, _ string, _ int) {}
|
||||
func (n *nopSessionStore) GetContextWindow(_ context.Context, _ string) int { return 0 }
|
||||
func (n *nopSessionStore) SetLastPromptTokens(_ context.Context, _ string, _, _ int) {}
|
||||
func (n *nopSessionStore) SetLastPromptTokens(_ context.Context, _ string, tokens, msgCount int) {
|
||||
n.setLastTokens = tokens
|
||||
n.setLastMsgCount = msgCount
|
||||
n.lastPromptTokens = tokens
|
||||
n.lastMsgCount = msgCount
|
||||
}
|
||||
func (n *nopSessionStore) GetLastPromptTokens(_ context.Context, _ string) (int, int) {
|
||||
return n.lastPromptTokens, n.lastMsgCount
|
||||
}
|
||||
@@ -74,6 +132,19 @@ func (n *nopSessionStore) LastUsedChannel(_ context.Context, _ string) (string,
|
||||
return "", ""
|
||||
}
|
||||
|
||||
// Thread-safe accessors for recording fields (maybeSummarize mutates them from a
|
||||
// background goroutine, so tests must not read the fields directly under -race).
|
||||
func (n *nopSessionStore) truncateCount() int {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
return n.truncateCalls
|
||||
}
|
||||
func (n *nopSessionStore) incrementCount() int {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
return n.incrementCalls
|
||||
}
|
||||
|
||||
// signallingProvider wraps capturingProvider and signals a channel when Chat is called.
|
||||
type signallingProvider struct {
|
||||
capturingProvider
|
||||
@@ -94,11 +165,17 @@ func (s *signallingProvider) Chat(ctx context.Context, req providers.ChatRequest
|
||||
func TestMaybeSummarize_MaxTokensDynamic(t *testing.T) {
|
||||
const contextWindow = 10000
|
||||
|
||||
// Build history large enough to exceed the compaction threshold.
|
||||
// Agent-only budget: the calling agent reserves a FIXED maxTokens from the
|
||||
// window, so compactionInputCap() = min(window-maxTokens, share*window-maxTokens).
|
||||
// With window=10000, maxTokens=1024 → inputCap = min(8976, 7476) = 7476.
|
||||
const maxTokens = 1024
|
||||
|
||||
// Build history large enough to exceed the compaction threshold while still
|
||||
// packing the to-summarize slice into a single chunk under the input cap.
|
||||
// threshold = contextWindow * DefaultHistoryShare = 10000 * 0.85 = 8500.
|
||||
// EstimateTokens uses ~4 chars/token; 9000 tokens * 4 = 36000 chars of content.
|
||||
// Use 5 user-assistant pairs each carrying ~9000 chars so EstimateTokens > threshold.
|
||||
longContent := makeLongString(9000)
|
||||
// estimateMessageTokens ≈ runes/3; 10 msgs × 3000 chars → ~10000 tokens > 8500.
|
||||
// The 6 to-summarize messages pack into one ~6520-token chunk (< 7476 cap).
|
||||
longContent := makeLongString(3000)
|
||||
history := make([]providers.Message, 10)
|
||||
for i := range history {
|
||||
if i%2 == 0 {
|
||||
@@ -124,6 +201,7 @@ func TestMaybeSummarize_MaxTokensDynamic(t *testing.T) {
|
||||
provider: sp,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: contextWindow,
|
||||
maxTokens: maxTokens,
|
||||
sessions: sessions,
|
||||
// hasMemory = false → shouldRunMemoryFlush returns false (skip memory flush)
|
||||
hasMemory: false,
|
||||
@@ -132,7 +210,7 @@ func TestMaybeSummarize_MaxTokensDynamic(t *testing.T) {
|
||||
// tokenCounter nil → estimateSummaryInputTokens uses rune/3 fallback
|
||||
}
|
||||
|
||||
loop.maybeSummarize(context.Background(), "test-session-key")
|
||||
loop.maybeSummarize(context.Background(), "test-session-key", false)
|
||||
|
||||
// Wait for background goroutine to call provider.Chat (up to 5s).
|
||||
select {
|
||||
@@ -151,7 +229,7 @@ func TestMaybeSummarize_MaxTokensDynamic(t *testing.T) {
|
||||
t.Fatal("Options[\"max_tokens\"] not set in ChatRequest from maybeSummarize")
|
||||
}
|
||||
|
||||
maxTokens, ok := maxTokensRaw.(int)
|
||||
gotMaxTokens, ok := maxTokensRaw.(int)
|
||||
if !ok {
|
||||
t.Fatalf("Options[\"max_tokens\"] type = %T, want int", maxTokensRaw)
|
||||
}
|
||||
@@ -162,8 +240,8 @@ func TestMaybeSummarize_MaxTokensDynamic(t *testing.T) {
|
||||
toSummarize := history[:len(history)-4]
|
||||
expectedIn := loop.estimateSummaryInputTokens(toSummarize)
|
||||
wantMax := dynamicSummaryMax(expectedIn)
|
||||
if maxTokens != wantMax {
|
||||
t.Errorf("max_tokens = %d, want %d (dynamicSummaryMax(%d))", maxTokens, wantMax, expectedIn)
|
||||
if gotMaxTokens != wantMax {
|
||||
t.Errorf("max_tokens = %d, want %d (dynamicSummaryMax(%d))", gotMaxTokens, wantMax, expectedIn)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,7 +264,7 @@ func TestMaybeSummarize_LogsSkipDecisionUnderThreshold(t *testing.T) {
|
||||
}
|
||||
|
||||
logs := captureSlog(t, func() {
|
||||
loop.maybeSummarize(context.Background(), "test-session-key")
|
||||
loop.maybeSummarize(context.Background(), "test-session-key", false)
|
||||
})
|
||||
|
||||
assertLogContains(t, logs,
|
||||
@@ -200,10 +278,49 @@ func TestMaybeSummarize_LogsSkipDecisionUnderThreshold(t *testing.T) {
|
||||
)
|
||||
}
|
||||
|
||||
func TestMaybeSummarize_IgnoresInvalidLastPromptCalibration(t *testing.T) {
|
||||
provider := &capturingProvider{response: "unused"}
|
||||
sessions := &nopSessionStore{
|
||||
history: []providers.Message{
|
||||
{Role: "user", Content: "short request"},
|
||||
{Role: "assistant", Content: "short response"},
|
||||
},
|
||||
lastPromptTokens: 264722,
|
||||
lastMsgCount: 28,
|
||||
}
|
||||
loop := &Loop{
|
||||
id: "test-agent",
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: 200000,
|
||||
sessions: sessions,
|
||||
hasMemory: false,
|
||||
}
|
||||
|
||||
logs := captureSlog(t, func() {
|
||||
loop.maybeSummarize(context.Background(), "test-session-key", false)
|
||||
})
|
||||
|
||||
if len(provider.captured) != 0 {
|
||||
t.Fatalf("provider.Chat called %d times, want 0 for invalid calibration under fallback threshold", len(provider.captured))
|
||||
}
|
||||
assertLogContains(t, logs,
|
||||
"compaction_decision",
|
||||
`"decision":"skip"`,
|
||||
`"skip_reason":"under_threshold"`,
|
||||
`"last_prompt_tokens":264722`,
|
||||
`"calibration_invalid":true`,
|
||||
`"invalid_last_prompt_tokens":264722`,
|
||||
)
|
||||
}
|
||||
|
||||
func TestMaybeSummarize_LogsTriggerDecisionOverThreshold(t *testing.T) {
|
||||
const contextWindow = 10000
|
||||
// Fixed agent reserve so compactionInputCap() leaves room to pack the
|
||||
// to-summarize slice into one chunk (see TestMaybeSummarize_MaxTokensDynamic).
|
||||
const maxTokens = 1024
|
||||
|
||||
longContent := makeLongString(9000)
|
||||
longContent := makeLongString(3000)
|
||||
history := make([]providers.Message, 10)
|
||||
for i := range history {
|
||||
if i%2 == 0 {
|
||||
@@ -228,12 +345,13 @@ func TestMaybeSummarize_LogsTriggerDecisionOverThreshold(t *testing.T) {
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: contextWindow,
|
||||
maxTokens: maxTokens,
|
||||
sessions: sessions,
|
||||
hasMemory: false,
|
||||
}
|
||||
|
||||
logs := captureSlog(t, func() {
|
||||
loop.maybeSummarize(context.Background(), "test-session-key")
|
||||
loop.maybeSummarize(context.Background(), "test-session-key", false)
|
||||
|
||||
// Wait for the background summarize goroutine to reach the point where it
|
||||
// calls provider.Chat (which happens strictly after the "compact_budget"
|
||||
|
||||
@@ -0,0 +1,258 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// Fixture math for these tests (contextWindow=200000, maxTokens=8192, nil cfg):
|
||||
//
|
||||
// effectiveMaxTokens = 8192
|
||||
// compactionInputCap = min(200000-8192, int(200000*0.85)-8192) = min(191808, 161808) = 161808
|
||||
// baseline threshold = int(200000*0.85) = 170000
|
||||
// overhead (no calib) = min(int(200000*0.2), 40000) = 40000
|
||||
// pressureThreshold = max(161808-40000, 0) = 121808
|
||||
// defensive floor = 170000/2 = 85000 (< 121808, so unused)
|
||||
//
|
||||
// Token estimate (no calibration) = sum(runes(content)/3) over all history messages.
|
||||
// Bands used below:
|
||||
//
|
||||
// ~100k tokens → below BOTH thresholds → skip even under mid-loop pressure
|
||||
// ~150k tokens → between pressure(121808) and baseline(170000) → persist ONLY under pressure
|
||||
const (
|
||||
pressureContextWindow = 200000
|
||||
pressureMaxTokens = 8192
|
||||
)
|
||||
|
||||
// buildTokenHistory returns n alternating user/assistant messages whose combined
|
||||
// rune/3 token estimate is approximately targetTokens. Each message carries
|
||||
// targetTokens/n * 3 runes so EstimateTokens(history) ≈ targetTokens.
|
||||
func buildTokenHistory(n, targetTokens int) []providers.Message {
|
||||
perMsgTokens := targetTokens / n
|
||||
perMsgRunes := perMsgTokens * 3
|
||||
content := makeLongString(perMsgRunes)
|
||||
history := make([]providers.Message, n)
|
||||
for i := range history {
|
||||
if i%2 == 0 {
|
||||
history[i] = providers.Message{Role: "user", Content: content}
|
||||
} else {
|
||||
history[i] = providers.Message{Role: "assistant", Content: content}
|
||||
}
|
||||
}
|
||||
return history
|
||||
}
|
||||
|
||||
// waitForCount polls get() until it reaches want or the deadline elapses.
|
||||
func waitForCount(t *testing.T, get func() int, want int) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if get() >= want {
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("timed out waiting for count to reach %d, last=%d", want, get())
|
||||
}
|
||||
|
||||
func newPressureLoop(sessions *nopSessionStore, provider providers.Provider) *Loop {
|
||||
return &Loop{
|
||||
id: "test-agent",
|
||||
provider: provider,
|
||||
model: "claude-3-5-sonnet",
|
||||
contextWindow: pressureContextWindow,
|
||||
maxTokens: pressureMaxTokens,
|
||||
sessions: sessions,
|
||||
hasMemory: false,
|
||||
compactionCfg: nil,
|
||||
}
|
||||
}
|
||||
|
||||
// V1-B / V2 core: a run that compacted mid-loop lowers the threshold to
|
||||
// pressureThreshold (121808). History at ~150k tokens is BELOW the baseline
|
||||
// (170000) but ABOVE the pressure threshold → summarize must run and PERSIST
|
||||
// (TruncateHistory + IncrementCompaction each exactly once).
|
||||
func TestMaybeSummarize_MidLoopPressurePersists(t *testing.T) {
|
||||
sessions := &nopSessionStore{
|
||||
history: buildTokenHistory(20, 150000), // ~150k tokens
|
||||
}
|
||||
loop := newPressureLoop(sessions, &capturingProvider{response: "compaction summary"})
|
||||
|
||||
loop.maybeSummarize(context.Background(), "sess-1", true)
|
||||
|
||||
waitForCount(t, sessions.incrementCount, 1)
|
||||
|
||||
if got := sessions.incrementCount(); got != 1 {
|
||||
t.Errorf("IncrementCompaction called %d times, want 1", got)
|
||||
}
|
||||
if got := sessions.truncateCount(); got != 1 {
|
||||
t.Errorf("TruncateHistory called %d times, want 1", got)
|
||||
}
|
||||
sessions.mu.Lock()
|
||||
keepLast := sessions.truncateKeepLast
|
||||
setSummary := sessions.setSummaryCalls
|
||||
sessions.mu.Unlock()
|
||||
if keepLast != 4 {
|
||||
t.Errorf("TruncateHistory keepLast = %d, want 4", keepLast)
|
||||
}
|
||||
if setSummary != 1 {
|
||||
t.Errorf("SetSummary called %d times, want 1", setSummary)
|
||||
}
|
||||
}
|
||||
|
||||
// V2 no-over-compaction (baseline case): the SAME ~150k history without mid-loop
|
||||
// pressure must SKIP (150000 < 170000 baseline). No summarize, no persist.
|
||||
func TestMaybeSummarize_NonMidLoopSkipsUnderBaseline(t *testing.T) {
|
||||
provider := &capturingProvider{response: "should not be called"}
|
||||
sessions := &nopSessionStore{
|
||||
history: buildTokenHistory(20, 150000), // ~150k tokens
|
||||
}
|
||||
loop := newPressureLoop(sessions, provider)
|
||||
|
||||
loop.maybeSummarize(context.Background(), "sess-1", false)
|
||||
|
||||
// Skip returns synchronously before spawning the goroutine; give any
|
||||
// (erroneously spawned) goroutine a brief window to prove it didn't run.
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if got := len(provider.captured); got != 0 {
|
||||
t.Errorf("provider.Chat called %d times, want 0 (under baseline, no pressure)", got)
|
||||
}
|
||||
if got := sessions.incrementCount(); got != 0 {
|
||||
t.Errorf("IncrementCompaction called %d times, want 0", got)
|
||||
}
|
||||
if got := sessions.truncateCount(); got != 0 {
|
||||
t.Errorf("TruncateHistory called %d times, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
// V2 no-over-compaction (the KEY case that distinguishes "lower threshold" from
|
||||
// "skip threshold"): mid-loop pressure is set, but the real history is small
|
||||
// (~100k tokens, below pressureThreshold=121808). The mid-loop compaction was
|
||||
// driven by transient tool-result bloat, not history — so we must STILL SKIP and
|
||||
// NOT truncate the session. This proves we lowered the threshold rather than
|
||||
// bypassing it.
|
||||
func TestMaybeSummarize_MidLoopSkipsWhenHistorySmall(t *testing.T) {
|
||||
provider := &capturingProvider{response: "should not be called"}
|
||||
sessions := &nopSessionStore{
|
||||
history: buildTokenHistory(20, 100000), // ~100k tokens < pressureThreshold
|
||||
}
|
||||
loop := newPressureLoop(sessions, provider)
|
||||
|
||||
loop.maybeSummarize(context.Background(), "sess-1", true)
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if got := len(provider.captured); got != 0 {
|
||||
t.Errorf("provider.Chat called %d times, want 0 (history below pressure threshold)", got)
|
||||
}
|
||||
if got := sessions.incrementCount(); got != 0 {
|
||||
t.Errorf("IncrementCompaction called %d times, want 0 (no over-compaction)", got)
|
||||
}
|
||||
if got := sessions.truncateCount(); got != 0 {
|
||||
t.Errorf("TruncateHistory called %d times, want 0 (no over-compaction)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// V2 threshold-sync boundary: with one fixed history between the pressure
|
||||
// threshold (121808) and the baseline (170000), the ONLY thing that flips the
|
||||
// decision is the mid-loop flag. Persist when true, skip when false — proving
|
||||
// the two thresholds bracket this history exactly as designed.
|
||||
func TestMaybeSummarize_ThresholdSyncBoundary(t *testing.T) {
|
||||
history := buildTokenHistory(20, 150000) // between 121808 and 170000
|
||||
|
||||
// mid-loop = false → skip
|
||||
skipProvider := &capturingProvider{response: "unused"}
|
||||
skipSessions := &nopSessionStore{history: cloneMessages(history)}
|
||||
newPressureLoop(skipSessions, skipProvider).maybeSummarize(context.Background(), "sess-1", false)
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
if got := skipSessions.incrementCount(); got != 0 {
|
||||
t.Errorf("non-mid-loop: IncrementCompaction = %d, want 0 (skip above pressure, below baseline)", got)
|
||||
}
|
||||
|
||||
// mid-loop = true → persist
|
||||
triggerSessions := &nopSessionStore{history: cloneMessages(history)}
|
||||
newPressureLoop(triggerSessions, &capturingProvider{response: "compaction summary"}).
|
||||
maybeSummarize(context.Background(), "sess-1", true)
|
||||
waitForCount(t, triggerSessions.incrementCount, 1)
|
||||
if got := triggerSessions.incrementCount(); got != 1 {
|
||||
t.Errorf("mid-loop: IncrementCompaction = %d, want 1 (persist under pressure)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// V2 anti-loop (characterization): turn 1 persists the mid-loop compaction,
|
||||
// shrinking the session history. Turn 2 loads the now-small history, so even
|
||||
// with mid-loop pressure still set it stays below the pressure threshold and
|
||||
// does NOT re-compact — the unbounded re-compaction loop is broken.
|
||||
func TestMaybeSummarize_AntiLoop_SecondTurnDoesNotRecompact(t *testing.T) {
|
||||
sessions := &nopSessionStore{
|
||||
history: buildTokenHistory(20, 150000),
|
||||
}
|
||||
loop := newPressureLoop(sessions, &capturingProvider{response: "compaction summary"})
|
||||
|
||||
// Turn 1: pressure present, history large → persist (truncate to keepLast=4).
|
||||
loop.maybeSummarize(context.Background(), "sess-1", true)
|
||||
waitForCount(t, sessions.incrementCount, 1)
|
||||
|
||||
sessions.mu.Lock()
|
||||
remaining := len(sessions.history)
|
||||
sessions.mu.Unlock()
|
||||
if remaining > 4 {
|
||||
t.Fatalf("after turn 1 persist, history len = %d, want <= 4 (keepLast)", remaining)
|
||||
}
|
||||
|
||||
// Turn 2: history is now tiny (4 messages). Even with pressure set, the
|
||||
// history-only estimate is far below pressureThreshold → skip.
|
||||
loop.maybeSummarize(context.Background(), "sess-1", true)
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if got := sessions.incrementCount(); got != 1 {
|
||||
t.Errorf("after turn 2, IncrementCompaction total = %d, want 1 (no re-compaction)", got)
|
||||
}
|
||||
if got := sessions.truncateCount(); got != 1 {
|
||||
t.Errorf("after turn 2, TruncateHistory total = %d, want 1 (no re-compaction)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// Media preservation regression: MediaRefs on messages about to be truncated
|
||||
// must be carried onto the first kept message so the agent doesn't lose track of
|
||||
// shared files across a pressure-driven compaction.
|
||||
func TestMaybeSummarize_MidLoopPreservesMediaRefs(t *testing.T) {
|
||||
history := buildTokenHistory(20, 150000)
|
||||
// Attach a media ref to an early message (will be in the to-summarize slice).
|
||||
history[0].MediaRefs = []providers.MediaRef{
|
||||
{ID: "img-1", MimeType: "image/png", Kind: "image", Path: "/tmp/img-1.png"},
|
||||
}
|
||||
sessions := &nopSessionStore{history: history}
|
||||
loop := newPressureLoop(sessions, &capturingProvider{response: "compaction summary"})
|
||||
|
||||
loop.maybeSummarize(context.Background(), "sess-1", true)
|
||||
waitForCount(t, sessions.incrementCount, 1)
|
||||
|
||||
sessions.mu.Lock()
|
||||
defer sessions.mu.Unlock()
|
||||
if len(sessions.history) == 0 {
|
||||
t.Fatal("history empty after truncation")
|
||||
}
|
||||
var found bool
|
||||
for _, ref := range sessions.history[0].MediaRefs {
|
||||
if ref.ID == "img-1" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("preserved media ref img-1 not found on first kept message; refs=%+v", sessions.history[0].MediaRefs)
|
||||
}
|
||||
}
|
||||
|
||||
// cloneMessages returns a shallow copy of a message slice so two Loops don't
|
||||
// share the same backing array (TruncateHistory mutates it in the mock).
|
||||
func cloneMessages(msgs []providers.Message) []providers.Message {
|
||||
out := make([]providers.Message, len(msgs))
|
||||
copy(out, msgs)
|
||||
return out
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/config"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/eventbus"
|
||||
@@ -18,8 +19,10 @@ func (l *Loop) runViaPipeline(ctx context.Context, req RunRequest) (*RunResult,
|
||||
input := convertRunInput(&req)
|
||||
// Bridge runState shares loop detection state between pipeline and agent.
|
||||
bridgeRS := &runState{}
|
||||
deps := l.buildPipelineDeps(&req, bridgeRS)
|
||||
|
||||
// Resolve the effective model + provider BEFORE building deps so the pre-call
|
||||
// budget estimate reserves reasoning output for the model that will actually
|
||||
// run (a ModelOverride can change the reasoning capability, hence the bump).
|
||||
model := l.model
|
||||
if req.ModelOverride != "" {
|
||||
model = req.ModelOverride
|
||||
@@ -33,6 +36,8 @@ func (l *Loop) runViaPipeline(ctx context.Context, req RunRequest) (*RunResult,
|
||||
}
|
||||
}
|
||||
|
||||
deps := l.buildPipelineDeps(&req, bridgeRS)
|
||||
|
||||
p := pipeline.NewDefaultPipeline(deps)
|
||||
state := pipeline.NewRunState(input, nil, model, provider)
|
||||
|
||||
@@ -44,6 +49,9 @@ func (l *Loop) runViaPipeline(ctx context.Context, req RunRequest) (*RunResult,
|
||||
}
|
||||
|
||||
// buildPipelineDeps maps Loop fields + methods to PipelineDeps callbacks.
|
||||
// effProvider/effModel are the resolved provider+model for THIS run (after any
|
||||
// ModelOverride/ProviderOverride) so reasoning-effort resolution matches the
|
||||
// request the pipeline will actually send.
|
||||
func (l *Loop) buildPipelineDeps(req *RunRequest, bridgeRS *runState) pipeline.PipelineDeps {
|
||||
maxIter := l.maxIterations
|
||||
if req.MaxIterations > 0 && req.MaxIterations < maxIter {
|
||||
@@ -69,9 +77,10 @@ func (l *Loop) buildPipelineDeps(req *RunRequest, bridgeRS *runState) pipeline.P
|
||||
}
|
||||
|
||||
return pipeline.PipelineDeps{
|
||||
TokenCounter: tokencount.NewTiktokenCounter(),
|
||||
EventBus: l.domainBus,
|
||||
Hooks: l.hookDispatcher,
|
||||
TokenCounter: tokencount.NewTiktokenCounter(),
|
||||
BudgetCounter: l.budgetCounter,
|
||||
EventBus: l.domainBus,
|
||||
Hooks: l.hookDispatcher,
|
||||
Config: pipeline.PipelineConfig{
|
||||
MaxIterations: maxIter,
|
||||
MaxToolCalls: l.maxToolCalls,
|
||||
@@ -82,19 +91,7 @@ func (l *Loop) buildPipelineDeps(req *RunRequest, bridgeRS *runState) pipeline.P
|
||||
Compaction: l.compactionCfg,
|
||||
// V3 memory/retrieval flags removed — always true at runtime.
|
||||
},
|
||||
// Resolve per-model context window once per run. Falls back to
|
||||
// Config.ContextWindow when registry/model is unknown (existing
|
||||
// behaviour unchanged for tests and lite edition).
|
||||
ResolveContextWindow: func(provider, model string) int {
|
||||
if l.modelRegistry == nil || model == "" {
|
||||
return 0
|
||||
}
|
||||
spec := l.modelRegistry.Resolve(provider, model)
|
||||
if spec == nil {
|
||||
return 0
|
||||
}
|
||||
return spec.ContextWindow
|
||||
},
|
||||
ResolveContextWindow: l.resolveEffectiveContextWindow,
|
||||
EmitEvent: func(event any) {
|
||||
if ae, ok := event.(AgentEvent); ok {
|
||||
l.emit(ae)
|
||||
@@ -184,30 +181,10 @@ func (l *Loop) buildPipelineDeps(req *RunRequest, bridgeRS *runState) pipeline.P
|
||||
StripMessageDirectives: StripMessageDirectives,
|
||||
DeduplicateMediaSuffix: deduplicateMediaSuffix,
|
||||
IsSilentReply: IsSilentReply,
|
||||
EmitSessionCompleted: func(ctx context.Context, sessionKey string, msgCount, tokensUsed, compactionCount int) {
|
||||
if l.domainBus != nil {
|
||||
// Include existing session summary (from previous compaction cycles).
|
||||
// Current cycle's compaction runs async so its summary isn't ready yet,
|
||||
// but previous summaries are available and useful for episodic creation.
|
||||
var summary string
|
||||
if compactionCount > 0 {
|
||||
summary = l.sessions.GetSummary(ctx, sessionKey)
|
||||
}
|
||||
l.domainBus.Publish(eventbus.DomainEvent{
|
||||
Type: eventbus.EventSessionCompleted,
|
||||
TenantID: l.tenantID.String(),
|
||||
AgentID: l.agentUUID.String(),
|
||||
UserID: req.UserID,
|
||||
SourceID: sessionKey,
|
||||
Payload: &eventbus.SessionCompletedPayload{
|
||||
SessionKey: sessionKey,
|
||||
MessageCount: msgCount,
|
||||
TokensUsed: tokensUsed,
|
||||
CompactionCount: compactionCount,
|
||||
Summary: summary,
|
||||
},
|
||||
})
|
||||
}
|
||||
EmitSessionCompleted: func(ctx context.Context, sessionKey string, msgCount, tokensUsed, _ int) {
|
||||
// The per-run count (5th arg) is intentionally ignored — emitSessionCompleted
|
||||
// reads the CUMULATIVE session count itself (Bug A). See method doc.
|
||||
l.emitSessionCompleted(ctx, sessionKey, req.UserID, msgCount, tokensUsed)
|
||||
},
|
||||
UpdateMetadata: cb.updateMetadata,
|
||||
BootstrapCleanup: cb.bootstrapCleanup,
|
||||
@@ -215,6 +192,51 @@ func (l *Loop) buildPipelineDeps(req *RunRequest, bridgeRS *runState) pipeline.P
|
||||
}
|
||||
}
|
||||
|
||||
// emitSessionCompleted publishes the session.completed domain event that drives
|
||||
// the consolidation pipeline (episodic → semantic → dreaming). No-op when no bus.
|
||||
//
|
||||
// Bug A: it reads the CUMULATIVE session compaction count via GetCompactionCount
|
||||
// (matching the legacy v2 emit path) rather than the per-run counter the pipeline
|
||||
// tracks — the per-run counter resets to 0 each run, which pinned source_id at
|
||||
// ":0" and made the episodic worker skip every cycle after the first.
|
||||
//
|
||||
// Bug C: SourceID embeds that count ("<sessionKey>:<count>") so the eventbus dedup
|
||||
// key (Type+":"+SourceID, 5m TTL) advances per compaction cycle instead of
|
||||
// swallowing every rapid same-session turn. Safe for the worker, which builds its
|
||||
// own idempotency key from payload.SessionKey+payload.CompactionCount and never
|
||||
// parses SourceID.
|
||||
func (l *Loop) emitSessionCompleted(ctx context.Context, sessionKey, userID string, msgCount, tokensUsed int) {
|
||||
if l.domainBus == nil {
|
||||
return
|
||||
}
|
||||
count := l.sessions.GetCompactionCount(ctx, sessionKey)
|
||||
// Attach the existing session summary (from a PREVIOUS compaction cycle) when
|
||||
// one exists. The current cycle's summary is produced asynchronously and isn't
|
||||
// ready yet, but a prior summary spares the worker an LLM re-summarize call.
|
||||
var summary string
|
||||
if count > 0 {
|
||||
summary = l.sessions.GetSummary(ctx, sessionKey)
|
||||
}
|
||||
l.domainBus.Publish(eventbus.DomainEvent{
|
||||
Type: eventbus.EventSessionCompleted,
|
||||
TenantID: l.tenantID.String(),
|
||||
AgentID: l.agentUUID.String(),
|
||||
UserID: userID,
|
||||
SourceID: fmt.Sprintf("%s:%d", sessionKey, count),
|
||||
Payload: &eventbus.SessionCompletedPayload{
|
||||
SessionKey: sessionKey,
|
||||
MessageCount: msgCount,
|
||||
TokensUsed: tokensUsed,
|
||||
CompactionCount: count,
|
||||
Summary: summary,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func (l *Loop) resolveEffectiveContextWindow() int {
|
||||
return l.contextWindow
|
||||
}
|
||||
|
||||
// convertRunInput converts agent.RunRequest to pipeline.RunInput.
|
||||
func convertRunInput(req *RunRequest) *pipeline.RunInput {
|
||||
return &pipeline.RunInput{
|
||||
@@ -260,6 +282,11 @@ func convertRunResult(pr *pipeline.RunResult) *RunResult {
|
||||
if pr == nil {
|
||||
return nil
|
||||
}
|
||||
var lastUsage *providers.Usage
|
||||
if pr.LastUsage.PromptTokens > 0 || pr.LastUsage.CompletionTokens > 0 || pr.LastUsage.TotalTokens > 0 {
|
||||
lu := pr.LastUsage
|
||||
lastUsage = &lu
|
||||
}
|
||||
media := make([]MediaResult, len(pr.MediaResults))
|
||||
for i, m := range pr.MediaResults {
|
||||
media[i] = MediaResult{
|
||||
@@ -276,6 +303,7 @@ func convertRunResult(pr *pipeline.RunResult) *RunResult {
|
||||
RunID: pr.RunID,
|
||||
Iterations: pr.Iterations,
|
||||
Usage: &pr.TotalUsage,
|
||||
LastUsage: lastUsage,
|
||||
Media: media,
|
||||
Deliverables: pr.Deliverables,
|
||||
BlockReplies: pr.BlockReplies,
|
||||
|
||||
@@ -95,7 +95,7 @@ type pipelineCallbackSet struct {
|
||||
flushMessages func(ctx context.Context, sessionKey string, msgs []providers.Message) error
|
||||
updateMetadata func(ctx context.Context, sessionKey string, usage, lastUsage providers.Usage, msgCount int) error
|
||||
bootstrapCleanup func(ctx context.Context, state *pipeline.RunState) error
|
||||
maybeSummarize func(ctx context.Context, sessionKey string)
|
||||
maybeSummarize func(ctx context.Context, sessionKey string, midLoopCompacted bool)
|
||||
}
|
||||
|
||||
func (l *Loop) makeResolveWorkspace(req *RunRequest) func(ctx context.Context, input *pipeline.RunInput) (*workspace.WorkspaceContext, error) {
|
||||
@@ -759,13 +759,25 @@ func (l *Loop) makeBootstrapCleanup() func(ctx context.Context, state *pipeline.
|
||||
}
|
||||
|
||||
func (l *Loop) reserveLLMUsage(ctx context.Context, req *RunRequest, state *pipeline.RunState, chatReq providers.ChatRequest, attempt string) (*usagecaps.Reservation, error) {
|
||||
if l.usageCaps == nil || state.Provider == nil {
|
||||
return nil, nil
|
||||
providerName := ""
|
||||
if state.Provider != nil {
|
||||
providerName = state.Provider.Name()
|
||||
}
|
||||
return l.reserveLLMUsageFor(ctx, req, state.Iteration, chatReq, attempt, state.Provider.Name(), state.Model)
|
||||
// reserveLLMUsageFor runs the mandatory hard-ceiling guard before any
|
||||
// reservation or transport, so the non-fallback path is guarded here even
|
||||
// when usage caps are disabled (Lite runtime).
|
||||
return l.reserveLLMUsageFor(ctx, req, state.Iteration, chatReq, attempt, providerName, state.Model)
|
||||
}
|
||||
|
||||
func (l *Loop) reserveLLMUsageFor(ctx context.Context, req *RunRequest, iteration int, chatReq providers.ChatRequest, attempt, providerName, model string) (*usagecaps.Reservation, error) {
|
||||
// Mandatory final hard-ceiling guard for the concrete request that is about
|
||||
// to be sent, after all directive/retry/reasoning mutations. This is the
|
||||
// shared pre-transport chokepoint for BOTH the fallback candidate path and
|
||||
// the non-fallback path (via reserveLLMUsage), and for every retry attempt.
|
||||
// It runs regardless of usage-cap configuration so the ceiling holds on Lite.
|
||||
if guardErr := l.guardCompleteModelRequest(chatReq, providerName, model, attempt); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
if l.usageCaps == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -103,6 +103,59 @@ func TestSupportsPromptCacheParams(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveEffectiveContextWindow_UsesAgentConfigOnly(t *testing.T) {
|
||||
registry := &panicModelRegistry{}
|
||||
loop := &Loop{contextWindow: 128_000, modelRegistry: registry}
|
||||
if got := loop.resolveEffectiveContextWindow(); got != 128_000 {
|
||||
t.Fatalf("resolveEffectiveContextWindow() = %d, want 128000", got)
|
||||
}
|
||||
}
|
||||
|
||||
type panicModelRegistry struct{}
|
||||
|
||||
func (*panicModelRegistry) Resolve(_, _ string) *providers.ModelSpec {
|
||||
panic("model registry must not participate in request budgeting")
|
||||
}
|
||||
|
||||
func (*panicModelRegistry) Register(providers.ModelSpec) {
|
||||
panic("model registry must not participate in request budgeting")
|
||||
}
|
||||
|
||||
func (*panicModelRegistry) Catalog(string) []providers.ModelSpec {
|
||||
panic("model registry must not participate in request budgeting")
|
||||
}
|
||||
|
||||
func TestMakeUpdateMetadataStoresLastUsagePromptTokens(t *testing.T) {
|
||||
sessions := &nopSessionStore{}
|
||||
loop := &Loop{
|
||||
model: "test-model",
|
||||
provider: finalThinkingStreamProvider{},
|
||||
sessions: sessions,
|
||||
}
|
||||
req := &RunRequest{Channel: "telegram"}
|
||||
|
||||
update := loop.makeUpdateMetadata(req)
|
||||
err := update(context.Background(), "sess-1",
|
||||
providers.Usage{PromptTokens: 225000, CompletionTokens: 3000},
|
||||
providers.Usage{PromptTokens: 70000, CompletionTokens: 100},
|
||||
12,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("update metadata error: %v", err)
|
||||
}
|
||||
if sessions.inputTokens != 225000 || sessions.outputTokens != 3000 {
|
||||
t.Fatalf("accumulated tokens = %d/%d, want total run 225000/3000", sessions.inputTokens, sessions.outputTokens)
|
||||
}
|
||||
// Upstream 503909d3 calibration: SetLastPromptTokens stores the final call's
|
||||
// full context size (Usage.ContextTokens(), which adds cached segments back
|
||||
// for Anthropic-style usage) PLUS the final completion — the reply joins
|
||||
// history so it occupies the next request's prompt. No cache tokens here, so
|
||||
// ContextTokens()=70000; +100 completion = 70100.
|
||||
if sessions.setLastTokens != 70100 || sessions.setLastMsgCount != 12 {
|
||||
t.Fatalf("last prompt calibration = %d/%d, want last request 70100/12", sessions.setLastTokens, sessions.setLastMsgCount)
|
||||
}
|
||||
}
|
||||
|
||||
// A Function-nil tool definition (e.g. the native image_generation sentinel,
|
||||
// providers.ToolDefinition{Type: "image_generation"}) must not panic the
|
||||
// mcp-def counter. Regression for the v3.14.0 nil-pointer crash.
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// RequestBudgetExceededError is returned before provider transport when the
|
||||
// COMPLETE final request (messages + tools + output reserve) exceeds the
|
||||
// receiving agent's configured context window. It is intentionally internal
|
||||
// to the agent loop: callers abort/reduce instead of falling through to model
|
||||
// fallback or sending an oversized payload.
|
||||
type RequestBudgetExceededError struct {
|
||||
AgentID string
|
||||
Provider string
|
||||
Model string
|
||||
Attempt string
|
||||
InputTokens int
|
||||
OutputReserve int
|
||||
ContextWindow int
|
||||
}
|
||||
|
||||
func (e *RequestBudgetExceededError) Error() string {
|
||||
return fmt.Sprintf(
|
||||
"context budget exceeded: input=%d + output_reserve=%d > agent_window=%d (agent=%s provider=%s model=%s attempt=%s)",
|
||||
e.InputTokens, e.OutputReserve, e.ContextWindow, e.AgentID, e.Provider, e.Model, e.Attempt,
|
||||
)
|
||||
}
|
||||
|
||||
// ContextBudgetExceeded marks this error as a request-level context-budget
|
||||
// overflow. Package-neutral consumers (pipeline.ThinkStage) detect it via an
|
||||
// interface check without importing the agent package, so they can re-enter
|
||||
// reduction (prune -> compact -> shrink) and rebuild before aborting.
|
||||
func (e *RequestBudgetExceededError) ContextBudgetExceeded() bool { return true }
|
||||
|
||||
// guardCompleteModelRequest is the mandatory final boundary check for EVERY
|
||||
// model request emitted by an agent loop. Call it only after all messages,
|
||||
// tools, retry instructions, model overrides and options have been finalized,
|
||||
// and immediately before provider.Chat/ChatStream.
|
||||
//
|
||||
// Hard invariant:
|
||||
//
|
||||
// completeInputTokens + outputReserveTokens <= agent.ContextWindow
|
||||
//
|
||||
// The configured agent context window and configured max_tokens are the only
|
||||
// budget authorities. Model and provider are diagnostic fields only.
|
||||
func (l *Loop) guardCompleteModelRequest(request providers.ChatRequest, providerName, model, attempt string) error {
|
||||
// Bare Loop literals used by unrelated unit tests opt out; NewLoop always
|
||||
// wires the fixed local budget counter.
|
||||
if l.budgetCounter == nil {
|
||||
return nil
|
||||
}
|
||||
contextWindow := l.resolveEffectiveContextWindow()
|
||||
if contextWindow <= 0 {
|
||||
return fmt.Errorf("context_window_unresolved: agent=%s provider=%s model=%s attempt=%s", l.id, providerName, model, attempt)
|
||||
}
|
||||
|
||||
inputTokens, err := l.budgetCounter.CountRequest(request)
|
||||
if err != nil {
|
||||
return fmt.Errorf("count complete request: %w", err)
|
||||
}
|
||||
outputReserve := l.effectiveMaxTokens()
|
||||
|
||||
if inputTokens+outputReserve <= contextWindow {
|
||||
slog.Debug("context_budget.guard",
|
||||
"agent", l.id,
|
||||
"provider", providerName,
|
||||
"model", model,
|
||||
"attempt", attempt,
|
||||
"input_tokens", inputTokens,
|
||||
"output_reserve_tokens", outputReserve,
|
||||
"context_window", contextWindow,
|
||||
"action", "allow")
|
||||
return nil
|
||||
}
|
||||
|
||||
slog.Info("context_budget.guard",
|
||||
"agent", l.id,
|
||||
"provider", providerName,
|
||||
"model", model,
|
||||
"attempt", attempt,
|
||||
"input_tokens", inputTokens,
|
||||
"output_reserve_tokens", outputReserve,
|
||||
"context_window", contextWindow,
|
||||
"action", "abort",
|
||||
"reason", "context_budget_exceeded")
|
||||
|
||||
return &RequestBudgetExceededError{
|
||||
AgentID: l.id,
|
||||
Provider: providerName,
|
||||
Model: model,
|
||||
Attempt: attempt,
|
||||
InputTokens: inputTokens,
|
||||
OutputReserve: outputReserve,
|
||||
ContextWindow: contextWindow,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/tokencount"
|
||||
)
|
||||
|
||||
func newBudgetTestLoop(window, maxTokens int) *Loop {
|
||||
return &Loop{
|
||||
id: "budget-test",
|
||||
contextWindow: window,
|
||||
maxTokens: maxTokens,
|
||||
budgetCounter: tokencount.NewBudgetCounter(),
|
||||
}
|
||||
}
|
||||
|
||||
func budgetMessageAtLeast(t *testing.T, counter tokencount.BudgetCounter, target int) providers.Message {
|
||||
t.Helper()
|
||||
low, high := 1, target*8
|
||||
for low < high {
|
||||
mid := low + (high-low)/2
|
||||
msg := providers.Message{Role: "user", Content: strings.Repeat("word ", mid)}
|
||||
count, err := counter.CountMessages([]providers.Message{msg})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if count < target {
|
||||
low = mid + 1
|
||||
} else {
|
||||
high = mid
|
||||
}
|
||||
}
|
||||
return providers.Message{Role: "user", Content: strings.Repeat("word ", low)}
|
||||
}
|
||||
|
||||
func TestGuardCompleteModelRequest_AgentWindowOnly(t *testing.T) {
|
||||
counter := tokencount.NewBudgetCounter()
|
||||
msg := budgetMessageAtLeast(t, counter, 130_000)
|
||||
req := providers.ChatRequest{Model: "model-a", Messages: []providers.Message{msg}}
|
||||
|
||||
loop128 := newBudgetTestLoop(128_000, 8_192)
|
||||
err := loop128.guardCompleteModelRequest(req, "provider-a", "model-a", "initial")
|
||||
if err == nil {
|
||||
t.Fatal("128k agent: expected request to be blocked")
|
||||
}
|
||||
var budgetErr *RequestBudgetExceededError
|
||||
if !errors.As(err, &budgetErr) {
|
||||
t.Fatalf("expected RequestBudgetExceededError, got %T: %v", err, err)
|
||||
}
|
||||
|
||||
loop200 := newBudgetTestLoop(200_000, 8_192)
|
||||
if err := loop200.guardCompleteModelRequest(req, "provider-a", "model-a", "initial"); err != nil {
|
||||
t.Fatalf("200k agent: expected request to fit, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGuardCompleteModelRequest_ModelProviderIndependent(t *testing.T) {
|
||||
loop := newBudgetTestLoop(128_000, 8_192)
|
||||
base := providers.ChatRequest{
|
||||
Model: "model-a",
|
||||
Messages: []providers.Message{{
|
||||
Role: "assistant",
|
||||
Content: "same content",
|
||||
Thinking: "same thinking",
|
||||
ToolCalls: []providers.ToolCall{{
|
||||
ID: "call-1",
|
||||
Name: "read_file",
|
||||
Arguments: map[string]any{"path": "/tmp/a"},
|
||||
Metadata: map[string]string{"thought_signature": "sig"},
|
||||
}},
|
||||
}},
|
||||
}
|
||||
other := base
|
||||
other.Model = "totally-different-model"
|
||||
|
||||
first, err := loop.budgetCounter.CountRequest(base)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, err := loop.budgetCounter.CountRequest(other)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if first != second {
|
||||
t.Fatalf("completeInput changed with model: %d != %d", first, second)
|
||||
}
|
||||
if errA, errB := loop.guardCompleteModelRequest(base, "provider-a", base.Model, "a"), loop.guardCompleteModelRequest(other, "provider-b", other.Model, "b"); (errA == nil) != (errB == nil) {
|
||||
t.Fatalf("allow/abort changed with model/provider: %v vs %v", errA, errB)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGuardCompleteModelRequest_MaxTokensChangesHardCapExactly(t *testing.T) {
|
||||
counter := tokencount.NewBudgetCounter()
|
||||
msg := budgetMessageAtLeast(t, counter, 7_500)
|
||||
req := providers.ChatRequest{Messages: []providers.Message{msg}}
|
||||
input, err := counter.CountRequest(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
window := input + 2_500
|
||||
|
||||
allow := newBudgetTestLoop(window, 2_500)
|
||||
if err := allow.guardCompleteModelRequest(req, "p", "m", "allow"); err != nil {
|
||||
t.Fatalf("input + 2500 should fit exactly: input=%d window=%d err=%v", input, window, err)
|
||||
}
|
||||
block := newBudgetTestLoop(window, 2_501)
|
||||
if err := block.guardCompleteModelRequest(req, "p", "m", "block"); err == nil {
|
||||
t.Fatalf("input + 2501 must exceed by one: input=%d window=%d", input, window)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGuardCompleteModelRequest_FailsClosedOnUnresolvedAgentWindow(t *testing.T) {
|
||||
loop := newBudgetTestLoop(0, 8_192)
|
||||
req := providers.ChatRequest{Messages: []providers.Message{{Role: "user", Content: "hi"}}}
|
||||
if err := loop.guardCompleteModelRequest(req, "provider", "model", "initial"); err == nil {
|
||||
t.Fatal("expected wiring failure without configured agent window")
|
||||
}
|
||||
}
|
||||
@@ -206,7 +206,13 @@ func (l *Loop) Run(ctx context.Context, req RunRequest) (*RunResult, error) {
|
||||
"iterations", result.Iterations,
|
||||
}
|
||||
if result.Usage != nil {
|
||||
logAttrs = append(logAttrs, "total_tokens", result.Usage.TotalTokens)
|
||||
logAttrs = append(logAttrs,
|
||||
"total_tokens", result.Usage.TotalTokens,
|
||||
"total_prompt_tokens", result.Usage.PromptTokens,
|
||||
)
|
||||
}
|
||||
if result.LastUsage != nil {
|
||||
logAttrs = append(logAttrs, "last_usage_prompt_tokens", result.LastUsage.PromptTokens)
|
||||
}
|
||||
slog.Info("v3.run.completed", logAttrs...)
|
||||
|
||||
|
||||
@@ -152,9 +152,11 @@ type Loop struct {
|
||||
// Context pruning config (trim old tool results in-memory)
|
||||
contextPruningCfg *config.ContextPruningConfig
|
||||
|
||||
// tokenCounter provides accurate per-model token counting for context pruning.
|
||||
// Nil means the legacy char-based heuristic is used.
|
||||
// tokenCounter is retained for legacy compaction/pruning estimates.
|
||||
tokenCounter tokencount.TokenCounter
|
||||
// budgetCounter is the fixed local, model-independent complete-input counter
|
||||
// used by the request-budget invariant.
|
||||
budgetCounter tokencount.BudgetCounter
|
||||
|
||||
// Sandbox info
|
||||
sandboxEnabled bool
|
||||
@@ -480,6 +482,12 @@ func (l *Loop) effectiveMaxTokens() int {
|
||||
return defaultMaxTokens
|
||||
}
|
||||
|
||||
// ContextWindow returns the operator-configured agent context window.
|
||||
func (l *Loop) ContextWindow() int { return l.contextWindow }
|
||||
|
||||
// MaxTokens returns the operator-configured effective agent max_tokens.
|
||||
func (l *Loop) MaxTokens() int { return l.effectiveMaxTokens() }
|
||||
|
||||
// resolveReserveTokens returns the reserve token buffer from compaction config.
|
||||
// Issue 958: Wire ReserveTokensFloor to prevent context overflow before compaction.
|
||||
func (l *Loop) resolveReserveTokens() int {
|
||||
@@ -559,6 +567,7 @@ func NewLoop(cfg LoopConfig) *Loop {
|
||||
compactionCfg: cfg.CompactionCfg,
|
||||
contextPruningCfg: cfg.ContextPruningCfg,
|
||||
tokenCounter: tokencount.NewTiktokenCounter(),
|
||||
budgetCounter: tokencount.NewBudgetCounter(),
|
||||
sandboxEnabled: cfg.SandboxEnabled,
|
||||
sandboxContainerDir: cfg.SandboxContainerDir,
|
||||
sandboxWorkspaceAccess: cfg.SandboxWorkspaceAccess,
|
||||
@@ -678,6 +687,7 @@ type RunResult struct {
|
||||
RunID string `json:"runId"`
|
||||
Iterations int `json:"iterations"`
|
||||
Usage *providers.Usage `json:"usage,omitempty"`
|
||||
LastUsage *providers.Usage `json:"lastUsage,omitempty"`
|
||||
Media []MediaResult `json:"media,omitempty"` // media files from tool results (MEDIA: prefix)
|
||||
Deliverables []string `json:"deliverables,omitempty"` // actual content from tool outputs (for team task results)
|
||||
BlockReplies int `json:"blockReplies,omitempty"` // number of block.reply events emitted
|
||||
@@ -703,10 +713,12 @@ type MediaResult struct {
|
||||
// on *runState without passing 20+ individual variables.
|
||||
type runState struct {
|
||||
// Loop control
|
||||
loopDetector toolLoopState
|
||||
totalUsage providers.Usage
|
||||
iteration int
|
||||
totalToolCalls int
|
||||
loopDetector toolLoopState
|
||||
totalUsage providers.Usage
|
||||
lastUsage providers.Usage
|
||||
lastUsageMsgCount int
|
||||
iteration int
|
||||
totalToolCalls int
|
||||
|
||||
// Output accumulators
|
||||
finalContent string
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
)
|
||||
|
||||
// Agent is the core abstraction for an AI agent execution loop.
|
||||
@@ -21,3 +22,20 @@ type Agent interface {
|
||||
ProviderName() string
|
||||
Provider() providers.Provider
|
||||
}
|
||||
|
||||
// BudgetedAgent exposes the configured agent-only request budget without adding
|
||||
// model/provider authority or expanding the broad Agent test interface.
|
||||
type BudgetedAgent interface {
|
||||
ContextWindow() int
|
||||
MaxTokens() int
|
||||
}
|
||||
|
||||
// WithAgentBudget propagates an agent's configured budget to nested model calls.
|
||||
func WithAgentBudget(ctx context.Context, a Agent) context.Context {
|
||||
budgeted, ok := a.(BudgetedAgent)
|
||||
if !ok {
|
||||
return ctx
|
||||
}
|
||||
ctx = store.WithAgentContextWindow(ctx, budgeted.ContextWindow())
|
||||
return store.WithAgentMaxTokens(ctx, budgeted.MaxTokens())
|
||||
}
|
||||
@@ -46,6 +46,12 @@ func (l *Loop) callInternalLLMWithUsage(ctx context.Context, chatReq providers.C
|
||||
if fallbackProvider, ok := l.provider.(*providers.ModelFallbackProvider); ok {
|
||||
before := func(callCtx context.Context, entry providers.FallbackCandidate, actualReq providers.ChatRequest) (providers.FallbackAfterCall, error) {
|
||||
candidatePurpose := fmt.Sprintf("%s:%s:%s", purpose, entry.ProviderName, actualReq.Model)
|
||||
// Mandatory hard-ceiling guard: internal calls (compaction, session
|
||||
// summarization, memory flush) obey the SAME per-agent window as
|
||||
// main-loop calls, per candidate, after all request mutations.
|
||||
if guardErr := l.guardCompleteModelRequest(actualReq, entry.ProviderName, actualReq.Model, candidatePurpose); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
reservation, err := l.reserveInternalLLMUsageFor(callCtx, actualReq, candidatePurpose, entry.ProviderName, actualReq.Model)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -58,6 +64,18 @@ func (l *Loop) callInternalLLMWithUsage(ctx context.Context, chatReq providers.C
|
||||
}
|
||||
return fallbackProvider.ChatWithHook(ctx, chatReq, before)
|
||||
}
|
||||
providerName := ""
|
||||
if l.provider != nil {
|
||||
providerName = l.provider.Name()
|
||||
}
|
||||
callModel := chatReq.Model
|
||||
if callModel == "" {
|
||||
callModel = l.model
|
||||
}
|
||||
// Mandatory hard-ceiling guard for the non-fallback internal call path.
|
||||
if guardErr := l.guardCompleteModelRequest(chatReq, providerName, callModel, purpose); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
reservation, reserveErr := l.reserveInternalLLMUsage(ctx, chatReq, purpose)
|
||||
if reserveErr != nil {
|
||||
return nil, reserveErr
|
||||
|
||||
@@ -330,7 +330,8 @@ type AgentDefaults struct {
|
||||
// Matching TS agents.defaults.compaction.
|
||||
type CompactionConfig struct {
|
||||
ReserveTokensFloor int `json:"reserveTokensFloor,omitempty"` // min reserve tokens (default 20000)
|
||||
MaxHistoryShare float64 `json:"maxHistoryShare,omitempty"` // max share of context for history (default 0.85)
|
||||
MaxHistoryShare float64 `json:"maxHistoryShare,omitempty"` // max share of context for history-only post-turn compaction (default 0.85)
|
||||
MaxRequestShare float64 `json:"maxRequestShare,omitempty"` // max share of context for the final request sent to the model (default 0.85)
|
||||
KeepLastMessages int `json:"keepLastMessages,omitempty"` // messages to keep after compaction (default 4)
|
||||
TimeoutSeconds int `json:"timeoutSeconds,omitempty"` // summarization timeout in seconds (default 120)
|
||||
MemoryFlush *MemoryFlushConfig `json:"memoryFlush,omitempty"` // pre-compaction flush
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
package consolidation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/config"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
)
|
||||
|
||||
// withAgentRequestBudget wires the calling agent's configured request
|
||||
// budget (context_window, max_tokens) into ctx so background LLM calls
|
||||
// (episodic summarization, dreaming synthesis) pass the agent-only preflight
|
||||
// guard instead of failing closed with an AgentBudgetWiringError.
|
||||
//
|
||||
// Background workers run off a session.completed / episodic.created event and
|
||||
// do not carry a live agent Loop, so the budget cannot come from
|
||||
// agent.WithAgentBudget. We load it straight from the agent row. If the store
|
||||
// is unavailable or the row cannot be read, we fall back to the operator
|
||||
// defaults rather than let a transient lookup miss silently kill memory
|
||||
// consolidation — the guard's job is to bound requests to the agent window,
|
||||
// and the defaults are the same window the agent would use unconfigured.
|
||||
func withAgentRequestBudget(ctx context.Context, agents store.AgentCRUDStore, agentID uuid.UUID, purpose string) context.Context {
|
||||
contextWindow := config.DefaultContextWindow
|
||||
maxTokens := config.DefaultMaxTokens
|
||||
|
||||
if agents != nil && agentID != uuid.Nil {
|
||||
if ag, err := agents.GetByIDUnscoped(ctx, agentID); err != nil {
|
||||
slog.Warn("consolidation: agent budget lookup failed, using defaults",
|
||||
"purpose", purpose, "agent", agentID, "err", err,
|
||||
"context_window", contextWindow, "max_tokens", maxTokens)
|
||||
} else {
|
||||
if ag.ContextWindow > 0 {
|
||||
contextWindow = ag.ContextWindow
|
||||
}
|
||||
if ag.MaxTokens > 0 {
|
||||
maxTokens = ag.MaxTokens
|
||||
}
|
||||
}
|
||||
} else {
|
||||
slog.Warn("consolidation: no agent store or agent id, using default budget",
|
||||
"purpose", purpose, "agent", agentID,
|
||||
"context_window", contextWindow, "max_tokens", maxTokens)
|
||||
}
|
||||
|
||||
ctx = store.WithAgentContextWindow(ctx, contextWindow)
|
||||
ctx = store.WithAgentMaxTokens(ctx, maxTokens)
|
||||
return ctx
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package consolidation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/config"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
)
|
||||
|
||||
// budgetAgentStore overrides only GetByIDUnscoped; the embedded nil interface
|
||||
// panics if any other method is called, which keeps the mock honest.
|
||||
type budgetAgentStore struct {
|
||||
store.AgentCRUDStore
|
||||
agent *store.AgentData
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *budgetAgentStore) GetByIDUnscoped(_ context.Context, _ uuid.UUID) (*store.AgentData, error) {
|
||||
if m.err != nil {
|
||||
return nil, m.err
|
||||
}
|
||||
return m.agent, nil
|
||||
}
|
||||
|
||||
// The whole point of the fix: background workers must carry the agent's
|
||||
// configured request budget into ctx, or the agent-only preflight guard fails
|
||||
// closed with AgentBudgetWiringError and silently kills memory consolidation.
|
||||
func TestWithAgentRequestBudget_WiresConfiguredAgentBudget(t *testing.T) {
|
||||
id := uuid.New()
|
||||
agents := &budgetAgentStore{agent: &store.AgentData{ContextWindow: 128000, MaxTokens: 4096}}
|
||||
|
||||
ctx := withAgentRequestBudget(context.Background(), agents, id, "episodic-summary")
|
||||
|
||||
if got := store.AgentContextWindowFromContext(ctx); got != 128000 {
|
||||
t.Fatalf("context window: want 128000, got %d", got)
|
||||
}
|
||||
if got := store.AgentMaxTokensFromContext(ctx); got != 4096 {
|
||||
t.Fatalf("max tokens: want 4096, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithAgentRequestBudget_NilStoreFallsBackToDefaults(t *testing.T) {
|
||||
ctx := withAgentRequestBudget(context.Background(), nil, uuid.New(), "dreaming-synthesis")
|
||||
|
||||
if got := store.AgentContextWindowFromContext(ctx); got != config.DefaultContextWindow {
|
||||
t.Fatalf("context window: want default %d, got %d", config.DefaultContextWindow, got)
|
||||
}
|
||||
if got := store.AgentMaxTokensFromContext(ctx); got != config.DefaultMaxTokens {
|
||||
t.Fatalf("max tokens: want default %d, got %d", config.DefaultMaxTokens, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithAgentRequestBudget_LookupErrorFallsBackToDefaults(t *testing.T) {
|
||||
agents := &budgetAgentStore{err: errors.New("db down")}
|
||||
|
||||
ctx := withAgentRequestBudget(context.Background(), agents, uuid.New(), "episodic-summary")
|
||||
|
||||
if got := store.AgentContextWindowFromContext(ctx); got != config.DefaultContextWindow {
|
||||
t.Fatalf("context window: want default %d, got %d", config.DefaultContextWindow, got)
|
||||
}
|
||||
if got := store.AgentMaxTokensFromContext(ctx); got != config.DefaultMaxTokens {
|
||||
t.Fatalf("max tokens: want default %d, got %d", config.DefaultMaxTokens, got)
|
||||
}
|
||||
}
|
||||
|
||||
// A zero/partial agent row must not zero out the budget — the setters ignore
|
||||
// non-positive values, and the helper must supply defaults for the missing half.
|
||||
func TestWithAgentRequestBudget_ZeroFieldsUseDefaults(t *testing.T) {
|
||||
agents := &budgetAgentStore{agent: &store.AgentData{ContextWindow: 0, MaxTokens: 0}}
|
||||
|
||||
ctx := withAgentRequestBudget(context.Background(), agents, uuid.New(), "episodic-summary")
|
||||
|
||||
if got := store.AgentContextWindowFromContext(ctx); got != config.DefaultContextWindow {
|
||||
t.Fatalf("context window: want default %d, got %d", config.DefaultContextWindow, got)
|
||||
}
|
||||
if got := store.AgentMaxTokensFromContext(ctx); got != config.DefaultMaxTokens {
|
||||
t.Fatalf("max tokens: want default %d, got %d", config.DefaultMaxTokens, got)
|
||||
}
|
||||
}
|
||||
@@ -33,6 +33,7 @@ type dreamingWorker struct {
|
||||
registry *providers.Registry // provider resolution
|
||||
alertDeps bgalert.AlertDeps
|
||||
usageCaps *usagecaps.Service
|
||||
agents store.AgentCRUDStore // for resolving per-agent request budget
|
||||
|
||||
// threshold/debounce are the global defaults. Per-agent overrides come
|
||||
// from resolveConfig which reads the agent's MemoryConfig.Dreaming JSONB.
|
||||
@@ -166,6 +167,14 @@ func (w *dreamingWorker) Handle(ctx context.Context, event eventbus.DomainEvent)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Wire the agent's configured request budget so the synthesis call passes
|
||||
// the agent-only preflight guard instead of failing closed.
|
||||
if agentUUID, perr := uuid.Parse(agentID); perr == nil {
|
||||
ctx = withAgentRequestBudget(ctx, w.agents, agentUUID, "dreaming-synthesis")
|
||||
} else {
|
||||
ctx = withAgentRequestBudget(ctx, w.agents, uuid.Nil, "dreaming-synthesis")
|
||||
}
|
||||
|
||||
// Build LLM prompt and call provider.
|
||||
synthesis, err := w.synthesize(ctx, provider, model, entries)
|
||||
if err != nil {
|
||||
|
||||
@@ -25,6 +25,7 @@ type episodicWorker struct {
|
||||
eventBus eventbus.DomainEventBus
|
||||
alertDeps bgalert.AlertDeps
|
||||
usageCaps *usagecaps.Service
|
||||
agents store.AgentCRUDStore // resolves per-agent request budget for the preflight guard
|
||||
}
|
||||
|
||||
// resolveProvider delegates to shared background provider resolution.
|
||||
@@ -56,6 +57,10 @@ func (w *episodicWorker) Handle(ctx context.Context, event eventbus.DomainEvent)
|
||||
return fmt.Errorf("episodic: invalid agent_id %q: %w", event.AgentID, err)
|
||||
}
|
||||
ctx = store.WithAgentID(ctx, agentUUID)
|
||||
// Wire the calling agent's request budget so nested LLM summarization calls
|
||||
// pass the agent-only preflight guard instead of failing closed with an
|
||||
// AgentBudgetWiringError (purpose=episodic-summary).
|
||||
ctx = withAgentRequestBudget(ctx, w.agents, agentUUID, "episodic-summary")
|
||||
|
||||
// Build source_id for idempotency
|
||||
sourceID := fmt.Sprintf("%s:%d", payload.SessionKey, payload.CompactionCount)
|
||||
|
||||
@@ -45,6 +45,7 @@ func Register(deps ConsolidationDeps) func() {
|
||||
eventBus: deps.EventBus,
|
||||
alertDeps: deps.AlertDeps,
|
||||
usageCaps: deps.UsageCaps,
|
||||
agents: deps.AgentStore,
|
||||
}
|
||||
semantic := &semanticWorker{
|
||||
kgStore: deps.KGStore,
|
||||
@@ -66,6 +67,7 @@ func Register(deps ConsolidationDeps) func() {
|
||||
threshold: dreamingDefaultThreshold,
|
||||
debounce: dreamingDefaultDebounce,
|
||||
resolveConfig: newAgentStoreResolver(deps.AgentStore),
|
||||
agents: deps.AgentStore,
|
||||
}
|
||||
|
||||
unsub1 := deps.EventBus.Subscribe(eventbus.EventSessionCompleted, episodic.Handle)
|
||||
|
||||
@@ -23,6 +23,18 @@ var sensitiveKeys = []string{
|
||||
"credential", "authorization", "cookie",
|
||||
}
|
||||
|
||||
var safeTokenTelemetryKeys = map[string]bool{
|
||||
"total_tokens": true,
|
||||
"prompt_tokens": true,
|
||||
"completion_tokens": true,
|
||||
"input_tokens": true,
|
||||
"output_tokens": true,
|
||||
"thinking_tokens": true,
|
||||
"last_prompt_tokens": true,
|
||||
"last_usage_prompt_tokens": true,
|
||||
"total_prompt_tokens": true,
|
||||
}
|
||||
|
||||
// LogTee is a slog.Handler that forwards log records to subscribed WS clients
|
||||
// while delegating to an underlying handler for normal output.
|
||||
type LogTee struct {
|
||||
@@ -347,6 +359,9 @@ func logLevelValue(v any) slog.Level {
|
||||
|
||||
func isSensitiveKey(key string) bool {
|
||||
lower := strings.ToLower(key)
|
||||
if safeTokenTelemetryKeys[lower] {
|
||||
return false
|
||||
}
|
||||
for _, s := range sensitiveKeys {
|
||||
if strings.Contains(lower, s) {
|
||||
return true
|
||||
|
||||
@@ -33,7 +33,12 @@ func TestLogTeeAggregateIncludesWithAttrsAndGroupEntries(t *testing.T) {
|
||||
func TestLogTeeAggregateRedactsSensitiveAttrs(t *testing.T) {
|
||||
tee := NewLogTee(slog.NewTextHandler(io.Discard, nil))
|
||||
rec := slog.NewRecord(time.Now(), slog.LevelInfo, "secret", 0)
|
||||
rec.AddAttrs(slog.String("api_token", "leak"), slog.String("source", "test"))
|
||||
rec.AddAttrs(
|
||||
slog.String("api_token", "leak"),
|
||||
slog.Int("total_tokens", 123),
|
||||
slog.Int("last_prompt_tokens", 100),
|
||||
slog.String("source", "test"),
|
||||
)
|
||||
if err := tee.Handle(context.Background(), rec); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -42,4 +47,10 @@ func TestLogTeeAggregateRedactsSensitiveAttrs(t *testing.T) {
|
||||
if attrs["api_token"] != redactedValue {
|
||||
t.Fatalf("attrs = %+v", attrs)
|
||||
}
|
||||
if attrs["total_tokens"] != "123" {
|
||||
t.Fatalf("total_tokens = %#v, want 123", attrs["total_tokens"])
|
||||
}
|
||||
if attrs["last_prompt_tokens"] != "100" {
|
||||
t.Fatalf("last_prompt_tokens = %#v, want 100", attrs["last_prompt_tokens"])
|
||||
}
|
||||
}
|
||||
@@ -340,6 +340,12 @@ func (m *ChatMethods) dispatchChatSends(requests []chatSendRequest) {
|
||||
if userID != "" {
|
||||
runCtxBase = store.WithUserID(runCtxBase, userID)
|
||||
}
|
||||
// Team Work gate runs before Loop.injectContext; inject the resolved agent
|
||||
// budget now so the classifier cannot fall back to model/provider guesses.
|
||||
runCtxBase = agent.WithAgentBudget(runCtxBase, loop)
|
||||
if uid := loop.UUID(); uid != uuid.Nil {
|
||||
runCtxBase = store.WithAgentID(runCtxBase, uid)
|
||||
}
|
||||
gate := m.applyTeamWorkGate(runCtxBase, params, loop, sessionKey)
|
||||
params.Message = gate.message
|
||||
// Inject team dispatch tracker: gates team_tasks create (must search/list first)
|
||||
|
||||
@@ -83,17 +83,10 @@ func (s *ContextStage) Execute(ctx context.Context, state *RunState) error {
|
||||
state.Ctx = ctx
|
||||
}
|
||||
|
||||
// 0.5. Resolve the effective context window for this run's provider/model.
|
||||
// Done once here so PruneStage reads a stable value on every iteration and
|
||||
// the budget can't drift if the model somehow changes mid-run. A zero
|
||||
// result from the resolver (unknown model, no registry) leaves the field
|
||||
// zero — PruneStage then falls back to Config.ContextWindow.
|
||||
if s.deps.ResolveContextWindow != nil && state.Model != "" {
|
||||
providerID := ""
|
||||
if state.Provider != nil {
|
||||
providerID = state.Provider.Name()
|
||||
}
|
||||
if cw := s.deps.ResolveContextWindow(providerID, state.Model); cw > 0 {
|
||||
// 0.5. Snapshot the configured agent window once for this run. Model and
|
||||
// provider do not participate in request-budget authority.
|
||||
if s.deps.ResolveContextWindow != nil {
|
||||
if cw := s.deps.ResolveContextWindow(); cw > 0 {
|
||||
state.Context.EffectiveContextWindow = cw
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,16 +24,16 @@ type PruneStats struct {
|
||||
// PipelineDeps bundles all external dependencies stages need.
|
||||
// Passed to NewDefaultPipeline; individual stages receive what they need via closure or direct field access.
|
||||
type PipelineDeps struct {
|
||||
TokenCounter tokencount.TokenCounter
|
||||
EventBus eventbus.DomainEventBus
|
||||
Config PipelineConfig
|
||||
TokenCounter tokencount.TokenCounter
|
||||
BudgetCounter tokencount.BudgetCounter
|
||||
EventBus eventbus.DomainEventBus
|
||||
Config PipelineConfig
|
||||
// Hooks is the hook dispatcher. nil = no hooks (zero-overhead fast path).
|
||||
Hooks hooks.Dispatcher
|
||||
|
||||
// ResolveContextWindow returns the effective context window (in tokens) for
|
||||
// a given provider/model pair. Nil = always use Config.ContextWindow.
|
||||
// Invoked ONCE per run by ContextStage and stored in RunState.Context.EffectiveContextWindow.
|
||||
ResolveContextWindow func(provider, model string) int
|
||||
// ResolveContextWindow returns this agent's configured context window.
|
||||
// Model/provider are intentionally absent from the budget authority surface.
|
||||
ResolveContextWindow func() int
|
||||
|
||||
// Callbacks from agent.Loop — Phase 8 adapter wires these.
|
||||
EmitEvent func(event any)
|
||||
@@ -125,9 +125,13 @@ type PipelineDeps struct {
|
||||
// total (billing/AccumulateTokens); lastUsage is the final iteration's own
|
||||
// usage — its ContextTokens() is the session's current context size
|
||||
// (SetLastPromptTokens → sessions context display + compaction calibration).
|
||||
// msgCount is the message count captured alongside lastUsage.
|
||||
UpdateMetadata func(ctx context.Context, sessionKey string, usage, lastUsage providers.Usage, msgCount int) error
|
||||
BootstrapCleanup func(ctx context.Context, state *RunState) error
|
||||
MaybeSummarize func(ctx context.Context, sessionKey string)
|
||||
// MaybeSummarize takes midLoopCompacted: when the final-request guard already
|
||||
// compacted mid-loop this run, post-turn summarization lowers its trigger
|
||||
// threshold so the compaction is PERSISTED (episodic Bug B / anti-loop).
|
||||
MaybeSummarize func(ctx context.Context, sessionKey string, midLoopCompacted bool)
|
||||
}
|
||||
|
||||
// FireHook is nil-safe. Returns FireResult{Decision: DecisionAllow} when no
|
||||
|
||||
@@ -0,0 +1,291 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/config"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// FinalRequestEstimate describes the complete pre-call context budget for the
|
||||
// request that will actually be sent to the model.
|
||||
type FinalRequestEstimate struct {
|
||||
MessageTokens int
|
||||
ToolTokens int
|
||||
InputTokens int
|
||||
OutputReserveTokens int
|
||||
HardInputCapTokens int
|
||||
CompactTargetTokens int
|
||||
ContextWindow int
|
||||
MaxRequestShare float64
|
||||
}
|
||||
|
||||
const defaultMaxRequestShare = 0.85
|
||||
|
||||
func effectiveMaxRequestShare(cfg *config.CompactionConfig) float64 {
|
||||
if cfg != nil && cfg.MaxRequestShare > 0 && cfg.MaxRequestShare <= 1 {
|
||||
return cfg.MaxRequestShare
|
||||
}
|
||||
return defaultMaxRequestShare
|
||||
}
|
||||
|
||||
func (s *ThinkStage) buildChatRequest(state *RunState, toolDefs []providers.ToolDefinition) providers.ChatRequest {
|
||||
options := map[string]any{
|
||||
providers.OptMaxTokens: s.deps.Config.MaxTokens,
|
||||
}
|
||||
return providers.ChatRequest{
|
||||
Messages: state.Messages.All(),
|
||||
Tools: toolDefs,
|
||||
Model: state.Model,
|
||||
Options: options,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ThinkStage) finalRequestEstimate(state *RunState, req providers.ChatRequest) (FinalRequestEstimate, error) {
|
||||
contextWindow := state.Context.EffectiveContextWindow
|
||||
if contextWindow == 0 {
|
||||
contextWindow = s.deps.Config.ContextWindow
|
||||
}
|
||||
if contextWindow <= 0 {
|
||||
return FinalRequestEstimate{}, nil
|
||||
}
|
||||
|
||||
messageTokens, toolTokens, err := countBudgetInput(s.deps, state.Model, req)
|
||||
if err != nil {
|
||||
return FinalRequestEstimate{}, err
|
||||
}
|
||||
outputReserve := s.deps.Config.MaxTokens
|
||||
if outputReserve < 0 {
|
||||
outputReserve = 0
|
||||
}
|
||||
share := effectiveMaxRequestShare(s.deps.Config.Compaction)
|
||||
inputTokens := messageTokens + toolTokens
|
||||
hardInputCap := contextWindow - outputReserve
|
||||
shareTarget := int(float64(contextWindow)*share) - outputReserve
|
||||
compactTarget := min(hardInputCap, shareTarget)
|
||||
return FinalRequestEstimate{
|
||||
MessageTokens: messageTokens,
|
||||
ToolTokens: toolTokens,
|
||||
InputTokens: inputTokens,
|
||||
OutputReserveTokens: outputReserve,
|
||||
HardInputCapTokens: hardInputCap,
|
||||
CompactTargetTokens: compactTarget,
|
||||
ContextWindow: contextWindow,
|
||||
MaxRequestShare: share,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func countBudgetInput(deps *PipelineDeps, model string, req providers.ChatRequest) (int, int, error) {
|
||||
if deps != nil && deps.BudgetCounter != nil {
|
||||
messages, err := deps.BudgetCounter.CountMessages(req.Messages)
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("count request messages: %w", err)
|
||||
}
|
||||
tools, err := deps.BudgetCounter.CountToolSchemas(req.Tools)
|
||||
if err != nil {
|
||||
return 0, 0, fmt.Errorf("count request tools: %w", err)
|
||||
}
|
||||
return messages, tools, nil
|
||||
}
|
||||
// Isolated pipeline tests may still wire only the legacy counter. Runtime
|
||||
// always provides BudgetCounter.
|
||||
return countRequestMessages(deps, model, req.Messages), countRequestTools(deps, model, req.Tools), nil
|
||||
}
|
||||
|
||||
func countRequestMessages(deps *PipelineDeps, model string, messages []providers.Message) int {
|
||||
if deps != nil && deps.TokenCounter != nil {
|
||||
return deps.TokenCounter.CountMessages(model, messages)
|
||||
}
|
||||
total := 0
|
||||
for _, msg := range messages {
|
||||
total += utf8.RuneCountInString(msg.Content)/3 + 4
|
||||
for _, tc := range msg.ToolCalls {
|
||||
total += utf8.RuneCountInString(tc.ID)/3 + utf8.RuneCountInString(tc.Name)/3
|
||||
}
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
func countRequestTools(deps *PipelineDeps, model string, tools []providers.ToolDefinition) int {
|
||||
if len(tools) == 0 {
|
||||
return 0
|
||||
}
|
||||
if deps != nil && deps.TokenCounter != nil {
|
||||
return deps.TokenCounter.CountToolSchemas(model, tools)
|
||||
}
|
||||
blob, err := json.Marshal(tools)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return utf8.RuneCountInString(string(blob)) / 3
|
||||
}
|
||||
|
||||
func (e FinalRequestEstimate) withinLimit() bool {
|
||||
return e.ContextWindow > 0 &&
|
||||
e.HardInputCapTokens > 0 &&
|
||||
e.CompactTargetTokens > 0 &&
|
||||
e.InputTokens <= e.HardInputCapTokens &&
|
||||
e.InputTokens <= e.CompactTargetTokens
|
||||
}
|
||||
|
||||
func (s *ThinkStage) prepareFinalRequest(ctx context.Context, state *RunState, toolDefs []providers.ToolDefinition) (providers.ChatRequest, FinalRequestEstimate, error) {
|
||||
req := s.buildChatRequest(state, toolDefs)
|
||||
estimate, err := s.finalRequestEstimate(state, req)
|
||||
if err != nil {
|
||||
return req, estimate, err
|
||||
}
|
||||
if estimate.ContextWindow <= 0 {
|
||||
// A nil counter means this lightweight pipeline instance did not opt into
|
||||
// request budgeting (primarily isolated stage tests). Runtime wiring always
|
||||
// provides a counter, so an unresolved runtime window still fails closed.
|
||||
if s.deps.TokenCounter == nil {
|
||||
return req, estimate, nil
|
||||
}
|
||||
return req, estimate, fmt.Errorf("context_window_unresolved: no configured agent context window")
|
||||
}
|
||||
if estimate.HardInputCapTokens <= 0 || estimate.CompactTargetTokens <= 0 {
|
||||
s.logFinalRequestGuard(state, estimate, "abort", "invalid_budget")
|
||||
return req, estimate, fmt.Errorf("final request context budget unavailable: context_window=%d output_reserve=%d hard_input_cap=%d compact_target=%d",
|
||||
estimate.ContextWindow, estimate.OutputReserveTokens, estimate.HardInputCapTokens, estimate.CompactTargetTokens)
|
||||
}
|
||||
if estimate.withinLimit() {
|
||||
s.logFinalRequestGuard(state, estimate, "allow", "initial")
|
||||
return req, estimate, nil
|
||||
}
|
||||
|
||||
s.logFinalRequestGuard(state, estimate, "reduce", "initial")
|
||||
steps := []string{"prune_history", "compact_history", "shrink_memory"}
|
||||
for _, step := range steps {
|
||||
changed, err := s.reduceFinalRequestContext(ctx, state, estimate, step)
|
||||
if err != nil {
|
||||
return req, estimate, err
|
||||
}
|
||||
if !changed {
|
||||
continue
|
||||
}
|
||||
req = s.buildChatRequest(state, toolDefs)
|
||||
estimate, err = s.finalRequestEstimate(state, req)
|
||||
if err != nil {
|
||||
return req, estimate, err
|
||||
}
|
||||
if estimate.withinLimit() {
|
||||
s.logFinalRequestGuard(state, estimate, "allow", step)
|
||||
return req, estimate, nil
|
||||
}
|
||||
s.logFinalRequestGuard(state, estimate, "reduce", step)
|
||||
}
|
||||
|
||||
s.logFinalRequestGuard(state, estimate, "abort", "exhausted")
|
||||
return req, estimate, fmt.Errorf("final request context budget exceeded: estimated_input=%d compact_target=%d hard_input_cap=%d context_window=%d max_request_share=%.2f",
|
||||
estimate.InputTokens, estimate.CompactTargetTokens, estimate.HardInputCapTokens, estimate.ContextWindow, estimate.MaxRequestShare)
|
||||
}
|
||||
|
||||
func (s *ThinkStage) reduceFinalRequestContext(ctx context.Context, state *RunState, estimate FinalRequestEstimate, step string) (bool, error) {
|
||||
switch step {
|
||||
case "prune_history":
|
||||
return s.pruneForFinalRequestBudget(state, estimate), nil
|
||||
case "compact_history":
|
||||
return s.compactForFinalRequestBudget(ctx, state)
|
||||
case "shrink_memory":
|
||||
return s.shrinkMemoryForFinalRequestBudget(state), nil
|
||||
default:
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ThinkStage) pruneForFinalRequestBudget(state *RunState, estimate FinalRequestEstimate) bool {
|
||||
if s.deps.PruneMessages == nil {
|
||||
return false
|
||||
}
|
||||
history := state.Messages.History()
|
||||
if len(history) == 0 {
|
||||
return false
|
||||
}
|
||||
fixedMessages := []providers.Message{state.Messages.System()}
|
||||
fixedMessages = append(fixedMessages, state.Messages.Pending()...)
|
||||
fixedMessageTokens := countRequestMessages(s.deps, state.Model, fixedMessages)
|
||||
budget := estimate.CompactTargetTokens - estimate.ToolTokens - fixedMessageTokens
|
||||
if budget <= 0 {
|
||||
budget = 1
|
||||
}
|
||||
pruned, stats := s.deps.PruneMessages(history, budget)
|
||||
changed := stats.ResultsTrimmed > 0 || stats.ResultsCleared > 0 || stats.Compacted || len(pruned) != len(history)
|
||||
if !changed {
|
||||
return false
|
||||
}
|
||||
if s.deps.SanitizeHistory != nil {
|
||||
pruned, _ = s.deps.SanitizeHistory(pruned)
|
||||
}
|
||||
state.Messages.SetHistory(pruned)
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *ThinkStage) compactForFinalRequestBudget(ctx context.Context, state *RunState) (bool, error) {
|
||||
if s.deps.CompactMessages == nil {
|
||||
return false, nil
|
||||
}
|
||||
history := state.Messages.History()
|
||||
if len(history) == 0 {
|
||||
return false, nil
|
||||
}
|
||||
savedPending := state.Messages.Pending()
|
||||
compacted, err := s.deps.CompactMessages(ctx, history, state.Model)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("compact final request context: %w", err)
|
||||
}
|
||||
state.Messages.ReplaceHistory(compacted)
|
||||
for _, msg := range savedPending {
|
||||
state.Messages.AppendPending(msg)
|
||||
}
|
||||
state.Prune.MidLoopCompacted = true
|
||||
state.Compact.CompactionCount++
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *ThinkStage) shrinkMemoryForFinalRequestBudget(state *RunState) bool {
|
||||
section := strings.TrimSpace(state.Context.MemorySection)
|
||||
if section == "" {
|
||||
return false
|
||||
}
|
||||
sys := state.Messages.System()
|
||||
content := sys.Content
|
||||
candidates := []string{"\n\n" + state.Context.MemorySection, state.Context.MemorySection, section}
|
||||
for _, candidate := range candidates {
|
||||
if candidate == "" {
|
||||
continue
|
||||
}
|
||||
if strings.Contains(content, candidate) {
|
||||
sys.Content = strings.Replace(content, candidate, "", 1)
|
||||
state.Messages.SetSystem(sys)
|
||||
state.Context.MemorySection = ""
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *ThinkStage) logFinalRequestGuard(state *RunState, estimate FinalRequestEstimate, action, step string) {
|
||||
if estimate.ContextWindow <= 0 {
|
||||
return
|
||||
}
|
||||
slog.Info("final_context.guard",
|
||||
"session_key", state.Input.SessionKey,
|
||||
"run_id", state.RunID,
|
||||
"model", state.Model,
|
||||
"context_window", estimate.ContextWindow,
|
||||
"max_request_share", estimate.MaxRequestShare,
|
||||
"hard_input_cap_tokens", estimate.HardInputCapTokens,
|
||||
"compact_target_input_tokens", estimate.CompactTargetTokens,
|
||||
"message_tokens", estimate.MessageTokens,
|
||||
"tool_tokens", estimate.ToolTokens,
|
||||
"input_tokens", estimate.InputTokens,
|
||||
"output_reserve_tokens", estimate.OutputReserveTokens,
|
||||
"action", action,
|
||||
"reduction_step", step,
|
||||
)
|
||||
}
|
||||
@@ -120,7 +120,9 @@ func (s *FinalizeStage) Execute(ctx context.Context, state *RunState) error {
|
||||
state.Tool.MediaResults = append(state.Tool.MediaResults, mr)
|
||||
}
|
||||
|
||||
// 4. Flush remaining pending messages to session store
|
||||
// 4. Flush remaining pending messages to session store.
|
||||
// Capture the pre-flush history length so metadata msgCount reflects
|
||||
// history + newly-persisted pending (matches upstream calibration).
|
||||
historyCountBeforeFlush := len(state.Messages.History())
|
||||
pending := state.Messages.FlushPending()
|
||||
persistablePending := persistableMessages(pending)
|
||||
@@ -145,9 +147,15 @@ func (s *FinalizeStage) Execute(ctx context.Context, state *RunState) error {
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Post-run summarization (async background)
|
||||
// 7. Post-run summarization (async background).
|
||||
// Pass the mid-loop pressure flag: when the guard had to compact mid-loop this
|
||||
// run, maybeSummarize uses a lower, unit-aligned threshold so the compaction is
|
||||
// PERSISTED to the session (TruncateHistory + IncrementCompaction) instead of
|
||||
// being thrown away — which both breaks the per-turn re-compaction loop (Việc 2)
|
||||
// and advances the cumulative compaction count so episodic can progress (Việc 1-B).
|
||||
// Both mid-loop paths (prune_stage + compactForFinalRequestBudget) set this flag.
|
||||
if s.deps.MaybeSummarize != nil {
|
||||
s.deps.MaybeSummarize(ctx, state.Input.SessionKey)
|
||||
s.deps.MaybeSummarize(ctx, state.Input.SessionKey, state.Prune.MidLoopCompacted)
|
||||
}
|
||||
|
||||
// 8. Emit session.completed for consolidation pipeline (episodic → semantic → dreaming).
|
||||
|
||||
@@ -400,6 +400,7 @@ func TestPipeline_BuildResultPopulatesRunID(t *testing.T) {
|
||||
state := buildMinimalRunState()
|
||||
state.Observe.FinalContent = "hello"
|
||||
state.Think.TotalUsage = providers.Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}
|
||||
state.Think.LastUsage = providers.Usage{PromptTokens: 7, CompletionTokens: 2, TotalTokens: 9}
|
||||
|
||||
result, err := p.Run(context.Background(), state)
|
||||
if err != nil {
|
||||
@@ -414,6 +415,9 @@ func TestPipeline_BuildResultPopulatesRunID(t *testing.T) {
|
||||
if result.TotalUsage.TotalTokens != 15 {
|
||||
t.Errorf("result.TotalUsage.TotalTokens = %d, want 15", result.TotalUsage.TotalTokens)
|
||||
}
|
||||
if result.LastUsage.TotalTokens != 9 {
|
||||
t.Errorf("result.LastUsage.TotalTokens = %d, want 9", result.LastUsage.TotalTokens)
|
||||
}
|
||||
if result.Duration <= 0 {
|
||||
t.Errorf("result.Duration = %v, want > 0", result.Duration)
|
||||
}
|
||||
@@ -506,6 +510,7 @@ func TestRunState_BuildResult_AllFields(t *testing.T) {
|
||||
state.Observe.FinalContent = "final"
|
||||
state.Observe.FinalThinking = "thinking"
|
||||
state.Think.TotalUsage = providers.Usage{PromptTokens: 100, CompletionTokens: 50, TotalTokens: 150}
|
||||
state.Think.LastUsage = providers.Usage{PromptTokens: 80, CompletionTokens: 20, TotalTokens: 100}
|
||||
state.Iteration = 7
|
||||
state.Tool.TotalToolCalls = 3
|
||||
state.Tool.LoopKilled = true
|
||||
@@ -528,6 +533,9 @@ func TestRunState_BuildResult_AllFields(t *testing.T) {
|
||||
if r.TotalUsage.TotalTokens != 150 {
|
||||
t.Errorf("TotalUsage.TotalTokens = %d", r.TotalUsage.TotalTokens)
|
||||
}
|
||||
if r.LastUsage.TotalTokens != 100 {
|
||||
t.Errorf("LastUsage.TotalTokens = %d", r.LastUsage.TotalTokens)
|
||||
}
|
||||
if r.Iterations != 7 {
|
||||
t.Errorf("Iterations = %d", r.Iterations)
|
||||
}
|
||||
|
||||
@@ -26,7 +26,7 @@ func NewPruneStage(deps *PipelineDeps, memFlush *MemoryFlushStage) *PruneStage {
|
||||
return &PruneStage{deps: deps, memoryFlush: memFlush, result: Continue}
|
||||
}
|
||||
|
||||
func (s *PruneStage) Name() string { return "prune" }
|
||||
func (s *PruneStage) Name() string { return "prune" }
|
||||
func (s *PruneStage) Result() StageResult { return s.result }
|
||||
|
||||
// defaultCachePruneTTL is used when cfg.TTL is empty or invalid.
|
||||
@@ -74,6 +74,17 @@ func (s *PruneStage) Execute(ctx context.Context, state *RunState) error {
|
||||
tokensBefore := historyTokens
|
||||
|
||||
softThreshold := budget * 70 / 100
|
||||
slog.Info("context.preflight_budget",
|
||||
"session_key", state.Input.SessionKey,
|
||||
"context_window", contextWindow,
|
||||
"effective_context_window", state.Context.EffectiveContextWindow,
|
||||
"history_tokens", historyTokens,
|
||||
"overhead_tokens", state.Context.OverheadTokens,
|
||||
"max_tokens", s.deps.Config.MaxTokens,
|
||||
"reserve_tokens", s.deps.Config.ReserveTokens,
|
||||
"budget", budget,
|
||||
"soft_threshold", softThreshold,
|
||||
)
|
||||
if historyTokens <= softThreshold {
|
||||
return nil // under budget, no action needed
|
||||
}
|
||||
|
||||
@@ -78,6 +78,7 @@ func (rs *RunState) BuildResult() *RunResult {
|
||||
Content: rs.Observe.FinalContent,
|
||||
Thinking: rs.Observe.FinalThinking,
|
||||
TotalUsage: rs.Think.TotalUsage,
|
||||
LastUsage: rs.Think.LastUsage,
|
||||
Iterations: rs.Iteration,
|
||||
ToolCalls: rs.Tool.TotalToolCalls,
|
||||
LoopKilled: rs.Tool.LoopKilled,
|
||||
|
||||
@@ -3,6 +3,7 @@ package pipeline
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
@@ -353,7 +354,18 @@ func TestThinkStage_UsageAccumulation(t *testing.T) {
|
||||
return &providers.ChatResponse{
|
||||
Content: "hello",
|
||||
FinishReason: "stop",
|
||||
Usage: &providers.Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15},
|
||||
Usage: &providers.Usage{
|
||||
PromptTokens: 10 * callCount,
|
||||
CompletionTokens: 5,
|
||||
TotalTokens: 10*callCount + 5,
|
||||
CacheCreationTokens: 2,
|
||||
CacheReadTokens: 3,
|
||||
PromptTokensIncludeCachedSegments: callCount == 2,
|
||||
ThinkingTokens: 4,
|
||||
RequestCount: 1,
|
||||
ImageCount: 1,
|
||||
WebSearchCount: 1,
|
||||
},
|
||||
}, nil
|
||||
},
|
||||
}
|
||||
@@ -364,14 +376,27 @@ func TestThinkStage_UsageAccumulation(t *testing.T) {
|
||||
_ = stage.Execute(context.Background(), state)
|
||||
_ = stage.Execute(context.Background(), state)
|
||||
|
||||
if state.Think.TotalUsage.PromptTokens != 20 {
|
||||
t.Errorf("PromptTokens = %d, want 20", state.Think.TotalUsage.PromptTokens)
|
||||
if state.Think.TotalUsage.PromptTokens != 30 {
|
||||
t.Errorf("PromptTokens = %d, want 30", state.Think.TotalUsage.PromptTokens)
|
||||
}
|
||||
if state.Think.TotalUsage.CompletionTokens != 10 {
|
||||
t.Errorf("CompletionTokens = %d, want 10", state.Think.TotalUsage.CompletionTokens)
|
||||
}
|
||||
if state.Think.TotalUsage.TotalTokens != 30 {
|
||||
t.Errorf("TotalTokens = %d, want 30", state.Think.TotalUsage.TotalTokens)
|
||||
if state.Think.TotalUsage.TotalTokens != 40 {
|
||||
t.Errorf("TotalTokens = %d, want 40", state.Think.TotalUsage.TotalTokens)
|
||||
}
|
||||
if state.Think.TotalUsage.CacheCreationTokens != 4 || state.Think.TotalUsage.CacheReadTokens != 6 {
|
||||
t.Errorf("cache tokens = %d/%d, want 4/6", state.Think.TotalUsage.CacheCreationTokens, state.Think.TotalUsage.CacheReadTokens)
|
||||
}
|
||||
if !state.Think.TotalUsage.PromptTokensIncludeCachedSegments {
|
||||
t.Error("PromptTokensIncludeCachedSegments = false, want true")
|
||||
}
|
||||
if state.Think.TotalUsage.ThinkingTokens != 8 || state.Think.TotalUsage.RequestCount != 2 ||
|
||||
state.Think.TotalUsage.ImageCount != 2 || state.Think.TotalUsage.WebSearchCount != 2 {
|
||||
t.Errorf("extended usage = %+v, want thinking/request/image/web = 8/2/2/2", state.Think.TotalUsage)
|
||||
}
|
||||
if state.Think.LastUsage.PromptTokens != 20 || state.Think.LastUsage.TotalTokens != 25 {
|
||||
t.Errorf("LastUsage = %+v, want last prompt=20 total=25", state.Think.LastUsage)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3045,7 +3070,7 @@ func TestFinalizeStage_MaybeSummarize_Called(t *testing.T) {
|
||||
t.Parallel()
|
||||
summarizeCalled := false
|
||||
deps := &PipelineDeps{
|
||||
MaybeSummarize: func(_ context.Context, sessionKey string) {
|
||||
MaybeSummarize: func(_ context.Context, sessionKey string, _ bool) {
|
||||
if sessionKey == "sess-1" {
|
||||
summarizeCalled = true
|
||||
}
|
||||
@@ -3408,3 +3433,251 @@ func TestToolStage_Parallel_DefersNonToolMessages(t *testing.T) {
|
||||
t.Errorf("pending[3].Content = %q, want nudge", pending[3].Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestThinkStage_FinalRequestGuard_CompactsBeforeCallLLM(t *testing.T) {
|
||||
t.Parallel()
|
||||
compacted := false
|
||||
called := false
|
||||
|
||||
deps := &PipelineDeps{
|
||||
TokenCounter: finalRequestBudgetCounter{},
|
||||
Config: PipelineConfig{ContextWindow: 100, MaxTokens: 10, Compaction: &config.CompactionConfig{MaxRequestShare: 0.85}},
|
||||
BuildFilteredTools: func(_ *RunState) ([]providers.ToolDefinition, error) {
|
||||
return []providers.ToolDefinition{{
|
||||
Type: "function",
|
||||
Function: &providers.ToolFunctionSchema{
|
||||
Name: "large_tool",
|
||||
Description: strings.Repeat("t", 50),
|
||||
Parameters: map[string]any{"type": "object"},
|
||||
},
|
||||
}}, nil
|
||||
},
|
||||
CompactMessages: func(_ context.Context, msgs []providers.Message, _ string) ([]providers.Message, error) {
|
||||
compacted = true
|
||||
if len(msgs) != 1 || !strings.Contains(msgs[0].Content, "long-history") {
|
||||
t.Fatalf("CompactMessages got history = %#v", msgs)
|
||||
}
|
||||
return []providers.Message{{Role: "user", Content: "short"}}, nil
|
||||
},
|
||||
CallLLM: func(_ context.Context, _ *RunState, req providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
called = true
|
||||
if len(req.Messages) == 0 || req.Messages[len(req.Messages)-1].Content != "short" {
|
||||
t.Fatalf("CallLLM received un-compacted request messages = %#v", req.Messages)
|
||||
}
|
||||
return &providers.ChatResponse{Content: "ok", FinishReason: "stop"}, nil
|
||||
},
|
||||
}
|
||||
stage := NewThinkStage(deps)
|
||||
state := defaultState()
|
||||
state.Messages.SetHistory([]providers.Message{{Role: "user", Content: strings.Repeat("long-history", 10)}})
|
||||
|
||||
if err := stage.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("Execute() error: %v", err)
|
||||
}
|
||||
if !compacted {
|
||||
t.Fatal("expected final request guard to compact before CallLLM")
|
||||
}
|
||||
if !called {
|
||||
t.Fatal("expected CallLLM after compaction")
|
||||
}
|
||||
}
|
||||
|
||||
type finalRequestBudgetCounter struct{}
|
||||
|
||||
func (finalRequestBudgetCounter) Count(_ string, text string) int { return len(text) }
|
||||
func (finalRequestBudgetCounter) CountMessages(_ string, msgs []providers.Message) int {
|
||||
total := 0
|
||||
for _, msg := range msgs {
|
||||
total += len(msg.Content)
|
||||
}
|
||||
return total
|
||||
}
|
||||
func (finalRequestBudgetCounter) CountToolSchemas(_ string, tools []providers.ToolDefinition) int {
|
||||
total := 0
|
||||
for _, tool := range tools {
|
||||
if tool.Function != nil {
|
||||
total += len(tool.Function.Name) + len(tool.Function.Description)
|
||||
}
|
||||
}
|
||||
return total
|
||||
}
|
||||
func (finalRequestBudgetCounter) ModelContextWindow(_ string) int { return 100 }
|
||||
|
||||
func TestThinkStage_FinalRequestGuard_AllowsRequestUnderLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
compacted := false
|
||||
called := false
|
||||
|
||||
deps := &PipelineDeps{
|
||||
TokenCounter: finalRequestBudgetCounter{},
|
||||
Config: PipelineConfig{ContextWindow: 100, MaxTokens: 10, Compaction: &config.CompactionConfig{MaxRequestShare: 0.85}},
|
||||
CompactMessages: func(_ context.Context, _ []providers.Message, _ string) ([]providers.Message, error) {
|
||||
compacted = true
|
||||
return nil, nil
|
||||
},
|
||||
CallLLM: func(_ context.Context, _ *RunState, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
called = true
|
||||
return &providers.ChatResponse{Content: "ok", FinishReason: "stop"}, nil
|
||||
},
|
||||
}
|
||||
stage := NewThinkStage(deps)
|
||||
state := defaultState()
|
||||
state.Messages.SetHistory([]providers.Message{{Role: "user", Content: "short"}})
|
||||
|
||||
if err := stage.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("Execute() error: %v", err)
|
||||
}
|
||||
if compacted {
|
||||
t.Fatal("did not expect compaction for request under final context limit")
|
||||
}
|
||||
if !called {
|
||||
t.Fatal("expected CallLLM for request under final context limit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestThinkStage_FinalRequestGuard_AbortsWhenStillOverLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
called := false
|
||||
compactCalls := 0
|
||||
|
||||
deps := &PipelineDeps{
|
||||
TokenCounter: finalRequestBudgetCounter{},
|
||||
Config: PipelineConfig{ContextWindow: 100, MaxTokens: 10, Compaction: &config.CompactionConfig{MaxRequestShare: 0.85}},
|
||||
CompactMessages: func(_ context.Context, _ []providers.Message, _ string) ([]providers.Message, error) {
|
||||
compactCalls++
|
||||
return []providers.Message{{Role: "user", Content: strings.Repeat("still-long", 20)}}, nil
|
||||
},
|
||||
CallLLM: func(_ context.Context, _ *RunState, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
called = true
|
||||
return &providers.ChatResponse{Content: "should not call", FinishReason: "stop"}, nil
|
||||
},
|
||||
}
|
||||
stage := NewThinkStage(deps)
|
||||
state := defaultState()
|
||||
state.Messages.SetHistory([]providers.Message{{Role: "user", Content: strings.Repeat("long-history", 10)}})
|
||||
|
||||
err := stage.Execute(context.Background(), state)
|
||||
if err == nil {
|
||||
t.Fatal("expected final request context budget error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "final request context budget exceeded") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if compactCalls != 1 {
|
||||
t.Fatalf("CompactMessages calls = %d, want 1", compactCalls)
|
||||
}
|
||||
if called {
|
||||
t.Fatal("CallLLM must not run while final request remains over limit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalRequestEstimate_UsesAgentWindowBudgetMath(t *testing.T) {
|
||||
t.Parallel()
|
||||
tests := []struct {
|
||||
name string
|
||||
window int
|
||||
hard int
|
||||
target int
|
||||
}{
|
||||
{name: "200k", window: 200_000, hard: 191_808, target: 161_808},
|
||||
{name: "128k", window: 128_000, hard: 119_808, target: 100_608},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
deps := &PipelineDeps{
|
||||
TokenCounter: finalRequestBudgetCounter{},
|
||||
Config: PipelineConfig{
|
||||
ContextWindow: tt.window,
|
||||
MaxTokens: 8192,
|
||||
Compaction: &config.CompactionConfig{MaxRequestShare: 0.85},
|
||||
},
|
||||
}
|
||||
stage := NewThinkStage(deps)
|
||||
estimate, err := stage.finalRequestEstimate(defaultState(), providers.ChatRequest{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if estimate.HardInputCapTokens != tt.hard || estimate.CompactTargetTokens != tt.target {
|
||||
t.Fatalf("budget = hard %d target %d, want %d/%d", estimate.HardInputCapTokens, estimate.CompactTargetTokens, tt.hard, tt.target)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestThinkStage_FinalRequestGuard_RejectsUnresolvedWindow(t *testing.T) {
|
||||
t.Parallel()
|
||||
called := false
|
||||
stage := NewThinkStage(&PipelineDeps{
|
||||
TokenCounter: finalRequestBudgetCounter{},
|
||||
CallLLM: func(context.Context, *RunState, providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
called = true
|
||||
return &providers.ChatResponse{}, nil
|
||||
},
|
||||
})
|
||||
err := stage.Execute(context.Background(), defaultState())
|
||||
if err == nil || !strings.Contains(err.Error(), "context_window_unresolved") {
|
||||
t.Fatalf("Execute() error = %v, want context_window_unresolved", err)
|
||||
}
|
||||
if called {
|
||||
t.Fatal("CallLLM must not run without a resolved context window")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFinalRequestEstimate_OutputReserveIsAgentMaxTokens locks the agent-only
|
||||
// output reserve: the reserve equals the configured agent max_tokens exactly and
|
||||
// does NOT change with any provider/model reasoning transform. Model/provider
|
||||
// thinking bumps must never widen or narrow the pre-transport budget.
|
||||
func TestFinalRequestEstimate_OutputReserveIsAgentMaxTokens(t *testing.T) {
|
||||
t.Parallel()
|
||||
stage := NewThinkStage(&PipelineDeps{
|
||||
TokenCounter: finalRequestBudgetCounter{},
|
||||
Config: PipelineConfig{
|
||||
ContextWindow: 200_000,
|
||||
MaxTokens: 8192,
|
||||
Compaction: &config.CompactionConfig{MaxRequestShare: 0.85},
|
||||
},
|
||||
})
|
||||
|
||||
estimate, err := stage.finalRequestEstimate(defaultState(), stage.buildChatRequest(defaultState(), nil))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Reserve == agent max_tokens (8192), independent of any reasoning transform.
|
||||
if estimate.OutputReserveTokens != 8192 {
|
||||
t.Fatalf("output reserve = %d, want 8192 (agent max_tokens)", estimate.OutputReserveTokens)
|
||||
}
|
||||
// Hard cap == window - agent max_tokens, exactly.
|
||||
if estimate.HardInputCapTokens != 200_000-8192 {
|
||||
t.Fatalf("hard cap = %d, want %d", estimate.HardInputCapTokens, 200_000-8192)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInputContextTokens_CacheSemantics(t *testing.T) {
|
||||
t.Parallel()
|
||||
inclusive := providers.Usage{
|
||||
PromptTokens: 100, CacheReadTokens: 80, CacheCreationTokens: 10,
|
||||
PromptTokensIncludeCachedSegments: true,
|
||||
}
|
||||
if got := InputContextTokens(inclusive); got != 100 {
|
||||
t.Fatalf("inclusive input = %d, want 100", got)
|
||||
}
|
||||
exclusive := inclusive
|
||||
exclusive.PromptTokensIncludeCachedSegments = false
|
||||
if got := InputContextTokens(exclusive); got != 190 {
|
||||
t.Fatalf("exclusive input = %d, want 190", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalContextActualLevel(t *testing.T) {
|
||||
t.Parallel()
|
||||
estimate := FinalRequestEstimate{InputTokens: 100, HardInputCapTokens: 120}
|
||||
if level, reason := finalContextActualLevel(estimate, 100); level != slog.LevelDebug || reason != "" {
|
||||
t.Fatalf("exact estimate = %v/%q, want debug/empty", level, reason)
|
||||
}
|
||||
if level, reason := finalContextActualLevel(estimate, 106); level != slog.LevelWarn || reason != "estimate_undercount" {
|
||||
t.Fatalf("undercount = %v/%q, want warn/estimate_undercount", level, reason)
|
||||
}
|
||||
if level, reason := finalContextActualLevel(estimate, 121); level != slog.LevelWarn || reason != "hard_cap_exceeded" {
|
||||
t.Fatalf("hard cap = %v/%q, want warn/hard_cap_exceeded", level, reason)
|
||||
}
|
||||
}
|
||||
@@ -118,6 +118,7 @@ type RunResult struct {
|
||||
Content string
|
||||
Thinking string
|
||||
TotalUsage providers.Usage
|
||||
LastUsage providers.Usage
|
||||
Iterations int
|
||||
ToolCalls int
|
||||
LoopKilled bool
|
||||
|
||||
@@ -2,6 +2,7 @@ package pipeline
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
@@ -53,14 +54,11 @@ func (s *ThinkStage) Execute(ctx context.Context, state *RunState) error {
|
||||
state.Tool.AllowedTools = nil
|
||||
}
|
||||
|
||||
// 3. Construct ChatRequest
|
||||
req := providers.ChatRequest{
|
||||
Messages: state.Messages.All(),
|
||||
Tools: toolDefs,
|
||||
Model: state.Model,
|
||||
Options: map[string]any{
|
||||
providers.OptMaxTokens: s.deps.Config.MaxTokens,
|
||||
},
|
||||
// 3. Construct the final ChatRequest and enforce the request-level
|
||||
// context budget before any provider call is attempted.
|
||||
req, estimate, err := s.prepareFinalRequest(ctx, state, toolDefs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 4. Call LLM (stream or sync — delegated to callback)
|
||||
@@ -69,35 +67,54 @@ func (s *ThinkStage) Execute(ctx context.Context, state *RunState) error {
|
||||
}
|
||||
resp, err := s.deps.CallLLM(ctx, state, req)
|
||||
if err != nil {
|
||||
// Central hard-ceiling guard rejected the concrete request AFTER the
|
||||
// callback appended its final directive/retry/reasoning mutations (which
|
||||
// prepareFinalRequest could not see). Re-enter the full reduction chain
|
||||
// against the current messages, rebuild, and retry this iteration. Only
|
||||
// when reduction cannot bring the request under the ceiling do we abort —
|
||||
// the guard already guaranteed zero transport calls were made.
|
||||
if isRequestBudgetExceededErr(err) {
|
||||
if state.Think.OverflowRetries >= maxBudgetReductionRetries {
|
||||
return fmt.Errorf("request context budget exceeded after reduction: %w", err)
|
||||
}
|
||||
if s.reduceForBudgetExceeded(ctx, state) {
|
||||
// This LLM call produced no response, so LastResponse still holds
|
||||
// the PRIOR iteration's response. Clear it before returning Continue
|
||||
// so ToolStage/ObserveStage in this same iteration do not re-execute
|
||||
// the previous iteration's tool calls. The retry rebuilds and calls
|
||||
// the model again on the next iteration.
|
||||
state.Think.LastResponse = nil
|
||||
return nil // Retry this iteration (Continue result) with reduced context.
|
||||
}
|
||||
return fmt.Errorf("request context budget exceeded, reduction exhausted: %w", err)
|
||||
}
|
||||
// Issue 958: Check for context overflow — attempt emergency compaction + retry
|
||||
if isContextOverflowErr(err) {
|
||||
if state.Think.OverflowRetries > 0 {
|
||||
return fmt.Errorf("context overflow after compaction: %w", err)
|
||||
}
|
||||
if s.tryEmergencyCompaction(ctx, state, "context_overflow_error") {
|
||||
// Same stale-response hazard as the budget path: no response was
|
||||
// produced this iteration, so drop the prior one before retrying.
|
||||
state.Think.LastResponse = nil
|
||||
return nil // Retry this iteration (Continue result)
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("llm call: %w", err)
|
||||
}
|
||||
|
||||
// 5. Accumulate usage across turns: base tokens, ThinkingTokens for reasoning
|
||||
// models, and cache tokens. Cache read/creation must be summed too — otherwise
|
||||
// the aggregated RunResult.Usage (and thus webhook usage + usage_events) reports
|
||||
// zero cache even when each turn hit the prompt cache heavily.
|
||||
// 5. Accumulate usage across turns and retain the usage snapshot for the
|
||||
// request that was just sent. AddUsage preserves provider telemetry fields.
|
||||
if resp.Usage != nil {
|
||||
state.Think.TotalUsage.PromptTokens += resp.Usage.PromptTokens
|
||||
state.Think.TotalUsage.CompletionTokens += resp.Usage.CompletionTokens
|
||||
state.Think.TotalUsage.TotalTokens += resp.Usage.TotalTokens
|
||||
state.Think.TotalUsage.ThinkingTokens += resp.Usage.ThinkingTokens
|
||||
state.Think.TotalUsage.CacheReadTokens += resp.Usage.CacheReadTokens
|
||||
state.Think.TotalUsage.CacheCreationTokens += resp.Usage.CacheCreationTokens
|
||||
if resp.Usage.PromptTokensIncludeCachedSegments {
|
||||
state.Think.TotalUsage.PromptTokensIncludeCachedSegments = true
|
||||
}
|
||||
// Log actual (post-call) usage vs the pre-call budget estimate (my guard
|
||||
// observability), then accumulate the run-cumulative total.
|
||||
s.logFinalRequestActual(ctx, state, estimate, *resp.Usage)
|
||||
AddUsage(&state.Think.TotalUsage, *resp.Usage)
|
||||
// Snapshot per-call usage: the LAST iteration's prompt size IS the
|
||||
// session's current context. Keep the previous snapshot when a response
|
||||
// carries no prompt tokens (e.g. providers that omit usage on some turns).
|
||||
// carries no prompt tokens (e.g. providers that omit usage on some turns)
|
||||
// — matches upstream 503909d3 behaviour, which never wipes a usable
|
||||
// snapshot with an empty final response.
|
||||
if resp.Usage.PromptTokens > 0 {
|
||||
state.Think.LastUsage = *resp.Usage
|
||||
}
|
||||
@@ -108,6 +125,10 @@ func (s *ThinkStage) Execute(ctx context.Context, state *RunState) error {
|
||||
return fmt.Errorf("llm response truncated before content after compaction")
|
||||
}
|
||||
if s.tryEmergencyCompaction(ctx, state, "empty_length_response") {
|
||||
// LastResponse has NOT yet been updated to this empty response, so it
|
||||
// still holds the prior iteration's result. Clear it so downstream
|
||||
// stages this iteration don't act on the stale response during retry.
|
||||
state.Think.LastResponse = nil
|
||||
return nil // Retry next iteration with compacted history.
|
||||
}
|
||||
return fmt.Errorf("llm response truncated before content")
|
||||
@@ -172,6 +193,46 @@ func (s *ThinkStage) Execute(ctx context.Context, state *RunState) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *ThinkStage) logFinalRequestActual(ctx context.Context, state *RunState, estimate FinalRequestEstimate, usage providers.Usage) {
|
||||
actualInput := InputContextTokens(usage)
|
||||
if actualInput <= 0 || estimate.InputTokens <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
providerName := ""
|
||||
if state.Provider != nil {
|
||||
providerName = state.Provider.Name()
|
||||
}
|
||||
ratio := float64(actualInput) / float64(estimate.InputTokens)
|
||||
level, reason := finalContextActualLevel(estimate, actualInput)
|
||||
args := []any{
|
||||
"run_id", state.RunID,
|
||||
"iteration", state.Iteration,
|
||||
"provider", providerName,
|
||||
"model", state.Model,
|
||||
"estimated_input", estimate.InputTokens,
|
||||
"actual_input", actualInput,
|
||||
"estimate_ratio", ratio,
|
||||
"effective_window", estimate.ContextWindow,
|
||||
"hard_input_cap", estimate.HardInputCapTokens,
|
||||
"compact_target", estimate.CompactTargetTokens,
|
||||
}
|
||||
if reason != "" {
|
||||
args = append(args, "reason", reason)
|
||||
}
|
||||
slog.Log(ctx, level, "final_context.actual", args...)
|
||||
}
|
||||
|
||||
func finalContextActualLevel(estimate FinalRequestEstimate, actualInput int) (slog.Level, string) {
|
||||
if actualInput > estimate.HardInputCapTokens {
|
||||
return slog.LevelWarn, "hard_cap_exceeded"
|
||||
}
|
||||
if actualInput*100 > estimate.InputTokens*105 {
|
||||
return slog.LevelWarn, "estimate_undercount"
|
||||
}
|
||||
return slog.LevelDebug, ""
|
||||
}
|
||||
|
||||
func (s *ThinkStage) tryEmergencyCompaction(ctx context.Context, state *RunState, reason string) bool {
|
||||
state.Think.OverflowRetries++
|
||||
if s.deps.CompactMessages == nil {
|
||||
@@ -297,3 +358,64 @@ func isContextOverflowErr(err error) bool {
|
||||
lower := strings.ToLower(err.Error())
|
||||
return providers.IsContextOverflowMessage(lower)
|
||||
}
|
||||
|
||||
// maxBudgetReductionRetries bounds how many times a single iteration re-enters
|
||||
// reduction after the central request-budget guard rejects the built request.
|
||||
// Each retry runs one reduction step; three covers prune -> compact -> shrink.
|
||||
const maxBudgetReductionRetries = 3
|
||||
|
||||
// contextBudgetExceededError is satisfied by the agent loop's
|
||||
// RequestBudgetExceededError. Detected via interface so this package does not
|
||||
// import the agent package (which would be a cycle).
|
||||
type contextBudgetExceededError interface {
|
||||
ContextBudgetExceeded() bool
|
||||
}
|
||||
|
||||
// isRequestBudgetExceededErr reports whether err (or anything it wraps) is a
|
||||
// request-level context-budget overflow raised by the central pre-transport
|
||||
// guard. Such an error guarantees no provider transport call was made.
|
||||
func isRequestBudgetExceededErr(err error) bool {
|
||||
var budgetErr contextBudgetExceededError
|
||||
if errors.As(err, &budgetErr) {
|
||||
return budgetErr.ContextBudgetExceeded()
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// reduceForBudgetExceeded runs one pass of the reduction chain
|
||||
// (prune_history -> compact_history -> shrink_memory) against the current
|
||||
// state, stopping at the first step that changes anything. It increments
|
||||
// OverflowRetries so a stuck request eventually aborts. Returns true when some
|
||||
// reduction was applied and the iteration should retry.
|
||||
func (s *ThinkStage) reduceForBudgetExceeded(ctx context.Context, state *RunState) bool {
|
||||
state.Think.OverflowRetries++
|
||||
|
||||
// Rebuild the estimate against current messages so pruning targets the
|
||||
// right budget. Tools are rebuilt lazily by the retried iteration.
|
||||
req := s.buildChatRequest(state, nil)
|
||||
estimate, err := s.finalRequestEstimate(state, req)
|
||||
if err != nil {
|
||||
slog.Warn("request_budget.count_failed", "run_id", state.RunID, "error", err)
|
||||
return false
|
||||
}
|
||||
|
||||
steps := []string{"prune_history", "compact_history", "shrink_memory"}
|
||||
for _, step := range steps {
|
||||
changed, err := s.reduceFinalRequestContext(ctx, state, estimate, step)
|
||||
if err != nil {
|
||||
slog.Warn("request_budget.reduce_failed",
|
||||
"run_id", state.RunID,
|
||||
"step", step,
|
||||
"error", err)
|
||||
continue
|
||||
}
|
||||
if changed {
|
||||
slog.Info("request_budget.reduced",
|
||||
"run_id", state.RunID,
|
||||
"step", step,
|
||||
"overflow_retries", state.Think.OverflowRetries)
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package pipeline
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// fakeBudgetErr satisfies the contextBudgetExceededError interface that
|
||||
// ThinkStage detects, standing in for agent.RequestBudgetExceededError without
|
||||
// importing the agent package (which would be an import cycle).
|
||||
type fakeBudgetErr struct{}
|
||||
|
||||
func (fakeBudgetErr) Error() string { return "context budget exceeded" }
|
||||
func (fakeBudgetErr) ContextBudgetExceeded() bool { return true }
|
||||
|
||||
// TestThinkStage_BudgetExceeded_ReducesThenRetries verifies that when CallLLM
|
||||
// returns a request-budget-exceeded error (the central pre-transport guard
|
||||
// fired), the stage runs a reduction pass and returns Continue to retry the
|
||||
// iteration rather than aborting.
|
||||
func TestThinkStage_BudgetExceeded_ReducesThenRetries(t *testing.T) {
|
||||
pruned := false
|
||||
deps := &PipelineDeps{
|
||||
Config: PipelineConfig{MaxIterations: 10, MaxTokens: 1000, ContextWindow: 200_000},
|
||||
CallLLM: func(_ context.Context, _ *RunState, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
return nil, fakeBudgetErr{}
|
||||
},
|
||||
PruneMessages: func(msgs []providers.Message, _ int) ([]providers.Message, PruneStats) {
|
||||
pruned = true
|
||||
// Report a trim so reduceFinalRequestContext sees a change.
|
||||
return msgs, PruneStats{ResultsTrimmed: 1}
|
||||
},
|
||||
}
|
||||
stage := NewThinkStage(deps)
|
||||
state := defaultState()
|
||||
state.Messages.SetHistory([]providers.Message{{Role: "user", Content: "big history"}})
|
||||
|
||||
if err := stage.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("Execute() should retry after reduction, got error: %v", err)
|
||||
}
|
||||
if stage.Result() != Continue {
|
||||
t.Errorf("Result() = %v, want Continue (retry after reduction)", stage.Result())
|
||||
}
|
||||
if !pruned {
|
||||
t.Error("expected PruneMessages to run during budget reduction")
|
||||
}
|
||||
if state.Think.OverflowRetries != 1 {
|
||||
t.Errorf("OverflowRetries = %d, want 1", state.Think.OverflowRetries)
|
||||
}
|
||||
}
|
||||
|
||||
// TestThinkStage_BudgetExceeded_AbortsWhenReductionExhausted verifies that when
|
||||
// no reduction step can change the request, the stage surfaces the budget error
|
||||
// instead of looping forever. The guard already guaranteed zero transport calls.
|
||||
func TestThinkStage_BudgetExceeded_AbortsWhenReductionExhausted(t *testing.T) {
|
||||
deps := &PipelineDeps{
|
||||
Config: PipelineConfig{MaxIterations: 10, MaxTokens: 1000, ContextWindow: 200_000},
|
||||
CallLLM: func(_ context.Context, _ *RunState, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
return nil, fakeBudgetErr{}
|
||||
},
|
||||
// No PruneMessages / CompactMessages wired and no memory section, so
|
||||
// every reduction step reports "no change".
|
||||
}
|
||||
stage := NewThinkStage(deps)
|
||||
state := defaultState()
|
||||
|
||||
err := stage.Execute(context.Background(), state)
|
||||
if err == nil {
|
||||
t.Fatal("expected abort when reduction is exhausted, got nil")
|
||||
}
|
||||
if !errors.As(err, new(interface{ ContextBudgetExceeded() bool })) {
|
||||
t.Fatalf("expected wrapped budget error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestThinkThenTool_BudgetRetry_DoesNotReexecuteStaleToolCalls is the
|
||||
// full-pipeline regression guard for the stale-LastResponse hazard: when the
|
||||
// budget guard rejects iteration N+1's request, ThinkStage retries the
|
||||
// iteration, but ToolStage runs next in the SAME iteration and reads
|
||||
// state.Think.LastResponse. If ThinkStage left the PRIOR iteration's response
|
||||
// in place, ToolStage would re-execute those tool calls a second time. This
|
||||
// asserts ThinkStage clears LastResponse so ToolStage short-circuits.
|
||||
func TestThinkThenTool_BudgetRetry_DoesNotReexecuteStaleToolCalls(t *testing.T) {
|
||||
executed := 0
|
||||
deps := &PipelineDeps{
|
||||
Config: PipelineConfig{MaxIterations: 10, MaxTokens: 1000, ContextWindow: 200_000},
|
||||
CallLLM: func(_ context.Context, _ *RunState, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
return nil, fakeBudgetErr{}
|
||||
},
|
||||
PruneMessages: func(msgs []providers.Message, _ int) ([]providers.Message, PruneStats) {
|
||||
return msgs, PruneStats{ResultsTrimmed: 1} // report a change so reduction "succeeds"
|
||||
},
|
||||
ExecuteToolCall: func(_ context.Context, _ *RunState, tc providers.ToolCall) ([]providers.Message, error) {
|
||||
executed++
|
||||
return []providers.Message{{Role: "tool", ToolCallID: tc.ID, Content: "ran"}}, nil
|
||||
},
|
||||
}
|
||||
think := NewThinkStage(deps)
|
||||
tool := NewToolStage(deps)
|
||||
state := defaultState()
|
||||
state.Messages.SetHistory([]providers.Message{{Role: "user", Content: "big history"}})
|
||||
|
||||
// Simulate a prior iteration having produced a tool-call response that was
|
||||
// already executed. It must NOT be executed again when this iteration retries.
|
||||
state.Think.LastResponse = &providers.ChatResponse{
|
||||
FinishReason: "tool_calls",
|
||||
ToolCalls: []providers.ToolCall{{ID: "stale-1", Name: "write_file", Arguments: map[string]any{"path": "x"}}},
|
||||
}
|
||||
|
||||
if err := think.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("ThinkStage.Execute() should retry, got error: %v", err)
|
||||
}
|
||||
if state.Think.LastResponse != nil {
|
||||
t.Fatal("ThinkStage must clear LastResponse on budget-reduction retry")
|
||||
}
|
||||
if err := tool.Execute(context.Background(), state); err != nil {
|
||||
t.Fatalf("ToolStage.Execute() error: %v", err)
|
||||
}
|
||||
if executed != 0 {
|
||||
t.Fatalf("stale tool call re-executed %d time(s); want 0", executed)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package pipeline
|
||||
|
||||
import "github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
|
||||
// InputContextTokens returns the provider-reported input occupancy for one
|
||||
// request. Anthropic-style usage reports cached segments separately, while
|
||||
// OpenAI-style prompt tokens already include them.
|
||||
func InputContextTokens(usage providers.Usage) int {
|
||||
if usage.PromptTokensIncludeCachedSegments {
|
||||
return usage.PromptTokens
|
||||
}
|
||||
return usage.PromptTokens + usage.CacheReadTokens + usage.CacheCreationTokens
|
||||
}
|
||||
|
||||
// AddUsage accumulates provider usage while preserving telemetry fields that
|
||||
// are easy to drop when callers hand-roll partial sums.
|
||||
func AddUsage(dst *providers.Usage, src providers.Usage) {
|
||||
if dst == nil {
|
||||
return
|
||||
}
|
||||
dst.PromptTokens += src.PromptTokens
|
||||
dst.CompletionTokens += src.CompletionTokens
|
||||
dst.TotalTokens += src.TotalTokens
|
||||
dst.CacheCreationTokens += src.CacheCreationTokens
|
||||
dst.CacheReadTokens += src.CacheReadTokens
|
||||
dst.PromptTokensIncludeCachedSegments = dst.PromptTokensIncludeCachedSegments || src.PromptTokensIncludeCachedSegments
|
||||
dst.ThinkingTokens += src.ThinkingTokens
|
||||
dst.RequestCount += src.RequestCount
|
||||
dst.ImageCount += src.ImageCount
|
||||
dst.WebSearchCount += src.WebSearchCount
|
||||
}
|
||||
@@ -37,6 +37,10 @@ const (
|
||||
ShellDenyGroupsKey contextKey = "goclaw_shell_deny_groups"
|
||||
// AgentKeyKey is the context key for the agent key/name (string identifier, e.g. "default").
|
||||
AgentKeyKey contextKey = "goclaw_agent_key"
|
||||
// AgentContextWindowKey carries the calling agent's configured context window.
|
||||
AgentContextWindowKey contextKey = "goclaw_agent_context_window"
|
||||
// AgentMaxTokensKey carries the calling agent's configured output reserve.
|
||||
AgentMaxTokensKey contextKey = "goclaw_agent_max_tokens"
|
||||
// TenantIDKey is the context key for the tenant UUID.
|
||||
TenantIDKey contextKey = "goclaw_tenant_id"
|
||||
// CrossTenantKey indicates the caller has cross-tenant access (owner/system admin).
|
||||
@@ -139,6 +143,38 @@ func CredentialUserIDFromContext(ctx context.Context) string {
|
||||
return UserIDFromContext(ctx)
|
||||
}
|
||||
|
||||
// WithAgentContextWindow returns a context carrying the configured agent window.
|
||||
func WithAgentContextWindow(ctx context.Context, window int) context.Context {
|
||||
if window <= 0 {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, AgentContextWindowKey, window)
|
||||
}
|
||||
|
||||
// AgentContextWindowFromContext returns the configured agent window, or zero.
|
||||
func AgentContextWindowFromContext(ctx context.Context) int {
|
||||
if v, ok := ctx.Value(AgentContextWindowKey).(int); ok && v > 0 {
|
||||
return v
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// WithAgentMaxTokens returns a context carrying the configured agent max_tokens.
|
||||
func WithAgentMaxTokens(ctx context.Context, maxTokens int) context.Context {
|
||||
if maxTokens <= 0 {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, AgentMaxTokensKey, maxTokens)
|
||||
}
|
||||
|
||||
// AgentMaxTokensFromContext returns the configured agent max_tokens, or zero.
|
||||
func AgentMaxTokensFromContext(ctx context.Context) int {
|
||||
if v, ok := ctx.Value(AgentMaxTokensKey).(int); ok && v > 0 {
|
||||
return v
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// WithAgentID returns a new context with the given agent UUID.
|
||||
func WithAgentID(ctx context.Context, id uuid.UUID) context.Context {
|
||||
return context.WithValue(ctx, AgentIDKey, id)
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
package tokencount
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"compress/gzip"
|
||||
_ "embed"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
tiktoken "github.com/pkoukk/tiktoken-go"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
const (
|
||||
budgetEncodingPattern = `(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+`
|
||||
budgetMessageOverhead = 4
|
||||
budgetInlineMediaUnit = 1600
|
||||
)
|
||||
|
||||
//go:embed cl100k_base.tiktoken.gz
|
||||
var budgetEncodingData []byte
|
||||
|
||||
// BudgetCounter counts the complete input side of a model request with one
|
||||
// fixed, bundled GoClaw encoding. Its API intentionally has no model or provider
|
||||
// parameter: agent configuration is the only request-budget authority.
|
||||
type BudgetCounter interface {
|
||||
CountText(text string) (int, error)
|
||||
CountMessages(messages []providers.Message) (int, error)
|
||||
CountToolSchemas(tools []providers.ToolDefinition) (int, error)
|
||||
CountRequest(request providers.ChatRequest) (int, error)
|
||||
}
|
||||
|
||||
type fixedBudgetCounter struct {
|
||||
once sync.Once
|
||||
encoder *tiktoken.Tiktoken
|
||||
err error
|
||||
}
|
||||
|
||||
// NewBudgetCounter returns GoClaw's fixed local complete-input counter.
|
||||
func NewBudgetCounter() BudgetCounter { return &fixedBudgetCounter{} }
|
||||
|
||||
func (c *fixedBudgetCounter) CountText(text string) (int, error) {
|
||||
encoder, err := c.load()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return len(encoder.Encode(text, nil, nil)), nil
|
||||
}
|
||||
|
||||
func (c *fixedBudgetCounter) CountMessages(messages []providers.Message) (int, error) {
|
||||
encoder, err := c.load()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
total := 0
|
||||
for _, message := range messages {
|
||||
total += encodeBudgetText(encoder, message.Role)
|
||||
total += encodeBudgetText(encoder, message.Content)
|
||||
total += encodeBudgetText(encoder, message.Thinking)
|
||||
total += encodeBudgetText(encoder, message.ToolCallID)
|
||||
total += encodeBudgetText(encoder, message.Phase)
|
||||
total += encodeBudgetText(encoder, string(message.RawAssistantContent))
|
||||
total += budgetMessageOverhead
|
||||
if message.IsError {
|
||||
total++
|
||||
}
|
||||
for _, call := range message.ToolCalls {
|
||||
total += encodeBudgetText(encoder, call.ID)
|
||||
total += encodeBudgetText(encoder, call.Name)
|
||||
total += encodeBudgetJSON(encoder, call.Arguments)
|
||||
total += encodeBudgetJSON(encoder, call.Metadata)
|
||||
total += encodeBudgetText(encoder, call.ParseError)
|
||||
}
|
||||
for _, image := range message.Images {
|
||||
if image.Data != "" || image.URL != "" {
|
||||
total += budgetInlineMediaUnit
|
||||
}
|
||||
}
|
||||
for _, video := range message.Videos {
|
||||
if video.Data != "" || video.URL != "" {
|
||||
total += budgetInlineMediaUnit
|
||||
}
|
||||
}
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func (c *fixedBudgetCounter) CountToolSchemas(tools []providers.ToolDefinition) (int, error) {
|
||||
if len(tools) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
encoder, err := c.load()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
blob, err := json.Marshal(tools)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("marshal tool schemas for budget: %w", err)
|
||||
}
|
||||
return len(encoder.Encode(string(blob), nil, nil)), nil
|
||||
}
|
||||
|
||||
func (c *fixedBudgetCounter) CountRequest(request providers.ChatRequest) (int, error) {
|
||||
messages, err := c.CountMessages(request.Messages)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
tools, err := c.CountToolSchemas(request.Tools)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return messages + tools, nil
|
||||
}
|
||||
|
||||
func (c *fixedBudgetCounter) load() (*tiktoken.Tiktoken, error) {
|
||||
c.once.Do(func() {
|
||||
c.encoder, c.err = buildBudgetEncoder(budgetEncodingData)
|
||||
})
|
||||
return c.encoder, c.err
|
||||
}
|
||||
|
||||
func buildBudgetEncoder(compressed []byte) (*tiktoken.Tiktoken, error) {
|
||||
reader, err := gzip.NewReader(bytes.NewReader(compressed))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open bundled budget encoding: %w", err)
|
||||
}
|
||||
defer reader.Close()
|
||||
|
||||
ranks := make(map[string]int, 100_256)
|
||||
scanner := bufio.NewScanner(reader)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
parts := strings.SplitN(line, " ", 2)
|
||||
if len(parts) != 2 {
|
||||
return nil, fmt.Errorf("invalid bundled budget encoding row")
|
||||
}
|
||||
token, decodeErr := base64.StdEncoding.DecodeString(parts[0])
|
||||
if decodeErr != nil {
|
||||
return nil, fmt.Errorf("decode bundled budget token: %w", decodeErr)
|
||||
}
|
||||
rank, parseErr := strconv.Atoi(parts[1])
|
||||
if parseErr != nil {
|
||||
return nil, fmt.Errorf("parse bundled budget rank: %w", parseErr)
|
||||
}
|
||||
ranks[string(token)] = rank
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
return nil, fmt.Errorf("read bundled budget encoding: %w", err)
|
||||
}
|
||||
|
||||
special := map[string]int{
|
||||
tiktoken.ENDOFTEXT: 100257,
|
||||
tiktoken.FIM_PREFIX: 100258,
|
||||
tiktoken.FIM_MIDDLE: 100259,
|
||||
tiktoken.FIM_SUFFIX: 100260,
|
||||
tiktoken.ENDOFPROMPT: 100276,
|
||||
}
|
||||
core, err := tiktoken.NewCoreBPE(ranks, special, budgetEncodingPattern)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build bundled budget encoding: %w", err)
|
||||
}
|
||||
encoding := &tiktoken.Encoding{
|
||||
Name: "goclaw_budget_cl100k",
|
||||
PatStr: budgetEncodingPattern,
|
||||
MergeableRanks: ranks,
|
||||
SpecialTokens: special,
|
||||
}
|
||||
specialSet := make(map[string]any, len(special))
|
||||
for token := range special {
|
||||
specialSet[token] = true
|
||||
}
|
||||
return tiktoken.NewTiktoken(core, encoding, specialSet), nil
|
||||
}
|
||||
|
||||
func encodeBudgetText(encoder *tiktoken.Tiktoken, text string) int {
|
||||
if text == "" {
|
||||
return 0
|
||||
}
|
||||
return len(encoder.Encode(text, nil, nil))
|
||||
}
|
||||
|
||||
func encodeBudgetJSON(encoder *tiktoken.Tiktoken, value any) int {
|
||||
blob, err := json.Marshal(value)
|
||||
if err != nil || string(blob) == "null" || string(blob) == "{}" {
|
||||
return 0
|
||||
}
|
||||
return len(encoder.Encode(string(blob), nil, nil))
|
||||
}
|
||||
Binary file not shown.
@@ -10,43 +10,76 @@ import (
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// FallbackCounter uses rune-count/3 heuristic (matches v2 behavior).
|
||||
// Used when tiktoken-go is unavailable or model is unknown.
|
||||
// FallbackCounter is the BEST-EFFORT heuristic counter used when tiktoken-go is
|
||||
// unavailable or the model is unknown (e.g. 9router brand models that match no
|
||||
// registry prefix). It is NOT a proven token bound: real tokenizers vary by
|
||||
// language, and dense scripts like Vietnamese tokenize at ~2 chars/token or
|
||||
// worse. Because the pre-transport guard treats this count as its ceiling, the
|
||||
// heuristic deliberately uses a CONSERVATIVE chars-per-token ratio so it errs
|
||||
// toward over-counting (compact/block early) rather than under-counting (send an
|
||||
// oversized request). The only way to a provable ceiling for an un-tokenizable
|
||||
// model is to register it with a real tokenizer or use the provider's own token
|
||||
// count.
|
||||
type FallbackCounter struct{}
|
||||
|
||||
func NewFallbackCounter() *FallbackCounter { return &FallbackCounter{} }
|
||||
|
||||
// fallbackCharsPerToken is the conservative chars-per-token divisor for the
|
||||
// guard-facing count methods. Lower than a naive ~3-4 chars/token so mixed
|
||||
// Vietnamese/code content is over-counted rather than under-counted. This is a
|
||||
// safety heuristic, not an exact tokenization.
|
||||
const fallbackCharsPerToken = 2
|
||||
|
||||
func (c *FallbackCounter) Count(_ string, text string) int {
|
||||
return utf8.RuneCountInString(text) / 3
|
||||
return utf8.RuneCountInString(text) / fallbackCharsPerToken
|
||||
}
|
||||
|
||||
func (c *FallbackCounter) CountMessages(_ string, msgs []providers.Message) int {
|
||||
total := 0
|
||||
for _, m := range msgs {
|
||||
total += utf8.RuneCountInString(m.Content)/3 + PerMessageOverhead
|
||||
total += utf8.RuneCountInString(m.Content)/fallbackCharsPerToken + PerMessageOverhead
|
||||
// Match tiktokenCounter.CountMessages coverage so the fallback path
|
||||
// applies the same best-effort ceiling (thinking, tool-result id, raw
|
||||
// blocks, tool args, media all count toward the wire payload).
|
||||
total += utf8.RuneCountInString(m.Thinking) / fallbackCharsPerToken
|
||||
total += utf8.RuneCountInString(m.ToolCallID) / fallbackCharsPerToken
|
||||
total += utf8.RuneCountInString(string(m.RawAssistantContent)) / fallbackCharsPerToken
|
||||
for _, tc := range m.ToolCalls {
|
||||
total += utf8.RuneCountInString(tc.ID)/3 + utf8.RuneCountInString(tc.Name)/3
|
||||
for k, v := range tc.Arguments {
|
||||
total += utf8.RuneCountInString(k) / 3
|
||||
if s, ok := v.(string); ok {
|
||||
total += utf8.RuneCountInString(s) / 3
|
||||
} else {
|
||||
total += 10
|
||||
}
|
||||
total += utf8.RuneCountInString(tc.ID)/fallbackCharsPerToken + utf8.RuneCountInString(tc.Name)/fallbackCharsPerToken
|
||||
if b, err := json.Marshal(tc.Arguments); err == nil {
|
||||
total += utf8.RuneCountInString(string(b)) / fallbackCharsPerToken
|
||||
}
|
||||
}
|
||||
total += fallbackMediaTokenCost(m)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
// CountToolSchemas returns rune/3 heuristic count for the JSON-serialised tool list.
|
||||
// Returns 0 for nil or empty slice.
|
||||
// fallbackMediaTokenCost mirrors mediaTokenCost for the heuristic counter.
|
||||
func fallbackMediaTokenCost(m providers.Message) int {
|
||||
const perInlineMediaItem = 1600
|
||||
cost := 0
|
||||
for _, img := range m.Images {
|
||||
if img.Data != "" || img.URL != "" {
|
||||
cost += perInlineMediaItem
|
||||
}
|
||||
}
|
||||
for _, vid := range m.Videos {
|
||||
if vid.Data != "" || vid.URL != "" {
|
||||
cost += perInlineMediaItem
|
||||
}
|
||||
}
|
||||
return cost
|
||||
}
|
||||
|
||||
// CountToolSchemas returns a conservative heuristic count for the JSON-serialised
|
||||
// tool list. Returns 0 for nil or empty slice.
|
||||
func (c *FallbackCounter) CountToolSchemas(_ string, tools []providers.ToolDefinition) int {
|
||||
if len(tools) == 0 {
|
||||
return 0
|
||||
}
|
||||
blob, _ := json.Marshal(tools)
|
||||
return utf8.RuneCountInString(string(blob)) / 3
|
||||
return utf8.RuneCountInString(string(blob)) / fallbackCharsPerToken
|
||||
}
|
||||
|
||||
// ModelContextWindow uses longest-prefix-match to avoid ambiguity
|
||||
|
||||
@@ -53,9 +53,17 @@ func (c *tiktokenCounter) CountMessages(model string, msgs []providers.Message)
|
||||
return c.fallback.CountMessages(model, msgs)
|
||||
}
|
||||
|
||||
// The per-message count depends on the tokenizer, so the cache key MUST
|
||||
// include the tokenizer identity. Otherwise a message counted with model A's
|
||||
// tokenizer (e.g. cl100k for Claude) would return that stale count when the
|
||||
// same message is re-counted for model B on a different tokenizer (e.g.
|
||||
// o200k for GPT-4o) — most visibly on fallback candidates. Models that share
|
||||
// a tokenizer correctly share cache entries.
|
||||
tokenizerID := resolveModelInfo(model).TokenizerID
|
||||
|
||||
total := 0
|
||||
for _, m := range msgs {
|
||||
hash := messageHash(m)
|
||||
hash := messageHash(tokenizerID, m)
|
||||
|
||||
c.mu.RLock()
|
||||
cached, ok := c.msgCache[hash]
|
||||
@@ -67,10 +75,29 @@ func (c *tiktokenCounter) CountMessages(model string, msgs []providers.Message)
|
||||
}
|
||||
|
||||
count := len(enc.Encode(m.Content, nil, nil)) + PerMessageOverhead
|
||||
// Thinking/reasoning is sent back to the provider on subsequent turns
|
||||
// (Anthropic requires it for tool-use passback), so it counts toward input.
|
||||
if m.Thinking != "" {
|
||||
count += len(enc.Encode(m.Thinking, nil, nil))
|
||||
}
|
||||
// Tool result correlation id (role="tool" messages).
|
||||
if m.ToolCallID != "" {
|
||||
count += len(enc.Encode(m.ToolCallID, nil, nil))
|
||||
}
|
||||
// Raw assistant content blocks (Anthropic thinking-block passback) are
|
||||
// serialized into the wire payload verbatim — count their bytes.
|
||||
if len(m.RawAssistantContent) > 0 {
|
||||
count += len(enc.Encode(string(m.RawAssistantContent), nil, nil))
|
||||
}
|
||||
for _, tc := range m.ToolCalls {
|
||||
count += len(enc.Encode(tc.Name, nil, nil))
|
||||
count += len(enc.Encode(tc.ID, nil, nil))
|
||||
// Tool call arguments are serialized as JSON into the request and can
|
||||
// dominate token cost; the previous counter ignored them entirely.
|
||||
count += encodeToolArgs(enc, tc.Arguments)
|
||||
}
|
||||
// Media (images/videos) carry a per-item token cost when sent inline.
|
||||
count += mediaTokenCost(m)
|
||||
|
||||
c.mu.Lock()
|
||||
c.msgCache[hash] = count
|
||||
@@ -167,19 +194,84 @@ func resolveModelInfo(model string) ModelInfo {
|
||||
}
|
||||
|
||||
// messageHash computes FNV-1a hash of message content for cache keying.
|
||||
func messageHash(m providers.Message) uint64 {
|
||||
// It MUST cover every field that CountMessages encodes; otherwise a payload
|
||||
// whose thinking/args/tool-result/raw-blocks/media changed but whose Content is
|
||||
// unchanged would return a stale cached count for a different wire payload.
|
||||
// The tokenizer ID is folded in first because the token count is
|
||||
// tokenizer-dependent: the same message yields different counts under cl100k vs
|
||||
// o200k, so counts must not be shared across tokenizers.
|
||||
func messageHash(tokenizerID TokenizerID, m providers.Message) uint64 {
|
||||
h := fnv.New64a()
|
||||
h.Write([]byte(tokenizerID))
|
||||
h.Write([]byte{0}) // separator
|
||||
h.Write([]byte(m.Role))
|
||||
h.Write([]byte{0}) // separator
|
||||
h.Write([]byte(m.Content))
|
||||
h.Write([]byte{0})
|
||||
h.Write([]byte(m.Thinking))
|
||||
h.Write([]byte{0})
|
||||
h.Write([]byte(m.ToolCallID))
|
||||
h.Write([]byte{0})
|
||||
h.Write(m.RawAssistantContent)
|
||||
for _, tc := range m.ToolCalls {
|
||||
h.Write([]byte{0})
|
||||
h.Write([]byte(tc.ID))
|
||||
h.Write([]byte(tc.Name))
|
||||
if b, err := json.Marshal(tc.Arguments); err == nil {
|
||||
h.Write(b)
|
||||
}
|
||||
}
|
||||
for _, img := range m.Images {
|
||||
h.Write([]byte{0})
|
||||
h.Write([]byte(img.MimeType))
|
||||
h.Write([]byte(img.URL))
|
||||
}
|
||||
for _, vid := range m.Videos {
|
||||
h.Write([]byte{0})
|
||||
h.Write([]byte(vid.MimeType))
|
||||
h.Write([]byte(vid.URL))
|
||||
}
|
||||
return h.Sum64()
|
||||
}
|
||||
|
||||
// encodeToolArgs returns the BPE token count of a tool call's arguments as they
|
||||
// are serialized into the request payload (JSON). Returns 0 when arguments are
|
||||
// empty or cannot be marshalled.
|
||||
func encodeToolArgs(enc *tiktoken.Tiktoken, args map[string]any) int {
|
||||
if len(args) == 0 {
|
||||
return 0
|
||||
}
|
||||
b, err := json.Marshal(args)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return len(enc.Encode(string(b), nil, nil))
|
||||
}
|
||||
|
||||
// mediaTokenCost approximates the token cost of inline media on a message.
|
||||
// This is BEST-EFFORT, not a proven bound: a flat per-item estimate can
|
||||
// under-count a large inline video whose true token cost far exceeds it. It
|
||||
// exists so inline media is not counted as zero. Tool-internal media that
|
||||
// travels out-of-band (provider file/transcription APIs) is measured separately
|
||||
// by size in the caps guard (EstimateOutOfBandMediaTokens); this path only
|
||||
// covers media actually embedded in the message. Callers that persist media as
|
||||
// MediaRefs (not inlined) incur no cost here.
|
||||
func mediaTokenCost(m providers.Message) int {
|
||||
const perInlineMediaItem = 1600 // best-effort estimate per inline image/video
|
||||
cost := 0
|
||||
for _, img := range m.Images {
|
||||
if img.Data != "" || img.URL != "" {
|
||||
cost += perInlineMediaItem
|
||||
}
|
||||
}
|
||||
for _, vid := range m.Videos {
|
||||
if vid.Data != "" || vid.URL != "" {
|
||||
cost += perInlineMediaItem
|
||||
}
|
||||
}
|
||||
return cost
|
||||
}
|
||||
|
||||
// NewTokenCounter creates the best available counter.
|
||||
// Uses tiktoken if requested, falls back to rune/3 heuristic.
|
||||
func NewTokenCounter(useTiktoken bool) TokenCounter {
|
||||
|
||||
@@ -59,6 +59,38 @@ func TestCountMessages_Cache(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestCountMessages_CacheIsTokenizerScoped is the regression guard for the
|
||||
// stale cross-tokenizer count bug: the same message counted under a cl100k
|
||||
// model then an o200k model must re-encode (different tokenizer) rather than
|
||||
// return the first tokenizer's cached count. This matters most on fallback
|
||||
// candidates that switch tokenizer mid-run.
|
||||
func TestCountMessages_CacheIsTokenizerScoped(t *testing.T) {
|
||||
c := NewTiktokenCounter()
|
||||
// A string whose cl100k and o200k token counts differ.
|
||||
msgs := []providers.Message{
|
||||
{Role: "user", Content: "internationalization tokenization pseudopseudohypoparathyroidism"},
|
||||
}
|
||||
|
||||
cl := c.CountMessages("claude-sonnet-4-5-20250929", msgs) // cl100k_base
|
||||
o2 := c.CountMessages("gpt-4o", msgs) // o200k_base
|
||||
|
||||
if cl <= 0 || o2 <= 0 {
|
||||
t.Fatalf("counts must be positive: cl=%d o2=%d", cl, o2)
|
||||
}
|
||||
// The two tokenizers should produce different counts for this string; if they
|
||||
// were sharing a cache entry, o2 would equal cl (the stale bug).
|
||||
if cl == o2 {
|
||||
t.Fatalf("cl100k and o200k returned identical count %d — cache is not tokenizer-scoped", cl)
|
||||
}
|
||||
// Two distinct tokenizers => two cache entries for the one message.
|
||||
c.mu.RLock()
|
||||
cacheLen := len(c.msgCache)
|
||||
c.mu.RUnlock()
|
||||
if cacheLen != 2 {
|
||||
t.Errorf("cache has %d entries, want 2 (one per tokenizer)", cacheLen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountMessages_Overhead(t *testing.T) {
|
||||
c := NewTiktokenCounter()
|
||||
msgs := []providers.Message{
|
||||
|
||||
@@ -98,7 +98,7 @@ func (t *ReadAudioTool) callProvider(ctx context.Context, cp credentialProvider,
|
||||
Model: model,
|
||||
Options: map[string]any{"max_tokens": 16384},
|
||||
}
|
||||
reservation, reserveErr := reserveToolLLMUsage(ctx, t.usageCaps, t.Name(), providerName, model, chatReq)
|
||||
reservation, reserveErr := reserveToolLLMUsageWithMedia(ctx, t.usageCaps, t.Name(), providerName, model, chatReq, mime, data)
|
||||
if reserveErr != nil {
|
||||
return nil, nil, reserveErr
|
||||
}
|
||||
@@ -120,7 +120,7 @@ func (t *ReadAudioTool) callProvider(ctx context.Context, cp credentialProvider,
|
||||
Model: model,
|
||||
Options: map[string]any{"max_tokens": 16384},
|
||||
}
|
||||
reservation, reserveErr := reserveToolLLMUsage(ctx, t.usageCaps, t.Name(), providerName, model, chatReq)
|
||||
reservation, reserveErr := reserveToolLLMUsageWithMedia(ctx, t.usageCaps, t.Name(), providerName, model, chatReq, mime, data)
|
||||
if reserveErr != nil {
|
||||
return nil, nil, reserveErr
|
||||
}
|
||||
@@ -142,7 +142,7 @@ func (t *ReadAudioTool) callProvider(ctx context.Context, cp credentialProvider,
|
||||
Model: model,
|
||||
Options: map[string]any{"max_tokens": 16384},
|
||||
}
|
||||
reservation, reserveErr := reserveToolLLMUsage(ctx, t.usageCaps, t.Name(), providerName, model, chatReq)
|
||||
reservation, reserveErr := reserveToolLLMUsageWithMedia(ctx, t.usageCaps, t.Name(), providerName, model, chatReq, mime, data)
|
||||
if reserveErr != nil {
|
||||
return nil, nil, reserveErr
|
||||
}
|
||||
|
||||
@@ -160,7 +160,7 @@ func (t *ReadDocumentTool) callProvider(ctx context.Context, cp credentialProvid
|
||||
Model: model,
|
||||
Options: map[string]any{"max_tokens": 16384},
|
||||
}
|
||||
reservation, reserveErr := reserveToolLLMUsage(ctx, t.usageCaps, t.Name(), providerName, model, chatReq)
|
||||
reservation, reserveErr := reserveToolLLMUsageWithMedia(ctx, t.usageCaps, t.Name(), providerName, model, chatReq, mime, data)
|
||||
if reserveErr != nil {
|
||||
return nil, nil, reserveErr
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps"
|
||||
)
|
||||
|
||||
// resolveVideoFile finds the video file path from context MediaRefs.
|
||||
@@ -89,7 +90,17 @@ func (t *ReadVideoTool) callProvider(ctx context.Context, cp credentialProvider,
|
||||
Model: model,
|
||||
Options: map[string]any{"max_tokens": 16384},
|
||||
}
|
||||
reservation, reserveErr := reserveToolLLMUsage(ctx, t.usageCaps, t.Name(), providerName, model, chatReq)
|
||||
// The reservation differs by transport. The base64 path has the real
|
||||
// bytes now and counts them. The URL-stream path never buffers the
|
||||
// payload, so it cannot prove completeInput and fails closed under an
|
||||
// agent budget rather than trusting Content-Length.
|
||||
var reservation *usagecaps.Reservation
|
||||
var reserveErr error
|
||||
if videoURL != "" {
|
||||
reservation, reserveErr = reserveToolLLMUsageUnverifiableMedia(ctx, t.usageCaps, t.Name(), providerName, model, chatReq)
|
||||
} else {
|
||||
reservation, reserveErr = reserveToolLLMUsageWithMedia(ctx, t.usageCaps, t.Name(), providerName, model, chatReq, mime, data)
|
||||
}
|
||||
if reserveErr != nil {
|
||||
return nil, nil, reserveErr
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/security"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
)
|
||||
|
||||
type mockCredentialProvider struct {
|
||||
@@ -52,72 +53,50 @@ func TestReadVideo_PrivateURL_Error(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadVideo_GeminiURL_Validation(t *testing.T) {
|
||||
// TestReadVideo_GeminiURL_FailsClosedUnderAgentBudget locks the fail-closed
|
||||
// contract for a streamed video URL. The stream is never buffered, so its bytes
|
||||
// cannot be counted into completeInput; under an agent budget the call refuses
|
||||
// before any transport rather than undercounting. Because every tool LLM call
|
||||
// is agent-scoped in production, this refusal — not the downstream
|
||||
// Content-Length / 2 GB / HTTP-status checks — is the reachable behavior. Those
|
||||
// transport validations remain in callProvider for defense in depth but are
|
||||
// unreachable once an agent budget is present.
|
||||
func TestReadVideo_GeminiURL_FailsClosedUnderAgentBudget(t *testing.T) {
|
||||
security.SetAllowLoopbackForTest(true)
|
||||
defer security.SetAllowLoopbackForTest(false)
|
||||
|
||||
// Missing Content-Length should fail before upload.
|
||||
ts1 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Transfer-Encoding", "chunked")
|
||||
w.Write([]byte("chunked data mock video"))
|
||||
// A server that would otherwise satisfy the static-streaming constraints
|
||||
// (valid Content-Length, 2xx). The call must still fail closed before it is
|
||||
// ever contacted, because the payload cannot be counted.
|
||||
hits := 0
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits++
|
||||
w.Header().Set("Content-Length", "1024")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write(make([]byte, 1024))
|
||||
}))
|
||||
defer ts1.Close()
|
||||
defer ts.Close()
|
||||
|
||||
tool := NewReadVideoTool(nil, nil)
|
||||
cp := &mockCredentialProvider{apiKey: "test-key"}
|
||||
|
||||
params1 := map[string]any{
|
||||
ctx := store.WithAgentContextWindow(context.Background(), 200_000)
|
||||
ctx = store.WithAgentMaxTokens(ctx, 32_000)
|
||||
|
||||
params := map[string]any{
|
||||
"prompt": "describe this video",
|
||||
"url": ts1.URL,
|
||||
"url": ts.URL,
|
||||
"_provider_type": "gemini",
|
||||
}
|
||||
|
||||
_, _, err := tool.callProvider(context.Background(), cp, "gemini", "gemini-2.5-flash", params1)
|
||||
_, _, err := tool.callProvider(ctx, cp, "gemini", "gemini-2.5-flash", params)
|
||||
if err == nil {
|
||||
t.Fatalf("expected error for missing Content-Length")
|
||||
t.Fatalf("expected fail-closed for an unverifiable streamed video URL under an agent budget")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "URL does not support static streaming") {
|
||||
t.Errorf("unexpected error for missing Content-Length: %v", err)
|
||||
if !strings.Contains(err.Error(), "cannot verify streamed native media") {
|
||||
t.Errorf("unexpected error, want fail-closed refusal: %v", err)
|
||||
}
|
||||
|
||||
// Content-Length over the Gemini File API limit should fail before upload.
|
||||
ts2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Length", "2147483649") // 2GB + 1 byte
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer ts2.Close()
|
||||
|
||||
params2 := map[string]any{
|
||||
"prompt": "describe this video",
|
||||
"url": ts2.URL,
|
||||
"_provider_type": "gemini",
|
||||
}
|
||||
|
||||
_, _, err = tool.callProvider(context.Background(), cp, "gemini", "gemini-2.5-flash", params2)
|
||||
if err == nil {
|
||||
t.Fatalf("expected error for Content-Length exceeding 2GB")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "exceeds the maximum limit of 2 GB") {
|
||||
t.Errorf("unexpected error for limit exceed: %v", err)
|
||||
}
|
||||
|
||||
// Non-2xx status should be reported before upload.
|
||||
ts3 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
}))
|
||||
defer ts3.Close()
|
||||
|
||||
params3 := map[string]any{
|
||||
"prompt": "describe this video",
|
||||
"url": ts3.URL,
|
||||
"_provider_type": "gemini",
|
||||
}
|
||||
|
||||
_, _, err = tool.callProvider(context.Background(), cp, "gemini", "gemini-2.5-flash", params3)
|
||||
if err == nil {
|
||||
t.Fatalf("expected error for HTTP 404 status code")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "video URL returned status code 404") {
|
||||
t.Errorf("unexpected error for HTTP status code: %v", err)
|
||||
if hits != 0 {
|
||||
t.Errorf("fail-closed must refuse before any transport, but the URL was contacted %d time(s)", hits)
|
||||
}
|
||||
}
|
||||
+42
-28
@@ -41,34 +41,38 @@ const (
|
||||
|
||||
// SubagentTask tracks a running or completed subagent.
|
||||
type SubagentTask struct {
|
||||
ID string `json:"id"`
|
||||
ParentID string `json:"parentId"`
|
||||
Task string `json:"task"`
|
||||
Label string `json:"label"`
|
||||
Status string `json:"status"` // "running", "completed", "failed", "cancelled"
|
||||
Result string `json:"result,omitempty"`
|
||||
Depth int `json:"depth"`
|
||||
Model string `json:"model,omitempty"` // model override for this subagent
|
||||
TotalInputTokens int64 `json:"totalInputTokens,omitempty"`
|
||||
TotalOutputTokens int64 `json:"totalOutputTokens,omitempty"`
|
||||
OriginChannel string `json:"originChannel,omitempty"`
|
||||
OriginChatID string `json:"originChatId,omitempty"`
|
||||
OriginPeerKind string `json:"originPeerKind,omitempty"` // "direct" or "group" (for session key building)
|
||||
OriginLocalKey string `json:"originLocalKey,omitempty"` // composite key with topic/thread suffix for routing
|
||||
OriginUserID string `json:"originUserId,omitempty"` // parent's userID for per-user scoping propagation
|
||||
OriginSenderID string `json:"originSenderId,omitempty"` // real acting sender; preserves permission attribution in announce re-ingress (#915)
|
||||
OriginRole string `json:"originRole,omitempty"` // parent's RBAC role; bypasses per-user grants for admin/operator/owner in re-ingress (#915)
|
||||
OriginSessionKey string `json:"originSessionKey,omitempty"` // exact parent session key for announce routing (WS uses non-standard format)
|
||||
CreatedAt int64 `json:"createdAt"`
|
||||
CompletedAt int64 `json:"completedAt,omitempty"`
|
||||
Media []bus.MediaFile `json:"-"` // media files from tool results
|
||||
OriginAgentID uuid.UUID `json:"-"` // parent agent UUID for usage caps and scoped tools
|
||||
OriginTenantID uuid.UUID `json:"-"` // parent's tenant for announce routing
|
||||
OriginTraceID uuid.UUID `json:"-"` // parent trace for announce linking
|
||||
OriginRootSpanID uuid.UUID `json:"-"` // parent agent's root span ID
|
||||
cancelFunc context.CancelFunc `json:"-"` // per-task context cancel
|
||||
spawnConfig SubagentConfig `json:"-"` // resolved config at spawn time (per-agent override merged)
|
||||
dbID uuid.UUID `json:"-"` // persistent DB UUID (zero if not persisted)
|
||||
ID string `json:"id"`
|
||||
ParentID string `json:"parentId"`
|
||||
Task string `json:"task"`
|
||||
Label string `json:"label"`
|
||||
Status string `json:"status"` // "running", "completed", "failed", "cancelled"
|
||||
Result string `json:"result,omitempty"`
|
||||
Depth int `json:"depth"`
|
||||
Model string `json:"model,omitempty"` // model override for this subagent
|
||||
TotalInputTokens int64 `json:"totalInputTokens,omitempty"`
|
||||
TotalOutputTokens int64 `json:"totalOutputTokens,omitempty"`
|
||||
OriginChannel string `json:"originChannel,omitempty"`
|
||||
OriginChatID string `json:"originChatId,omitempty"`
|
||||
OriginPeerKind string `json:"originPeerKind,omitempty"` // "direct" or "group" (for session key building)
|
||||
OriginLocalKey string `json:"originLocalKey,omitempty"` // composite key with topic/thread suffix for routing
|
||||
OriginUserID string `json:"originUserId,omitempty"` // parent's userID for per-user scoping propagation
|
||||
OriginSenderID string `json:"originSenderId,omitempty"` // real acting sender; preserves permission attribution in announce re-ingress (#915)
|
||||
OriginRole string `json:"originRole,omitempty"` // parent's RBAC role; bypasses per-user grants for admin/operator/owner in re-ingress (#915)
|
||||
OriginSessionKey string `json:"originSessionKey,omitempty"` // exact parent session key for announce routing (WS uses non-standard format)
|
||||
CreatedAt int64 `json:"createdAt"`
|
||||
CompletedAt int64 `json:"completedAt,omitempty"`
|
||||
Media []bus.MediaFile `json:"-"` // media files from tool results
|
||||
OriginAgentID uuid.UUID `json:"-"` // parent agent UUID for usage caps and scoped tools
|
||||
OriginTenantID uuid.UUID `json:"-"` // parent's tenant for announce routing
|
||||
OriginTraceID uuid.UUID `json:"-"` // parent trace for announce linking
|
||||
OriginRootSpanID uuid.UUID `json:"-"` // parent agent's root span ID
|
||||
// OriginContextWindow and OriginMaxTokens are captured from the caller at
|
||||
// spawn so a shared manager cannot mix budgets between agents.
|
||||
OriginContextWindow int `json:"-"`
|
||||
OriginMaxTokens int `json:"-"`
|
||||
cancelFunc context.CancelFunc `json:"-"` // per-task context cancel
|
||||
spawnConfig SubagentConfig `json:"-"` // resolved config at spawn time (per-agent override merged)
|
||||
dbID uuid.UUID `json:"-"` // persistent DB UUID (zero if not persisted)
|
||||
}
|
||||
|
||||
// SubagentManager manages the lifecycle of spawned subagents.
|
||||
@@ -86,6 +90,9 @@ type SubagentManager struct {
|
||||
announceQueue *AnnounceQueue // optional: batches announces with debounce
|
||||
taskStore store.SubagentTaskStore // optional: persists tasks to DB (fire-and-forget)
|
||||
usageCaps *usagecaps.Service
|
||||
// Default agent budget used only when a task was created without caller context.
|
||||
contextWindow int
|
||||
maxTokens int
|
||||
}
|
||||
|
||||
// NewSubagentManager creates a new subagent manager.
|
||||
@@ -123,6 +130,13 @@ func (sm *SubagentManager) SetUsageCapService(s *usagecaps.Service) {
|
||||
sm.usageCaps = s
|
||||
}
|
||||
|
||||
// SetAgentBudget records the manager's fallback agent budget. Spawned tasks
|
||||
// normally use their caller-specific values captured from context.
|
||||
func (sm *SubagentManager) SetAgentBudget(window, maxTokens int) {
|
||||
sm.contextWindow = window
|
||||
sm.maxTokens = maxTokens
|
||||
}
|
||||
|
||||
// effectiveConfig returns the per-agent context override merged with defaults,
|
||||
// or falls back to sm.config when no override is present.
|
||||
func (sm *SubagentManager) effectiveConfig(ctx context.Context) SubagentConfig {
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps"
|
||||
)
|
||||
|
||||
// callTrackingProvider records whether Chat/ChatStream ran, so a test can prove
|
||||
// the pre-transport guard blocked BEFORE any transport call.
|
||||
type callTrackingProvider struct {
|
||||
called bool
|
||||
}
|
||||
|
||||
func (p *callTrackingProvider) Name() string { return "tracking" }
|
||||
func (p *callTrackingProvider) DefaultModel() string { return "claude-sonnet-4-5-20250929" }
|
||||
func (p *callTrackingProvider) Chat(_ context.Context, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
p.called = true
|
||||
return &providers.ChatResponse{Content: "ok", FinishReason: "stop"}, nil
|
||||
}
|
||||
func (p *callTrackingProvider) ChatStream(_ context.Context, _ providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) {
|
||||
p.called = true
|
||||
return &providers.ChatResponse{Content: "ok", FinishReason: "stop"}, nil
|
||||
}
|
||||
|
||||
func subagentBigMessage(model string, targetChars int) providers.Message {
|
||||
return providers.Message{Role: "user", Content: strings.Repeat("word ", targetChars)}
|
||||
}
|
||||
|
||||
// TestSubagentSpawn_CapturesOriginContextWindow proves the Spawn path captures
|
||||
// the CALLING agent's context window (set in ctx by the agent loop's
|
||||
// injectContext) into SubagentTask.OriginContextWindow — the value the guard
|
||||
// later uses. This locks the plumbing the previous round only claimed.
|
||||
func TestSubagentSpawn_CapturesOriginContextWindow(t *testing.T) {
|
||||
provider := &recordingSubagentProvider{}
|
||||
manager := NewSubagentManager(provider, nil, "manager-default", nil, NewRegistry, SubagentConfig{
|
||||
MaxConcurrent: 4, MaxSpawnDepth: 3, MaxChildrenPerAgent: 8,
|
||||
})
|
||||
manager.SetAgentBudget(200_000, 32_000) // manager default (would be wrong for a 128k caller)
|
||||
|
||||
// Calling agent configured at 128k/8192, propagated via ctx (as injectContext does).
|
||||
ctx := store.WithAgentContextWindow(context.Background(), 128_000)
|
||||
ctx = store.WithAgentMaxTokens(ctx, 8_192)
|
||||
|
||||
_, _, err := manager.RunSync(ctx, "parent", 0, "task", "label", "", "chan", "chat")
|
||||
if err != nil {
|
||||
t.Fatalf("RunSync error: %v", err)
|
||||
}
|
||||
|
||||
// Find the task the manager created and assert it captured the CALLER's
|
||||
// 128k/8192 budget, not the manager default 200k/32000.
|
||||
manager.mu.RLock()
|
||||
defer manager.mu.RUnlock()
|
||||
if len(manager.tasks) == 0 {
|
||||
t.Fatal("no subagent task recorded")
|
||||
}
|
||||
for _, task := range manager.tasks {
|
||||
if task.OriginContextWindow != 128_000 {
|
||||
t.Fatalf("OriginContextWindow = %d, want 128000 (calling agent, not manager default)", task.OriginContextWindow)
|
||||
}
|
||||
if task.OriginMaxTokens != 8_192 {
|
||||
t.Fatalf("OriginMaxTokens = %d, want 8192 (calling agent, not manager default)", task.OriginMaxTokens)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestChatSubagentWithUsageCap_HonorsOriginWindow is the production-path guard
|
||||
// for finding #4: a subagent request that fits the 200k model window but exceeds
|
||||
// the calling agent's 128k cap must be blocked before transport.
|
||||
func TestChatSubagentWithUsageCap_HonorsOriginWindow(t *testing.T) {
|
||||
model := "claude-sonnet-4-5-20250929" // 200k model window, cl100k tokenizer
|
||||
manager := NewSubagentManager(nil, nil, "manager-default", nil, NewRegistry, SubagentConfig{})
|
||||
// Manager default is 200k; the per-task origin window must override it.
|
||||
manager.SetAgentBudget(200_000, 32_000)
|
||||
|
||||
// ~140k tokens: fits 200k model window, exceeds a 128k agent cap.
|
||||
req := providers.ChatRequest{
|
||||
Model: model,
|
||||
Messages: []providers.Message{subagentBigMessage(model, 140_000)},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
|
||||
// Task from a 128k-configured calling agent.
|
||||
task128 := &SubagentTask{ID: "t128", OriginContextWindow: 128_000, OriginMaxTokens: 8_192}
|
||||
prov := &callTrackingProvider{}
|
||||
_, err := manager.chatSubagentWithUsageCap(context.Background(), task128, prov, model, req, 0, 0)
|
||||
if err == nil {
|
||||
t.Fatal("expected abort: 140k request exceeds 128k origin window")
|
||||
}
|
||||
if prov.called {
|
||||
t.Fatal("provider must NOT be called when the guard aborts (transport calls must be 0)")
|
||||
}
|
||||
var ctxErr *usagecaps.ContextWindowExceededError
|
||||
if !errors.As(err, &ctxErr) {
|
||||
t.Fatalf("expected *ContextWindowExceededError, got %T: %v", err, err)
|
||||
}
|
||||
if ctxErr.ContextWindow != 128_000 {
|
||||
t.Fatalf("guard used window %d, want 128000 (origin, not model/manager)", ctxErr.ContextWindow)
|
||||
}
|
||||
}
|
||||
|
||||
// TestChatSubagentWithUsageCap_TwoAgentsIsolated proves two subagents from
|
||||
// agents configured at 128k and 200k are guarded independently by ONE shared
|
||||
// manager: the same 140k request is blocked for the 128k caller but allowed for
|
||||
// the 200k caller.
|
||||
func TestChatSubagentWithUsageCap_TwoAgentsIsolated(t *testing.T) {
|
||||
model := "claude-sonnet-4-5-20250929"
|
||||
manager := NewSubagentManager(nil, nil, "manager-default", nil, NewRegistry, SubagentConfig{})
|
||||
manager.SetAgentBudget(200_000, 32_000)
|
||||
|
||||
req := providers.ChatRequest{
|
||||
Model: model,
|
||||
Messages: []providers.Message{subagentBigMessage(model, 140_000)},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
|
||||
// 128k caller: blocked.
|
||||
prov128 := &callTrackingProvider{}
|
||||
if _, err := manager.chatSubagentWithUsageCap(context.Background(), &SubagentTask{ID: "a128", OriginContextWindow: 128_000, OriginMaxTokens: 8_192}, prov128, model, req, 0, 0); err == nil {
|
||||
t.Fatal("128k caller: expected abort for 140k request")
|
||||
}
|
||||
if prov128.called {
|
||||
t.Fatal("128k caller: provider must not be called")
|
||||
}
|
||||
|
||||
// 200k caller: allowed (140k < 200k), provider reached.
|
||||
prov200 := &callTrackingProvider{}
|
||||
if _, err := manager.chatSubagentWithUsageCap(context.Background(), &SubagentTask{ID: "a200", OriginContextWindow: 200_000, OriginMaxTokens: 8_192}, prov200, model, req, 0, 0); err != nil {
|
||||
t.Fatalf("200k caller: expected allow for 140k request, got %v", err)
|
||||
}
|
||||
if !prov200.called {
|
||||
t.Fatal("200k caller: provider should have been called")
|
||||
}
|
||||
}
|
||||
@@ -344,8 +344,23 @@ func (sm *SubagentManager) executeTask(ctx context.Context, task *SubagentTask)
|
||||
}
|
||||
|
||||
func (sm *SubagentManager) chatSubagentWithUsageCap(ctx context.Context, task *SubagentTask, activeProvider providers.Provider, model string, chatReq providers.ChatRequest, iteration, attempt int) (*providers.ChatResponse, error) {
|
||||
budget := usagecaps.AgentBudget{
|
||||
ContextWindow: task.OriginContextWindow,
|
||||
MaxTokens: task.OriginMaxTokens,
|
||||
}
|
||||
if budget.ContextWindow <= 0 {
|
||||
budget.ContextWindow = sm.contextWindow
|
||||
}
|
||||
if budget.MaxTokens <= 0 {
|
||||
budget.MaxTokens = sm.maxTokens
|
||||
}
|
||||
chatReq = clampToolRequestMaxTokens(chatReq, budget.MaxTokens)
|
||||
if fallbackProvider, ok := activeProvider.(*providers.ModelFallbackProvider); ok {
|
||||
before := func(callCtx context.Context, entry providers.FallbackCandidate, actualReq providers.ChatRequest) (providers.FallbackAfterCall, error) {
|
||||
// Guard this fallback request with the calling agent budget.
|
||||
if guardErr := usagecaps.GuardContextWindow(clampToolRequestMaxTokens(actualReq, budget.MaxTokens), entry.ProviderName, actualReq.Model, "subagent:"+task.ID, budget); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
reservation, err := sm.reserveSubagentUsage(callCtx, task, entry.ProviderName, actualReq.Model, actualReq, fmt.Sprintf("%d:%d:%s:%s", iteration, attempt, entry.ProviderName, actualReq.Model))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -358,6 +373,10 @@ func (sm *SubagentManager) chatSubagentWithUsageCap(ctx context.Context, task *S
|
||||
}
|
||||
return fallbackProvider.ChatWithHook(ctx, chatReq, before)
|
||||
}
|
||||
// Guard the non-fallback request with the calling agent budget.
|
||||
if guardErr := usagecaps.GuardContextWindow(chatReq, activeProvider.Name(), model, "subagent:"+task.ID, budget); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
reservation, err := sm.reserveSubagentUsage(ctx, task, activeProvider.Name(), model, chatReq, fmt.Sprintf("%d:%d", iteration, attempt))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -72,27 +72,29 @@ func (sm *SubagentManager) Spawn(
|
||||
}
|
||||
|
||||
subTask := &SubagentTask{
|
||||
ID: id,
|
||||
ParentID: parentID,
|
||||
Task: task,
|
||||
Label: label,
|
||||
Status: "running",
|
||||
Depth: depth + 1,
|
||||
Model: modelOverride,
|
||||
OriginChannel: channel,
|
||||
OriginChatID: chatID,
|
||||
OriginPeerKind: peerKind,
|
||||
OriginLocalKey: ToolLocalKeyFromCtx(ctx),
|
||||
OriginUserID: store.UserIDFromContext(ctx),
|
||||
OriginSenderID: store.SenderIDFromContext(ctx),
|
||||
OriginRole: store.RoleFromContext(ctx),
|
||||
OriginSessionKey: ToolSessionKeyFromCtx(ctx),
|
||||
OriginAgentID: store.AgentIDFromContext(ctx),
|
||||
OriginTenantID: store.TenantIDFromContext(ctx),
|
||||
OriginTraceID: tracing.TraceIDFromContext(ctx),
|
||||
OriginRootSpanID: tracing.ParentSpanIDFromContext(ctx),
|
||||
CreatedAt: time.Now().UnixMilli(),
|
||||
spawnConfig: cfg,
|
||||
ID: id,
|
||||
ParentID: parentID,
|
||||
Task: task,
|
||||
Label: label,
|
||||
Status: "running",
|
||||
Depth: depth + 1,
|
||||
Model: modelOverride,
|
||||
OriginChannel: channel,
|
||||
OriginChatID: chatID,
|
||||
OriginPeerKind: peerKind,
|
||||
OriginLocalKey: ToolLocalKeyFromCtx(ctx),
|
||||
OriginUserID: store.UserIDFromContext(ctx),
|
||||
OriginSenderID: store.SenderIDFromContext(ctx),
|
||||
OriginRole: store.RoleFromContext(ctx),
|
||||
OriginSessionKey: ToolSessionKeyFromCtx(ctx),
|
||||
OriginAgentID: store.AgentIDFromContext(ctx),
|
||||
OriginTenantID: store.TenantIDFromContext(ctx),
|
||||
OriginTraceID: tracing.TraceIDFromContext(ctx),
|
||||
OriginRootSpanID: tracing.ParentSpanIDFromContext(ctx),
|
||||
OriginContextWindow: store.AgentContextWindowFromContext(ctx),
|
||||
OriginMaxTokens: store.AgentMaxTokensFromContext(ctx),
|
||||
CreatedAt: time.Now().UnixMilli(),
|
||||
spawnConfig: cfg,
|
||||
}
|
||||
// Detach from parent's cancellation chain so subagent survives after parent run completes.
|
||||
// WithoutCancel preserves all context values (agent ID, workspace, trace info, etc.)
|
||||
@@ -154,26 +156,28 @@ func (sm *SubagentManager) RunSync(
|
||||
}
|
||||
|
||||
subTask := &SubagentTask{
|
||||
ID: id,
|
||||
ParentID: parentID,
|
||||
Task: task,
|
||||
Label: label,
|
||||
Status: "running",
|
||||
Depth: depth + 1,
|
||||
Model: modelOverride,
|
||||
OriginChannel: channel,
|
||||
OriginChatID: chatID,
|
||||
OriginLocalKey: ToolLocalKeyFromCtx(ctx),
|
||||
OriginUserID: store.UserIDFromContext(ctx),
|
||||
OriginSenderID: store.SenderIDFromContext(ctx),
|
||||
OriginRole: store.RoleFromContext(ctx),
|
||||
OriginSessionKey: ToolSessionKeyFromCtx(ctx),
|
||||
OriginAgentID: store.AgentIDFromContext(ctx),
|
||||
OriginTenantID: store.TenantIDFromContext(ctx),
|
||||
OriginTraceID: tracing.TraceIDFromContext(ctx),
|
||||
OriginRootSpanID: tracing.ParentSpanIDFromContext(ctx),
|
||||
CreatedAt: time.Now().UnixMilli(),
|
||||
spawnConfig: cfg,
|
||||
ID: id,
|
||||
ParentID: parentID,
|
||||
Task: task,
|
||||
Label: label,
|
||||
Status: "running",
|
||||
Depth: depth + 1,
|
||||
Model: modelOverride,
|
||||
OriginChannel: channel,
|
||||
OriginChatID: chatID,
|
||||
OriginLocalKey: ToolLocalKeyFromCtx(ctx),
|
||||
OriginUserID: store.UserIDFromContext(ctx),
|
||||
OriginSenderID: store.SenderIDFromContext(ctx),
|
||||
OriginRole: store.RoleFromContext(ctx),
|
||||
OriginSessionKey: ToolSessionKeyFromCtx(ctx),
|
||||
OriginAgentID: store.AgentIDFromContext(ctx),
|
||||
OriginTenantID: store.TenantIDFromContext(ctx),
|
||||
OriginTraceID: tracing.TraceIDFromContext(ctx),
|
||||
OriginRootSpanID: tracing.ParentSpanIDFromContext(ctx),
|
||||
OriginContextWindow: store.AgentContextWindowFromContext(ctx),
|
||||
OriginMaxTokens: store.AgentMaxTokensFromContext(ctx),
|
||||
CreatedAt: time.Now().UnixMilli(),
|
||||
spawnConfig: cfg,
|
||||
}
|
||||
if sm.taskStore != nil {
|
||||
subTask.dbID = store.GenNewID()
|
||||
|
||||
@@ -34,6 +34,10 @@ func TestRunSyncHonorsPerTaskModelOverride(t *testing.T) {
|
||||
|
||||
ctx := store.WithTenantID(context.Background(), uuid.New())
|
||||
ctx = WithParentModel(ctx, "parent-model")
|
||||
// The subagent's internal LLM call is agent-scoped; the guard requires a
|
||||
// budget in ctx (propagated from the calling agent via injectContext).
|
||||
ctx = store.WithAgentContextWindow(ctx, 200_000)
|
||||
ctx = store.WithAgentMaxTokens(ctx, 32_000)
|
||||
|
||||
result, _, err := manager.RunSync(ctx, "parent", 0, "test task", "test", "requested-model", "test", "chat")
|
||||
if err != nil {
|
||||
|
||||
+121
-13
@@ -2,7 +2,9 @@ package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
@@ -10,7 +12,42 @@ import (
|
||||
usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps"
|
||||
)
|
||||
|
||||
func agentBudgetFromContext(ctx context.Context) usagecaps.AgentBudget {
|
||||
return usagecaps.AgentBudget{
|
||||
ContextWindow: store.AgentContextWindowFromContext(ctx),
|
||||
MaxTokens: store.AgentMaxTokensFromContext(ctx),
|
||||
}
|
||||
}
|
||||
|
||||
// reserveToolLLMUsage guards and reserves a tool-internal model call whose full
|
||||
// input already lives in the ChatRequest (text plus the fixed inline-media
|
||||
// convention). The fixed local BudgetCounter counts the request itself.
|
||||
func reserveToolLLMUsage(ctx context.Context, svc *usagecaps.Service, toolName, providerName, model string, req providers.ChatRequest) (*usagecaps.Reservation, error) {
|
||||
return reserveToolLLMUsageWithMedia(ctx, svc, toolName, providerName, model, req, "", nil)
|
||||
}
|
||||
|
||||
// reserveToolLLMUsageWithMedia guards a native-media tool call whose payload is
|
||||
// sent OUT-OF-BAND (native provider JSON body or File API upload) and is thus
|
||||
// invisible to the ChatRequest the counter would otherwise see. To keep the
|
||||
// budget authority model/provider-independent, the fixed local BudgetCounter
|
||||
// counts the standard-base64 representation of the raw bytes as one extra
|
||||
// synthetic input message. That synthetic message is added ONLY to the
|
||||
// guard/reservation copy — it is never transported, because native paths build
|
||||
// their own provider payload from the raw bytes separately.
|
||||
//
|
||||
// mediaData must be the real bytes that will be sent. A native-media call whose
|
||||
// payload cannot be buffered here (a streamed remote URL) must NOT route through
|
||||
// this helper with nil data — use reserveToolLLMUsageUnverifiableMedia, which
|
||||
// fails closed under an agent budget instead of undercounting.
|
||||
func reserveToolLLMUsageWithMedia(ctx context.Context, svc *usagecaps.Service, toolName, providerName, model string, req providers.ChatRequest, mediaMIME string, mediaData []byte) (*usagecaps.Reservation, error) {
|
||||
if mediaMIME != "" || len(mediaData) > 0 {
|
||||
req = appendNativeMediaBudget(req, mediaMIME, mediaData)
|
||||
}
|
||||
budget := agentBudgetFromContext(ctx)
|
||||
req = clampToolRequestMaxTokens(req, budget.MaxTokens)
|
||||
if guardErr := usagecaps.GuardContextWindow(req, providerName, model, "tool:"+toolName, budget); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
if svc == nil {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -21,24 +58,95 @@ func reserveToolLLMUsage(ctx context.Context, svc *usagecaps.Service, toolName,
|
||||
ModelID: model,
|
||||
ReservationKey: fmt.Sprintf("tool:%s:%s", toolName, uuid.NewString()),
|
||||
Messages: req.Messages,
|
||||
MaxOutputTokens: maxOutputTokensFromOptions(req.Options),
|
||||
MaxOutputTokens: budget.MaxTokens,
|
||||
})
|
||||
}
|
||||
|
||||
func maxOutputTokensFromOptions(options map[string]any) int {
|
||||
maxTokens := 1024
|
||||
// reserveToolLLMUsageUnverifiableMedia handles a native-media call whose payload
|
||||
// cannot be counted before transport (e.g. a remote video streamed straight to
|
||||
// the provider without buffering). The complete-input invariant cannot be proven
|
||||
// for such a call, so under an agent budget it fails closed with an explicit
|
||||
// streaming error rather than trusting a byte-size/Content-Length estimate.
|
||||
// Without a propagated agent budget it falls through to the text path, which
|
||||
// itself fails closed with an AgentBudgetWiringError — every tool LLM call is
|
||||
// agent-scoped, so there is no path here that reaches transport uncounted.
|
||||
func reserveToolLLMUsageUnverifiableMedia(ctx context.Context, svc *usagecaps.Service, toolName, providerName, model string, req providers.ChatRequest) (*usagecaps.Reservation, error) {
|
||||
budget := agentBudgetFromContext(ctx)
|
||||
if budget.ContextWindow > 0 || budget.MaxTokens > 0 {
|
||||
return nil, fmt.Errorf(
|
||||
"tool:%s: cannot verify streamed native media against the agent context budget (no in-memory payload to count); refusing to send",
|
||||
toolName,
|
||||
)
|
||||
}
|
||||
return reserveToolLLMUsage(ctx, svc, toolName, providerName, model, req)
|
||||
}
|
||||
|
||||
// appendNativeMediaBudget returns a copy of req with one extra synthetic user
|
||||
// message carrying the media MIME and the standard-base64 encoding of the raw
|
||||
// payload, so the fixed BudgetCounter counts the real out-of-band input. The
|
||||
// original message slice is not mutated; the returned request is guard-only.
|
||||
func appendNativeMediaBudget(req providers.ChatRequest, mime string, data []byte) providers.ChatRequest {
|
||||
var b strings.Builder
|
||||
b.WriteString(mime)
|
||||
if len(data) > 0 {
|
||||
b.WriteByte('\n')
|
||||
b.WriteString(base64.StdEncoding.EncodeToString(data))
|
||||
}
|
||||
msgs := make([]providers.Message, len(req.Messages), len(req.Messages)+1)
|
||||
copy(msgs, req.Messages)
|
||||
msgs = append(msgs, providers.Message{Role: "user", Content: b.String()})
|
||||
req.Messages = msgs
|
||||
return req
|
||||
}
|
||||
|
||||
// clampToolRequestMaxTokens enforces max_tokens <= agentMaxTokens for every
|
||||
// agent-originated tool call. A request that does not declare max_tokens is set
|
||||
// to the agent's max_tokens rather than left to a provider default, so the
|
||||
// window invariant (completeInput + agentMaxTokens <= window) holds for the
|
||||
// value actually sent.
|
||||
func clampToolRequestMaxTokens(req providers.ChatRequest, agentMaxTokens int) providers.ChatRequest {
|
||||
if agentMaxTokens <= 0 {
|
||||
return req
|
||||
}
|
||||
if req.Options == nil {
|
||||
req.Options = map[string]any{}
|
||||
}
|
||||
current, ok := maxOutputTokensDeclared(req.Options)
|
||||
if !ok || current > agentMaxTokens {
|
||||
req.Options[providers.OptMaxTokens] = agentMaxTokens
|
||||
}
|
||||
return req
|
||||
}
|
||||
|
||||
// maxOutputTokensDeclared reports the max_tokens declared in the request options
|
||||
// and whether it was present at all. The bool distinguishes a genuinely missing
|
||||
// option (ok == false) from an explicit zero, which the clamp needs so it can
|
||||
// set the agent's max_tokens when the caller declared nothing.
|
||||
func maxOutputTokensDeclared(options map[string]any) (int, bool) {
|
||||
if options == nil {
|
||||
return maxTokens
|
||||
return 0, false
|
||||
}
|
||||
if v, ok := options[providers.OptMaxTokens]; ok {
|
||||
switch n := v.(type) {
|
||||
case int:
|
||||
maxTokens = n
|
||||
case int64:
|
||||
maxTokens = int(n)
|
||||
case float64:
|
||||
maxTokens = int(n)
|
||||
}
|
||||
v, ok := options[providers.OptMaxTokens]
|
||||
if !ok {
|
||||
return 0, false
|
||||
}
|
||||
switch n := v.(type) {
|
||||
case int:
|
||||
return n, true
|
||||
case int64:
|
||||
return int(n), true
|
||||
case int32:
|
||||
return int(n), true
|
||||
case float64:
|
||||
return int(n), true
|
||||
case float32:
|
||||
return int(n), true
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func maxOutputTokensFromOptions(options map[string]any) int {
|
||||
maxTokens, _ := maxOutputTokensDeclared(options)
|
||||
return maxTokens
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps"
|
||||
)
|
||||
|
||||
// bigContent builds a user message string that exceeds `targetChars` words under
|
||||
// the caps guard's fixed BudgetCounter, so the guard math is predictable here.
|
||||
func bigContent(targetChars int) string {
|
||||
return strings.Repeat("word ", targetChars)
|
||||
}
|
||||
|
||||
// withToolAgentBudget mirrors what the agent loop's injectContext does before a
|
||||
// tool runs: it propagates the CALLING agent's context window and max_tokens.
|
||||
func withToolAgentBudget(window, maxTokens int) context.Context {
|
||||
ctx := store.WithAgentContextWindow(context.Background(), window)
|
||||
return store.WithAgentMaxTokens(ctx, maxTokens)
|
||||
}
|
||||
|
||||
// TestReserveToolLLMUsage_HonorsAgentWindowFromContext is the production-path
|
||||
// regression guard: reserveToolLLMUsage reads the CALLING agent's budget from
|
||||
// ctx (set by injectContext via store.WithAgentContextWindow /
|
||||
// store.WithAgentMaxTokens) and enforces completeInput + max_tokens <= window.
|
||||
// Model/provider are never budget authorities.
|
||||
func TestReserveToolLLMUsage_HonorsAgentWindowFromContext(t *testing.T) {
|
||||
model := "claude-sonnet-4-5-20250929" // 200k model window — must NOT matter
|
||||
req := providers.ChatRequest{
|
||||
Model: model,
|
||||
Messages: []providers.Message{{Role: "user", Content: bigContent(40_000)}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
|
||||
// A 200k agent window admits ~40k tokens of prompt.
|
||||
ctxBig := withToolAgentBudget(200_000, 8_192)
|
||||
if _, err := reserveToolLLMUsage(ctxBig, nil, "read_document", "anthropic", model, req); err != nil {
|
||||
t.Fatalf("expected allow under 200k agent window, got %v", err)
|
||||
}
|
||||
|
||||
// A 20k agent window must block the SAME request before transport.
|
||||
ctxSmall := withToolAgentBudget(20_000, 8_192)
|
||||
_, err := reserveToolLLMUsage(ctxSmall, nil, "read_document", "anthropic", model, req)
|
||||
if err == nil {
|
||||
t.Fatal("expected abort when agent window (20k) is below the request size")
|
||||
}
|
||||
var ctxErr *usagecaps.ContextWindowExceededError
|
||||
if !errors.As(err, &ctxErr) {
|
||||
t.Fatalf("expected *ContextWindowExceededError, got %T: %v", err, err)
|
||||
}
|
||||
if ctxErr.ContextWindow != 20_000 {
|
||||
t.Fatalf("guard used window %d, want 20000 (agent cap, not model window)", ctxErr.ContextWindow)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReserveToolLLMUsage_FailsClosedWithoutAgentBudget proves an agent-scoped
|
||||
// tool call reaching the model gate WITHOUT a propagated budget fails closed
|
||||
// with a wiring error before any transport — it must never silently guess a
|
||||
// model window.
|
||||
func TestReserveToolLLMUsage_FailsClosedWithoutAgentBudget(t *testing.T) {
|
||||
req := providers.ChatRequest{
|
||||
Model: "gpt-4o",
|
||||
Messages: []providers.Message{{Role: "user", Content: "tiny prompt"}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
_, err := reserveToolLLMUsage(context.Background(), nil, "read_document", "openai", "gpt-4o", req)
|
||||
if err == nil {
|
||||
t.Fatal("expected wiring error without a propagated agent budget")
|
||||
}
|
||||
var wiringErr *usagecaps.AgentBudgetWiringError
|
||||
if !errors.As(err, &wiringErr) {
|
||||
t.Fatalf("expected *AgentBudgetWiringError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReserveToolLLMUsageWithMedia_CountsNativeMediaBytes proves the out-of-band
|
||||
// media payload is part of completeInput. The fixed BudgetCounter counts the
|
||||
// standard-base64 representation of the raw bytes (a model/provider-independent
|
||||
// rule), so a small prompt with a large native-media payload is blocked when it
|
||||
// no longer fits the CALLING agent's window. The bytes are counted, never
|
||||
// transported — the guard copy carries them, the real request does not.
|
||||
func TestReserveToolLLMUsageWithMedia_CountsNativeMediaBytes(t *testing.T) {
|
||||
model := "gpt-4o"
|
||||
req := func() providers.ChatRequest {
|
||||
return providers.ChatRequest{
|
||||
Model: model,
|
||||
Messages: []providers.Message{{Role: "user", Content: "Transcribe this."}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
}
|
||||
|
||||
// A small media payload under a large agent window is allowed.
|
||||
small := bytes.Repeat([]byte{0xAB}, 1024)
|
||||
ctxBig := withToolAgentBudget(128_000, 8_192)
|
||||
if _, err := reserveToolLLMUsageWithMedia(ctxBig, nil, "read_document", "openai", model, req(), "application/pdf", small); err != nil {
|
||||
t.Fatalf("small native media must fit a 128k window, got %v", err)
|
||||
}
|
||||
|
||||
// A large media payload against a small agent window is blocked BEFORE
|
||||
// transport: base64(256 KiB) is well over the 20k window's input cap.
|
||||
big := bytes.Repeat([]byte{0xCD}, 256*1024)
|
||||
ctxSmall := withToolAgentBudget(20_000, 8_192)
|
||||
_, err := reserveToolLLMUsageWithMedia(ctxSmall, nil, "read_document", "openai", model, req(), "application/pdf", big)
|
||||
if err == nil {
|
||||
t.Fatal("expected abort: large native media must exceed a 20k agent window")
|
||||
}
|
||||
var ctxErr *usagecaps.ContextWindowExceededError
|
||||
if !errors.As(err, &ctxErr) {
|
||||
t.Fatalf("expected *ContextWindowExceededError, got %T: %v", err, err)
|
||||
}
|
||||
if ctxErr.ContextWindow != 20_000 {
|
||||
t.Fatalf("guard used window %d, want 20000 (agent cap)", ctxErr.ContextWindow)
|
||||
}
|
||||
}
|
||||
|
||||
// TestReserveToolLLMUsageUnverifiableMedia_FailsClosedUnderAgentBudget proves a
|
||||
// native-media call whose payload cannot be buffered (a streamed remote URL)
|
||||
// fails closed under an agent budget rather than undercounting — the
|
||||
// complete-input invariant cannot be proven, so it refuses to send.
|
||||
func TestReserveToolLLMUsageUnverifiableMedia_FailsClosedUnderAgentBudget(t *testing.T) {
|
||||
model := "gemini-2.0-flash"
|
||||
req := providers.ChatRequest{
|
||||
Model: model,
|
||||
Messages: []providers.Message{{Role: "user", Content: "Analyze this video."}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
|
||||
// Under an agent budget: no in-memory payload to count -> refuse with an
|
||||
// explicit streaming error (not a wiring error).
|
||||
ctx := withToolAgentBudget(200_000, 8_192)
|
||||
_, err := reserveToolLLMUsageUnverifiableMedia(ctx, nil, "read_video", "gemini", model, req)
|
||||
if err == nil {
|
||||
t.Fatal("expected fail-closed for unverifiable streamed media under an agent budget")
|
||||
}
|
||||
var wiringErr *usagecaps.AgentBudgetWiringError
|
||||
if errors.As(err, &wiringErr) {
|
||||
t.Fatalf("under a valid agent budget the failure must be the streaming refusal, not a wiring error: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "refusing to send") {
|
||||
t.Fatalf("expected an explicit streaming refusal, got %v", err)
|
||||
}
|
||||
|
||||
// Without a propagated agent budget it falls through to the text path, which
|
||||
// itself fails closed with a wiring error — every tool LLM call is
|
||||
// agent-scoped, so no path here reaches transport uncounted.
|
||||
_, err = reserveToolLLMUsageUnverifiableMedia(context.Background(), nil, "read_video", "gemini", model, req)
|
||||
if err == nil {
|
||||
t.Fatal("expected fail-closed without a propagated agent budget")
|
||||
}
|
||||
if !errors.As(err, &wiringErr) {
|
||||
t.Fatalf("expected *AgentBudgetWiringError without a budget, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
package caps
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
)
|
||||
|
||||
// A background/utility LLM call (carries an agent ID but no wired per-agent
|
||||
// budget — e.g. vault.classify / vault.batch_summarize) must NOT fail closed
|
||||
// with an AgentBudgetWiringError. It should fall through to a safe default
|
||||
// budget and reach the provider. Regression guard for the vault-enrichment
|
||||
// break introduced when the agent-only guard was added.
|
||||
func TestServiceChat_BackgroundBudgetFallback_ReachesProvider(t *testing.T) {
|
||||
// nil store service: falls back to a direct provider.Chat after the guard,
|
||||
// isolating the budget-wiring behaviour from usage-cap policy plumbing.
|
||||
var svc *Service
|
||||
provider := &fakeChatProvider{name: "openrouter", model: "token/model"}
|
||||
|
||||
ctx := store.WithAgentID(context.Background(), uuid.New()) // agent-scoped, but NO window/max_tokens wired
|
||||
_, err := svc.Chat(ctx, provider, providers.ChatRequest{
|
||||
Model: "token/model",
|
||||
Messages: []providers.Message{{Role: "user", Content: "classify this short doc"}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}, ChatOptions{
|
||||
AgentID: uuid.New(),
|
||||
ProviderName: "openrouter",
|
||||
Purpose: "vault.classify",
|
||||
MaxOutputTokens: 4096,
|
||||
})
|
||||
if err != nil {
|
||||
var wiringErr *AgentBudgetWiringError
|
||||
if errors.As(err, &wiringErr) {
|
||||
t.Fatalf("background call must not fail closed with a wiring error: %v", err)
|
||||
}
|
||||
t.Fatalf("background call returned unexpected error: %v", err)
|
||||
}
|
||||
if provider.calls != 1 {
|
||||
t.Fatalf("provider calls = %d, want 1 (background call must reach transport)", provider.calls)
|
||||
}
|
||||
}
|
||||
|
||||
// The fallback is a SAFE default, not a bypass: a background call whose real
|
||||
// input exceeds the default window is still aborted before transport by the
|
||||
// window guard. Proves fallback != "guess a model window and wave it through".
|
||||
func TestServiceChat_BackgroundBudgetFallback_StillGuardsOversizedInput(t *testing.T) {
|
||||
var svc *Service
|
||||
provider := &fakeChatProvider{name: "openrouter", model: "token/model"}
|
||||
|
||||
// Force a tiny window via opts so the guard has a small ceiling; leave
|
||||
// max_tokens unwired so the fallback fills only that half.
|
||||
huge := strings.Repeat("token ", 60_000) // well over a small window once counted
|
||||
ctx := store.WithAgentID(context.Background(), uuid.New())
|
||||
_, err := svc.Chat(ctx, provider, providers.ChatRequest{
|
||||
Model: "token/model",
|
||||
Messages: []providers.Message{{Role: "user", Content: huge}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 1024},
|
||||
}, ChatOptions{
|
||||
AgentID: uuid.New(),
|
||||
ProviderName: "openrouter",
|
||||
Purpose: "vault.batch_summarize",
|
||||
AgentContextWindow: 8_000, // operator-provided window is preserved by the fallback
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected the window guard to abort an oversized background request")
|
||||
}
|
||||
var ctxErr *ContextWindowExceededError
|
||||
if !errors.As(err, &ctxErr) {
|
||||
t.Fatalf("expected *ContextWindowExceededError, got %T: %v", err, err)
|
||||
}
|
||||
if ctxErr.ContextWindow != 8_000 {
|
||||
t.Fatalf("guard used window %d, want 8000 (operator value preserved, not defaulted)", ctxErr.ContextWindow)
|
||||
}
|
||||
if provider.calls != 0 {
|
||||
t.Fatalf("provider calls = %d, want 0 (oversized request must not reach transport)", provider.calls)
|
||||
}
|
||||
}
|
||||
|
||||
// A non-agent-scoped call (no agent ID anywhere) must be untouched by the
|
||||
// fallback path — it is neither guarded nor defaulted, matching prior behaviour.
|
||||
func TestServiceChat_NonAgentScoped_NoFallbackNoGuard(t *testing.T) {
|
||||
var svc *Service
|
||||
provider := &fakeChatProvider{name: "openrouter", model: "token/model"}
|
||||
|
||||
huge := strings.Repeat("token ", 60_000)
|
||||
_, err := svc.Chat(context.Background(), provider, providers.ChatRequest{
|
||||
Model: "token/model",
|
||||
Messages: []providers.Message{{Role: "user", Content: huge}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 1024},
|
||||
}, ChatOptions{
|
||||
ProviderName: "openrouter",
|
||||
Purpose: "system.utility",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("non-agent-scoped call must pass through untouched, got %v", err)
|
||||
}
|
||||
if provider.calls != 1 {
|
||||
t.Fatalf("provider calls = %d, want 1", provider.calls)
|
||||
}
|
||||
}
|
||||
|
||||
// backgroundBudgetFallback fills only the missing halves and preserves any
|
||||
// operator-provided value; it never lets the reserve meet/exceed the window.
|
||||
func TestBackgroundBudgetFallback_FillsOnlyMissing(t *testing.T) {
|
||||
// Both missing: window defaulted, max_tokens taken from opts.MaxOutputTokens.
|
||||
got := backgroundBudgetFallback("vault.classify", AgentBudget{}, ChatOptions{MaxOutputTokens: 4096})
|
||||
if got.ContextWindow != defaultBackgroundContextWindow {
|
||||
t.Errorf("ContextWindow = %d, want %d", got.ContextWindow, defaultBackgroundContextWindow)
|
||||
}
|
||||
if got.MaxTokens != 4096 {
|
||||
t.Errorf("MaxTokens = %d, want 4096 (from opts.MaxOutputTokens)", got.MaxTokens)
|
||||
}
|
||||
|
||||
// opts has no max_tokens either: falls back to the package default.
|
||||
got = backgroundBudgetFallback("vault.classify", AgentBudget{}, ChatOptions{})
|
||||
if got.MaxTokens != defaultBackgroundMaxTokens {
|
||||
t.Errorf("MaxTokens = %d, want %d (package default)", got.MaxTokens, defaultBackgroundMaxTokens)
|
||||
}
|
||||
|
||||
// Operator window preserved; only max_tokens defaulted.
|
||||
got = backgroundBudgetFallback("vault.classify", AgentBudget{ContextWindow: 32_000}, ChatOptions{MaxOutputTokens: 2048})
|
||||
if got.ContextWindow != 32_000 {
|
||||
t.Errorf("ContextWindow = %d, want 32000 (preserved)", got.ContextWindow)
|
||||
}
|
||||
if got.MaxTokens != 2048 {
|
||||
t.Errorf("MaxTokens = %d, want 2048", got.MaxTokens)
|
||||
}
|
||||
|
||||
// Reserve must never meet/exceed the window.
|
||||
got = backgroundBudgetFallback("vault.classify", AgentBudget{ContextWindow: 1000, MaxTokens: 4000}, ChatOptions{})
|
||||
if got.MaxTokens >= got.ContextWindow {
|
||||
t.Errorf("MaxTokens %d must be < ContextWindow %d", got.MaxTokens, got.ContextWindow)
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
@@ -19,6 +20,10 @@ type ChatOptions struct {
|
||||
ReservationKey string
|
||||
Purpose string
|
||||
MaxOutputTokens int
|
||||
// AgentContextWindow and AgentMaxTokens are the caller agent's configured
|
||||
// request budget. Both are required for an agent-scoped call.
|
||||
AgentContextWindow int
|
||||
AgentMaxTokens int
|
||||
}
|
||||
|
||||
// Chat wraps Provider.Chat with the same usage-cap preflight and reconciliation
|
||||
@@ -28,7 +33,44 @@ func (s *Service) Chat(ctx context.Context, provider providers.Provider, req pro
|
||||
if provider == nil {
|
||||
return nil, errors.New("usage cap chat: provider is nil")
|
||||
}
|
||||
agentID := opts.AgentID
|
||||
if agentID == uuid.Nil {
|
||||
agentID = store.AgentIDFromContext(ctx)
|
||||
}
|
||||
budget := AgentBudget{
|
||||
ContextWindow: opts.AgentContextWindow,
|
||||
MaxTokens: opts.AgentMaxTokens,
|
||||
}
|
||||
if budget.ContextWindow <= 0 {
|
||||
budget.ContextWindow = store.AgentContextWindowFromContext(ctx)
|
||||
}
|
||||
if budget.MaxTokens <= 0 {
|
||||
budget.MaxTokens = store.AgentMaxTokensFromContext(ctx)
|
||||
}
|
||||
agentScoped := agentID != uuid.Nil
|
||||
if agentScoped && !budget.valid() {
|
||||
// A call carrying an agent ID but no wired request budget is a
|
||||
// background/utility LLM call (e.g. vault.classify, vault.batch_summarize)
|
||||
// that runs OUTSIDE an agent turn, where no per-agent window/max_tokens is
|
||||
// propagated. Rather than fail closed — which silently degrades enrichment
|
||||
// to the extractive fallback — fill a safe default budget so the window
|
||||
// guard still protects the provider transport. Logged at WARN so a genuine
|
||||
// agent-run wiring regression remains observable instead of masked.
|
||||
budget = backgroundBudgetFallback(opts.Purpose, budget, opts)
|
||||
}
|
||||
if agentScoped {
|
||||
req = clampRequestMaxTokens(req, budget.MaxTokens)
|
||||
}
|
||||
if s == nil || s.store == nil {
|
||||
guardName := opts.ProviderName
|
||||
if guardName == "" {
|
||||
guardName = provider.Name()
|
||||
}
|
||||
if agentScoped {
|
||||
if guardErr := GuardContextWindow(req, guardName, req.Model, opts.Purpose, budget); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
}
|
||||
return provider.Chat(ctx, req)
|
||||
}
|
||||
if fallback, ok := provider.(*providers.ModelFallbackProvider); ok {
|
||||
@@ -40,6 +82,12 @@ func (s *Service) Chat(ctx context.Context, provider providers.Provider, req pro
|
||||
}
|
||||
callOpts.ModelID = actualReq.Model
|
||||
callOpts.ReservationKey = ""
|
||||
if agentScoped {
|
||||
actualReq = clampRequestMaxTokens(actualReq, budget.MaxTokens)
|
||||
if guardErr := GuardContextWindow(actualReq, callOpts.ProviderName, actualReq.Model, opts.Purpose, budget); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
}
|
||||
usageReq := s.chatRequest(callCtx, entry.Provider, actualReq, callOpts)
|
||||
scopedCtx := scopedRequestContext(callCtx, usageReq)
|
||||
reservation, err := s.Preflight(scopedCtx, usageReq)
|
||||
@@ -54,6 +102,15 @@ func (s *Service) Chat(ctx context.Context, provider providers.Provider, req pro
|
||||
})
|
||||
}
|
||||
|
||||
guardName := opts.ProviderName
|
||||
if guardName == "" {
|
||||
guardName = provider.Name()
|
||||
}
|
||||
if agentScoped {
|
||||
if guardErr := GuardContextWindow(req, guardName, req.Model, opts.Purpose, budget); guardErr != nil {
|
||||
return nil, guardErr
|
||||
}
|
||||
}
|
||||
usageReq := s.chatRequest(ctx, provider, req, opts)
|
||||
scopedCtx := scopedRequestContext(ctx, usageReq)
|
||||
reservation, err := s.Preflight(scopedCtx, usageReq)
|
||||
@@ -122,6 +179,91 @@ func reservationKey(opts ChatOptions) string {
|
||||
return fmt.Sprintf("%s:%s", purpose, uuid.NewString())
|
||||
}
|
||||
|
||||
// defaultBackgroundContextWindow bounds background/utility LLM calls that carry
|
||||
// an agent ID but no wired per-agent budget. It is intentionally conservative:
|
||||
// large enough for enrichment prompts (classify/summarize send truncated docs),
|
||||
// small enough that a runaway prompt is still caught by the window guard.
|
||||
const defaultBackgroundContextWindow = 200_000
|
||||
|
||||
// defaultBackgroundMaxTokens is the output reserve used when a background call
|
||||
// declares no max_tokens of its own.
|
||||
const defaultBackgroundMaxTokens = 4_096
|
||||
|
||||
// backgroundBudgetFallback fills a safe budget for an agent-scoped call whose
|
||||
// window/max_tokens were never wired (a background/utility call outside an agent
|
||||
// turn). Any operator-provided value is preserved; only the missing halves are
|
||||
// defaulted. The window guard still runs against the result, so an oversized
|
||||
// request fails closed on real input rather than on missing wiring.
|
||||
func backgroundBudgetFallback(purpose string, budget AgentBudget, opts ChatOptions) AgentBudget {
|
||||
out := budget
|
||||
if out.MaxTokens <= 0 {
|
||||
out.MaxTokens = opts.MaxOutputTokens
|
||||
}
|
||||
if out.MaxTokens <= 0 {
|
||||
out.MaxTokens = defaultBackgroundMaxTokens
|
||||
}
|
||||
if out.ContextWindow <= 0 {
|
||||
out.ContextWindow = defaultBackgroundContextWindow
|
||||
}
|
||||
// Keep the invariant satisfiable: never let the reserve meet/exceed the window.
|
||||
if out.MaxTokens >= out.ContextWindow {
|
||||
out.MaxTokens = out.ContextWindow / 2
|
||||
}
|
||||
slog.Warn("caps.background_budget_default",
|
||||
"purpose", purpose,
|
||||
"context_window", out.ContextWindow,
|
||||
"max_tokens", out.MaxTokens,
|
||||
"reason", "agent-scoped call missing wired budget; using safe default instead of fail-closed",
|
||||
)
|
||||
return out
|
||||
}
|
||||
|
||||
func clampRequestMaxTokens(req providers.ChatRequest, agentMaxTokens int) providers.ChatRequest {
|
||||
if agentMaxTokens <= 0 {
|
||||
return req
|
||||
}
|
||||
if req.Options == nil {
|
||||
req.Options = map[string]any{}
|
||||
}
|
||||
// A request that declares no max_tokens is SET to the agent's max_tokens
|
||||
// rather than left to a provider default, so the window invariant holds for
|
||||
// the value actually sent. A declared value is clamped down when it exceeds
|
||||
// the agent cap (or is a non-positive explicit value).
|
||||
current, ok := maxOutputTokensDeclared(req.Options)
|
||||
if !ok || current <= 0 || current > agentMaxTokens {
|
||||
req.Options[providers.OptMaxTokens] = agentMaxTokens
|
||||
}
|
||||
return req
|
||||
}
|
||||
|
||||
// maxOutputTokensDeclared reports the declared max_tokens and whether the option
|
||||
// was present at all. The bool lets the clamp distinguish a genuinely missing
|
||||
// option from an explicit value, so a caller that declared nothing is pinned to
|
||||
// the agent's max_tokens.
|
||||
func maxOutputTokensDeclared(options map[string]any) (int, bool) {
|
||||
if options == nil {
|
||||
return 0, false
|
||||
}
|
||||
v, ok := options[providers.OptMaxTokens]
|
||||
if !ok {
|
||||
return 0, false
|
||||
}
|
||||
switch n := v.(type) {
|
||||
case int:
|
||||
return n, true
|
||||
case int64:
|
||||
return int(n), true
|
||||
case int32:
|
||||
return int(n), true
|
||||
case float64:
|
||||
return int(n), true
|
||||
case float32:
|
||||
return int(n), true
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
|
||||
func maxOutputTokens(req providers.ChatRequest, fallback int) int {
|
||||
if fallback <= 0 {
|
||||
fallback = 1024
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
package caps
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/tokencount"
|
||||
)
|
||||
|
||||
var contextGuardCounter tokencount.BudgetCounter = tokencount.NewBudgetCounter()
|
||||
|
||||
// AgentBudget is the complete request budget configured for one agent.
|
||||
// Model and provider are intentionally absent: they are not budget authorities.
|
||||
type AgentBudget struct {
|
||||
ContextWindow int
|
||||
MaxTokens int
|
||||
}
|
||||
|
||||
func (b AgentBudget) valid() bool { return b.ContextWindow > 0 && b.MaxTokens > 0 }
|
||||
|
||||
// AgentBudgetWiringError means an agent-scoped call reached a model-call gate
|
||||
// without the configured window and max_tokens that must have been propagated.
|
||||
type AgentBudgetWiringError struct {
|
||||
Purpose string
|
||||
Missing string
|
||||
}
|
||||
|
||||
func (e *AgentBudgetWiringError) Error() string {
|
||||
return fmt.Sprintf("agent budget wiring error: missing %s (purpose=%s)", e.Missing, e.Purpose)
|
||||
}
|
||||
|
||||
// ContextWindowExceededError is returned before provider transport when a
|
||||
// complete agent request cannot fit its configured budget.
|
||||
type ContextWindowExceededError struct {
|
||||
Provider string
|
||||
Model string
|
||||
Purpose string
|
||||
InputTokens int
|
||||
OutputReserve int
|
||||
ContextWindow int
|
||||
}
|
||||
|
||||
func (e *ContextWindowExceededError) Error() string {
|
||||
return fmt.Sprintf(
|
||||
"context window exceeded: input=%d + output_reserve=%d > window=%d (provider=%s model=%s purpose=%s)",
|
||||
e.InputTokens, e.OutputReserve, e.ContextWindow, e.Provider, e.Model, e.Purpose,
|
||||
)
|
||||
}
|
||||
|
||||
func (e *ContextWindowExceededError) ContextBudgetExceeded() bool { return true }
|
||||
|
||||
// GuardContextWindow enforces completeInput + agentMaxTokens <= agentWindow.
|
||||
// Provider/model are retained only for diagnostics.
|
||||
func GuardContextWindow(req providers.ChatRequest, providerName, model, purpose string, budget AgentBudget) error {
|
||||
if !budget.valid() {
|
||||
missing := "agent_context_window,agent_max_tokens"
|
||||
switch {
|
||||
case budget.ContextWindow > 0:
|
||||
missing = "agent_max_tokens"
|
||||
case budget.MaxTokens > 0:
|
||||
missing = "agent_context_window"
|
||||
}
|
||||
return &AgentBudgetWiringError{Purpose: purpose, Missing: missing}
|
||||
}
|
||||
input, err := contextGuardCounter.CountRequest(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("count complete request: %w", err)
|
||||
}
|
||||
if input+budget.MaxTokens <= budget.ContextWindow {
|
||||
return nil
|
||||
}
|
||||
slog.Info("caps.context_guard",
|
||||
"provider", providerName,
|
||||
"model", model,
|
||||
"purpose", purpose,
|
||||
"input_tokens", input,
|
||||
"output_reserve_tokens", budget.MaxTokens,
|
||||
"context_window", budget.ContextWindow,
|
||||
"action", "abort",
|
||||
)
|
||||
return &ContextWindowExceededError{
|
||||
Provider: providerName,
|
||||
Model: model,
|
||||
Purpose: purpose,
|
||||
InputTokens: input,
|
||||
OutputReserve: budget.MaxTokens,
|
||||
ContextWindow: budget.ContextWindow,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package caps
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/providers"
|
||||
)
|
||||
|
||||
// stubGuardProvider records whether Chat was invoked and returns a canned reply.
|
||||
type stubGuardProvider struct {
|
||||
called bool
|
||||
}
|
||||
|
||||
func (p *stubGuardProvider) Name() string { return "stubguard" }
|
||||
func (p *stubGuardProvider) DefaultModel() string { return "claude-3-5-sonnet" }
|
||||
func (p *stubGuardProvider) Chat(_ context.Context, _ providers.ChatRequest) (*providers.ChatResponse, error) {
|
||||
p.called = true
|
||||
return &providers.ChatResponse{Content: "ok"}, nil
|
||||
}
|
||||
func (p *stubGuardProvider) ChatStream(_ context.Context, _ providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) {
|
||||
p.called = true
|
||||
return &providers.ChatResponse{Content: "ok"}, nil
|
||||
}
|
||||
|
||||
// bigMessage returns a user message whose complete-input count is at least
|
||||
// target tokens, using the fixed model/provider-independent budget counter.
|
||||
func bigMessage(t *testing.T, target int) providers.Message {
|
||||
t.Helper()
|
||||
low, high := 1, target*8
|
||||
for low < high {
|
||||
mid := low + (high-low)/2
|
||||
msg := providers.Message{Role: "user", Content: strings.Repeat("word ", mid)}
|
||||
count, err := contextGuardCounter.CountMessages([]providers.Message{msg})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if count < target {
|
||||
low = mid + 1
|
||||
} else {
|
||||
high = mid
|
||||
}
|
||||
}
|
||||
return providers.Message{Role: "user", Content: strings.Repeat("word ", low)}
|
||||
}
|
||||
|
||||
// TestCapsChat_NonAgentCallSkipsWindowGuard proves a call with no AgentID and no
|
||||
// agent budget is NOT forced against any model window: it passes straight to the
|
||||
// provider even when huge. Model window is not a budget authority.
|
||||
func TestCapsChat_NonAgentCallSkipsWindowGuard(t *testing.T) {
|
||||
var svc *Service // nil service => Lite path
|
||||
prov := &stubGuardProvider{}
|
||||
// ~200k tokens: would blow any model window, but this is not agent-scoped.
|
||||
msg := bigMessage(t, 200_000)
|
||||
req := providers.ChatRequest{
|
||||
Model: "gpt-4o",
|
||||
Messages: []providers.Message{msg},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
if _, err := svc.Chat(context.Background(), prov, req, ChatOptions{Purpose: "non-agent"}); err != nil {
|
||||
t.Fatalf("non-agent call must not be window-guarded, got %v", err)
|
||||
}
|
||||
if !prov.called {
|
||||
t.Fatal("provider.Chat should have been called for a non-agent call")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCapsChat_AgentMissingBudgetUsesSafeDefault proves an agent-scoped call via
|
||||
// Service.Chat (AgentID set) that reaches the gate WITHOUT a wired budget does
|
||||
// NOT fail closed with a wiring error. These are background/utility calls (e.g.
|
||||
// vault.classify) that run outside an agent turn; failing closed silently
|
||||
// degrades enrichment. Instead a safe default budget is filled and the call
|
||||
// reaches transport. NOTE: the agent-RUN tool path guards via GuardContextWindow
|
||||
// directly (see internal/tools/usage_caps.go) and still fails closed — this
|
||||
// relaxation is scoped to Service.Chat only.
|
||||
func TestCapsChat_AgentMissingBudgetUsesSafeDefault(t *testing.T) {
|
||||
var svc *Service
|
||||
prov := &stubGuardProvider{}
|
||||
req := providers.ChatRequest{
|
||||
Model: "gpt-4o",
|
||||
Messages: []providers.Message{{Role: "user", Content: "hello"}},
|
||||
Options: map[string]any{providers.OptMaxTokens: 1024},
|
||||
}
|
||||
_, err := svc.Chat(context.Background(), prov, req, ChatOptions{
|
||||
AgentID: uuid.New(),
|
||||
Purpose: "vault.classify",
|
||||
})
|
||||
if err != nil {
|
||||
var wiringErr *AgentBudgetWiringError
|
||||
if errors.As(err, &wiringErr) {
|
||||
t.Fatalf("background call must not fail closed with a wiring error: %v", err)
|
||||
}
|
||||
t.Fatalf("unexpected error from safe-default background call: %v", err)
|
||||
}
|
||||
if !prov.called {
|
||||
t.Fatal("provider.Chat should have been called after filling a safe default budget")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCapsChat_AgentBudgetSameRequestDiffersByWindow is the core agent-only
|
||||
// contract test: the SAME request is blocked for a 128k agent and allowed for a
|
||||
// 200k agent, decided only by the agent's configured window, never the model's.
|
||||
func TestCapsChat_AgentBudgetSameRequestDiffersByWindow(t *testing.T) {
|
||||
var svc *Service
|
||||
agentID := uuid.New()
|
||||
// ~140k tokens: fits a 200k window, exceeds a 128k window.
|
||||
msg := bigMessage(t, 140_000)
|
||||
req := providers.ChatRequest{
|
||||
Model: "gpt-4o", // 128k model window — irrelevant to the decision
|
||||
Messages: []providers.Message{msg},
|
||||
Options: map[string]any{providers.OptMaxTokens: 4096},
|
||||
}
|
||||
|
||||
// 128k agent: blocked.
|
||||
prov128 := &stubGuardProvider{}
|
||||
_, err := svc.Chat(context.Background(), prov128, req, ChatOptions{
|
||||
AgentID: agentID,
|
||||
Purpose: "cap-128k",
|
||||
AgentContextWindow: 128_000,
|
||||
AgentMaxTokens: 8_192,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("128k agent: expected abort for 140k request")
|
||||
}
|
||||
if prov128.called {
|
||||
t.Fatal("128k agent: provider must NOT be called when the guard aborts (transport calls must be 0)")
|
||||
}
|
||||
var ctxErr *ContextWindowExceededError
|
||||
if !errors.As(err, &ctxErr) {
|
||||
t.Fatalf("expected *ContextWindowExceededError, got %T: %v", err, err)
|
||||
}
|
||||
if ctxErr.ContextWindow != 128_000 {
|
||||
t.Fatalf("guard reported window %d, want 128000 (agent window, not model)", ctxErr.ContextWindow)
|
||||
}
|
||||
|
||||
// 200k agent: allowed, provider reached.
|
||||
prov200 := &stubGuardProvider{}
|
||||
if _, err := svc.Chat(context.Background(), prov200, req, ChatOptions{
|
||||
AgentID: agentID,
|
||||
Purpose: "cap-200k",
|
||||
AgentContextWindow: 200_000,
|
||||
AgentMaxTokens: 8_192,
|
||||
}); err != nil {
|
||||
t.Fatalf("200k agent: expected allow for 140k request, got %v", err)
|
||||
}
|
||||
if !prov200.called {
|
||||
t.Fatal("200k agent: provider should have been called")
|
||||
}
|
||||
}
|
||||
@@ -812,15 +812,17 @@
|
||||
},
|
||||
"compaction": {
|
||||
"title": "Compaction",
|
||||
"description": "Context window compaction and memory flush settings",
|
||||
"description": "Final request context guard, history compaction, and memory flush settings",
|
||||
"maxHistoryShare": "Max History Share (0-1)",
|
||||
"maxHistoryShareTip": "Maximum fraction of context window for conversation history (e.g. 0.85 = 85%). Compaction triggers when exceeded.",
|
||||
"maxHistoryShareTip": "History-only threshold for post-turn history compaction. This is not the final request limit.",
|
||||
"keepLastMessages": "Keep Last Messages",
|
||||
"keepLastMessagesTip": "Recent messages kept after compaction. Older messages are replaced by a summary.",
|
||||
"timeoutSeconds": "Timeout Seconds",
|
||||
"timeoutSecondsTip": "Maximum time to wait for a compaction summary before keeping the original history.",
|
||||
"memoryFlush": "Memory Flush",
|
||||
"memoryFlushTip": "Before compaction, the agent gets a turn to save important context to memory files. Also triggers Knowledge Graph extraction."
|
||||
"memoryFlushTip": "Before compaction, the agent gets a turn to save important context to memory files. Also triggers Knowledge Graph extraction.",
|
||||
"maxRequestShare": "Max Request Share (0-1)",
|
||||
"maxRequestShareTip": "Maximum fraction of the agent context window allowed for the final request sent to the model. If exceeded, the system reduces context before calling the model."
|
||||
},
|
||||
"inboundDebounce": {
|
||||
"title": "Inbound Debounce",
|
||||
|
||||
@@ -782,13 +782,15 @@
|
||||
},
|
||||
"compaction": {
|
||||
"title": "압축",
|
||||
"description": "컨텍스트 창 압축 및 메모리 플러시 설정",
|
||||
"description": "최종 요청 컨텍스트 보호, 기록 압축, 메모리 플러시 설정",
|
||||
"maxHistoryShare": "최대 기록 비율 (0-1)",
|
||||
"maxHistoryShareTip": "대화 기록을 위한 컨텍스트 창의 최대 비율 (예: 0.85 = 85%). 초과 시 압축이 트리거됩니다.",
|
||||
"maxHistoryShareTip": "턴 이후 기록 압축에만 사용하는 기록 전용 임계값입니다. 최종 요청 전체 제한이 아닙니다.",
|
||||
"keepLastMessages": "마지막 메시지 유지",
|
||||
"keepLastMessagesTip": "압축 후 유지되는 최근 메시지입니다. 오래된 메시지는 요약으로 대체됩니다.",
|
||||
"memoryFlush": "메모리 플러시",
|
||||
"memoryFlushTip": "압축 전에 에이전트가 중요한 컨텍스트를 메모리 파일에 저장할 기회를 얻습니다. 지식 그래프 추출도 트리거합니다."
|
||||
"memoryFlushTip": "압축 전에 에이전트가 중요한 컨텍스트를 메모리 파일에 저장할 기회를 얻습니다. 지식 그래프 추출도 트리거합니다.",
|
||||
"maxRequestShare": "최대 요청 비율 (0-1)",
|
||||
"maxRequestShareTip": "모델에 보내는 최종 요청이 에이전트 컨텍스트 창에서 사용할 수 있는 최대 비율입니다. 초과하면 시스템이 모델 호출 전에 컨텍스트를 줄입니다."
|
||||
},
|
||||
"contextPruning": {
|
||||
"title": "컨텍스트 정리",
|
||||
|
||||
@@ -797,15 +797,17 @@
|
||||
},
|
||||
"compaction": {
|
||||
"title": "Nén ngữ cảnh",
|
||||
"description": "Cài đặt nén cửa sổ ngữ cảnh và ghi nhớ trước nén",
|
||||
"description": "Cài đặt ngưỡng request cuối, nén lịch sử và flush trí nhớ",
|
||||
"maxHistoryShare": "Phần lịch sử tối đa (0-1)",
|
||||
"maxHistoryShareTip": "Tỷ lệ tối đa của cửa sổ ngữ cảnh dành cho lịch sử hội thoại (vd: 0.85 = 85%). Nén kích hoạt khi vượt quá.",
|
||||
"maxHistoryShareTip": "Ngưỡng chỉ dành cho nén lịch sử sau lượt chạy. Đây không phải giới hạn tổng request cuối gửi model.",
|
||||
"keepLastMessages": "Giữ tin nhắn cuối",
|
||||
"keepLastMessagesTip": "Số tin nhắn gần nhất giữ lại sau khi nén. Các tin nhắn cũ hơn được thay thế bằng bản tóm tắt.",
|
||||
"timeoutSeconds": "Thời gian chờ (giây)",
|
||||
"timeoutSecondsTip": "Thời gian tối đa chờ bản tóm tắt nén ngữ cảnh trước khi giữ nguyên lịch sử cũ.",
|
||||
"memoryFlush": "Ghi nhớ trước nén",
|
||||
"memoryFlushTip": "Trước khi nén, agent được một lượt để lưu ngữ cảnh quan trọng vào file bộ nhớ. Cũng kích hoạt trích xuất Knowledge Graph."
|
||||
"memoryFlushTip": "Trước khi nén, agent được một lượt để lưu ngữ cảnh quan trọng vào file bộ nhớ. Cũng kích hoạt trích xuất Knowledge Graph.",
|
||||
"maxRequestShare": "Phần request tối đa (0-1)",
|
||||
"maxRequestShareTip": "Tỷ lệ tối đa của cửa sổ ngữ cảnh agent được phép dùng cho request cuối gửi model. Nếu vượt, hệ thống sẽ giảm context trước khi gọi model."
|
||||
},
|
||||
"inboundDebounce": {
|
||||
"title": "Debounce đầu vào",
|
||||
|
||||
@@ -797,15 +797,17 @@
|
||||
},
|
||||
"compaction": {
|
||||
"title": "压缩",
|
||||
"description": "上下文窗口压缩和压缩前记忆设置",
|
||||
"description": "最终请求上下文保护、历史压缩和记忆刷新设置",
|
||||
"maxHistoryShare": "最大历史份额 (0-1)",
|
||||
"maxHistoryShareTip": "上下文窗口用于对话历史的最大比例(如0.85 = 85%)。超过时触发压缩。",
|
||||
"maxHistoryShareTip": "仅用于回合后的历史压缩阈值,不是最终请求的总限制。",
|
||||
"keepLastMessages": "保留最后消息数",
|
||||
"keepLastMessagesTip": "压缩后保留的最近消息数。旧消息被摘要替换。",
|
||||
"timeoutSeconds": "超时秒数",
|
||||
"timeoutSecondsTip": "等待上下文压缩摘要的最长时间,超时后保留原始历史。",
|
||||
"memoryFlush": "压缩前记忆",
|
||||
"memoryFlushTip": "压缩前,Agent可以将重要上下文保存到记忆文件。同时触发知识图谱提取。"
|
||||
"memoryFlushTip": "压缩前,Agent可以将重要上下文保存到记忆文件。同时触发知识图谱提取。",
|
||||
"maxRequestShare": "最大请求份额 (0-1)",
|
||||
"maxRequestShareTip": "最终发送给模型的请求最多可占用智能体上下文窗口的比例。超过时,系统会先缩减上下文再调用模型。"
|
||||
},
|
||||
"inboundDebounce": {
|
||||
"title": "入站防抖",
|
||||
|
||||
@@ -20,6 +20,16 @@ export function CompactionSection({ value, onChange }: CompactionSectionProps) {
|
||||
</div>
|
||||
<div className="rounded-lg border p-3 space-y-4 sm:p-4">
|
||||
<div className="grid grid-cols-1 gap-4 sm:grid-cols-2">
|
||||
<div className="space-y-2">
|
||||
<InfoLabel tip={t(`${s}.maxRequestShareTip`)}>{t(`${s}.maxRequestShare`)}</InfoLabel>
|
||||
<Input
|
||||
type="number"
|
||||
step="0.05"
|
||||
placeholder="0.85"
|
||||
value={value.maxRequestShare ?? ""}
|
||||
onChange={(e) => onChange({ ...value, maxRequestShare: numOrUndef(e.target.value) })}
|
||||
/>
|
||||
</div>
|
||||
<div className="space-y-2">
|
||||
<InfoLabel tip={t(`${s}.maxHistoryShareTip`)}>{t(`${s}.maxHistoryShare`)}</InfoLabel>
|
||||
<Input
|
||||
|
||||
@@ -28,6 +28,7 @@ export interface SubagentsConfig {
|
||||
export interface CompactionConfig {
|
||||
reserveTokensFloor?: number;
|
||||
maxHistoryShare?: number;
|
||||
maxRequestShare?: number;
|
||||
keepLastMessages?: number;
|
||||
timeoutSeconds?: number;
|
||||
memoryFlush?: {
|
||||
|
||||
Reference in new issue
Block a user