fix(pipeline): include tool-schema tokens in overhead + dynamic compact max_tokens

- Add TokenCounter.CountToolSchemas() to measure JSON schema size for all tools
- Include tool schemas in OverheadTokens calculation for accurate context usage
- Implement dynamic max_tokens: in/25 clamp [1024, 8192] for compaction
- Add characterization tests: count_tool_schemas_test.go
- Add overhead verification tests: context_stage_overhead_test.go, context_stage_tool_overhead_test.go
- Add integration tests: context_stage_integration_test.go
- Add compact tests: loop_compact_dynamic_max_test.go, loop_compact_max_tokens_test.go
- Add sanitize tests: loop_history_sanitize_max_tokens_test.go
- Add integration test: loop_compact_integration_test.go
This commit is contained in:
viettranx committed 2026-04-23 08:31:53 +07:00
1 parent 04a9938f4f
commit eb6723d674
16 files changed
+976 -4

No files matched your search

+28 -1
View File
@@ -87,13 +87,15 @@ func (l *Loop) compactMessagesInPlace(ctx context.Context, messages []providers.
sctx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
inTokens := l.estimateSummaryInputTokens(toSummarize)
slog.Info("compact_budget", "agent", l.id, "in_tokens", inTokens, "out_tokens", dynamicSummaryMax(inTokens))
resp, err := l.provider.Chat(sctx, providers.ChatRequest{
Messages: []providers.Message{{
Role: "user",
Content: compactionSummaryPrompt + sb.String(),
}},
Model: l.model,
Options: map[string]any{"max_tokens": 1024, "temperature": 0.3},
Options: map[string]any{"max_tokens": dynamicSummaryMax(inTokens), "temperature": 0.3},
})
if err != nil {
slog.Warn("mid_loop_compaction_failed", "agent", l.id, "error", err)
@@ -129,3 +131,28 @@ func (l *Loop) compactMessagesInPlace(ctx context.Context, messages []providers.
return result
}
// 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
// runaway output billing on pathological inputs.
func dynamicSummaryMax(inputTokens int) int {
out := max(inputTokens/25, 1024)
if out > 8192 {
out = 8192
}
return out
}
// estimateSummaryInputTokens returns a best-effort input-token count. Prefers
// TokenCounter when attached; else rune/3 fallback (~±15% for UTF-8).
func (l *Loop) estimateSummaryInputTokens(messages []providers.Message) int {
if l.tokenCounter != nil {
return l.tokenCounter.CountMessages(l.model, messages)
}
total := 0
for _, m := range messages {
total += len([]rune(m.Content)) / 3
}
return total
}
@@ -0,0 +1,26 @@
package agent
import "testing"
// TestDynamicSummaryMax validates boundary cases for dynamicSummaryMax.
// Formula: out = in/25, clamped to [1024, 8192].
func TestDynamicSummaryMax(t *testing.T) {
cases := []struct {
input int
want int
}{
{0, 1024}, // zero → floor
{20000, 1024}, // 20000/25=800 → below floor, clamped
{25000, 1024}, // 25000/25=1000 → below floor, clamped
{26000, 1040}, // 26000/25=1040 → just above floor
{100000, 4000}, // 100000/25=4000 → mid-range
{204800, 8192}, // 204800/25=8192 → exactly at cap
{500000, 8192}, // 500000/25=20000 → above cap, clamped
}
for _, tc := range cases {
got := dynamicSummaryMax(tc.input)
if got != tc.want {
t.Errorf("dynamicSummaryMax(%d) = %d, want %d", tc.input, got, tc.want)
}
}
}
@@ -0,0 +1,99 @@
package agent
import (
"context"
"strings"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/tokencount"
)
// buildVietnameseMsgs constructs n alternating user/assistant messages with
// Vietnamese UTF-8 content. Each message is ~viRunes runes to hit a realistic
// total token budget (~100k input tokens for 600 messages).
func buildVietnameseMsgs(n, viRunes int) []providers.Message {
// ~viRunes-rune Vietnamese segment (3-byte UTF-8 per diacritic char).
segment := strings.Repeat(
"Xin chào! Đây là nội dung kiểm tra với ký tự tiếng Việt đặc biệt: ắ ặ ầ ẩ ậ ề ể ệ ọ ộ. ",
(viRunes/80)+1,
)
runes := []rune(segment)
if len(runes) > viRunes {
segment = string(runes[:viRunes])
}
msgs := make([]providers.Message, n)
for i := range msgs {
role := "user"
if i%2 != 0 {
role = "assistant"
}
msgs[i] = providers.Message{Role: role, Content: segment}
}
return msgs
}
// TestLoopCompact_Integration_DynamicMaxTokens_VietnameseFixture verifies the
// end-to-end composition of Phase 03 (FallbackCounter) + Phase 04 (dynamicSummaryMax):
//
// 1. Loop with real FallbackCounter estimates ~100k input tokens from 600 Vietnamese messages.
// 2. compactMessagesInPlace passes max_tokens in [2000, 8192] to the provider.
// 3. The formula dynamicSummaryMax(in) = in/25 holds: for ~100k input → ~4000 output budget.
//
// Tolerance: FallbackCounter uses rune/3 heuristic so exact input count varies;
// we assert >= 2000 && <= 8192 rather than == 4000.
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(),
}
// 600 messages × ~500 runes each ≈ 300k runes ÷ 3 ≈ 100k tokens total.
// keepCount defaults to 4; splitIdx = 600-4 = 596 msgs to summarise.
// FallbackCounter on 596 msgs × ~500 runes ÷ 3 ≈ ~99k tokens → dynamicSummaryMax(99000) = 3960 (floor 1024).
msgs := buildVietnameseMsgs(600, 500)
result := loop.compactMessagesInPlace(context.Background(), msgs)
if result == nil {
t.Fatal("compactMessagesInPlace returned nil; expected compaction to succeed with 600 messages")
}
if len(cap.captured) != 1 {
t.Fatalf("provider.Chat called %d time(s), want 1", len(cap.captured))
}
req := cap.captured[0]
maxTokensRaw, ok := req.Options["max_tokens"]
if !ok {
t.Fatal("Options[\"max_tokens\"] not set in ChatRequest")
}
maxTokens, ok := maxTokensRaw.(int)
if !ok {
t.Fatalf("Options[\"max_tokens\"] type = %T, want int", maxTokensRaw)
}
// Tolerance: FallbackCounter rune/3 varies slightly by content.
// For ~100k token input: dynamicSummaryMax → ~4000 (formula in/25).
// Assert range [2000, 8192] to accommodate counter variance.
const minExpected = 2000
const maxExpected = 8192
if maxTokens < minExpected || maxTokens > maxExpected {
t.Errorf("max_tokens = %d, want in [%d, %d]; formula dynamicSummaryMax(estimatedInput)",
maxTokens, minExpected, maxExpected)
}
// Log actual observed value for diagnostics.
keepCount := 4
if minKeep := len(msgs) * 3 / 10; minKeep > keepCount {
keepCount = minKeep
}
splitIdx := len(msgs) - keepCount
estimatedIn := loop.estimateSummaryInputTokens(msgs[:splitIdx])
t.Logf("observed: msgs=%d splitIdx=%d estimatedIn=%d max_tokens=%d dynamicSummaryMax=%d",
len(msgs), splitIdx, estimatedIn, maxTokens, dynamicSummaryMax(estimatedIn))
}
@@ -0,0 +1,77 @@
package agent
import (
"context"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
)
// capturingProvider records every ChatRequest passed to Chat.
// Distinct from stubProvider in intent_classify_test.go (that one ignores the request).
type capturingProvider struct {
captured []providers.ChatRequest
response string
}
func (c *capturingProvider) Chat(_ context.Context, req providers.ChatRequest) (*providers.ChatResponse, error) {
c.captured = append(c.captured, req)
return &providers.ChatResponse{Content: c.response}, nil
}
func (c *capturingProvider) ChatStream(_ context.Context, req providers.ChatRequest, _ func(providers.StreamChunk)) (*providers.ChatResponse, error) {
c.captured = append(c.captured, req)
return &providers.ChatResponse{Content: c.response}, nil
}
func (c *capturingProvider) DefaultModel() string { return "capturing-model" }
func (c *capturingProvider) Name() string { return "capturing" }
// TestCompactMessagesInPlace_MaxTokensDynamic verifies that compactMessagesInPlace
// passes max_tokens == dynamicSummaryMax(estimatedInputTokens) to the provider.
func TestCompactMessagesInPlace_MaxTokensDynamic(t *testing.T) {
cap := &capturingProvider{response: "Summary of conversation."}
loop := &Loop{
provider: cap,
model: "claude-3-5-sonnet",
// tokenCounter nil → estimateSummaryInputTokens uses rune/3 fallback
}
// Build 10 dummy messages (>= 6 required by compactMessagesInPlace).
msgs := make([]providers.Message, 10)
for i := range msgs {
if i%2 == 0 {
msgs[i] = providers.Message{Role: "user", Content: "user message"}
} else {
msgs[i] = providers.Message{Role: "assistant", Content: "assistant reply"}
}
}
result := loop.compactMessagesInPlace(context.Background(), msgs)
if result == nil {
t.Fatal("compactMessagesInPlace returned nil; expected compaction to succeed")
}
if len(cap.captured) != 1 {
t.Fatalf("provider.Chat called %d time(s), want 1", len(cap.captured))
}
req := cap.captured[0]
maxTokensRaw, ok := req.Options["max_tokens"]
if !ok {
t.Fatal("Options[\"max_tokens\"] not set in ChatRequest")
}
maxTokens, ok := maxTokensRaw.(int)
if !ok {
t.Fatalf("Options[\"max_tokens\"] type = %T, want int", maxTokensRaw)
}
// Compute expected using the same formula the implementation uses.
// With keepCount=4 and 10 messages, splitIdx=6 (first 6 messages summarised).
// tokenCounter nil → rune/3 fallback.
expectedIn := loop.estimateSummaryInputTokens(msgs[:6])
wantMax := dynamicSummaryMax(expectedIn)
if maxTokens != wantMax {
t.Errorf("max_tokens = %d, want %d (dynamicSummaryMax(%d))", maxTokens, wantMax, expectedIn)
}
}
+3 -1
View File
@@ -282,10 +282,12 @@ func (l *Loop) maybeSummarize(ctx context.Context, sessionKey string) {
}
prompt.WriteString(sb.String())
inTokens := l.estimateSummaryInputTokens(toSummarize)
slog.Info("compact_budget", "agent", l.id, "in_tokens", inTokens, "out_tokens", dynamicSummaryMax(inTokens))
resp, err := l.provider.Chat(sctx, providers.ChatRequest{
Messages: []providers.Message{{Role: "user", Content: prompt.String()}},
Model: l.model,
Options: map[string]any{"max_tokens": 1024, "temperature": 0.3},
Options: map[string]any{"max_tokens": dynamicSummaryMax(inTokens), "temperature": 0.3},
})
if err != nil {
slog.Warn("summarization failed", "session", sessionKey, "error", err)
@@ -0,0 +1,174 @@
package agent
import (
"context"
"testing"
"time"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/store"
)
// nopSessionStore is a minimal no-op implementation of store.SessionStore
// for testing maybeSummarize without a real database.
// All methods return zero values except GetHistory and GetLastPromptTokens,
// which return controlled fixture data.
type nopSessionStore struct {
history []providers.Message
lastPromptTokens int
lastMsgCount int
}
// SessionCoreStore methods
func (n *nopSessionStore) GetOrCreate(_ context.Context, _ string) *store.SessionData {
return &store.SessionData{}
}
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 {
return n.history
}
func (n *nopSessionStore) GetSummary(_ context.Context, _ string) string { return "" }
func (n *nopSessionStore) SetSummary(_ context.Context, _, _ string) {}
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) 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) GetMemoryFlushCompactionCount(_ context.Context, _ string) int { return 0 }
func (n *nopSessionStore) SetMemoryFlushDone(_ context.Context, _ string) {}
func (n *nopSessionStore) GetSessionMetadata(_ context.Context, _ string) map[string]string {
return nil
}
func (n *nopSessionStore) SetSessionMetadata(_ context.Context, _ string, _ map[string]string) {}
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) GetLastPromptTokens(_ context.Context, _ string) (int, int) {
return n.lastPromptTokens, n.lastMsgCount
}
// SessionListingStore methods
func (n *nopSessionStore) List(_ context.Context, _ string) []store.SessionInfo { return nil }
func (n *nopSessionStore) ListPaged(_ context.Context, _ store.SessionListOpts) store.SessionListResult {
return store.SessionListResult{Sessions: []store.SessionInfo{}}
}
func (n *nopSessionStore) ListPagedRich(_ context.Context, _ store.SessionListOpts) store.SessionListRichResult {
return store.SessionListRichResult{Sessions: []store.SessionInfoRich{}}
}
func (n *nopSessionStore) LastUsedChannel(_ context.Context, _ string) (string, string) {
return "", ""
}
// signallingProvider wraps capturingProvider and signals a channel when Chat is called.
type signallingProvider struct {
capturingProvider
done chan struct{}
}
func (s *signallingProvider) Chat(ctx context.Context, req providers.ChatRequest) (*providers.ChatResponse, error) {
resp, err := s.capturingProvider.Chat(ctx, req)
select {
case s.done <- struct{}{}:
default:
}
return resp, err
}
// TestMaybeSummarize_MaxTokensDynamic verifies that maybeSummarize passes
// max_tokens == dynamicSummaryMax(estimatedInputTokens) to the provider.
func TestMaybeSummarize_MaxTokensDynamic(t *testing.T) {
const contextWindow = 10000
// Build history large enough to exceed the compaction threshold.
// 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)
history := make([]providers.Message, 10)
for i := range history {
if i%2 == 0 {
history[i] = providers.Message{Role: "user", Content: longContent}
} else {
history[i] = providers.Message{Role: "assistant", Content: longContent}
}
}
done := make(chan struct{}, 1)
sp := &signallingProvider{
capturingProvider: capturingProvider{response: "compaction summary"},
done: done,
}
sessions := &nopSessionStore{
history: history,
lastPromptTokens: 0, // no calibration → falls back to EstimateTokens
lastMsgCount: 0,
}
loop := &Loop{
provider: sp,
model: "claude-3-5-sonnet",
contextWindow: contextWindow,
sessions: sessions,
// hasMemory = false → shouldRunMemoryFlush returns false (skip memory flush)
hasMemory: false,
// compactionCfg nil → uses DefaultHistoryShare (0.85), keepLast=4
compactionCfg: nil,
// tokenCounter nil → estimateSummaryInputTokens uses rune/3 fallback
}
loop.maybeSummarize(context.Background(), "test-session-key")
// Wait for background goroutine to call provider.Chat (up to 5s).
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for maybeSummarize to call provider.Chat")
}
if len(sp.captured) == 0 {
t.Fatal("provider.Chat was not called")
}
req := sp.captured[0]
maxTokensRaw, ok := req.Options["max_tokens"]
if !ok {
t.Fatal("Options[\"max_tokens\"] not set in ChatRequest from maybeSummarize")
}
maxTokens, ok := maxTokensRaw.(int)
if !ok {
t.Fatalf("Options[\"max_tokens\"] type = %T, want int", maxTokensRaw)
}
// Compute expected using the same formula the implementation uses.
// keepLast=4, history has 10 messages → toSummarize = history[:6].
// tokenCounter nil → rune/3 fallback on the fixture content.
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)
}
}
// makeLongString returns a string of n ASCII characters ('a').
func makeLongString(n int) string {
b := make([]byte, n)
for i := range b {
b[i] = 'a'
}
return string(b)
}
+15 -1
View File
@@ -129,10 +129,24 @@ func (s *ContextStage) Execute(ctx context.Context, state *RunState) error {
}
}
// 5. Compute overhead tokens via TokenCounter (replaces heuristic estimateOverhead)
// 4.5. Build filtered tools early so OverheadTokens includes tool-schema tokens.
// ThinkStage still calls BuildFilteredTools every iteration (tool list is
// iteration-dependent; final iteration strips all tools). This call is
// best-effort: errors are silently swallowed and the tool slice stays nil,
// which means overhead will under-count but remains safe/conservative.
if s.deps.BuildFilteredTools != nil {
if tools, err := s.deps.BuildFilteredTools(state); err == nil {
state.Think.Tools = tools
}
}
// 5. Compute overhead tokens via TokenCounter (replaces heuristic estimateOverhead).
// Includes both system-prompt tokens and tool-schema tokens so PruneStage
// budget shrinks correctly when tools are large.
if s.deps.TokenCounter != nil {
system := state.Messages.System()
overhead := s.deps.TokenCounter.CountMessages(state.Model, []providers.Message{system})
overhead += s.deps.TokenCounter.CountToolSchemas(state.Model, state.Think.Tools)
state.Context.OverheadTokens = overhead
}
@@ -0,0 +1,131 @@
package pipeline
import (
"context"
"encoding/json"
"strings"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/tokencount"
)
// buildRealisticToolDefinitions returns n ToolDefinitions with ~3KB JSON each,
// matching the trace-019dab16 scenario (agent with 10 realistic tools).
// Each tool has 8 uniquely-named parameters with long descriptions to reach ~3KB.
func buildRealisticToolDefinitions(n int) []providers.ToolDefinition {
// Long description filler (~200 chars) repeated per property.
descFiller := "This parameter controls an important aspect of the tool behaviour. " +
"Provide a valid value according to the schema constraints documented above. "
paramNames := []string{
"source_file_path", "destination_path", "encoding_format",
"compression_level", "output_template", "max_retry_count",
"timeout_seconds", "verbose_logging",
}
tools := make([]providers.ToolDefinition, n)
for i := range tools {
properties := map[string]any{}
required := make([]string, 0, 2)
for j, name := range paramNames {
properties[name] = map[string]any{
"type": "string",
"description": strings.Repeat(descFiller, 2),
}
if j < 2 {
required = append(required, name)
}
}
tools[i] = providers.ToolDefinition{
Type: "function",
Function: providers.ToolFunctionSchema{
Name: "realistic_tool",
Description: strings.Repeat(
"A realistic tool that performs complex file and system operations. "+
"It accepts multiple parameters and returns structured JSON output. "+
"Use this tool when you need to process, transform, or analyse data. ",
4,
),
Parameters: map[string]any{
"type": "object",
"properties": properties,
"required": required,
},
},
}
}
return tools
}
// TestContextStage_Integration_ToolOverhead_RealCounter verifies the end-to-end
// composition of Phase 03 (CountToolSchemas) with the real FallbackCounter:
// 1. state.Think.Tools is populated (len == numTools).
// 2. OverheadTokens > system-prompt-only count (tools add non-zero overhead).
// 3. OverheadTokens > 5000 when 10 tools each ~3KB JSON are provided.
//
// Uses real tokencount.FallbackCounter (no spy) for deterministic non-zero counts.
func TestContextStage_Integration_ToolOverhead_RealCounter(t *testing.T) {
t.Parallel()
const numTools = 10
counter := tokencount.NewFallbackCounter()
// Build system prompt ~1500 chars.
systemPrompt := strings.Repeat(
"You are a capable AI assistant with access to many tools. "+
"Use them wisely to help the user accomplish their goals. ",
10,
)
fixture := buildRealisticToolDefinitions(numTools)
// Sanity: verify fixtures are actually ~3KB each.
toolJSON, _ := json.Marshal(fixture[0])
if len(toolJSON) < 1000 {
t.Logf("WARNING: single tool JSON only %d bytes; fixture may be smaller than expected", len(toolJSON))
}
deps := &PipelineDeps{
TokenCounter: counter,
BuildMessages: func(_ context.Context, _ *RunInput, _ []providers.Message, _ string) ([]providers.Message, error) {
return []providers.Message{
{Role: "system", Content: systemPrompt},
}, nil
},
BuildFilteredTools: func(_ *RunState) ([]providers.ToolDefinition, error) {
return fixture, nil
},
}
stage := NewContextStage(deps)
state := defaultState()
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() error: %v", err)
}
// Assert 1: Tools populated.
if len(state.Think.Tools) != numTools {
t.Errorf("state.Think.Tools len = %d, want %d", len(state.Think.Tools), numTools)
}
// Compute system-only overhead for comparison.
sysMsg := providers.Message{Role: "system", Content: systemPrompt}
systemOnly := counter.CountMessages("claude-3", []providers.Message{sysMsg})
// Assert 2: OverheadTokens strictly greater than system-only (tools counted).
if state.Context.OverheadTokens <= systemOnly {
t.Errorf("OverheadTokens = %d, want > %d (system-only=%d); tool schemas not contributing",
state.Context.OverheadTokens, systemOnly, systemOnly)
}
// Assert 3: OverheadTokens > 5000 (system ~500 + 10 tools × 3KB JSON ÷ 3 ≈ 10000+).
const wantMinOverhead = 5000
if state.Context.OverheadTokens <= wantMinOverhead {
t.Errorf("OverheadTokens = %d, want > %d; 10 tools with ~3KB JSON each should contribute significantly",
state.Context.OverheadTokens, wantMinOverhead)
}
t.Logf("observed: systemOnly=%d, overhead=%d, tools=%d", systemOnly, state.Context.OverheadTokens, numTools)
}
@@ -0,0 +1,145 @@
package pipeline
import (
"context"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
)
// spyTokenCounter records all CountMessages invocations for assertion.
type spyTokenCounter struct {
calls [][]providers.Message // each element is the msgs slice from one CountMessages call
toolCounts int // number of CountToolSchemas calls
fixed int // tokens returned per CountMessages call
toolFixed int // tokens returned per CountToolSchemas call
}
func (s *spyTokenCounter) Count(_ string, _ string) int { return s.fixed }
func (s *spyTokenCounter) CountMessages(_ string, msgs []providers.Message) int {
// Deep-copy the slice so later mutations don't affect recorded state.
cp := make([]providers.Message, len(msgs))
copy(cp, msgs)
s.calls = append(s.calls, cp)
return len(msgs) * s.fixed
}
func (s *spyTokenCounter) CountToolSchemas(_ string, tools []providers.ToolDefinition) int {
s.toolCounts++
return len(tools) * s.toolFixed
}
func (s *spyTokenCounter) ModelContextWindow(_ string) int { return 200_000 }
// fixtureTools returns a slice of n minimal ToolDefinitions for testing.
func fixtureTools(n int) []providers.ToolDefinition {
tools := make([]providers.ToolDefinition, n)
for i := range tools {
tools[i] = providers.ToolDefinition{
Type: "function",
Function: providers.ToolFunctionSchema{
Name: "tool_fixture",
Description: "A fixture tool for testing overhead calculation.",
Parameters: map[string]any{"type": "object", "properties": map[string]any{}},
},
}
}
return tools
}
// TestContextStage_OverheadSystemPlusTools_PostFix verifies the POST-fix overhead
// calculation: OverheadTokens = system-message tokens + tool-schema tokens.
// Both CountMessages and CountToolSchemas are called exactly once.
func TestContextStage_OverheadSystemPlusTools_PostFix(t *testing.T) {
t.Parallel()
const systemFixed = 100
const toolFixed = 50
const numTools = 5
spy := &spyTokenCounter{fixed: systemFixed, toolFixed: toolFixed}
fixture := fixtureTools(numTools)
deps := &PipelineDeps{
TokenCounter: spy,
// BuildMessages seeds a system message so the counter has content.
BuildMessages: func(_ context.Context, _ *RunInput, _ []providers.Message, _ string) ([]providers.Message, error) {
return []providers.Message{
{Role: "system", Content: "You are a helpful assistant with many capabilities."},
}, nil
},
// BuildFilteredTools returns fixture tools so ContextStage can count them.
BuildFilteredTools: func(_ *RunState) ([]providers.ToolDefinition, error) {
return fixture, nil
},
}
stage := NewContextStage(deps)
state := defaultState()
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() error: %v", err)
}
// POST-fix: OverheadTokens = system (1 msg × 100) + tools (5 × 50) = 350.
wantOverhead := systemFixed + numTools*toolFixed
if state.Context.OverheadTokens != wantOverhead {
t.Errorf("OverheadTokens = %d, want %d (system=%d + tools=%d)",
state.Context.OverheadTokens, wantOverhead, systemFixed, numTools*toolFixed)
}
// Assert: exactly 1 call to CountMessages (system msg).
if len(spy.calls) != 1 {
t.Errorf("CountMessages called %d time(s), want exactly 1", len(spy.calls))
}
// Assert: CountToolSchemas called exactly once.
if spy.toolCounts != 1 {
t.Errorf("CountToolSchemas called %d time(s), want exactly 1", spy.toolCounts)
}
// Assert: state.Think.Tools populated by ContextStage.
if len(state.Think.Tools) != numTools {
t.Errorf("state.Think.Tools len = %d, want %d", len(state.Think.Tools), numTools)
}
}
// TestContextStage_OverheadSystemOnly_NoToolsCallback verifies that when
// BuildFilteredTools is nil, OverheadTokens = system tokens only (no panic).
// CountToolSchemas IS called with a nil slice (returns 0) — that's correct behavior.
func TestContextStage_OverheadSystemOnly_NoToolsCallback(t *testing.T) {
t.Parallel()
spy := &spyTokenCounter{fixed: 100, toolFixed: 0}
deps := &PipelineDeps{
TokenCounter: spy,
BuildMessages: func(_ context.Context, _ *RunInput, _ []providers.Message, _ string) ([]providers.Message, error) {
return []providers.Message{
{Role: "system", Content: "You are a helpful assistant with many capabilities."},
}, nil
},
// BuildFilteredTools intentionally nil.
}
stage := NewContextStage(deps)
state := defaultState()
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() error: %v", err)
}
// No tools → overhead = system only (1 msg × 100 = 100).
// CountToolSchemas(nil) = 0, so overhead is unchanged.
wantOverhead := 100
if state.Context.OverheadTokens != wantOverhead {
t.Errorf("OverheadTokens = %d, want %d", state.Context.OverheadTokens, wantOverhead)
}
}
// roleList returns a slice of role strings for error messages.
func roleList(msgs []providers.Message) []string {
roles := make([]string, len(msgs))
for i, m := range msgs {
roles[i] = m.Role
}
return roles
}
@@ -0,0 +1,98 @@
package pipeline
import (
"context"
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/tokencount"
)
// TestContextStage_ToolOverhead_ThinkToolsPopulated verifies that:
// 1. state.Think.Tools is populated by BuildFilteredTools called in ContextStage.
// 2. state.Context.OverheadTokens > CountMessages(system) when tools are present.
func TestContextStage_ToolOverhead_ThinkToolsPopulated(t *testing.T) {
t.Parallel()
const numTools = 5
// Use real FallbackCounter so we get deterministic non-zero tool counts.
counter := tokencount.NewFallbackCounter()
fixture := fixtureTools(numTools)
deps := &PipelineDeps{
TokenCounter: counter,
BuildMessages: func(_ context.Context, _ *RunInput, _ []providers.Message, _ string) ([]providers.Message, error) {
return []providers.Message{
{Role: "system", Content: "You are a capable AI assistant."},
}, nil
},
BuildFilteredTools: func(_ *RunState) ([]providers.ToolDefinition, error) {
return fixture, nil
},
}
stage := NewContextStage(deps)
state := defaultState()
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() error: %v", err)
}
// Assert: state.Think.Tools populated.
if len(state.Think.Tools) != numTools {
t.Errorf("state.Think.Tools len = %d, want %d", len(state.Think.Tools), numTools)
}
// Compute expected system-only overhead to compare.
sysMsg := providers.Message{Role: "system", Content: "You are a capable AI assistant."}
systemOnly := counter.CountMessages("claude-3", []providers.Message{sysMsg})
// Assert: overhead includes tool tokens — strictly greater than system-only.
if state.Context.OverheadTokens <= systemOnly {
t.Errorf("OverheadTokens = %d, want > %d (system-only); tool schemas not counted",
state.Context.OverheadTokens, systemOnly)
}
}
// TestContextStage_ToolOverhead_BuildFilteredToolsError_FallsBackToSystemOnly verifies
// that a BuildFilteredTools error is silently swallowed and overhead = system only.
func TestContextStage_ToolOverhead_BuildFilteredToolsError_FallsBackToSystemOnly(t *testing.T) {
t.Parallel()
counter := tokencount.NewFallbackCounter()
deps := &PipelineDeps{
TokenCounter: counter,
BuildMessages: func(_ context.Context, _ *RunInput, _ []providers.Message, _ string) ([]providers.Message, error) {
return []providers.Message{
{Role: "system", Content: "You are a capable AI assistant."},
}, nil
},
BuildFilteredTools: func(_ *RunState) ([]providers.ToolDefinition, error) {
return nil, context.DeadlineExceeded // simulate error
},
}
stage := NewContextStage(deps)
state := defaultState()
// Should not return an error even though BuildFilteredTools failed.
if err := stage.Execute(context.Background(), state); err != nil {
t.Fatalf("Execute() error: %v", err)
}
// state.Think.Tools should remain nil/empty.
if len(state.Think.Tools) != 0 {
t.Errorf("state.Think.Tools len = %d, want 0 on BuildFilteredTools error", len(state.Think.Tools))
}
// Overhead = system only (no tool penalty).
sysMsg := providers.Message{Role: "system", Content: "You are a capable AI assistant."}
wantOverhead := counter.CountMessages("claude-3", []providers.Message{sysMsg})
if state.Context.OverheadTokens != wantOverhead {
t.Errorf("OverheadTokens = %d, want %d (system-only on tool-build error)",
state.Context.OverheadTokens, wantOverhead)
}
}
+2 -1
View File
@@ -45,7 +45,8 @@ func (m *mockTokenCounter) Count(_ string, _ string) int { return m.countPerMess
func (m *mockTokenCounter) CountMessages(_ string, msgs []providers.Message) int {
return len(msgs) * m.countPerMessage
}
func (m *mockTokenCounter) ModelContextWindow(_ string) int { return 200_000 }
func (m *mockTokenCounter) CountToolSchemas(_ string, _ []providers.ToolDefinition) int { return 0 }
func (m *mockTokenCounter) ModelContextWindow(_ string) int { return 200_000 }
// --- ThinkStage tests ---
+7
View File
@@ -34,6 +34,13 @@ type ThinkState struct {
TruncRetries int // consecutive truncation retries (max 3)
OverflowRetries int // context overflow compact+retry attempts (max 1)
StreamingActive bool // true during active stream
// Tools is populated by ContextStage (iteration=0) for overhead calculation.
// It holds the best-effort tool list at run start and is used exclusively by
// the overhead counter in ContextStage. ThinkStage does NOT consume this field —
// it always calls BuildFilteredTools per iteration because the tool list is
// iteration-dependent (final iteration strips all tools).
Tools []providers.ToolDefinition
}
// PruneState: owned by PruneStage.
@@ -0,0 +1,136 @@
package tokencount_test
import (
"testing"
"github.com/nextlevelbuilder/goclaw/internal/providers"
"github.com/nextlevelbuilder/goclaw/internal/tokencount"
)
const testModel = "claude-sonnet-4-5-20250929"
// smallTool returns a minimal ToolDefinition with a short description.
func smallTool() providers.ToolDefinition {
return providers.ToolDefinition{
Type: "function",
Function: providers.ToolFunctionSchema{
Name: "get_time",
Description: "Returns the current UTC time.",
Parameters: map[string]any{"type": "object", "properties": map[string]any{}},
},
}
}
// largeTool returns a ToolDefinition with a longer description and parameters.
func largeTool(name string) providers.ToolDefinition {
return providers.ToolDefinition{
Type: "function",
Function: providers.ToolFunctionSchema{
Name: name,
Description: "Reads, writes, and appends content to files in the workspace. " +
"Supports binary and text modes. Path must be relative to the active workspace root. " +
"Returns byte count on success. Errors on path traversal attempts.",
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"path": map[string]any{"type": "string", "description": "Relative file path"},
"content": map[string]any{"type": "string", "description": "Content to write"},
"mode": map[string]any{"type": "string", "enum": []string{"read", "write", "append"}},
},
"required": []string{"path", "mode"},
},
},
}
}
// fiveLargeTools returns 5 distinct large tool definitions.
func fiveLargeTools() []providers.ToolDefinition {
names := []string{"write_file", "read_file", "exec_command", "web_search", "create_image"}
tools := make([]providers.ToolDefinition, len(names))
for i, n := range names {
tools[i] = largeTool(n)
}
return tools
}
func TestCountToolSchemas_NilSlice_ReturnsZero(t *testing.T) {
t.Parallel()
tc := tokencount.NewTiktokenCounter()
fc := tokencount.NewFallbackCounter()
if got := tc.CountToolSchemas(testModel, nil); got != 0 {
t.Errorf("tiktokenCounter.CountToolSchemas(nil) = %d, want 0", got)
}
if got := fc.CountToolSchemas(testModel, nil); got != 0 {
t.Errorf("FallbackCounter.CountToolSchemas(nil) = %d, want 0", got)
}
}
func TestCountToolSchemas_EmptySlice_ReturnsZero(t *testing.T) {
t.Parallel()
tc := tokencount.NewTiktokenCounter()
fc := tokencount.NewFallbackCounter()
if got := tc.CountToolSchemas(testModel, []providers.ToolDefinition{}); got != 0 {
t.Errorf("tiktokenCounter.CountToolSchemas([]) = %d, want 0", got)
}
if got := fc.CountToolSchemas(testModel, []providers.ToolDefinition{}); got != 0 {
t.Errorf("FallbackCounter.CountToolSchemas([]) = %d, want 0", got)
}
}
func TestCountToolSchemas_OneSmallTool_PositiveCount(t *testing.T) {
t.Parallel()
tools := []providers.ToolDefinition{smallTool()}
tc := tokencount.NewTiktokenCounter()
fc := tokencount.NewFallbackCounter()
if got := tc.CountToolSchemas(testModel, tools); got <= 0 {
t.Errorf("tiktokenCounter.CountToolSchemas(1 small tool) = %d, want > 0", got)
}
if got := fc.CountToolSchemas(testModel, tools); got <= 0 {
t.Errorf("FallbackCounter.CountToolSchemas(1 small tool) = %d, want > 0", got)
}
}
func TestCountToolSchemas_FiveLargeToolsGtOneSmall(t *testing.T) {
t.Parallel()
one := []providers.ToolDefinition{smallTool()}
five := fiveLargeTools()
tc := tokencount.NewTiktokenCounter()
fc := tokencount.NewFallbackCounter()
tcOne := tc.CountToolSchemas(testModel, one)
tcFive := tc.CountToolSchemas(testModel, five)
if tcFive <= tcOne {
t.Errorf("tiktokenCounter: 5 large tools (%d) should produce more tokens than 1 small tool (%d)", tcFive, tcOne)
}
fcOne := fc.CountToolSchemas(testModel, one)
fcFive := fc.CountToolSchemas(testModel, five)
if fcFive <= fcOne {
t.Errorf("FallbackCounter: 5 large tools (%d) should produce more tokens than 1 small tool (%d)", fcFive, fcOne)
}
}
func TestCountToolSchemas_FallbackModel_UsesRuneHeuristic(t *testing.T) {
t.Parallel()
// Unknown model forces tiktoken to use FallbackCounter path.
const unknownModel = "unknown-model-xyz"
tools := fiveLargeTools()
tc := tokencount.NewTiktokenCounter()
fc := tokencount.NewFallbackCounter()
tcCount := tc.CountToolSchemas(unknownModel, tools)
fcCount := fc.CountToolSchemas(unknownModel, tools)
// Both should return same value since tiktoken falls back to FallbackCounter.
if tcCount != fcCount {
t.Errorf("unknown model: tiktokenCounter(%d) != FallbackCounter(%d), expected same fallback path", tcCount, fcCount)
}
if tcCount <= 0 {
t.Errorf("unknown model: CountToolSchemas = %d, want > 0", tcCount)
}
}
+11
View File
@@ -2,6 +2,7 @@ package tokencount
import (
"cmp"
"encoding/json"
"slices"
"strings"
"unicode/utf8"
@@ -38,6 +39,16 @@ func (c *FallbackCounter) CountMessages(_ string, msgs []providers.Message) int
return total
}
// CountToolSchemas returns rune/3 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
}
// ModelContextWindow uses longest-prefix-match to avoid ambiguity
// (e.g., "gpt-4o" must match before "gpt-4").
func (c *FallbackCounter) ModelContextWindow(model string) int {
+19
View File
@@ -1,6 +1,7 @@
package tokencount
import (
"encoding/json"
"hash/fnv"
"log/slog"
"sync"
@@ -80,6 +81,24 @@ func (c *tiktokenCounter) CountMessages(model string, msgs []providers.Message)
return total
}
// CountToolSchemas returns BPE token count for the JSON-serialised tool list.
// Falls back to FallbackCounter if the encoder is unavailable.
// Returns 0 for nil or empty slice.
func (c *tiktokenCounter) CountToolSchemas(model string, tools []providers.ToolDefinition) int {
if len(tools) == 0 {
return 0
}
enc := c.encoderForModel(model)
if enc == nil {
return c.fallback.CountToolSchemas(model, tools)
}
blob, err := json.Marshal(tools)
if err != nil {
return 0
}
return len(enc.Encode(string(blob), nil, nil))
}
// ModelContextWindow delegates to FallbackCounter (same prefix-match logic).
func (c *tiktokenCounter) ModelContextWindow(model string) int {
return c.fallback.ModelContextWindow(model)
+5
View File
@@ -16,6 +16,11 @@ type TokenCounter interface {
// including per-message overhead (role tokens, separators).
CountMessages(model string, msgs []providers.Message) int
// CountToolSchemas returns token count for a slice of tool definitions
// serialised as JSON (the form sent to the LLM provider).
// Returns 0 for nil or empty slice.
CountToolSchemas(model string, tools []providers.ToolDefinition) int
// ModelContextWindow returns max context tokens for a model.
// Falls back to provider default if model unknown.
ModelContextWindow(model string) int