mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/agent: calibrate soft-compact with last-request prompt_tokens
After the first model call, scale char-based token estimates by API/resolved prompt size so budget checks track near-real tokens. /compact uses the same calibration when available.
This commit is contained in:
@@ -54,7 +54,7 @@ Async SQLAlchemy + aiosqlite. `MemoryStore`: schema init (+ lightweight SQLite `
|
|||||||
|
|
||||||
- **`ChatClient`** Protocol for `chat_completions`.
|
- **`ChatClient`** Protocol for `chat_completions`.
|
||||||
- **`@tool` / `ToolRegistry`**: decorator infers JSON Schema from type hints; execute tools by name.
|
- **`@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()`.
|
- **`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.
|
- **`/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`).
|
- Events: text_delta, assistant_message, tool_call/result, max_rounds, **error** (`retryable`/`source`), **cancelled** (`reason`), **usage** (`TokenUsage`).
|
||||||
|
|||||||
@@ -11,12 +11,14 @@ from plyngent.lmproto.openai_compatible.model import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import Sequence
|
from collections.abc import Callable, Sequence
|
||||||
|
|
||||||
from plyngent.lmproto.openai_compatible.model import AnyChatMessage
|
from plyngent.lmproto.openai_compatible.model import AnyChatMessage
|
||||||
|
|
||||||
|
type TokenMeasure = Callable[[Sequence[AnyChatMessage]], int]
|
||||||
|
|
||||||
DEFAULT_TOOL_RESULT_MAX_CHARS = 32_000
|
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_CONTEXT_MAX_TOKENS = 200_000
|
||||||
DEFAULT_OLD_TOOL_RESULT_CHARS = 800
|
DEFAULT_OLD_TOOL_RESULT_CHARS = 800
|
||||||
DEFAULT_RECENT_TOOL_RESULTS = 4
|
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:
|
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
|
from plyngent.agent.usage import chars_to_tokens
|
||||||
|
|
||||||
return chars_to_tokens(estimate_messages_chars(messages))
|
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:
|
def _shrink_tool(message: ToolChatMessage, max_chars: int) -> ToolChatMessage:
|
||||||
if len(message.content) <= max_chars:
|
if len(message.content) <= max_chars:
|
||||||
return message
|
return message
|
||||||
@@ -109,13 +136,14 @@ def _shrink_largest(
|
|||||||
*,
|
*,
|
||||||
max_tokens: int,
|
max_tokens: int,
|
||||||
shrink_cap: int,
|
shrink_cap: int,
|
||||||
|
measure: TokenMeasure,
|
||||||
) -> None:
|
) -> None:
|
||||||
def tool_len(i: int) -> int:
|
def tool_len(i: int) -> int:
|
||||||
msg = messages[i]
|
msg = messages[i]
|
||||||
return len(msg.content) if isinstance(msg, ToolChatMessage) else 0
|
return len(msg.content) if isinstance(msg, ToolChatMessage) else 0
|
||||||
|
|
||||||
for idx in sorted(tool_indices, key=tool_len, reverse=True):
|
for idx in sorted(tool_indices, key=tool_len, reverse=True):
|
||||||
if estimate_messages_tokens(messages) <= max_tokens:
|
if measure(messages) <= max_tokens:
|
||||||
return
|
return
|
||||||
tool_msg = messages[idx]
|
tool_msg = messages[idx]
|
||||||
if isinstance(tool_msg, ToolChatMessage):
|
if isinstance(tool_msg, ToolChatMessage):
|
||||||
@@ -128,17 +156,29 @@ def compact_messages_for_request(
|
|||||||
max_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS,
|
max_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS,
|
||||||
old_tool_result_chars: int = DEFAULT_OLD_TOOL_RESULT_CHARS,
|
old_tool_result_chars: int = DEFAULT_OLD_TOOL_RESULT_CHARS,
|
||||||
keep_recent_tool_results: int = DEFAULT_RECENT_TOOL_RESULTS,
|
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,
|
max_chars: int | None = None,
|
||||||
) -> list[AnyChatMessage]:
|
) -> list[AnyChatMessage]:
|
||||||
"""Return a request-time copy with older tool dumps shrunk if over budget.
|
"""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.
|
``max_tokens < 1`` disables compacting.
|
||||||
"""
|
"""
|
||||||
budget = max_tokens if max_chars is None else max_chars
|
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)
|
out: list[AnyChatMessage] = list(messages)
|
||||||
if budget < 1 or estimate_messages_tokens(out) <= budget:
|
if budget < 1 or measure(out) <= budget:
|
||||||
return out
|
return out
|
||||||
|
|
||||||
indices = _tool_indices(out)
|
indices = _tool_indices(out)
|
||||||
@@ -147,7 +187,7 @@ def compact_messages_for_request(
|
|||||||
|
|
||||||
protect = _protect_indices(indices, keep_recent_tool_results)
|
protect = _protect_indices(indices, keep_recent_tool_results)
|
||||||
_shrink_except(out, indices, protect, old_tool_result_chars)
|
_shrink_except(out, indices, protect, old_tool_result_chars)
|
||||||
if estimate_messages_tokens(out) <= budget:
|
if measure(out) <= budget:
|
||||||
return out
|
return out
|
||||||
|
|
||||||
_shrink_largest(
|
_shrink_largest(
|
||||||
@@ -155,5 +195,6 @@ def compact_messages_for_request(
|
|||||||
indices,
|
indices,
|
||||||
max_tokens=budget,
|
max_tokens=budget,
|
||||||
shrink_cap=max(64, old_tool_result_chars // 2),
|
shrink_cap=max(64, old_tool_result_chars // 2),
|
||||||
|
measure=measure,
|
||||||
)
|
)
|
||||||
return out
|
return out
|
||||||
|
|||||||
@@ -68,9 +68,16 @@ def soft_compact_transcript(
|
|||||||
messages: Sequence[AnyChatMessage],
|
messages: Sequence[AnyChatMessage],
|
||||||
*,
|
*,
|
||||||
max_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS,
|
max_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS,
|
||||||
|
prompt_tokens_hint: int | None = None,
|
||||||
|
sent_estimate_tokens: int | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Soft-compact tool dumps then format as transcript text."""
|
"""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)
|
return format_transcript(compacted)
|
||||||
|
|
||||||
|
|
||||||
@@ -81,12 +88,19 @@ async def summarize_messages(
|
|||||||
model: str,
|
model: str,
|
||||||
max_context_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS,
|
max_context_tokens: int = DEFAULT_CONTEXT_MAX_TOKENS,
|
||||||
temperature: float | None = 0.2,
|
temperature: float | None = 0.2,
|
||||||
|
prompt_tokens_hint: int | None = None,
|
||||||
|
sent_estimate_tokens: int | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Soft-compact history and ask the model for a dense summary (no tools)."""
|
"""Soft-compact history and ask the model for a dense summary (no tools)."""
|
||||||
if not messages:
|
if not messages:
|
||||||
msg = "nothing to compact"
|
msg = "nothing to compact"
|
||||||
raise ValueError(msg)
|
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():
|
if not transcript.strip():
|
||||||
msg = "nothing to compact"
|
msg = "nothing to compact"
|
||||||
raise ValueError(msg)
|
raise ValueError(msg)
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from .budget import (
|
|||||||
DEFAULT_CONTEXT_MAX_TOKENS,
|
DEFAULT_CONTEXT_MAX_TOKENS,
|
||||||
DEFAULT_TOOL_RESULT_MAX_CHARS,
|
DEFAULT_TOOL_RESULT_MAX_CHARS,
|
||||||
compact_messages_for_request,
|
compact_messages_for_request,
|
||||||
|
estimate_messages_tokens,
|
||||||
truncate_tool_result,
|
truncate_tool_result,
|
||||||
)
|
)
|
||||||
from .events import (
|
from .events import (
|
||||||
@@ -240,6 +241,9 @@ async def run_chat_loop(
|
|||||||
|
|
||||||
rounds_used = 0
|
rounds_used = 0
|
||||||
allowance = max_rounds
|
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 True:
|
||||||
while rounds_used < allowance:
|
while rounds_used < allowance:
|
||||||
@@ -247,7 +251,10 @@ async def run_chat_loop(
|
|||||||
request_messages = compact_messages_for_request(
|
request_messages = compact_messages_for_request(
|
||||||
messages,
|
messages,
|
||||||
max_tokens=max_context_tokens,
|
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(
|
param = ChatCompletionsParam(
|
||||||
messages=request_messages,
|
messages=request_messages,
|
||||||
model=model,
|
model=model,
|
||||||
@@ -257,6 +264,10 @@ async def run_chat_loop(
|
|||||||
|
|
||||||
pre_len = len(messages)
|
pre_len = len(messages)
|
||||||
async for event in _assistant_round(client, param, messages, stream=stream):
|
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
|
yield event
|
||||||
assistant = _last_assistant(messages, pre_len)
|
assistant = _last_assistant(messages, pre_len)
|
||||||
tool_calls = assistant.tool_calls
|
tool_calls = assistant.tool_calls
|
||||||
|
|||||||
@@ -170,11 +170,22 @@ class ReplState:
|
|||||||
msg = "nothing to compact (empty history)"
|
msg = "nothing to compact (empty history)"
|
||||||
raise ValueError(msg)
|
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(
|
summary = await summarize_messages(
|
||||||
self.client,
|
self.client,
|
||||||
messages,
|
messages,
|
||||||
model=self.model,
|
model=self.model,
|
||||||
max_context_tokens=self.agent.max_context_tokens,
|
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}"
|
session_name = name or f"compact-from-{old_id}"
|
||||||
await self.new_session(name=session_name)
|
await self.new_session(name=session_name)
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from msgspec import UNSET
|
|||||||
from plyngent.agent.budget import (
|
from plyngent.agent.budget import (
|
||||||
compact_messages_for_request,
|
compact_messages_for_request,
|
||||||
estimate_messages_tokens,
|
estimate_messages_tokens,
|
||||||
|
measure_messages_tokens,
|
||||||
truncate_tool_result,
|
truncate_tool_result,
|
||||||
)
|
)
|
||||||
from plyngent.agent.loop import run_chat_loop
|
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:
|
def _response(message: AssistantChatMessage) -> ChatCompletionResponse:
|
||||||
return ChatCompletionResponse(
|
return ChatCompletionResponse(
|
||||||
id="1",
|
id="1",
|
||||||
|
|||||||
Reference in New Issue
Block a user