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:
Clark Cant committed 2026-07-25 00:39:02 +07:00
1 parent 9957a94868
commit 5ca433b8ac
69 files changed
+4110 -371

No files matched your search

+1
View File
@@ -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
}
+6
View File
@@ -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
View File
@@ -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
+106 -3
View File
@@ -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
}
+9 -6
View File
@@ -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,
},
+3
View File
@@ -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)
}
}
+19 -7
View File
@@ -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,
+73 -32
View File
@@ -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
}
+69 -41
View File
@@ -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,
+16 -4
View File
@@ -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.
+99
View File
@@ -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,
}
}
+122
View File
@@ -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")
}
}
+7 -1
View File
@@ -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...)
+18 -6
View File
@@ -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
+18
View File
@@ -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())
}
+18
View File
@@ -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
+2 -1
View File
@@ -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
+50
View File
@@ -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)
+2
View File
@@ -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)
+15
View File
@@ -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
+12 -1
View File
@@ -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"])
}
}
+6
View File
@@ -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)
+4 -11
View File
@@ -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
}
}
+12 -8
View File
@@ -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
+291
View File
@@ -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,
)
}
+11 -3
View File
@@ -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).
+8
View File
@@ -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)
}
+12 -1
View File
@@ -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
}
+1
View File
@@ -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,
+279 -6
View File
@@ -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)
}
}
+1
View File
@@ -118,6 +118,7 @@ type RunResult struct {
Content string
Thinking string
TotalUsage providers.Usage
LastUsage providers.Usage
Iterations int
ToolCalls int
LoopKilled bool
+144 -22
View File
@@ -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)
}
}
+31
View File
@@ -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
}
+36
View File
@@ -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)
+198
View File
@@ -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.
+48 -15
View File
@@ -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
+94 -2
View File
@@ -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{
+3 -3
View File
@@ -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
}
+1 -1
View File
@@ -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
}
+12 -1
View File
@@ -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
}
+31 -52
View File
@@ -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
View File
@@ -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")
}
}
+19
View File
@@ -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
+45 -41
View File
@@ -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()
+4
View File
@@ -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
View File
@@ -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
}
+159
View File
@@ -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)
}
}
+142
View File
@@ -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
+90
View File
@@ -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,
}
}
+152
View File
@@ -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")
}
}
+5 -3
View File
@@ -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",
+5 -3
View File
@@ -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": "컨텍스트 정리",
+5 -3
View File
@@ -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",
+5 -3
View File
@@ -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
+1
View File
@@ -28,6 +28,7 @@ export interface SubagentsConfig {
export interface CompactionConfig {
reserveTokensFloor?: number;
maxHistoryShare?: number;
maxRequestShare?: number;
keepLastMessages?: number;
timeoutSeconds?: number;
memoryFlush?: {