diff --git a/CLAUDE.md b/CLAUDE.md index 0c5d3b4..4117d33 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -54,7 +54,7 @@ Async SQLAlchemy + aiosqlite. `MemoryStore`: schema init (+ lightweight SQLite ` - **`ChatClient`** Protocol for `chat_completions`. - **`@tool` / `ToolRegistry`**: decorator infers JSON Schema from type hints; execute tools by name. -- **`run_chat_loop`**: multi-round tool loop; default **streaming** text deltas + stream tool-call merge; parallel tools; tool-result char budget; soft context compact on request; cooperative cancel points; optional `on_limit`. +- **`run_chat_loop`**: multi-round tool loop; default **streaming** text deltas + stream tool-call merge; parallel tools; tool-result char budget; soft context compact on request (**API-calibrated** after first usage when available); cooperative cancel points; optional `on_limit`. - **`ChatAgent`**: optional `MemoryStore` (persist on success only); `stream`; system prompt; `pending_retry_text` + `retry()`. - **`/compact`**: soft-compact tool dumps → model summary (no tools) → **new** session seeded with summary message. - Events: text_delta, assistant_message, tool_call/result, max_rounds, **error** (`retryable`/`source`), **cancelled** (`reason`), **usage** (`TokenUsage`). diff --git a/src/plyngent/agent/budget.py b/src/plyngent/agent/budget.py index 04b3c7a..c41b021 100644 --- a/src/plyngent/agent/budget.py +++ b/src/plyngent/agent/budget.py @@ -11,12 +11,14 @@ from plyngent.lmproto.openai_compatible.model import ( ) if TYPE_CHECKING: - from collections.abc import Sequence + from collections.abc import Callable, Sequence from plyngent.lmproto.openai_compatible.model import AnyChatMessage + type TokenMeasure = Callable[[Sequence[AnyChatMessage]], int] + DEFAULT_TOOL_RESULT_MAX_CHARS = 32_000 -# Soft context budget in estimated tokens (~4 chars/token); not a hard model limit. +# Soft context budget in tokens (API-calibrated when possible; else ~4 chars/token). DEFAULT_CONTEXT_MAX_TOKENS = 200_000 DEFAULT_OLD_TOOL_RESULT_CHARS = 800 DEFAULT_RECENT_TOOL_RESULTS = 4 @@ -64,12 +66,37 @@ def estimate_messages_chars(messages: Sequence[AnyChatMessage]) -> int: def estimate_messages_tokens(messages: Sequence[AnyChatMessage]) -> int: - """Char-based token estimate for soft context budget checks.""" + """Char-based token estimate (fallback when no API calibration is available).""" from plyngent.agent.usage import chars_to_tokens return chars_to_tokens(estimate_messages_chars(messages)) +def measure_messages_tokens( + messages: Sequence[AnyChatMessage], + *, + prompt_tokens_hint: int | None = None, + sent_estimate_tokens: int | None = None, +) -> int: + """Token size for budget checks. + + When ``prompt_tokens_hint`` is the last request's API (or resolved) prompt + size and ``sent_estimate_tokens`` is the char-estimate of that same payload, + scale the current char-estimate by ``hint / sent_estimate`` so soft-compact + tracks near-real tokens after the first model call. + """ + est = estimate_messages_tokens(messages) + if ( + prompt_tokens_hint is not None + and prompt_tokens_hint > 0 + and sent_estimate_tokens is not None + and sent_estimate_tokens > 0 + ): + scaled = round(est * (prompt_tokens_hint / sent_estimate_tokens)) + return max(1, scaled) if est > 0 else 0 + return est + + def _shrink_tool(message: ToolChatMessage, max_chars: int) -> ToolChatMessage: if len(message.content) <= max_chars: return message @@ -109,13 +136,14 @@ def _shrink_largest( *, max_tokens: int, shrink_cap: int, + measure: TokenMeasure, ) -> None: def tool_len(i: int) -> int: msg = messages[i] return len(msg.content) if isinstance(msg, ToolChatMessage) else 0 for idx in sorted(tool_indices, key=tool_len, reverse=True): - if estimate_messages_tokens(messages) <= max_tokens: + if measure(messages) <= max_tokens: return tool_msg = messages[idx] if isinstance(tool_msg, ToolChatMessage): @@ -128,17 +156,29 @@ def compact_messages_for_request( max_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS, old_tool_result_chars: int = DEFAULT_OLD_TOOL_RESULT_CHARS, keep_recent_tool_results: int = DEFAULT_RECENT_TOOL_RESULTS, - # Deprecated alias: treated as token budget if max_tokens not overridden via callers. + prompt_tokens_hint: int | None = None, + sent_estimate_tokens: int | None = None, + # Deprecated alias: treated as token budget. max_chars: int | None = None, ) -> list[AnyChatMessage]: """Return a request-time copy with older tool dumps shrunk if over budget. - Budget is in **estimated tokens** (char/4). Does not mutate the original history. + Budget is in tokens. Prefer API-calibrated measurement via + ``prompt_tokens_hint`` / ``sent_estimate_tokens`` (last request); otherwise + fall back to char/4. Does not mutate the original history. ``max_tokens < 1`` disables compacting. """ budget = max_tokens if max_chars is None else max_chars + + def measure(msgs: Sequence[AnyChatMessage]) -> int: + return measure_messages_tokens( + msgs, + prompt_tokens_hint=prompt_tokens_hint, + sent_estimate_tokens=sent_estimate_tokens, + ) + out: list[AnyChatMessage] = list(messages) - if budget < 1 or estimate_messages_tokens(out) <= budget: + if budget < 1 or measure(out) <= budget: return out indices = _tool_indices(out) @@ -147,7 +187,7 @@ def compact_messages_for_request( protect = _protect_indices(indices, keep_recent_tool_results) _shrink_except(out, indices, protect, old_tool_result_chars) - if estimate_messages_tokens(out) <= budget: + if measure(out) <= budget: return out _shrink_largest( @@ -155,5 +195,6 @@ def compact_messages_for_request( indices, max_tokens=budget, shrink_cap=max(64, old_tool_result_chars // 2), + measure=measure, ) return out diff --git a/src/plyngent/agent/compact.py b/src/plyngent/agent/compact.py index 4f5084d..4d6efe0 100644 --- a/src/plyngent/agent/compact.py +++ b/src/plyngent/agent/compact.py @@ -68,9 +68,16 @@ def soft_compact_transcript( messages: Sequence[AnyChatMessage], *, max_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS, + prompt_tokens_hint: int | None = None, + sent_estimate_tokens: int | None = None, ) -> str: """Soft-compact tool dumps then format as transcript text.""" - compacted = compact_messages_for_request(messages, max_tokens=max_tokens) + compacted = compact_messages_for_request( + messages, + max_tokens=max_tokens, + prompt_tokens_hint=prompt_tokens_hint, + sent_estimate_tokens=sent_estimate_tokens, + ) return format_transcript(compacted) @@ -81,12 +88,19 @@ async def summarize_messages( model: str, max_context_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS, temperature: float | None = 0.2, + prompt_tokens_hint: int | None = None, + sent_estimate_tokens: int | None = None, ) -> str: """Soft-compact history and ask the model for a dense summary (no tools).""" if not messages: msg = "nothing to compact" raise ValueError(msg) - transcript = soft_compact_transcript(messages, max_tokens=max_context_tokens) + transcript = soft_compact_transcript( + messages, + max_tokens=max_context_tokens, + prompt_tokens_hint=prompt_tokens_hint, + sent_estimate_tokens=sent_estimate_tokens, + ) if not transcript.strip(): msg = "nothing to compact" raise ValueError(msg) diff --git a/src/plyngent/agent/loop.py b/src/plyngent/agent/loop.py index 8364e2a..fde8f43 100644 --- a/src/plyngent/agent/loop.py +++ b/src/plyngent/agent/loop.py @@ -23,6 +23,7 @@ from .budget import ( DEFAULT_CONTEXT_MAX_TOKENS, DEFAULT_TOOL_RESULT_MAX_CHARS, compact_messages_for_request, + estimate_messages_tokens, truncate_tool_result, ) from .events import ( @@ -240,6 +241,9 @@ async def run_chat_loop( rounds_used = 0 allowance = max_rounds + # Calibrate soft-compact from last model call's prompt_tokens (API preferred). + prompt_tokens_hint: int | None = None + sent_estimate_tokens: int | None = None while True: while rounds_used < allowance: @@ -247,7 +251,10 @@ async def run_chat_loop( request_messages = compact_messages_for_request( messages, max_tokens=max_context_tokens, + prompt_tokens_hint=prompt_tokens_hint, + sent_estimate_tokens=sent_estimate_tokens, ) + sent_est = estimate_messages_tokens(request_messages) param = ChatCompletionsParam( messages=request_messages, model=model, @@ -257,6 +264,10 @@ async def run_chat_loop( pre_len = len(messages) async for event in _assistant_round(client, param, messages, stream=stream): + if isinstance(event, UsageEvent): + # Next rounds scale char-estimates by real/resolved prompt size. + prompt_tokens_hint = event.usage.prompt_tokens + sent_estimate_tokens = sent_est yield event assistant = _last_assistant(messages, pre_len) tool_calls = assistant.tool_calls diff --git a/src/plyngent/cli/state.py b/src/plyngent/cli/state.py index 023c4d7..79d7259 100644 --- a/src/plyngent/cli/state.py +++ b/src/plyngent/cli/state.py @@ -170,11 +170,22 @@ class ReplState: msg = "nothing to compact (empty history)" raise ValueError(msg) + # Prefer last API prompt_tokens to drive soft-compact toward real size. + hint: int | None = None + sent_est: int | None = None + if not self.agent.last_request_usage.is_zero(): + from plyngent.agent.budget import estimate_messages_tokens + + hint = self.agent.last_request_usage.prompt_tokens + # Approximate: calibrate against current full history char-est. + sent_est = estimate_messages_tokens(messages) summary = await summarize_messages( self.client, messages, model=self.model, max_context_tokens=self.agent.max_context_tokens, + prompt_tokens_hint=hint, + sent_estimate_tokens=sent_est, ) session_name = name or f"compact-from-{old_id}" await self.new_session(name=session_name) diff --git a/tests/test_agent/test_compact.py b/tests/test_agent/test_compact.py index dea3ad4..0a88c40 100644 --- a/tests/test_agent/test_compact.py +++ b/tests/test_agent/test_compact.py @@ -7,6 +7,7 @@ from msgspec import UNSET from plyngent.agent.budget import ( compact_messages_for_request, estimate_messages_tokens, + measure_messages_tokens, truncate_tool_result, ) from plyngent.agent.loop import run_chat_loop @@ -100,6 +101,45 @@ def test_compact_disabled_when_max_tokens_zero() -> None: ) +def test_measure_messages_tokens_calibrates_to_api_hint() -> None: + messages: list[AnyChatMessage] = [UserChatMessage(content="a" * 40)] + raw = estimate_messages_tokens(messages) + # If char-est was 10 and API said 100, scale 2x content → ~200 + calibrated = measure_messages_tokens( + messages, + prompt_tokens_hint=raw * 10, + sent_estimate_tokens=raw, + ) + assert calibrated == raw * 10 + + +def test_compact_uses_api_calibration() -> None: + """With a high API hint scale, compact triggers earlier than raw char-est.""" + messages: list[AnyChatMessage] = [ + UserChatMessage(content="start"), + ToolChatMessage(content="OLD" * 400, tool_call_id="1"), + ToolChatMessage(content="NEW" * 20, tool_call_id="2"), + ] + est = estimate_messages_tokens(messages) + # Raw est under budget → no compact + no_api = compact_messages_for_request(messages, max_tokens=est + 100) + assert isinstance(no_api[1], ToolChatMessage) + assert "truncated" not in no_api[1].content + # Calibrate so measured size is 5x → over a mid budget → shrink old tool + compacted = compact_messages_for_request( + messages, + max_tokens=max(1, est * 2), + prompt_tokens_hint=est * 5, + sent_estimate_tokens=est, + old_tool_result_chars=40, + keep_recent_tool_results=1, + ) + assert isinstance(compacted[1], ToolChatMessage) + old = messages[1] + assert isinstance(old, ToolChatMessage) + assert "truncated" in compacted[1].content or len(compacted[1].content) < len(old.content) + + def _response(message: AssistantChatMessage) -> ChatCompletionResponse: return ChatCompletionResponse( id="1",