mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-11 03:13:24 +00:00
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:
1 parent
04a9938f4f
commit
eb6723d674
16 files changed
+976
-4
No files matched your search
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 ---
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in new issue
Block a user