refactor(agent): compact memory/compaction/prompt-cache modules (pass 2, corpus parity)

This commit is contained in:
Teknium
2026-09-02 19:07:30 -07:00
parent 4f20954c5f
commit 21cbb27d89
7 changed files with 542 additions and 1191 deletions
+146 -343
View File
@@ -1,9 +1,7 @@
"""MemoryManager — orchestrates memory providers for the agent.
"""MemoryManager — fans the agent's memory hooks out to registered providers.
Single integration point (run_agent.py) that fans out to registered providers.
The builtin provider is always allowed; only ONE external plugin provider may be
registered at a time — a second is rejected with a warning to prevent tool
schema bloat and conflicting memory backends.
registered at a time (tool-schema bloat, conflicting backends).
"""
from __future__ import annotations
@@ -24,23 +22,19 @@ from tools.registry import tool_error
logger = logging.getLogger(__name__)
# Providers that predate the checkpoint-API attribute are implicitly on the
# historical best-effort contract (API v1).
# Providers that predate the checkpoint-API attribute are on the best-effort v1 contract.
_LEGACY_PRE_COMPRESS_API_VERSION = 1
# How long shutdown_all() waits for in-flight background work to drain before
# abandoning it. Workers are daemon threads, so a wedged provider never blocks
# interpreter exit.
# shutdown_all() drain bound; workers are daemon threads so a wedged provider never
# blocks interpreter exit.
_SYNC_DRAIN_TIMEOUT_S = 5.0
_EXTERNAL_PREFETCH_TIMEOUT_S = 8.0
# ---------------------------------------------------------------------------
# Signature introspection (providers are duck-typed; call shapes vary)
# ---------------------------------------------------------------------------
# -- Signature introspection (providers are duck-typed; call shapes vary) -----
def _signature_params(fn: Callable[..., Any]):
"""Return ``fn``'s parameter mapping, or None when uninspectable (C callables, exotic proxies)."""
"""``fn``'s parameter mapping, or None when uninspectable (C callables, exotic proxies)."""
try:
return inspect.signature(fn).parameters
except (TypeError, ValueError):
@@ -54,9 +48,8 @@ def _has_var_kwargs(params) -> bool:
def _accepts_require_checkpoint(fn: Callable[..., Any]) -> bool:
"""True if ``fn`` can receive the ``require_checkpoint`` keyword.
Checkpoint (v2) providers written against the original docs example use the
bare ``on_pre_compress(self, messages)`` shape; passing the keyword would
raise TypeError, which under ``require_checkpoint=True`` the host would
v2 providers written against the docs example use the bare ``on_pre_compress(self,
messages)`` shape; passing the keyword would raise TypeError, which the host would
re-raise as a checkpoint failure even though the durable write succeeded.
Unreadable signatures conservatively report False.
"""
@@ -67,41 +60,34 @@ def _accepts_require_checkpoint(fn: Callable[..., Any]) -> bool:
return True
param = params.get("require_checkpoint")
return param is not None and param.kind in (
inspect.Parameter.KEYWORD_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
def _ctx_bound(fn: Callable[[], Any]) -> Callable[[], Any]:
"""Bind ``fn`` to the CALLER's contextvars for execution on another thread.
Profile isolation is a ContextVar-scoped HERMES_HOME override; worker threads
start with empty contexts, so an unbound provider resolving config paths or
secrets from a worker would silently land on the default profile.
Profile isolation is a ContextVar-scoped HERMES_HOME override; an unbound provider
on a worker thread would silently resolve paths/secrets against the default profile.
"""
return partial(contextvars.copy_context().run, fn)
# ---------------------------------------------------------------------------
# Tool-schema plumbing
# ---------------------------------------------------------------------------
# -- Tool-schema plumbing -----------------------------------------------------
def normalize_tool_schema(schema: Any) -> Optional[Dict[str, Any]]:
"""Return a bare function-tool dict with a resolvable top-level ``name``, else None.
Providers should return ``{"name", "description", "parameters"}``; some return
the already wrapped OpenAI form. Wrapping that twice yields a ``function`` with
no ``name`` and strict providers (e.g. DeepSeek) reject the ENTIRE request,
so both shapes are normalized and nameless entries can be skipped.
Providers should return ``{"name", "description", "parameters"}`` but some return the
wrapped OpenAI form; wrapping that twice yields a nameless ``function`` and strict
providers (DeepSeek) reject the ENTIRE request, so both shapes are normalized here.
"""
if not isinstance(schema, dict):
return None
if schema.get("type") == "function" and isinstance(schema.get("function"), dict):
schema = schema["function"]
name = schema.get("name", "")
if not name or not isinstance(name, str):
return None
return schema
return schema if name and isinstance(name, str) else None
def memory_provider_tools_enabled(
@@ -119,7 +105,6 @@ def memory_provider_tools_enabled(
return False
if "memory" in enabled_toolsets:
return True
try:
from toolsets import resolve_toolset
@@ -129,22 +114,21 @@ def memory_provider_tools_enabled(
return False
def _tool_name(tool: Any) -> Any:
return tool.get("function", {}).get("name") if isinstance(tool, dict) else None
def memory_provider_tools_exposed(agent: Any) -> bool:
"""Whether external memory-provider tools are exposed on ``agent``.
Same gate as ``inject_memory_provider_tools`` so a provider's
``system_prompt_block()`` and its tool schemas are presented together —
the system prompt must never advertise tools absent from the tool surface.
Same gate as ``inject_memory_provider_tools`` so a provider's ``system_prompt_block()``
never advertises tools absent from the tool surface.
"""
tools = getattr(agent, "tools", None)
memory_tool_present = isinstance(tools, (list, tuple)) and any(
isinstance(tool, dict) and tool.get("function", {}).get("name") == "memory"
for tool in tools
)
return memory_provider_tools_enabled(
getattr(agent, "enabled_toolsets", None),
getattr(agent, "disabled_toolsets", None),
memory_tool_present=memory_tool_present,
memory_tool_present=isinstance(tools, (list, tuple)) and any(_tool_name(t) == "memory" for t in tools),
)
@@ -155,11 +139,7 @@ def inject_memory_provider_tools(agent: Any) -> int:
if not memory_manager or tools is None:
return 0
existing_tool_names = {
tool.get("function", {}).get("name")
for tool in tools
if isinstance(tool, dict)
}
existing_tool_names = {_tool_name(tool) for tool in tools if isinstance(tool, dict)}
if not memory_provider_tools_exposed(agent):
# Say so once: a silent 0 leaves the provider looking "half on" with no
# clue which config key (platform_toolsets / disabled_toolsets) gated it.
@@ -183,8 +163,7 @@ def inject_memory_provider_tools(agent: Any) -> int:
valid_tool_names = getattr(agent, "valid_tool_names", None)
if valid_tool_names is None:
valid_tool_names = set()
agent.valid_tool_names = valid_tool_names
valid_tool_names = agent.valid_tool_names = set()
added = 0
for raw_schema in get_schemas():
@@ -192,8 +171,7 @@ def inject_memory_provider_tools(agent: Any) -> int:
if schema is None:
logger.warning(
"Memory provider returned a tool schema with no resolvable "
"name; skipping to avoid poisoning the request (%r)",
raw_schema,
"name; skipping to avoid poisoning the request (%r)", raw_schema,
)
continue
tool_name = schema["name"]
@@ -203,13 +181,10 @@ def inject_memory_provider_tools(agent: Any) -> int:
valid_tool_names.add(tool_name)
existing_tool_names.add(tool_name)
added += 1
return added
# ---------------------------------------------------------------------------
# Context fencing helpers
# ---------------------------------------------------------------------------
# -- Context fencing helpers --------------------------------------------------
_FENCE_TAG_RE = re.compile(r'</?\s*memory-context\s*>', re.IGNORECASE)
_INTERNAL_CONTEXT_RE = re.compile(
@@ -232,11 +207,10 @@ def sanitize_context(text: str) -> str:
class StreamingContextScrubber:
"""Stateful scrubber for streaming text whose memory-context spans may straddle deltas.
The one-shot ``sanitize_context`` regex needs both tags in one string, so a
span opened in one delta and closed in a later one would leak to the UI.
This state machine holds back partial-tag tails between ``feed()`` calls and
drops everything inside a span. Create a fresh scrubber (or ``reset()``) per
top-level response; call ``flush()`` at end of stream.
``sanitize_context`` needs both tags in one string, so a span split across deltas
would leak to the UI. This holds back partial-tag tails between ``feed()`` calls and
drops everything inside a span. One scrubber (or ``reset()``) per top-level
response; call ``flush()`` at end of stream.
"""
_OPEN_TAG = "<memory-context>"
@@ -267,32 +241,22 @@ class StreamingContextScrubber:
self._buf = buf[-held:] if held else ""
break
buf = buf[idx + len(self._CLOSE_TAG):]
self._in_span = False
else:
idx = self._find_boundary_open_tag(buf)
if idx == -1:
held = (
self._max_pending_open_suffix(buf)
or self._max_partial_suffix(buf, self._OPEN_TAG)
)
held = self._max_pending_open_suffix(buf) or self._max_partial_suffix(buf, self._OPEN_TAG)
self._append_visible(out, buf[:-held] if held else buf)
if held:
self._buf = buf[-held:]
self._buf = buf[-held:] if held else ""
break
if idx > 0:
self._append_visible(out, buf[:idx])
self._append_visible(out, buf[:idx])
buf = buf[idx + len(self._OPEN_TAG):]
self._in_span = True
self._in_span = not self._in_span
return "".join(out)
def flush(self) -> str:
"""Emit the held-back tail at end-of-stream.
Inside an unterminated span the remainder is discarded — leaking partial
memory context is worse than a truncated answer. Otherwise the held tail
was not a real tag and is emitted verbatim.
"""
"""Emit the held-back tail at end-of-stream; inside an unterminated span it is discarded
(leaking partial memory context is worse than a truncated answer)."""
tail = "" if self._in_span else self._buf
self._buf = ""
self._in_span = False
@@ -370,10 +334,10 @@ def _nonblank(text: Any) -> Any:
class MemoryManager:
"""Orchestrates the built-in provider plus at most one external provider.
"""Builtin provider (always first) plus at most one external provider.
The builtin provider is always first. Failures in one provider never block
the other: every fan-out hook logs and swallows per-provider exceptions.
Failures in one provider never block the other: every fan-out hook logs and
swallows per-provider exceptions.
"""
def __init__(self, *, external_prefetch_timeout: Optional[float] = None) -> None:
@@ -381,40 +345,29 @@ class MemoryManager:
self._tool_to_provider: Dict[str, MemoryProvider] = {}
self._has_external: bool = False
self._external_prefetch_timeout = (
_EXTERNAL_PREFETCH_TIMEOUT_S
if external_prefetch_timeout is None
else float(external_prefetch_timeout)
_EXTERNAL_PREFETCH_TIMEOUT_S if external_prefetch_timeout is None else float(external_prefetch_timeout)
)
if self._external_prefetch_timeout <= 0:
raise ValueError("external_prefetch_timeout must be positive")
self._external_prefetch_threads: Dict[str, threading.Thread] = {}
self._external_prefetch_lock = threading.Lock()
# Single-worker background executor for end-of-turn sync/prefetch,
# created lazily so the builtin-only path spawns no threads. One worker
# serializes a provider's writes (turn N lands before turn N+1).
# Single-worker background executor for end-of-turn sync/prefetch, created lazily so
# the builtin-only path spawns no threads; one worker serializes a provider's writes.
self._sync_executor: Optional[ThreadPoolExecutor] = None
self._sync_executor_lock = threading.Lock()
# Futures tracked by durability class ("write" / "prefetch") so shutdown
# can drain FIFO within a bound, then report exactly what it abandoned.
# Futures by durability class ("write" / "prefetch") so shutdown can drain FIFO
# within a bound, then report exactly what it abandoned.
self._background_futures: Dict[Future, str] = {}
self._shutting_down = False
self._shutdown_drain_state: Dict[str, Any] = {
"status": "not_started",
"abandoned_writes": 0,
"abandoned_prefetches": 0,
"active_tasks": 0,
"status": "not_started", "abandoned_writes": 0, "abandoned_prefetches": 0, "active_tasks": 0,
}
# -- Fan-out helper ------------------------------------------------------
def _each_provider(
self,
label: str,
call: Callable[[MemoryProvider], Any],
*,
level: int = logging.DEBUG,
providers: Optional[List[MemoryProvider]] = None,
exc_info: bool = False,
self, label: str, call: Callable[[MemoryProvider], Any], *,
level: int = logging.DEBUG, providers: Optional[List[MemoryProvider]] = None, exc_info: bool = False,
) -> List[Any]:
"""Call ``call(provider)`` for each provider, logging and swallowing failures.
@@ -426,10 +379,7 @@ class MemoryManager:
try:
results.append(call(provider))
except Exception as e:
logger.log(
level, "Memory provider '%s' %s: %s", provider.name, label, e,
exc_info=exc_info,
)
logger.log(level, "Memory provider '%s' %s: %s", provider.name, label, e, exc_info=exc_info)
return results
# -- Registration --------------------------------------------------------
@@ -438,24 +388,20 @@ class MemoryManager:
"""Register a provider; builtin always accepted, only ONE external allowed."""
if provider.name != "builtin":
if self._has_external:
existing = next(
(p.name for p in self._providers if p.name != "builtin"), "unknown"
)
existing = next((p.name for p in self._providers if p.name != "builtin"), "unknown")
logger.warning(
"Rejected memory provider '%s' — external provider '%s' is "
"already registered. Only one external memory provider is "
"allowed at a time. Configure which one via memory.provider "
"in config.yaml.",
provider.name, existing,
"in config.yaml.", provider.name, existing,
)
return
self._has_external = True
self._providers.append(provider)
# Core tool names are reserved: built-ins always win at agent init, so a
# shadowing provider tool would linger in ``_tool_to_provider`` and
# hijack dispatch. Reject it at the door.
# Core tool names are reserved: built-ins always win at agent init, so a shadowing
# provider tool would linger in ``_tool_to_provider`` and hijack dispatch.
from toolsets import _HERMES_CORE_TOOLS
for raw_schema in provider.get_tool_schemas():
@@ -467,25 +413,17 @@ class MemoryManager:
logger.warning(
"Memory provider '%s' tool '%s' shadows a reserved core "
"tool name; registration ignored. Core tools always win — "
"rename the provider's tool to something unique.",
provider.name, tool_name,
"rename the provider's tool to something unique.", provider.name, tool_name,
)
elif tool_name in self._tool_to_provider:
logger.warning(
"Memory tool name conflict: '%s' already registered by %s, "
"ignoring from %s",
tool_name,
self._tool_to_provider[tool_name].name,
provider.name,
"ignoring from %s", tool_name, self._tool_to_provider[tool_name].name, provider.name,
)
else:
self._tool_to_provider[tool_name] = provider
logger.info(
"Memory provider '%s' registered (%d tools)",
provider.name,
len(provider.get_tool_schemas()),
)
logger.info("Memory provider '%s' registered (%d tools)", provider.name, len(provider.get_tool_schemas()))
@property
def providers(self) -> List[MemoryProvider]:
@@ -501,23 +439,15 @@ class MemoryManager:
def build_system_prompt(self) -> str:
"""Join every provider's non-empty ``system_prompt_block()`` with blank lines."""
blocks = self._each_provider(
"system_prompt_block() failed",
lambda p: _nonblank(p.system_prompt_block()),
level=logging.WARNING,
"system_prompt_block() failed", lambda p: _nonblank(p.system_prompt_block()), level=logging.WARNING,
)
return "\n\n".join(b for b in blocks if b)
# -- Prefetch / recall ---------------------------------------------------
@staticmethod
def _strip_skill_scaffolding(text: str) -> Optional[str]:
"""Return memory-worthy user text, or None to skip the turn.
A /skill or /bundle turn expands into a model-facing message embedding the
whole skill body; feeding that to providers pollutes stores/embeddings with
prompt scaffolding. A bare invocation (no instruction) yields None.
"""
return extract_user_instruction_from_skill_message(text)
# A /skill or /bundle turn embeds the whole skill body in the model-facing message;
# providers get just the user's instruction (None for a bare invocation).
_strip_skill_scaffolding = staticmethod(extract_user_instruction_from_skill_message)
def prefetch_all(self, query: str, *, session_id: str = "") -> str:
"""Merge non-empty prefetch context from all providers (failures are non-fatal)."""
@@ -530,13 +460,11 @@ class MemoryManager:
)
return "\n\n".join(p for p in parts if p)
def _prefetch_provider(
self, provider: MemoryProvider, query: str, *, session_id: str = ""
) -> str:
def _prefetch_provider(self, provider: MemoryProvider, query: str, *, session_id: str = "") -> str:
"""Run one provider's prefetch; external providers are bounded by a timeout.
A stuck external call is left running on its daemon thread and the
provider is skipped on subsequent turns until that call returns.
A stuck external call keeps running on its daemon thread and the provider is
skipped on later turns until it returns.
"""
if provider.name == "builtin":
return provider.prefetch(query, session_id=session_id)
@@ -550,19 +478,12 @@ class MemoryManager:
except Exception as exc: # pragma: no cover - re-raised by caller
error_box["value"] = exc
thread = threading.Thread(
target=_ctx_bound(_run),
daemon=True,
name=f"memory-prefetch-{provider.name}",
)
thread = threading.Thread(target=_ctx_bound(_run), daemon=True, name=f"memory-prefetch-{provider.name}")
with self._external_prefetch_lock:
existing = self._external_prefetch_threads.get(provider.name)
if existing is not None:
if existing.is_alive():
logger.debug(
"Memory provider '%s' prefetch is still running; skipping this turn",
provider.name,
)
logger.debug("Memory provider '%s' prefetch is still running; skipping this turn", provider.name)
return ""
self._external_prefetch_threads.pop(provider.name, None)
self._external_prefetch_threads[provider.name] = thread
@@ -572,9 +493,7 @@ class MemoryManager:
if thread.is_alive():
logger.warning(
"Memory provider '%s' prefetch timed out after %.1fs; skipping it until "
"the stuck call returns",
provider.name,
self._external_prefetch_timeout,
"the stuck call returns", provider.name, self._external_prefetch_timeout,
)
return ""
@@ -588,22 +507,18 @@ class MemoryManager:
def describe_recall(self) -> str:
"""Deterministic recall indicator line (e.g. ``"🧠 Provider — recalled 3 memories"``).
Call right after :meth:`prefetch_all` so the user SEES memory was used
regardless of whether the model mentions it. Returns ``""`` when no
provider injected memory this turn, so callers can emit unconditionally.
Call right after :meth:`prefetch_all` so the user SEES memory was used regardless
of whether the model mentions it; ``""`` when no provider injected memory.
"""
segments: List[str] = []
for status in self._each_provider(
"recall_status failed (non-fatal)", lambda p: p.recall_status()
):
for status in self._each_provider("recall_status failed (non-fatal)", lambda p: p.recall_status()):
if status is None:
continue
if status.count == 1:
detail = "recalled 1 memory"
elif status.count > 1:
detail = f"recalled {status.count} memories"
else:
# count <= 0 → content injected but no discrete count (reflect).
else: # content injected but no discrete count (reflect)
detail = "recalled relevant memory"
segments.append(f"{status.glyph} {status.provider_label} — {detail}")
return " ".join(segments)
@@ -632,19 +547,14 @@ class MemoryManager:
return params is None or _has_var_kwargs(params) or "messages" in params
def sync_all(
self,
user_content: str,
assistant_content: str,
*,
session_id: str = "",
messages: Optional[List[Dict[str, Any]]] = None,
self, user_content: str, assistant_content: str, *,
session_id: str = "", messages: Optional[List[Dict[str, Any]]] = None,
) -> None:
"""Sync a completed turn to all providers on the background worker.
Never inline: a provider's ``sync_turn`` may block on a network/daemon
call for minutes, which kept ``run_conversation`` open after the user saw
the response. The single worker also serializes writes so turn N lands
before turn N+1 without provider-side ordering logic.
Never inline: a provider's ``sync_turn`` may block for minutes, which kept
``run_conversation`` open after the user saw the response. The single worker
also serializes writes so turn N lands before turn N+1.
"""
providers = list(self._providers)
clean_user_content = self._strip_skill_scaffolding(user_content) if providers else None
@@ -658,9 +568,7 @@ class MemoryManager:
provider.sync_turn(clean_user_content, assistant_content, **kwargs)
self._submit_background(
lambda: self._each_provider(
"sync_turn failed", _sync, level=logging.WARNING, providers=providers
)
lambda: self._each_provider("sync_turn failed", _sync, level=logging.WARNING, providers=providers)
)
# -- Background dispatch -------------------------------------------------
@@ -668,41 +576,33 @@ class MemoryManager:
def _submit_background(self, fn, *, kind: str = "write") -> None:
"""Queue ``fn`` on the serialized worker and track its durability class.
The callable runs under the caller's contextvars (see ``_ctx_bound``).
If the executor is unavailable outside shutdown, fall back to running
inline — the historical fail-safe (slow but correct).
Runs under the caller's contextvars (``_ctx_bound``). If the executor is
unavailable outside shutdown, run inline — the historical fail-safe.
"""
fn = _ctx_bound(fn)
def _run_inline() -> None:
try:
fn()
except Exception as e: # pragma: no cover - fn guards internally
logger.debug("Inline memory background task failed: %s", e)
executor = self._get_sync_executor()
if executor is None:
if self._shutting_down:
logger.warning("Memory manager is shutting down; rejecting late %s task", kind)
return
_run_inline()
return
future = None
try:
# Submit+track atomically with the shutdown snapshot. The callback is
# attached outside the lock: an already-completed future invokes
# callbacks synchronously.
# Submit+track atomically with the shutdown snapshot. The callback is attached
# outside the lock: an already-completed future invokes callbacks synchronously.
with self._sync_executor_lock:
if self._shutting_down:
logger.warning("Memory manager is shutting down; rejecting late %s task", kind)
return
future = executor.submit(fn)
self._background_futures[future] = kind
future.add_done_callback(self._forget_background_future)
if executor is not None:
future = executor.submit(fn)
self._background_futures[future] = kind
except RuntimeError:
if self._shutting_down:
logger.warning("Memory manager shut down during %s submission; task rejected", kind)
return
_run_inline()
if future is not None:
future.add_done_callback(self._forget_background_future)
return
try:
fn()
except Exception as e: # pragma: no cover - fn guards internally
logger.debug("Inline memory background task failed: %s", e)
def _forget_background_future(self, future: Future) -> None:
with self._sync_executor_lock:
@@ -722,21 +622,16 @@ class MemoryManager:
# Daemon workers: a provider wedged on a network call must
# never block interpreter exit.
from tools.daemon_pool import DaemonThreadPoolExecutor
self._sync_executor = DaemonThreadPoolExecutor(
max_workers=1,
thread_name_prefix="mem-sync",
)
self._sync_executor = DaemonThreadPoolExecutor(max_workers=1, thread_name_prefix="mem-sync")
except Exception as e: # pragma: no cover - resource exhaustion
logger.warning("Failed to create memory sync executor: %s", e)
return None
return self._sync_executor
def flush_pending(self, timeout: Optional[float] = None) -> bool:
"""Block until queued sync/prefetch work has drained.
"""Block until queued sync/prefetch work has drained (False on timeout).
With a single worker, a sentinel task completing proves every earlier
task ran. Returns True when drained within ``timeout`` (or no executor
exists), False on timeout.
With a single worker, a sentinel task completing proves every earlier task ran.
"""
executor = self._sync_executor
if executor is None:
@@ -744,8 +639,7 @@ class MemoryManager:
try:
fut = executor.submit(lambda: None)
except RuntimeError:
# Executor already shut down — nothing pending.
return True
return True # executor already shut down — nothing pending
try:
fut.result(timeout=timeout)
return True
@@ -757,8 +651,7 @@ class MemoryManager:
def get_all_tool_schemas(self) -> List[Dict[str, Any]]:
"""Collect deduplicated tool schemas from all providers.
Reserved core tool names are skipped: :meth:`add_provider` refuses to
route them, so the manager must not advertise a schema it never routes.
Reserved core tool names are skipped: :meth:`add_provider` refuses to route them.
"""
from toolsets import _HERMES_CORE_TOOLS
@@ -771,8 +664,7 @@ class MemoryManager:
if schema is None:
logger.warning(
"Memory provider '%s' returned a tool schema with "
"no resolvable name; skipping (%r)",
provider.name, raw_schema,
"no resolvable name; skipping (%r)", provider.name, raw_schema,
)
continue
name = schema["name"]
@@ -791,9 +683,7 @@ class MemoryManager:
"""Check if any provider handles this tool."""
return tool_name in self._tool_to_provider
def handle_tool_call(
self, tool_name: str, args: Dict[str, Any], **kwargs
) -> str:
def handle_tool_call(self, tool_name: str, args: Dict[str, Any], **kwargs) -> str:
"""Route a tool call to its provider; returns a JSON string (tool_error on failure)."""
provider = self._tool_to_provider.get(tool_name)
if provider is None:
@@ -801,45 +691,30 @@ class MemoryManager:
try:
return provider.handle_tool_call(tool_name, args, **kwargs)
except Exception as e:
logger.error(
"Memory provider '%s' handle_tool_call(%s) failed: %s",
provider.name, tool_name, e,
)
logger.error("Memory provider '%s' handle_tool_call(%s) failed: %s", provider.name, tool_name, e)
return tool_error(f"Memory tool '{tool_name}' failed: {e}")
# -- Lifecycle hooks -----------------------------------------------------
def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None:
"""Notify all providers of a new turn (kwargs: remaining_tokens, model, platform, tool_count)."""
self._each_provider(
"on_turn_start failed",
lambda p: p.on_turn_start(turn_number, message, **kwargs),
)
self._each_provider("on_turn_start failed", lambda p: p.on_turn_start(turn_number, message, **kwargs))
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
"""Notify all providers of session end."""
self._each_provider(
"on_session_end failed",
lambda p: p.on_session_end(messages),
level=logging.WARNING,
exc_info=True,
"on_session_end failed", lambda p: p.on_session_end(messages), level=logging.WARNING, exc_info=True,
)
def commit_session_boundary_async(
self,
messages: List[Dict[str, Any]],
*,
new_session_id: str,
parent_session_id: str = "",
reason: str = "new_session",
self, messages: List[Dict[str, Any]], *,
new_session_id: str, parent_session_id: str = "", reason: str = "new_session",
) -> None:
"""Queue old-session extraction + provider rebinding as ONE serialized task.
``on_session_end`` (LLM-bound extraction, seconds) must run strictly
BEFORE ``on_session_switch`` rebinds provider-internal session state; an
ad-hoc thread raced the inline switch and misattributed transcripts. One
task on the single FIFO worker gives both an immediate return and
ordering against every other provider write.
``on_session_end`` (LLM-bound, seconds) must run strictly BEFORE ``on_session_switch``
rebinds provider session state; an ad-hoc thread raced the inline switch and
misattributed transcripts. One FIFO task gives immediate return plus ordering.
"""
if not self._providers:
return
@@ -851,63 +726,38 @@ class MemoryManager:
except Exception as e: # pragma: no cover - on_session_end guards per-provider
logger.warning("Session-boundary extraction failed: %s", e)
try:
self.on_session_switch(
new_session_id,
parent_session_id=parent_session_id,
reset=True,
reason=reason,
)
self.on_session_switch(new_session_id, parent_session_id=parent_session_id, reset=True, reason=reason)
except Exception as e: # pragma: no cover - on_session_switch guards per-provider
logger.warning("Session-boundary switch failed: %s", e)
self._submit_background(_run)
def on_session_switch(
self,
new_session_id: str,
*,
parent_session_id: str = "",
reset: bool = False,
rewound: bool = False,
**kwargs,
self, new_session_id: str, *,
parent_session_id: str = "", reset: bool = False, rewound: bool = False, **kwargs,
) -> None:
"""Notify providers that ``AIAgent.session_id`` rotated without teardown.
Fires on ``/resume``, ``/branch``, ``/reset``, ``/new`` and compression.
``rewound=True`` (``/undo``) means the id is unchanged but the transcript
was truncated.
"""
"""Notify providers that ``AIAgent.session_id`` rotated without teardown
(``/resume``, ``/branch``, ``/reset``, ``/new``, compression). ``rewound=True``
(``/undo``): same id, truncated transcript."""
if not new_session_id:
return
# Forward ``rewound`` only when set: an unconditional ``rewound=False``
# would pollute every provider's **kwargs on the common paths.
# Forward ``rewound`` only when set so it never pollutes providers' **kwargs.
if rewound:
kwargs["rewound"] = True
self._each_provider(
"on_session_switch failed",
lambda p: p.on_session_switch(
new_session_id, parent_session_id=parent_session_id, reset=reset, **kwargs
),
lambda p: p.on_session_switch(new_session_id, parent_session_id=parent_session_id, reset=reset, **kwargs),
)
@staticmethod
def _checkpoint_api_version(provider: MemoryProvider) -> Optional[int]:
"""Provider's advertised pre-compress checkpoint API version; None if unparseable."""
try:
return int(
getattr(
provider,
"pre_compress_checkpoint_api_version",
_LEGACY_PRE_COMPRESS_API_VERSION,
)
)
return int(getattr(provider, "pre_compress_checkpoint_api_version", _LEGACY_PRE_COMPRESS_API_VERSION))
except (TypeError, ValueError):
return None
def supports_pre_compress_checkpoint(
self,
api_version: int = PRE_COMPRESS_CHECKPOINT_API_VERSION,
) -> bool:
def supports_pre_compress_checkpoint(self, api_version: int = PRE_COMPRESS_CHECKPOINT_API_VERSION) -> bool:
"""Return whether an active provider guarantees checkpoint API support."""
return any(
(version := self._checkpoint_api_version(p)) is not None and version >= api_version
@@ -915,21 +765,17 @@ class MemoryManager:
)
def on_pre_compress(
self,
messages: List[Dict[str, Any]],
*,
self, messages: List[Dict[str, Any]], *,
evidence_messages: Optional[List[Dict[str, Any]]] = None,
require_checkpoint: bool = False,
checkpoint_api_version: int = PRE_COMPRESS_CHECKPOINT_API_VERSION,
) -> str:
"""Notify providers before compression; return their combined summary-prompt text.
``messages`` is the raw transcript (the API v1 contract every provider
gets). ``evidence_messages`` is the host-normalized evidence list handed
only to checkpoint (v2+) providers; when omitted they get the raw list.
With ``require_checkpoint``, at least one checkpoint provider must
succeed — its exception propagates so the caller can keep the
uncompressed transcript.
``messages`` is the raw transcript (the v1 contract). ``evidence_messages`` is the
host-normalized list handed only to checkpoint (v2+) providers. With
``require_checkpoint``, at least one checkpoint provider must succeed — its
exception propagates so the caller keeps the uncompressed transcript.
"""
parts = []
checkpoint_succeeded = False
@@ -942,24 +788,14 @@ class MemoryManager:
if is_checkpoint_provider and evidence_messages is not None:
provider_messages = evidence_messages
try:
if is_checkpoint_provider and _accepts_require_checkpoint(
provider.on_pre_compress
):
result = provider.on_pre_compress(
provider_messages,
require_checkpoint=require_checkpoint,
)
else:
# v1 providers, and v2 providers with the bare one-argument
# shape, never see the requirement signal.
if is_checkpoint_provider and _accepts_require_checkpoint(provider.on_pre_compress):
result = provider.on_pre_compress(provider_messages, require_checkpoint=require_checkpoint)
else: # v1 providers and bare-shape v2 providers never see the signal
result = provider.on_pre_compress(provider_messages)
if result and result.strip():
parts.append(result)
except Exception as e:
logger.debug(
"Memory provider '%s' on_pre_compress failed: %s",
provider.name, e,
)
logger.debug("Memory provider '%s' on_pre_compress failed: %s", provider.name, e)
if require_checkpoint and is_checkpoint_provider:
raise
else:
@@ -982,11 +818,7 @@ class MemoryManager:
return "positional" if accepted >= 4 else "legacy"
def on_memory_write(
self,
action: str,
target: str,
content: str,
metadata: Optional[Dict[str, Any]] = None,
self, action: str, target: str, content: str, metadata: Optional[Dict[str, Any]] = None,
) -> None:
"""Notify external providers when the built-in memory tool writes (skips builtin, the source)."""
@@ -1000,23 +832,19 @@ class MemoryManager:
provider.on_memory_write(action, target, content)
self._each_provider(
"on_memory_write failed",
_notify,
providers=[p for p in self._providers if p.name != "builtin"],
"on_memory_write failed", _notify, providers=[p for p in self._providers if p.name != "builtin"],
)
# Actions the bridge mirrors to external providers. Non-mutating tool result
# shapes (errors, staged-for-approval) are filtered by
# ``notify_memory_tool_write`` before reaching a provider.
# Actions mirrored to external providers; non-mutating results (errors, staged) are
# filtered by ``notify_memory_tool_write`` first.
_MIRRORED_MEMORY_ACTIONS = {"add", "replace", "remove"}
@staticmethod
def _memory_tool_result_succeeded(result: Any) -> bool:
"""True only when the built-in memory tool actually committed a write.
Fails closed: non-JSON, non-dict, missing ``success``, or a write staged
for approval all return False so providers never mirror a write that
did not land.
Fails closed (non-JSON, non-dict, missing ``success``, staged for approval) so
providers never mirror a write that did not land.
"""
if isinstance(result, str):
try:
@@ -1028,18 +856,14 @@ class MemoryManager:
return result.get("success") is True and result.get("staged") is not True
def notify_memory_tool_write(
self,
tool_result: Any,
tool_args: Dict[str, Any],
*,
self, tool_result: Any, tool_args: Dict[str, Any], *,
build_metadata: Optional[Callable[[], Dict[str, Any]]] = None,
) -> None:
"""Mirror a built-in memory tool call to external providers.
Gates on a committed write, expands single-op and batched ``operations``
shapes, keeps only mutating actions, and forwards ``old_text`` plus the
per-op provenance from ``build_metadata`` (the loop knows session/task/
tool-call identity the manager does not).
Gates on a committed write, expands single-op and batched ``operations`` shapes,
keeps only mutating actions, and forwards ``old_text`` plus per-op provenance from
``build_metadata`` (the loop knows session/task/tool-call identity; we do not).
"""
if not self._memory_tool_result_succeeded(tool_result):
return
@@ -1047,11 +871,7 @@ class MemoryManager:
target = str(tool_args.get("target") or "memory")
operations = tool_args.get("operations")
if not (isinstance(operations, list) and operations):
operations = [{
"action": tool_args.get("action"),
"content": tool_args.get("content"),
"old_text": tool_args.get("old_text"),
}]
operations = [{k: tool_args.get(k) for k in ("action", "content", "old_text")}]
for op in operations:
if not isinstance(op, dict):
@@ -1064,17 +884,11 @@ class MemoryManager:
old_text = op.get("old_text")
if old_text:
metadata["old_text"] = str(old_text)
self.on_memory_write(
action,
target,
str(op.get("content") or ""),
metadata=metadata,
)
self.on_memory_write(action, target, str(op.get("content") or ""), metadata=metadata)
except Exception as e:
logger.debug("notify_memory_tool_write failed for op %s: %s", action, e)
def on_delegation(self, task: str, result: str, *,
child_session_id: str = "", **kwargs) -> None:
def on_delegation(self, task: str, result: str, *, child_session_id: str = "", **kwargs) -> None:
"""Notify all providers that a subagent completed."""
self._each_provider(
"on_delegation failed",
@@ -1085,10 +899,7 @@ class MemoryManager:
"""Drain the background executor (bounded), then shut providers down in reverse order."""
self._drain_sync_executor()
self._each_provider(
"shutdown failed",
lambda p: p.shutdown(),
level=logging.WARNING,
providers=list(reversed(self._providers)),
"shutdown failed", lambda p: p.shutdown(), level=logging.WARNING, providers=list(reversed(self._providers)),
)
@property
@@ -1113,9 +924,8 @@ class MemoryManager:
if executor is None:
return
# shutdown(wait=False) closes submission without touching the FIFO;
# waiting on the tracked futures lets the worker run every queued
# write/boundary task in order up to the deadline.
# shutdown(wait=False) closes submission without touching the FIFO; waiting on the
# tracked futures lets the worker run every queued task in order up to the deadline.
executor.shutdown(wait=False, cancel_futures=False)
_, pending = wait(tuple(tracked), timeout=_SYNC_DRAIN_TIMEOUT_S)
if not pending:
@@ -1134,18 +944,13 @@ class MemoryManager:
with self._sync_executor_lock:
self._shutdown_drain_state.update(
status="timed_out",
abandoned_writes=abandoned_writes,
abandoned_prefetches=abandoned_prefetches,
active_tasks=active_tasks,
status="timed_out", abandoned_writes=abandoned_writes,
abandoned_prefetches=abandoned_prefetches, active_tasks=active_tasks,
)
logger.warning(
"Memory shutdown drain timed out after %.2fs; abandoning %d queued "
"memory write(s) and %d queued prefetch(es); %d active task(s) remain detached",
_SYNC_DRAIN_TIMEOUT_S,
abandoned_writes,
abandoned_prefetches,
active_tasks,
_SYNC_DRAIN_TIMEOUT_S, abandoned_writes, abandoned_prefetches, active_tasks,
)
def initialize_all(self, session_id: str, **kwargs) -> None:
@@ -1154,7 +959,5 @@ class MemoryManager:
from hermes_constants import get_hermes_home
kwargs["hermes_home"] = str(get_hermes_home())
self._each_provider(
"initialize failed",
lambda p: p.initialize(session_id=session_id, **kwargs),
level=logging.WARNING,
"initialize failed", lambda p: p.initialize(session_id=session_id, **kwargs), level=logging.WARNING,
)
+26 -46
View File
@@ -1,9 +1,8 @@
"""Abstract base class for pluggable memory providers.
Plugins ship in ``plugins/memory/<name>/`` and are activated via ``memory.provider``;
MemoryManager allows only ONE external provider at a time. Lifecycle is driven by
MemoryManager: initialize -> system_prompt_block / prefetch / sync_turn per turn ->
tool dispatch -> shutdown, plus the optional ``on_*`` hooks below.
Plugins ship in ``plugins/memory/<name>/``, activated via ``memory.provider`` (ONE external
provider at a time). Lifecycle, driven by MemoryManager: initialize -> system_prompt_block /
prefetch / sync_turn per turn -> tool dispatch -> shutdown, plus optional ``on_*`` hooks.
"""
from __future__ import annotations
@@ -16,9 +15,8 @@ from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
# v1 = historical implicit contract (best-effort on_pre_compress() with the raw
# message list); v2 = opt-in fail-closed checkpoint (normalized evidence handoff
# + strict-mode failure propagation).
# v1 = best-effort on_pre_compress() with the raw message list; v2 = opt-in fail-closed
# checkpoint (normalized evidence handoff + strict-mode failure propagation).
PRE_COMPRESS_CHECKPOINT_API_VERSION = 2
# Default glyph for recall indicators; providers may use their own brand mark.
@@ -36,10 +34,9 @@ class RecallStatus:
glyph: str = INDICATOR_GLYPH
# Prompts with no semantic signal. Single source of truth for the core prefetch
# gate (turn_context.py, run_agent.py) and provider-side classifiers (honcho).
# Anchored and followed only by whitespace/punctuation, so "k8s"/"yolo"/"note"
# do NOT match while "hi!"/"thanks :)"/"done???" do.
# Prompts with no semantic signal; single source of truth for the core prefetch gate and
# provider-side classifiers. Anchored and followed only by whitespace/punctuation, so
# "k8s"/"yolo"/"note" do NOT match while "hi!"/"thanks :)"/"done???" do.
TRIVIAL_PROMPT_RE = re.compile(
r'^(yes|no|ok|okay|sure|thanks|thank you|y|n|yep|nope|yeah|nah|'
r'hi|hey|hello|yo|sup|'
@@ -50,11 +47,8 @@ TRIVIAL_PROMPT_RE = re.compile(
def is_trivial_prompt(text: Optional[str]) -> bool:
"""True for empty input, slash commands and bare greetings/acknowledgements.
Skipping recall on these saves a blocking network round-trip and keeps
stale user-model context from derailing one-word replies.
"""
"""True for empty input, slash commands and bare greetings/acknowledgements (skipping
recall saves a round-trip and keeps stale context from derailing one-word replies)."""
stripped = (text or "").strip()
if not stripped or stripped.startswith("/"):
return True
@@ -64,8 +58,8 @@ def is_trivial_prompt(text: Optional[str]) -> bool:
class MemoryProvider(ABC):
"""Abstract base class for memory providers."""
# Providers that durably checkpoint every successful on_pre_compress() opt
# in by setting PRE_COMPRESS_CHECKPOINT_API_VERSION; 1 = best-effort legacy.
# Providers that durably checkpoint every successful on_pre_compress() set this to
# PRE_COMPRESS_CHECKPOINT_API_VERSION; 1 = best-effort legacy.
pre_compress_checkpoint_api_version = 1
@property
@@ -84,12 +78,10 @@ class MemoryProvider(ABC):
def initialize(self, session_id: str, **kwargs) -> None:
"""Initialize once at agent startup (connections, resources, threads).
kwargs always include ``hermes_home`` (use it for profile-scoped storage,
never hardcode ``~/.hermes``) and ``platform``. May include
``agent_context`` ("primary" | "subagent" | "cron" | "flush" — skip
writes for non-primary contexts, cron prompts would corrupt user
representations), ``agent_identity`` (profile name), ``agent_workspace``,
``parent_session_id``, ``user_id``, ``user_id_alt``.
kwargs always include ``hermes_home`` (profile-scoped storage; never hardcode
``~/.hermes``) and ``platform``; may include ``agent_context`` ("primary" |
"subagent" | "cron" | "flush" — skip writes for non-primary contexts),
``agent_identity``, ``agent_workspace``, ``parent_session_id``, ``user_id``, ``user_id_alt``.
"""
def unavailable_reason(self) -> str:
@@ -103,11 +95,8 @@ class MemoryProvider(ABC):
return ""
def prefetch(self, query: str, *, session_id: str = "") -> str:
"""Formatted recall context for the upcoming turn ("" if none).
Must be fast — do the recall in the background and return cached
results. ``session_id`` scopes concurrent sessions (gateway, cached agents).
"""
"""Formatted recall context for the upcoming turn ("" if none). Must be fast — recall
in the background and return cached results; ``session_id`` scopes concurrent sessions."""
return ""
def queue_prefetch(self, query: str, *, session_id: str = "") -> None:
@@ -161,15 +150,10 @@ class MemoryProvider(ABC):
rewound: bool = False,
**kwargs,
) -> None:
"""session_id reassigned mid-process (/resume, /branch, /reset, /new,
gateway equivalents, context compression) without a provider teardown.
Update or reset per-session state cached in ``initialize()`` so later
writes land in the right record. ``parent_session_id`` carries lineage
("" when none). ``reset`` is True only for a genuinely new conversation
(/reset, /new) — flush per-session buffers. ``rewound``: same id but the
transcript was truncated, so invalidate per-turn document state.
"""
"""session_id reassigned mid-process (/resume, /branch, /reset, /new, compression)
without teardown: rebind per-session state so later writes land in the right record.
``reset`` is True only for a genuinely new conversation (flush buffers); ``rewound``:
same id but the transcript was truncated."""
def on_pre_compress(self, messages: List[Dict[str, Any]]) -> str:
"""Extract insights from ``messages`` about to be compressed; the returned
@@ -182,14 +166,10 @@ class MemoryProvider(ABC):
result); the subagent itself has no provider session (skip_memory=True)."""
def get_config_schema(self) -> List[Dict[str, Any]]:
"""Setup fields for ``hermes memory setup`` ([] if none).
Each field: ``key``, ``description``, optional ``secret`` (goes to .env),
``required``, ``default``, ``choices``, ``type`` (text | integer |
number | boolean), ``minimum`` / ``maximum`` / ``step`` (numeric),
``url`` (where to get the credential), ``env_var`` (explicit secret env
var; default auto-generated).
"""
"""Setup fields for ``hermes memory setup`` ([] if none): ``key``, ``description``,
optional ``secret`` (goes to .env), ``required``, ``default``, ``choices``, ``type``
(text | integer | number | boolean), ``minimum``/``maximum``/``step``, ``url``,
``env_var`` (explicit secret env var; default auto-generated)."""
return []
def save_config(self, values: Dict[str, Any], hermes_home: str) -> None:
+139 -266
View File
@@ -1,8 +1,7 @@
"""Message and tool-payload sanitization helpers.
"""Message and tool-payload sanitization helpers (pure; documented in-place mutation).
Pure functions that walk OpenAI-format message lists and structured payloads,
repairing or stripping characters that would crash ``json.dumps`` in the OpenAI
SDK or be rejected upstream. Stateless except for documented in-place mutation;
Walk OpenAI-format message lists and structured payloads, repairing or stripping
characters that would crash ``json.dumps`` in the OpenAI SDK or be rejected upstream.
``run_agent`` re-exports them for old imports.
"""
@@ -16,12 +15,11 @@ from typing import Any, Callable
logger = logging.getLogger(__name__)
# Lone surrogate code points are invalid in UTF-8 and crash json.dumps inside
# the OpenAI SDK. Also used by run_agent and the CLI for paste scrubbing.
# Lone surrogates are invalid UTF-8 and crash json.dumps in the OpenAI SDK; also used for
# CLI paste scrubbing.
_SURROGATE_RE = re.compile(r'[\ud800-\udfff]')
# Message keys handled explicitly by _sanitize_messages; every OTHER key is
# swept generically (reasoning, reasoning_content, reasoning_details, ...).
# Keys handled explicitly by _sanitize_messages; every OTHER key is swept generically.
_MESSAGE_CORE_KEYS = frozenset({"content", "name", "tool_calls", "role"})
@@ -73,11 +71,9 @@ def _sanitize_structure(payload: Any, fix: Callable[[str], str]) -> bool:
def _sanitize_messages(messages: list, fix: Callable[[str], str], *, deep: bool) -> bool:
"""Apply ``fix`` to the string fields of every message dict in-place.
Covers content / content-part text, name, tool_call function arguments, and
every non-core top-level str field (reasoning_content etc.) so retries don't
fail on a non-content field. ``deep=True`` additionally covers tool_call ids,
function names, and NESTED non-core fields (``reasoning_details`` arrays from
byte-level reasoning models such as xiaomi/mimo, kimi, glm).
Covers content / content-part text, name, tool_call arguments, and every non-core
top-level str field. ``deep=True`` additionally covers tool_call ids, function names,
and NESTED non-core fields (``reasoning_details`` arrays from byte-level reasoning models).
"""
found = False
for msg in messages:
@@ -139,11 +135,8 @@ def _sanitize_tools_non_ascii(tools: list) -> bool:
def _escape_invalid_chars_in_json_strings(raw: str) -> str:
"""Escape literal control chars (0x00-0x1F) inside JSON string values as ``\\uXXXX``.
Complements ``json.loads(strict=False)`` in ``_repair_tool_call_arguments``
for llama.cpp-style output that mixes control chars with other malformations.
"""
"""Escape literal control chars (0x00-0x1F) inside JSON string values as ``\\uXXXX``
(for llama.cpp-style output mixing control chars with other malformations)."""
out: list[str] = []
in_string = False
i = 0
@@ -165,9 +158,8 @@ def _escape_invalid_chars_in_json_strings(raw: str) -> str:
return "".join(out)
# When a repair rewrites arguments to "{}", the WARNING log is the last surviving
# copy of content that can hold real user data (e.g. a truncated write_file's
# streamed file content). Bound it here rather than at a short preview.
# When a repair rewrites arguments to "{}", the WARNING log is the last surviving copy of
# content that can hold real user data (a truncated write_file), so bound it generously.
_FULL_ARGS_LOG_BOUND = 100_000
@@ -180,10 +172,8 @@ def _loads_ok(text: str) -> bool:
def _repair_tool_call_arguments(raw_args: str, tool_name: str = "?") -> str:
"""Repair malformed tool_call argument JSON (truncation, trailing commas,
Python ``None``, literal control chars); returns ``"{}"`` if unrepairable so
the request succeeds instead of crashing the session. Repairs log at WARNING.
"""
"""Repair malformed tool_call argument JSON (truncation, trailing commas, Python ``None``,
control chars); ``"{}"`` if unrepairable so the request succeeds. Repairs log at WARNING."""
raw_stripped = raw_args.strip() if isinstance(raw_args, str) else ""
if not raw_stripped:
@@ -194,22 +184,17 @@ def _repair_tool_call_arguments(raw_args: str, tool_name: str = "?") -> str:
logger.warning("Sanitized Python-None tool_call arguments for %s", tool_name)
return "{}"
# Pass 0: strict=False accepts literal control chars inside strings (the
# most common local-model case) and re-serialises to wire-valid JSON.
# Pass 0: strict=False accepts literal control chars inside strings (the most common
# local-model case) and re-serialises to wire-valid JSON.
try:
parsed = json.loads(raw_stripped, strict=False)
reserialised = json.dumps(parsed, separators=(",", ":"))
reserialised = json.dumps(json.loads(raw_stripped, strict=False), separators=(",", ":"))
if reserialised != raw_stripped:
logger.warning(
"Repaired unescaped control chars in tool_call arguments for %s",
tool_name,
)
logger.warning("Repaired unescaped control chars in tool_call arguments for %s", tool_name)
return reserialised
except (json.JSONDecodeError, TypeError, ValueError):
pass
# Passes 1-3: strip trailing commas, close unclosed structures, then trim
# excess closers (bounded).
# Passes 1-3: strip trailing commas, close unclosed structures, trim excess closers (bounded).
fixed = re.sub(r',\s*([}\]])', r'\1', raw_stripped)
fixed += '}' * max(0, fixed.count('{') - fixed.count('}'))
fixed += ']' * max(0, fixed.count('[') - fixed.count(']'))
@@ -224,25 +209,20 @@ def _repair_tool_call_arguments(raw_args: str, tool_name: str = "?") -> str:
break
if _loads_ok(fixed):
logger.warning(
"Repaired malformed tool_call arguments for %s: %s → %s",
tool_name, raw_stripped[:80], fixed[:80],
)
logger.warning("Repaired malformed tool_call arguments for %s: %s → %s", tool_name, raw_stripped[:80], fixed[:80])
return fixed
# Pass 4: escape control chars inside strings (strict=False alone fails
# when other malformations are present too), then retry.
# Pass 4: escape control chars inside strings (strict=False alone fails when other
# malformations are present too), then retry.
escaped = _escape_invalid_chars_in_json_strings(fixed)
if escaped != fixed and _loads_ok(escaped):
logger.warning(
"Repaired control-char-laced tool_call arguments for %s: %s → %s",
tool_name, raw_stripped[:80], escaped[:80],
"Repaired control-char-laced tool_call arguments for %s: %s → %s", tool_name, raw_stripped[:80], escaped[:80],
)
return escaped
logger.warning(
"Unrepairable tool_call arguments for %s — "
"replaced with empty object (was: %s)",
"Unrepairable tool_call arguments for %s — replaced with empty object (was: %s)",
tool_name, raw_stripped[:_FULL_ARGS_LOG_BOUND],
)
return "{}"
@@ -251,43 +231,31 @@ def _repair_tool_call_arguments(raw_args: str, tool_name: str = "?") -> str:
def close_interrupted_tool_sequence(messages: list, final_response: Any = None) -> bool:
"""Append a synthetic assistant turn when an interrupted tail is a tool result.
A transcript ending on a raw ``tool`` message makes the next user message
land as ``tool → user`` — a role-alternation violation strict providers
(Gemini, Claude) answer by hallucinating a continuation and dropping prior
context. Mutates in place; returns True if a closing turn was appended.
A transcript ending on a raw ``tool`` message makes the next user message land as
``tool → user`` — an alternation violation strict providers (Gemini, Claude) answer by
hallucinating a continuation. Mutates in place; True if a closing turn was appended.
"""
if not messages:
return False
last = messages[-1]
last = messages[-1] if messages else None
if not isinstance(last, dict) or last.get("role") != "tool":
return False
text = final_response if isinstance(final_response, str) else ""
from agent.message_metadata import append_message
append_message(messages, {
"role": "assistant",
"content": text.strip() or "Operation interrupted.",
})
append_message(messages, {"role": "assistant", "content": text.strip() or "Operation interrupted."})
return True
def serialized_messages_bytes(messages: list) -> int:
"""Exact serialized byte size of the ``messages`` payload (HTTP 413 recovery).
A 413 is a BYTE-size error, but the token estimator prices an image at a
flat per-image cost, so it cannot score recovery from an image-dominated
413. This measures what the provider actually rejected, identically before
and after each pass. Non-serializable values fall back to ``str()`` so a
malformed message can never crash recovery.
A 413 is a BYTE-size error, but the token estimator prices images at a flat cost, so
it cannot score recovery from an image-dominated 413. Non-serializable values fall
back to ``str()`` so a malformed message can never crash recovery.
"""
if not isinstance(messages, list) or not messages:
return 0
try:
return len(
json.dumps(
messages, ensure_ascii=False, separators=(",", ":"), default=str
).encode("utf-8")
)
return len(json.dumps(messages, ensure_ascii=False, separators=(",", ":"), default=str).encode("utf-8"))
except (TypeError, ValueError):
return sum(len(str(m)) for m in messages)
@@ -298,12 +266,10 @@ _IMAGE_PART_TYPES = {"image_url", "image", "input_image"}
def _strip_images_from_messages(messages: list) -> bool:
"""Remove image content parts from all messages in-place (server rejected images).
Preserves alternation invariants: ``tool`` messages and assistant messages
carrying ``tool_calls`` whose content was entirely images are replaced with a
placeholder, NOT deleted (deleting orphans the paired ``tool_call_id`` →
HTTP 400); other now-empty messages are dropped. Any rewritten message also
loses its ``api_content`` sidecar — it carries the exact bytes previously
sent, i.e. the images being removed. Returns True if any image parts were removed.
``tool`` messages and assistant messages carrying ``tool_calls`` whose content was
entirely images get a placeholder, NOT deleted (deleting orphans the paired
``tool_call_id`` → HTTP 400); other now-empty messages are dropped. Rewritten messages
lose their ``api_content`` sidecar (it carries the images being removed).
"""
from agent.turn_context import drop_stale_api_content
@@ -315,10 +281,7 @@ def _strip_images_from_messages(messages: list) -> bool:
content = msg.get("content")
if not isinstance(content, list):
continue
new_parts = [
part for part in content
if not (isinstance(part, dict) and part.get("type") in _IMAGE_PART_TYPES)
]
new_parts = [p for p in content if not (isinstance(p, dict) and p.get("type") in _IMAGE_PART_TYPES)]
if len(new_parts) < len(content):
found = True
if new_parts:
@@ -333,36 +296,26 @@ def _strip_images_from_messages(messages: list) -> bool:
return found
# Provider error bodies (lowercased substring match) meaning "image/multimodal
# input unsupported" — the loop then strips images and retries text-only instead
# of cascading into compression / context-too-large recovery or wedging on retries.
# Provider error bodies (lowercased substring match) meaning "image/multimodal input
# unsupported" — the loop then strips images and retries text-only instead of cascading
# into compression / context-too-large recovery or wedging on retries.
_IMAGE_REJECTION_PHRASES = (
"only 'text' content type is supported",
"only text content type is supported",
"image_url is not supported",
"image content is not supported",
"multimodal is not supported",
"multimodal content is not supported",
"multimodal input is not supported",
"vision is not supported",
"vision input is not supported",
"does not support images",
"does not support image input",
"does not support multimodal",
"does not support vision",
"model does not support image",
"only 'text' content type is supported", "only text content type is supported",
"image_url is not supported", "image content is not supported",
"multimodal is not supported", "multimodal content is not supported", "multimodal input is not supported",
"vision is not supported", "vision input is not supported",
"does not support images", "does not support image input", "does not support multimodal",
"does not support vision", "model does not support image",
# DashScope-style gateways reject non-text blocks with this generic body.
"unexpected item type in content",
# ChatGPT-account Codex backend rejects data:image URLs in input_image;
# keyed on the field-path apostrophe so other URL errors don't false-trip.
"image_url'. expected",
# ChatGPT-account Codex wording for corrupt/unsupported native image payloads.
"image data you provided does not represent a valid image",
# ChatGPT-account Codex backend rejects data:image URLs in input_image; keyed on the
# field-path apostrophe so other URL errors don't false-trip. Second: its wording for
# corrupt/unsupported native image payloads.
"image_url'. expected", "image data you provided does not represent a valid image",
# DeepSeek's text-only request-body variant error.
"unknown variant `image_url`, expected `text`",
"unknown variant image_url, expected text",
# OpenRouter HTTP 404 when no upstream endpoint accepts image input (passes
# the 4xx gate; without this the gateway queue wedges behind the stuck turn).
"unknown variant `image_url`, expected `text`", "unknown variant image_url, expected text",
# OpenRouter HTTP 404 when no upstream endpoint accepts image input (passes the 4xx
# gate; without this the gateway queue wedges behind the stuck turn).
"no endpoints found that support image input",
# Kimi/Moonshot et al. reject truncated/corrupt image bytes baked into history.
"failed to decode image",
@@ -376,45 +329,25 @@ def _looks_like_image_content_rejection(error_body: str) -> bool:
__all__ = [
"_SURROGATE_RE",
"close_interrupted_tool_sequence",
"_sanitize_surrogates",
"_sanitize_structure_surrogates",
"_sanitize_messages_surrogates",
"_escape_invalid_chars_in_json_strings",
"_repair_tool_call_arguments",
"_strip_non_ascii",
"_sanitize_messages_non_ascii",
"_sanitize_tools_non_ascii",
"_strip_images_from_messages",
"_sanitize_structure_non_ascii",
"_SURROGATE_RE", "close_interrupted_tool_sequence",
"_sanitize_surrogates", "_sanitize_structure_surrogates", "_sanitize_messages_surrogates",
"_escape_invalid_chars_in_json_strings", "_repair_tool_call_arguments",
"_strip_non_ascii", "_sanitize_messages_non_ascii", "_sanitize_tools_non_ascii",
"_strip_images_from_messages", "_sanitize_structure_non_ascii",
# call_id policy owners
"deterministic_call_id",
"coalesce_tool_call_id",
"tool_call_id_variants",
"tool_result_id_variants",
"uniquify_tool_call_ids",
"deterministic_call_id", "coalesce_tool_call_id", "tool_call_id_variants",
"tool_result_id_variants", "uniquify_tool_call_ids",
# reasoning_content policy owners
"reasoning_echo_family",
"matches_reasoning_echo_family",
"needs_reasoning_echo",
"stale_thinking_reaches_wire",
"apply_reasoning_content_policy",
"reapply_reasoning_echo",
"reasoning_echo_family", "matches_reasoning_echo_family", "needs_reasoning_echo",
"stale_thinking_reaches_wire", "apply_reasoning_content_policy", "reapply_reasoning_echo",
]
# ---------------------------------------------------------------------------
# call_id policy — single owner for hash synthesis, ``call_id or id``
# coalescing, and duplicate-id repair.
#
# NOT consolidated on purpose: agent/transports/codex_event_projector's
# _deterministic_call_id maps codex app-server ITEM ids, not chat tool-call
# content; merging would change ids and invalidate caches.
#
# HARD INVARIANT: everything here stays deterministic (never uuid4) and
# byte-identical for existing inputs — these ids feed prompt-cache prefixes.
# ---------------------------------------------------------------------------
# -- call_id policy: hash synthesis, ``call_id or id`` coalescing, duplicate-id repair ----
# NOT merged with codex_event_projector._deterministic_call_id (maps app-server ITEM ids,
# not chat tool-call content; merging would change ids and invalidate caches).
# HARD INVARIANT: deterministic (never uuid4) and byte-identical for existing inputs —
# these ids feed prompt-cache prefixes.
def _tc_field(tc: Any, key: str) -> Any:
@@ -423,19 +356,14 @@ def _tc_field(tc: Any, key: str) -> Any:
def deterministic_call_id(fn_name: str, arguments: str, index: int = 0) -> str:
"""Deterministic call_id fallback when the API omits one (random ids would
make every prefix unique and break prompt caching)."""
"""Deterministic call_id fallback when the API omits one (random ids would break caching)."""
seed = f"{fn_name}:{arguments}:{index}"
digest = hashlib.sha256(seed.encode("utf-8", errors="replace")).hexdigest()[:12]
return f"call_{digest}"
return f"call_{hashlib.sha256(seed.encode('utf-8', errors='replace')).hexdigest()[:12]}"
def _expand_tool_id_variants(values: tuple[Any, ...]) -> frozenset[str]:
"""Every wire spelling of one tool-call identifier.
Responses bridges may expose the pairing id and response-item id separately
or encode both as ``call_id|response_item_id``; all are aliases for ONE call.
"""
"""Every wire spelling of one tool-call identifier: Responses bridges may expose the pairing
id and response-item id separately or as ``call_id|response_item_id``; all alias ONE call."""
variants: set[str] = set()
for raw in values:
value = raw.strip() if isinstance(raw, str) else ""
@@ -449,9 +377,7 @@ def _expand_tool_id_variants(values: tuple[Any, ...]) -> frozenset[str]:
def tool_call_id_variants(tc: Any) -> frozenset[str]:
"""Return all pairing-id variants carried by a tool-call entry."""
return _expand_tool_id_variants(
(_tc_field(tc, "call_id"), _tc_field(tc, "id"), _tc_field(tc, "response_item_id"))
)
return _expand_tool_id_variants(tuple(_tc_field(tc, k) for k in ("call_id", "id", "response_item_id")))
def tool_result_id_variants(tool_call_id: Any) -> frozenset[str]:
@@ -460,11 +386,10 @@ def tool_result_id_variants(tool_call_id: Any) -> frozenset[str]:
def coalesce_tool_call_id(tc: Any) -> str:
"""Effective call id of a tool_call entry (dict or object).
"""Effective call id of a tool_call entry (dict or object); ``""`` when none.
Codex Responses calls carry ``call_id`` (authoritative pairing key), Chat
Completions carry ``id`` only, and bridge ids may be ``call_id|response_item_id``.
Returns ``""`` when neither is set.
Codex Responses calls carry ``call_id`` (authoritative pairing key), Chat Completions
carry ``id`` only, and bridge ids may be ``call_id|response_item_id``.
"""
for raw in (_tc_field(tc, "call_id"), _tc_field(tc, "id")):
value = raw.strip() if isinstance(raw, str) else ""
@@ -476,21 +401,18 @@ def coalesce_tool_call_id(tc: Any) -> str:
def uniquify_tool_call_ids(tool_calls: list) -> list:
"""Ensure every tool call in one assistant turn has a distinct id.
Some providers reuse one id across calls in a batch; the pre-API sanitizer
then keeps only the first call/result pair per id and strict providers
reject duplicates outright. First occurrence keeps its id; later collisions
get a deterministic ``<id>_d<n>`` suffix (never uuid4 — cache-prefix
stability). Mutates entries in place (SDK models / SimpleNamespace / dicts)
and returns the same list. Blank ids are left for the deterministic fallback
in ``build_assistant_message``.
Some providers reuse one id across a batch; the pre-API sanitizer then keeps only the
first call/result pair per id and strict providers reject duplicates. First occurrence
keeps its id; later collisions get a deterministic ``<id>_d<n>`` suffix (never uuid4 —
cache-prefix stability). Mutates entries in place (SDK models / SimpleNamespace /
dicts). Blank ids are left for the deterministic fallback in ``build_assistant_message``.
"""
seen: set = set()
for tc in tool_calls or []:
# Same coalescing rule as coalesce_tool_call_id, tolerant of non-string ids.
raw = _tc_field(tc, "call_id") or _tc_field(tc, "id") or ""
raw = raw.strip() if isinstance(raw, str) else ""
# Composite Responses ids ("call_x|fc_y") collide on the call half —
# that's the pairing key providers enforce per turn.
# Composite Responses ids ("call_x|fc_y") collide on the call half — the pairing key.
cid = raw.split("|", 1)[0]
if not cid:
continue
@@ -504,8 +426,7 @@ def uniquify_tool_call_ids(tool_calls: list) -> list:
seen.add(new_id)
def _renamed(value):
# Keep a composite id's response-item half so the provider's real
# fc_/item id survives the rename.
# Keep a composite id's response-item half so the provider's fc_/item id survives.
if isinstance(value, str) and "|" in value:
return f"{new_id}|{value.split('|', 1)[1]}"
return new_id
@@ -520,12 +441,9 @@ def uniquify_tool_call_ids(tool_calls: list) -> list:
if getattr(tc, "call_id", None):
tc.call_id = new_id
except Exception:
logger.warning(
"Could not uniquify duplicate tool call id %s", cid
)
logger.warning("Could not uniquify duplicate tool call id %s", cid)
continue
_fn = _tc_field(tc, "function")
_fn_name = (_fn.get("name") if isinstance(_fn, dict) else getattr(_fn, "name", None)) or "?"
_fn_name = _tc_field(_tc_field(tc, "function"), "name") or "?"
logger.warning(
"Model reused tool call id %s within one turn; renamed the "
"duplicate to %s (tool=%s) to keep call/result pairing "
@@ -534,44 +452,28 @@ def uniquify_tool_call_ids(tool_calls: list) -> list:
return tool_calls
# ---------------------------------------------------------------------------
# reasoning_content policy — single owner. The POLICY (which provider direction
# gets strip vs re-pad) lives here as one rule table + apply functions; adapters
# keep only SYNTAX mapping (e.g. anthropic_adapter → thinking block).
#
# -- reasoning_content policy: single owner of strip-vs-re-pad; adapters keep only SYNTAX --
# require-side (echo-back enforced; replays 400 without the field):
# kimi — provider kimi-coding/kimi-coding-cn, or host api.kimi.com /
# moonshot.ai / moonshot.cn. Host-driven on purpose: aggregators
# re-exporting kimi models reject the echo.
# deepseek — provider "deepseek", model contains "deepseek", or host
# api.deepseek.com. V4 rejects empty-string pads → " " single space.
# kimi — provider kimi-coding/kimi-coding-cn, or host api.kimi.com / moonshot.ai /
# moonshot.cn. Host-driven on purpose: aggregators re-exporting kimi reject it.
# deepseek — provider "deepseek", model contains "deepseek", or host api.deepseek.com.
# V4 rejects empty-string pads → " " single space.
# mimo — provider "xiaomi", model contains "mimo", or host *.xiaomimimo.com.
# strict side (field rejected 400/422 "Extra inputs are not permitted"):
# everyone else — Mistral, Cerebras, Groq, SambaNova, … Strip the key
# entirely, even a single-space pad.
# ---------------------------------------------------------------------------
# strict side (field rejected 400/422 "Extra inputs are not permitted"): everyone else —
# Mistral, Cerebras, Groq, SambaNova, … Strip the key entirely, even a one-space pad.
_REASONING_ECHO_RULES: tuple = (
# (family, exact providers (raw), exact providers (lowered),
# model substrings (lowered), base_url hosts)
("kimi", frozenset({"kimi-coding", "kimi-coding-cn"}), frozenset(), (),
("api.kimi.com", "moonshot.ai", "moonshot.cn")),
("deepseek", frozenset(), frozenset({"deepseek"}), ("deepseek",),
("api.deepseek.com",)),
("mimo", frozenset(), frozenset({"xiaomi"}), ("mimo",),
("api.xiaomimimo.com", "xiaomimimo.com")),
# (family, exact providers (raw), exact providers (lowered), model substrings (lowered), hosts)
("kimi", frozenset({"kimi-coding", "kimi-coding-cn"}), frozenset(), (), ("api.kimi.com", "moonshot.ai", "moonshot.cn")),
("deepseek", frozenset(), frozenset({"deepseek"}), ("deepseek",), ("api.deepseek.com",)),
("mimo", frozenset(), frozenset({"xiaomi"}), ("mimo",), ("api.xiaomimimo.com", "xiaomimimo.com")),
)
_REASONING_ECHO_RULE_BY_FAMILY = {rule[0]: rule for rule in _REASONING_ECHO_RULES}
def matches_reasoning_echo_family(
family: str, provider: Any, model: Any, base_url: Any
) -> bool:
"""True when (provider, model, base_url) matches one echo-back family.
Families can overlap (a deepseek-named model on a kimi host); membership is
tested independently per family. Raises KeyError for an unknown family.
"""
def matches_reasoning_echo_family(family: str, provider: Any, model: Any, base_url: Any) -> bool:
"""True when (provider, model, base_url) matches one echo-back family (families can overlap;
membership is tested independently). Raises KeyError for an unknown family."""
from utils import base_url_host_matches
_, raw_providers, lowered_providers, model_subs, hosts = _REASONING_ECHO_RULE_BY_FAMILY[family]
@@ -587,10 +489,10 @@ def matches_reasoning_echo_family(
def reasoning_echo_family(provider: Any, model: Any, base_url: Any) -> "str | None":
"""``"kimi"`` / ``"deepseek"`` / ``"mimo"`` (first match in table order) when the
endpoint enforces reasoning_content echo-back, else ``None`` (strip side)."""
for rule in _REASONING_ECHO_RULES:
if matches_reasoning_echo_family(rule[0], provider, model, base_url):
return rule[0]
return None
return next(
(rule[0] for rule in _REASONING_ECHO_RULES if matches_reasoning_echo_family(rule[0], provider, model, base_url)),
None,
)
def needs_reasoning_echo(provider: Any, model: Any, base_url: Any) -> bool:
@@ -598,80 +500,54 @@ def needs_reasoning_echo(provider: Any, model: Any, base_url: Any) -> bool:
return reasoning_echo_family(provider, model, base_url) is not None
def stale_thinking_reaches_wire(
api_mode: Any, provider: Any, model: Any, base_url: Any
) -> bool:
"""True when stale assistant ``reasoning``/``reasoning_content`` text is
actually replayed on the wire for the active route.
def stale_thinking_reaches_wire(api_mode: Any, provider: Any, model: Any, base_url: Any) -> bool:
"""True when stale assistant ``reasoning``/``reasoning_content`` text is actually replayed
on the wire for the active route.
The single wire-truth predicate the compaction TRIGGER estimator and the
tail-budget walks must share: if they disagree, a reasoning-heavy session
can look over-threshold to preflight yet fully tail-protected to the walk —
an infinite ineffective compaction loop.
* ``codex_responses``: the Responses input builder never reads the text keys
(continuity rides the encrypted ``codex_reasoning_items`` sidecar) → False.
* echo-back families: stored ``reasoning_content`` is replayed verbatim → True.
* everything else: stripped or one-space-padded at send time → False.
The single wire-truth predicate the compaction TRIGGER estimator and the tail-budget
walks must share: if they disagree, a reasoning-heavy session can look over-threshold
to preflight yet fully tail-protected to the walk — an infinite compaction loop.
``codex_responses`` never reads the text keys (continuity rides the encrypted sidecar);
echo-back families replay stored ``reasoning_content`` verbatim; everyone else strips.
"""
if (api_mode or "") == "codex_responses":
return False
return needs_reasoning_echo(provider, model, base_url)
return (api_mode or "") != "codex_responses" and needs_reasoning_echo(provider, model, base_url)
def apply_reasoning_content_policy(
source_msg: dict, api_msg: dict, needs_thinking_pad: bool
) -> None:
"""Copy provider-facing reasoning fields onto an API replay message.
``needs_thinking_pad`` is the require-side flag (``needs_reasoning_echo``).
Mutates ``api_msg`` in place.
"""
def apply_reasoning_content_policy(source_msg: dict, api_msg: dict, needs_thinking_pad: bool) -> None:
"""Copy provider-facing reasoning fields onto an API replay message (mutates ``api_msg``).
``needs_thinking_pad`` is the require-side flag (``needs_reasoning_echo``)."""
if source_msg.get("role") != "assistant":
return
# 1. Explicit reasoning_content set. Require-side: preserve verbatim,
# upgrading legacy "" placeholders to " " (DeepSeek V4 400s on ""). Strict
# side: strip entirely — a reasoning primary pads history with " ", then a
# fallback to Mistral/Cerebras/Groq replays the pad and 422s.
if not needs_thinking_pad:
# Strict side: never carry the field — a reasoning primary pads history with " ",
# then a fallback to Mistral/Cerebras/Groq replays the pad and 422s. Also drops a
# non-string value (None after compaction): never pass null to the API.
api_msg.pop("reasoning_content", None)
return
existing = source_msg.get("reasoning_content")
if isinstance(existing, str):
if not needs_thinking_pad:
api_msg.pop("reasoning_content", None)
else:
api_msg["reasoning_content"] = existing or " "
# Explicit value: preserve verbatim, upgrading legacy "" to " " (DeepSeek V4 400s on "").
api_msg["reasoning_content"] = existing or " "
return
normalized_reasoning = source_msg.get("reasoning")
has_reasoning = isinstance(normalized_reasoning, str) and bool(normalized_reasoning)
if needs_thinking_pad:
# 2. Cross-provider poisoned history: tool_calls + 'reasoning' but no
# 'reasoning_content' key means the reasoning text came from ANOTHER
# provider (DeepSeek's own build pins reasoning_content for tool-call
# turns). Pad with " " to satisfy the API without leaking foreign CoT.
# 3. Healthy session: promote internal 'reasoning' → 'reasoning_content'.
# 4. No reasoning at all: every assistant turn still needs the field;
# " " (not "") because DeepSeek V4 rejects empty string.
if has_reasoning and not source_msg.get("tool_calls"):
api_msg["reasoning_content"] = normalized_reasoning
else:
api_msg["reasoning_content"] = " "
reasoning = source_msg.get("reasoning")
if isinstance(reasoning, str) and reasoning and not source_msg.get("tool_calls"):
# Healthy session: promote internal 'reasoning' → 'reasoning_content'.
api_msg["reasoning_content"] = reasoning
return
# 5. Strict side: never carry the field (incl. a non-string value such as
# None after compaction — never pass null to the API).
api_msg.pop("reasoning_content", None)
# tool_calls + 'reasoning' but no 'reasoning_content' means the reasoning came from
# ANOTHER provider (DeepSeek's own build pins reasoning_content for tool-call turns):
# pad without leaking foreign CoT. No reasoning at all: every assistant turn still needs
# the field; " " (not "") because DeepSeek V4 rejects empty string.
api_msg["reasoning_content"] = " "
def reapply_reasoning_echo(api_messages: list, needs_thinking_pad: bool) -> int:
"""Re-pad (or strip) assistant turns' reasoning_content for the ACTIVE provider.
``api_messages`` is built once before the retry loop under the primary
provider; a mid-conversation fallback can switch providers, so the baked-in
reasoning fields must be reconciled: switching TO a require-side provider
needs the pad re-applied (else 400), switching TO a strict provider needs
the stale pad stripped (else 422). Idempotent; call every iteration.
Returns the number of assistant turns whose reasoning_content changed.
``api_messages`` is built once under the primary provider; a mid-conversation
fallback can switch providers, so baked-in fields must be reconciled: TO a
require-side provider re-applies the pad (else 400), TO a strict provider strips it
(else 422). Idempotent. Returns the number of assistant turns changed.
"""
changed = 0
for api_msg in api_messages:
@@ -689,8 +565,5 @@ def reapply_reasoning_echo(api_messages: list, needs_thinking_pad: bool) -> int:
return changed
# Image / multimodal parts are deliberately NOT consolidated here: per-adapter
# handling (anthropic base64 source blocks, Responses input_image items) is
# format-specific SYNTAX. The one shared image POLICY — removing images when a
# server rejects them while preserving tool_call_id pairing — is
# ``_strip_images_from_messages`` above.
# Image / multimodal parts are deliberately NOT consolidated here: per-adapter handling is
# format-specific SYNTAX. The one shared image POLICY is ``_strip_images_from_messages``.
+65 -170
View File
@@ -13,37 +13,34 @@ from typing import Any, Dict, List, Optional
from agent.model_metadata import estimate_messages_tokens_rough, estimate_tokens_rough
# Origin-module constants/helpers (and call_llm) are imported lazily inside methods: it avoids the
# import cycle and keeps tests that patch ``agent.context_compressor.X`` effective.
# Log name parity with the origin module.
logger = logging.getLogger("agent.context_compressor")
def _cc():
"""The origin module, resolved lazily: avoids the import cycle and keeps tests that patch
``agent.context_compressor.X`` effective (attributes are read at call time)."""
from agent import context_compressor
return context_compressor
def _is_summary_marker(entry: Any) -> bool:
return isinstance(entry, dict) and bool(entry.get(_cc().COMPRESSED_SUMMARY_METADATA_KEY))
def _is_micro_marker(entry: Any) -> bool:
"""True for a summary marker provably absorbed into the rolling summary (micro, not batch)."""
from agent.context_compressor import COMPRESSED_SUMMARY_METADATA_KEY, MICRO_COMPACT_MARKER_KEY
return (
isinstance(entry, dict)
and bool(entry.get(COMPRESSED_SUMMARY_METADATA_KEY))
and bool(entry.get(MICRO_COMPACT_MARKER_KEY))
)
return _is_summary_marker(entry) and bool(entry.get(_cc().MICRO_COMPACT_MARKER_KEY))
class MicroCompactionMixin:
"""Rolling micro-compaction; host must be a ``ContextCompressor``."""
def _resolve_compact_cursor(
self,
messages: List[Dict[str, Any]],
head_end: int,
tail_start: int,
) -> int:
"""Return the index of the first message not yet absorbed into the rolling summary.
def _resolve_compact_cursor(self, messages: List[Dict[str, Any]], head_end: int, tail_start: int) -> int:
"""Index of the first message not yet absorbed into the rolling summary.
Uses the in-memory cursor when valid; otherwise scans for the last summary marker.
"""
from agent.context_compressor import MICRO_COMPACT_MARKER_KEY
if head_end < self._micro_compact_cursor < tail_start:
return self._micro_compact_cursor
last_summary_idx = -1
@@ -56,14 +53,12 @@ class MicroCompactionMixin:
# Resumed session: rehydrate the rolling summary from the surviving marker so the next
# pass merges, not replaces.
if not self._micro_compact_rolling_summary.strip():
recovered = self._rolling_summary_from_marker(
messages[last_summary_idx].get("content")
)
recovered = self._rolling_summary_from_marker(messages[last_summary_idx].get("content"))
if recovered:
self._micro_compact_rolling_summary = recovered
# Rehydration proves containment: this marker (batch or micro) becomes
# supersede/defrag-eligible; unabsorbed markers never get the key.
messages[last_summary_idx][MICRO_COMPACT_MARKER_KEY] = True
messages[last_summary_idx][_cc().MICRO_COMPACT_MARKER_KEY] = True
logger.info(
"Micro-compaction: recovered rolling summary from "
"transcript (%d chars)", len(recovered),
@@ -72,10 +67,7 @@ class MicroCompactionMixin:
return cursor
def _find_one_exchange(
self,
messages: List[Dict[str, Any]],
start: int,
tail_start: int,
self, messages: List[Dict[str, Any]], start: int, tail_start: int,
) -> Optional[tuple[int, int]]:
"""Find the next complete exchange (full agent turn) starting at *start*.
@@ -111,23 +103,13 @@ class MicroCompactionMixin:
return None
return (exchange_start, idx)
def _serialize_one_exchange(
self,
messages: List[Dict[str, Any]],
start: int,
end: int,
) -> str:
def _serialize_one_exchange(self, messages: List[Dict[str, Any]], start: int, end: int) -> str:
"""Serialize a single exchange for the micro-summarizer via ``_serialize_for_summary``."""
return self._serialize_for_summary(messages[start:end])
def _build_micro_summary_prompt(
self,
existing_summary: str,
exchange_text: str,
) -> List[Dict[str, str]]:
def _build_micro_summary_prompt(self, existing_summary: str, exchange_text: str) -> List[Dict[str, str]]:
"""Build the prompt messages for a single-exchange micro-summary."""
summary_block = existing_summary if existing_summary.strip() else "(No previous summary yet.)"
user_prompt = (
"You are a summarization agent creating a compact record of an "
"ongoing conversation. You are given a running summary and the "
@@ -144,28 +126,18 @@ class MicroCompactionMixin:
"Return ONLY the updated summary text, no preamble or explanation. "
"Do not include this instruction block in your output."
)
return [
{"role": "system", "content": "You are a conversation summarization assistant."},
{"role": "user", "content": user_prompt},
]
def _micro_summarize_one(
self,
exchange_text: str,
) -> Optional[str]:
"""Micro-summarize one exchange into the rolling summary via the aux LLM.
Returns the updated summary text, or ``None`` on failure.
"""
def _micro_summarize_one(self, exchange_text: str) -> Optional[str]:
"""Micro-summarize one exchange into the rolling summary via the aux LLM (None on failure)."""
from agent.auxiliary_client import aux_interrupt_protection, call_llm
from agent.context_compressor import _response_finish_reason
call_kwargs = {
"task": "compression",
"messages": self._build_micro_summary_prompt(
self._micro_compact_rolling_summary, exchange_text,
),
"messages": self._build_micro_summary_prompt(self._micro_compact_rolling_summary, exchange_text),
"max_tokens": min(1500, self.max_summary_tokens or 1500),
"temperature": 0.1,
}
@@ -188,7 +160,7 @@ class MicroCompactionMixin:
return None
# A length stop means a partial merge; leave the exchange unabsorbed so a later pass retries.
if _response_finish_reason(response) == "length":
if _cc()._response_finish_reason(response) == "length":
logger.warning(
"micro-summarization output hit the token cap "
"(finish_reason=length) — discarding partial summary",
@@ -196,10 +168,7 @@ class MicroCompactionMixin:
return None
message = response.choices[0].message
if isinstance(message, dict):
content = message.get("content")
else:
content = getattr(message, "content", message)
content = message.get("content") if isinstance(message, dict) else getattr(message, "content", message)
if not isinstance(content, str):
content = str(content) if content else ""
content = content.strip()
@@ -212,18 +181,13 @@ class MicroCompactionMixin:
def _needs_defrag(self) -> bool:
"""Return True when the rolling summary is large enough to defrag."""
content_tokens = estimate_tokens_rough(self._micro_compact_rolling_summary)
return content_tokens >= self._micro_compact_defrag_threshold_tokens
return estimate_tokens_rough(self._micro_compact_rolling_summary) >= self._micro_compact_defrag_threshold_tokens
def _defrag_rolling_summary(
self,
messages: List[Dict[str, Any]],
) -> bool:
def _defrag_rolling_summary(self, messages: List[Dict[str, Any]]) -> bool:
"""Re-summarize the rolling summary text and rewrite the marker in place.
Transcript-shape-neutral (no splice, no cursor move). Returns True when it rewrote.
"""
from agent.context_compressor import _DB_PERSISTED_MARKER
old_summary = self._micro_compact_rolling_summary
if not old_summary.strip():
return False
@@ -239,10 +203,10 @@ class MicroCompactionMixin:
for entry in reversed(messages):
if _is_micro_marker(entry):
entry["content"] = self._render_micro_marker_content(fresh_summary)
# Content changed: clear the persisted stamp so the DB sync rewrites the row.
entry.pop(_DB_PERSISTED_MARKER, None)
# In-place pop on a live dict would be identity-skipped by the bounded flush scan;
# Content changed: clear the persisted stamp so the DB sync rewrites the row. An
# in-place pop on a live dict would be identity-skipped by the bounded flush scan;
# flag the finalizer.
entry.pop(_cc()._DB_PERSISTED_MARKER, None)
self._flush_scan_cursor_invalidated = True
break
logger.info(
@@ -255,16 +219,12 @@ class MicroCompactionMixin:
self._micro_compact_consecutive_failures = 0
self._micro_compact_last_failure_cursor = -1
def _micro_compact(
self,
messages: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
def _micro_compact(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Run one round of micro-compaction; public entry point from ``finalize_turn()``.
Returns the (possibly modified) list and syncs the session DB via ``archive_and_compact``
(the append-only flush alone would double-load on resume).
"""
from agent.context_compressor import _MICRO_COMPACT_MAX_CONSECUTIVE_FAILURES
if not self._micro_compact_enabled:
return messages
@@ -281,16 +241,13 @@ class MicroCompactionMixin:
if n_messages < 4:
return messages
head_size = self._protect_head_size(messages)
compress_start = self._align_boundary_forward(messages, head_size)
compress_start = self._align_boundary_forward(messages, self._protect_head_size(messages))
compress_end = self._find_tail_cut_by_tokens(messages, compress_start)
if compress_start >= compress_end:
return messages
cursor = self._resolve_compact_cursor(messages, compress_start, compress_end)
if cursor >= compress_end:
return messages
exchange = self._find_one_exchange(messages, cursor, compress_end)
if exchange is None:
return messages
@@ -302,12 +259,8 @@ class MicroCompactionMixin:
def _telemetry(outcome: str, result: List[Dict[str, Any]], **extra: Any) -> None:
self._emit_micro_compaction_telemetry(
outcome=outcome,
messages_before=n_messages,
messages_after=len(result),
tokens_before=_tokens_before,
duration_ms=int((time.monotonic() - _started_at) * 1000),
**extra,
outcome=outcome, messages_before=n_messages, messages_after=len(result),
tokens_before=_tokens_before, duration_ms=int((time.monotonic() - _started_at) * 1000), **extra,
)
# Defrag rewrites summary text/marker in place (no splice, no cursor move) instead of
@@ -338,7 +291,7 @@ class MicroCompactionMixin:
self._micro_compact_last_failure_cursor = exchange_start
_outcome = "summarize_failed"
if self._micro_compact_consecutive_failures >= _MICRO_COMPACT_MAX_CONSECUTIVE_FAILURES:
if self._micro_compact_consecutive_failures >= _cc()._MICRO_COMPACT_MAX_CONSECUTIVE_FAILURES:
logger.info(
"Micro-compaction: skipping exchange at cursor %d "
"after %d consecutive failures",
@@ -348,18 +301,14 @@ class MicroCompactionMixin:
self._micro_compact_cursor = exchange_end
self._reset_micro_failure_tracking()
_outcome = "exchange_skipped"
_telemetry(
_outcome, messages, tokens_after=_tokens_before, exchange_tokens=_exchange_tokens,
)
_telemetry(_outcome, messages, tokens_after=_tokens_before, exchange_tokens=_exchange_tokens)
return messages
self._micro_compact_rolling_summary = updated_summary
self._micro_compact_cursor = exchange_end
self._reset_micro_failure_tracking()
result = self._splice_micro_compact_result(
messages, exchange_start, exchange_end, supersede=_cumulative,
)
result = self._splice_micro_compact_result(messages, exchange_start, exchange_end, supersede=_cumulative)
self._micro_compact_cursor = self._cursor_after_splice(result, exchange_start + 1)
self._sync_micro_compact_to_db(result)
_telemetry(
@@ -371,56 +320,41 @@ class MicroCompactionMixin:
@staticmethod
def _rolling_summary_from_marker(content: Any) -> str:
"""Recover the rolling-summary text from a summary marker (resume rehydration)."""
from agent.context_compressor import (
_SUMMARY_END_MARKER,
HISTORICAL_TASK_HEADING,
)
cc = _cc()
if not isinstance(content, str) or not content.strip():
return ""
body = content
# rfind: SUMMARY_PREFIX itself mentions the heading, so the first hit is in the preamble.
idx = body.rfind(HISTORICAL_TASK_HEADING)
idx = body.rfind(cc.HISTORICAL_TASK_HEADING)
if idx != -1:
body = body[idx + len(HISTORICAL_TASK_HEADING):]
end = body.find(_SUMMARY_END_MARKER)
body = body[idx + len(cc.HISTORICAL_TASK_HEADING):]
end = body.find(cc._SUMMARY_END_MARKER)
if end != -1:
body = body[:end]
return body.strip()
def _cursor_after_splice(
self,
result: List[Dict[str, Any]],
fallback: int,
) -> int:
def _cursor_after_splice(self, result: List[Dict[str, Any]], fallback: int) -> int:
"""Cursor position just past the summary marker in *result*.
Must derive from the SPLICED list: a splice collapses several rows into one marker
(and may drop a superseded one), so pre-splice indices land inside a later exchange
and silently skip it.
"""
from agent.context_compressor import COMPRESSED_SUMMARY_METADATA_KEY
for idx in range(len(result) - 1, -1, -1):
entry = result[idx]
if isinstance(entry, dict) and entry.get(COMPRESSED_SUMMARY_METADATA_KEY):
if _is_summary_marker(result[idx]):
return idx + 1
return fallback
def _emit_micro_compaction_telemetry(
self,
*,
outcome: str,
messages_before: int,
messages_after: int,
tokens_before: int | None,
tokens_after: int | None,
exchange_tokens: int | None = None,
duration_ms: int | None = None,
self, *, outcome: str, messages_before: int, messages_after: int,
tokens_before: int | None, tokens_after: int | None,
exchange_tokens: int | None = None, duration_ms: int | None = None,
) -> None:
"""Emit one content-free JSON log line for a micro-compaction pass.
``tokens_delta`` < 0 means the pass shrank the transcript; ``*_total`` fields accumulate.
"""
from agent.context_compressor import _safe_int
_safe_int = _cc()._safe_int
try:
delta = None
if tokens_before is not None and tokens_after is not None:
@@ -429,7 +363,6 @@ class MicroCompactionMixin:
self._micro_compact_passes += 1
# Cached reads only: the lazy properties can fire a synchronous /models probe.
threshold = self._threshold_tokens
context_limit = self._resolved_context_length
occupancy = None
if threshold and tokens_after is not None and threshold > 0:
occupancy = round(tokens_after / threshold * 100, 1)
@@ -443,37 +376,28 @@ class MicroCompactionMixin:
"tokens_after": _safe_int(tokens_after),
"tokens_delta": _safe_int(delta),
"exchange_tokens": _safe_int(exchange_tokens),
"rolling_summary_tokens": estimate_tokens_rough(
self._micro_compact_rolling_summary
),
"rolling_summary_tokens": estimate_tokens_rough(self._micro_compact_rolling_summary),
"cursor": _safe_int(self._micro_compact_cursor),
"passes_total": self._micro_compact_passes,
"tokens_saved_total": self._micro_compact_tokens_saved_total,
"duration_ms": _safe_int(duration_ms),
# Headroom: how full the window is being kept.
"threshold_tokens": _safe_int(threshold),
"context_limit": _safe_int(context_limit),
"context_limit": _safe_int(self._resolved_context_length),
"occupancy_pct": occupancy,
"main_model": self.model or "",
"aux_model": self.summary_model or "",
}
logger.info(
"micro compaction telemetry: %s",
json.dumps(payload, sort_keys=True, separators=(",", ":")),
)
logger.info("micro compaction telemetry: %s", json.dumps(payload, sort_keys=True, separators=(",", ":")))
except Exception as exc:
logger.debug("failed to emit micro-compaction telemetry: %s", exc)
def _sync_micro_compact_to_db(
self,
compacted_messages: List[Dict[str, Any]],
) -> None:
def _sync_micro_compact_to_db(self, compacted_messages: List[Dict[str, Any]]) -> None:
"""Persist the micro-compacted set to the session DB atomically and stamp rows persisted.
Without this the old exchange rows stay ``active=1`` and a resume double-loads
both the summary and the originals.
"""
from agent.context_compressor import stamp_db_persisted_markers
session_db = getattr(self, "_session_db", None)
session_id = getattr(self, "_session_id", "")
if not session_db or not session_id:
@@ -482,12 +406,10 @@ class MicroCompactionMixin:
# Every row except the marker is a carried-forward original: archive pre-splice
# originals rewind-style.
session_db.archive_and_compact(
session_id,
compacted_messages,
tail_count=max(0, len(compacted_messages) - 1),
session_id, compacted_messages, tail_count=max(0, len(compacted_messages) - 1),
)
# Shared post-commit stamp site with batch commit and proactive prune.
stamp_db_persisted_markers(compacted_messages)
_cc().stamp_db_persisted_markers(compacted_messages)
except Exception:
logger.info(
"Micro-compaction DB sync failed — resume will double-load "
@@ -495,21 +417,13 @@ class MicroCompactionMixin:
)
def _splice_micro_compact_result(
self,
messages: List[Dict[str, Any]],
splice_start: int,
splice_end: int,
supersede: bool = True,
self, messages: List[Dict[str, Any]], splice_start: int, splice_end: int, supersede: bool = True,
) -> List[Dict[str, Any]]:
"""Replace *messages[splice_start:splice_end]* with an assistant-role summary marker.
Merges user turns left adjacent by a superseded marker so the result is alternation-valid.
"""
from agent.context_compressor import (
COMPRESSED_SUMMARY_HAS_USER_TURN_KEY,
COMPRESSED_SUMMARY_METADATA_KEY,
MICRO_COMPACT_MARKER_KEY,
)
cc = _cc()
summary_text = self._micro_compact_rolling_summary
if not summary_text.strip():
return messages
@@ -517,13 +431,12 @@ class MicroCompactionMixin:
summary_msg = {
"role": "assistant",
"content": self._render_micro_marker_content(summary_text),
COMPRESSED_SUMMARY_METADATA_KEY: True,
cc.COMPRESSED_SUMMARY_METADATA_KEY: True,
# Micro marker: eligible for supersede/defrag; batch markers never carry this key.
MICRO_COMPACT_MARKER_KEY: True,
cc.MICRO_COMPACT_MARKER_KEY: True,
# Micro markers absorb only assistant/tool content; user turns stay in the transcript.
COMPRESSED_SUMMARY_HAS_USER_TURN_KEY: False,
cc.COMPRESSED_SUMMARY_HAS_USER_TURN_KEY: False,
}
result = messages[:splice_start] + [summary_msg] + messages[splice_end:]
# Cumulative summary: keep only the newest marker. Drop an older one only if supersede AND
@@ -532,8 +445,7 @@ class MicroCompactionMixin:
marker_idxs = [i for i, m in enumerate(result) if _is_micro_marker(m)]
if len(marker_idxs) > 1:
superseded = set(marker_idxs[:-1])
result = [m for i, m in enumerate(result) if i not in superseded]
result = self._merge_adjacent_user_turns(result)
result = self._merge_adjacent_user_turns([m for i, m in enumerate(result) if i not in superseded])
# Deliberately no _strip_persistence_markers: micro archives in place under the same session
# id, so stamps stay accurate and a failed archive keeps the append-only flush idempotent.
@@ -542,48 +454,31 @@ class MicroCompactionMixin:
@staticmethod
def _render_micro_marker_content(summary_text: str) -> str:
"""Assemble the marker content wrapper around *summary_text*."""
from agent.context_compressor import (
_SUMMARY_END_MARKER,
HISTORICAL_TASK_HEADING,
SUMMARY_PREFIX,
)
return (
f"{SUMMARY_PREFIX}\n\n"
f"{HISTORICAL_TASK_HEADING}\n"
f"{summary_text.strip()}"
f"\n\n{_SUMMARY_END_MARKER}"
)
cc = _cc()
return f"{cc.SUMMARY_PREFIX}\n\n{cc.HISTORICAL_TASK_HEADING}\n{summary_text.strip()}\n\n{cc._SUMMARY_END_MARKER}"
@staticmethod
def _merge_adjacent_user_turns(
result: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
def _merge_adjacent_user_turns(result: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Merge consecutive plain-text real user turns left by a supersede.
Same ``\\n\\n`` join as ``repair_message_sequence`` pass 2, done here so the marker
and cursor are never collateral damage of the downstream repair. Lists untouched.
"""
from agent.context_compressor import COMPRESSED_SUMMARY_METADATA_KEY
from agent.turn_context import drop_stale_api_content
def _plain_user(m: Any) -> bool:
return (
isinstance(m, dict)
and m.get("role") == "user"
and not m.get(COMPRESSED_SUMMARY_METADATA_KEY)
and isinstance(m.get("content"), str)
isinstance(m, dict) and m.get("role") == "user"
and not _is_summary_marker(m) and isinstance(m.get("content"), str)
)
merged: List[Dict[str, Any]] = []
for msg in result:
prev = merged[-1] if merged else None
if _plain_user(msg) and _plain_user(prev):
prev_content = prev["content"]
new_content = msg["content"]
prev_content, new_content = prev["content"], msg["content"]
prev["content"] = (
(prev_content + "\n\n" + new_content)
if prev_content and new_content
else (prev_content or new_content)
(prev_content + "\n\n" + new_content) if prev_content and new_content else (prev_content or new_content)
)
# Merged content invalidates the api_content sidecar.
drop_stale_api_content(prev)
+58 -112
View File
@@ -1,12 +1,11 @@
"""Native OpenAI Responses server-side compaction — gpt-5.6 on direct OpenAI routes only.
``context_management=[{"type": "compaction", "compact_threshold": N}]`` makes the server
summarize older context into an opaque ``compaction`` item (sealed to the issuing
endpoint) once the input crosses N tokens. Support is deliberately narrow (live-verified):
gpt-5.6 family only (5.1/5.2 fail server-side with no structured rejection) on direct
OpenAI routes (api.openai.com or the ChatGPT Codex backend). Hermes' local compressor stays
armed as fallback: the native threshold is clamped below the local trigger, and captured
compaction items ride the existing ``codex_reasoning_items`` sidecar. No transport imports.
summarize older context into an opaque ``compaction`` item once the input crosses N tokens.
Deliberately narrow (live-verified): gpt-5.6 only (5.1/5.2 fail server-side with no
structured rejection) on api.openai.com or the ChatGPT Codex backend. The local compressor
stays armed as fallback (native threshold clamped below the local trigger); compaction items
ride the ``codex_reasoning_items`` sidecar. No transport imports (shared gate, no cycles).
"""
from __future__ import annotations
@@ -20,8 +19,7 @@ from agent.message_content import flatten_message_text
logger = logging.getLogger(__name__)
# Native compaction fires this many tokens below the local compressor's
# trigger so the server always gets the first shot.
# Native compaction fires this far below the local trigger so the server gets the first shot.
LOCAL_TRIGGER_SAFETY_MARGIN = 8_192
# Fallback when automatic mode has no local trigger to follow.
DEFAULT_COMPACT_THRESHOLD = 200_000
@@ -35,30 +33,18 @@ def is_native_compaction_model(model: Optional[str]) -> bool:
def resolve_native_compaction_capabilities(
*,
model: Optional[str],
base_url: Optional[str],
provider: Optional[str] = None,
is_codex_backend: bool = False,
*, model: Optional[str], base_url: Optional[str], provider: Optional[str] = None, is_codex_backend: bool = False,
) -> Dict[str, bool]:
"""Resolve the native-compaction capability for a runtime destination.
A resolved ``False`` is distinct from "unresolved" and must survive model
switches unchanged.
"""
"""Resolve the native-compaction capability for a runtime destination (a resolved ``False``
is distinct from "unresolved" and must survive model switches unchanged)."""
direct_default = (provider or "").strip().lower() == "openai" and not base_url
eligible = is_native_compaction_model(model) and (
direct_default
or is_direct_openai_route(base_url, is_codex_backend=is_codex_backend)
direct_default or is_direct_openai_route(base_url, is_codex_backend=is_codex_backend)
)
return {"native_compaction": eligible}
def is_direct_openai_route(
base_url: Optional[str],
*,
is_codex_backend: bool = False,
) -> bool:
def is_direct_openai_route(base_url: Optional[str], *, is_codex_backend: bool = False) -> bool:
"""True for api.openai.com or the ChatGPT Codex backend — nothing else."""
if is_codex_backend:
return True
@@ -80,24 +66,17 @@ def _positive_int(value: Any, *, reject: tuple = (bool,)) -> Optional[int]:
return parsed if parsed > 0 else None
def resolve_compact_threshold(
configured_threshold: Any,
local_trigger_tokens: Any = None,
) -> int:
def resolve_compact_threshold(configured_threshold: Any, local_trigger_tokens: Any = None) -> int:
"""Resolve automatic mode or clamp an explicit native threshold.
An omitted/invalid setting follows the local compressor trigger
(``ContextCompressor.threshold_tokens``) minus the safety margin. An
explicit positive integer is absolute unless it must be clamped so native
compaction fires first. Booleans are never thresholds.
Omitted/invalid follows the local compressor trigger minus the safety margin. An
explicit positive integer is absolute unless it must be clamped so native compaction
fires first. Booleans are never thresholds.
"""
local = _positive_int(local_trigger_tokens)
upper = None
if local is not None:
if local > LOCAL_TRIGGER_SAFETY_MARGIN:
upper = max(1_024, local - LOCAL_TRIGGER_SAFETY_MARGIN)
else:
upper = max(1_024, int(local * 0.8))
upper = max(1_024, local - LOCAL_TRIGGER_SAFETY_MARGIN if local > LOCAL_TRIGGER_SAFETY_MARGIN else int(local * 0.8))
configured = _positive_int(configured_threshold, reject=(bool, float))
if configured is None:
@@ -124,43 +103,29 @@ def _warn_native_compaction_suppressed_by_checkpoint_gate() -> None:
def native_compaction_context_management(
agent: Any,
*,
is_codex_backend: bool,
is_xai_responses: bool = False,
is_github_responses: bool = False,
agent: Any, *, is_codex_backend: bool, is_xai_responses: bool = False, is_github_responses: bool = False,
) -> Optional[List[Dict[str, Any]]]:
"""Return the ``context_management`` payload for this request, or None.
"""Return the ``context_management`` payload for this request, or None ("do not send").
None means "do not send the field" (request byte-identical to pre-feature).
Every gate is re-checked per request so a mid-session model switch or the
in-session kill switch (``agent.codex_responses_native_compaction = False``,
set by rejection recovery) takes effect on the next call.
Every gate is re-checked per request so a mid-session model switch or the in-session
kill switch (``agent.codex_responses_native_compaction = False``) takes effect next call.
"""
capabilities = getattr(agent, "runtime_capabilities", None)
if isinstance(capabilities, dict) and not capabilities.get("native_compaction", False):
return None
if not getattr(agent, "codex_responses_native_compaction", False):
return None
# compression.enabled: false disables ALL automatic compaction, native included.
if not getattr(agent, "compression_enabled", True):
if not getattr(agent, "codex_responses_native_compaction", False) or not getattr(agent, "compression_enabled", True):
return None
# Server-side compaction is a lossy boundary the provider owns — no
# pre-compress checkpoint can run first — so the checkpoint-aware Hermes
# compressor stays authoritative. Explicit-True matches compress_context().
# Server-side compaction is a lossy boundary the provider owns (no pre-compress checkpoint
# can run first), so the checkpoint-aware compressor stays authoritative. Explicit-True
# matches compress_context().
if getattr(agent, "compression_checkpoint_required", False) is True:
_warn_native_compaction_suppressed_by_checkpoint_gate()
return None
if is_xai_responses or is_github_responses:
if is_xai_responses or is_github_responses or not is_native_compaction_model(getattr(agent, "model", None)):
return None
if not is_native_compaction_model(getattr(agent, "model", None)):
return None
trusted_proxy = bool(
getattr(agent, "capabilities", {}).get("openai_native_compaction", False)
)
if not trusted_proxy and not is_direct_openai_route(
getattr(agent, "base_url", None), is_codex_backend=is_codex_backend
):
trusted_proxy = bool(getattr(agent, "capabilities", {}).get("openai_native_compaction", False))
if not trusted_proxy and not is_direct_openai_route(getattr(agent, "base_url", None), is_codex_backend=is_codex_backend):
return None
compressor = getattr(agent, "context_compressor", None)
@@ -171,9 +136,8 @@ def native_compaction_context_management(
return [{"type": "compaction", "compact_threshold": threshold}]
# Retention budgets for plaintext user messages / local compression summaries
# carried across a native compaction boundary (mirrors Codex CLI's
# RETAINED_MESSAGE_TOKEN_BUDGET; the summary budget prevents summary inflation).
# Retention budgets for plaintext user messages / local summaries carried across a native
# compaction boundary (mirrors Codex CLI's RETAINED_MESSAGE_TOKEN_BUDGET).
RETAINED_USER_MESSAGE_TOKEN_BUDGET = 64_000
RETAINED_SUMMARY_TOKEN_BUDGET = 32_000
@@ -216,11 +180,8 @@ def _extract_item_text(item: Any) -> Optional[str]:
def _has_retainable_image_content(item: Any) -> bool:
"""True for a converted Responses message with a valid ``input_image`` part.
Only the adapter-owned ``input_image`` shape counts: unknown or empty
multipart placeholders must not become durable history for being non-empty.
"""
"""True for a converted Responses message with a valid ``input_image`` part (only the
adapter-owned shape counts, so empty multipart placeholders never become durable history)."""
if not isinstance(item, dict):
return False
content = item.get("content")
@@ -235,10 +196,8 @@ def _has_retainable_image_content(item: Any) -> bool:
)
# Canonical provenance check (metadata marker, then canonical prefix classifier).
# Deliberately NOT a second heuristic: no underscore-key scan, no matching on
# ad-hoc headings — either could promote ordinary or adversarial content to
# durable retained history.
# Canonical provenance check. Deliberately NOT a second heuristic (no underscore-key scan,
# no ad-hoc headings) — either could promote adversarial content to durable history.
_is_summary_item = is_compaction_summary_message
@@ -255,29 +214,23 @@ def prune_pre_checkpoint_items(
) -> List[Dict[str, Any]]:
"""Restructure Responses input around the newest compaction checkpoint.
The server drops every input item preceding a replayed ``compaction`` item,
which silently erases the user's plaintext asks and any local-compression
summary (``role="assistant"``). With a checkpoint present, rebuild as::
The server drops every input item preceding a replayed ``compaction`` item, erasing the
user's plaintext asks and any local-compression summary. Rebuild as::
[checkpoint run] + [retained user & summary messages (newest-first budget)] + [post]
- The NEWEST contiguous run of checkpoints wins.
- User messages are kept verbatim within ``retained_user_token_budget``;
the boundary message is head-truncated when it only partially fits
(string content only — goals are stated up front). A recognized
image-only user message is retained whole at one-token cost.
- Summaries are retained whole within ``retained_summary_token_budget`` and
never sliced (their structural framing would corrupt); one that doesn't
fit is dropped. Identical summary text is never retained twice.
- Relative order between user messages and summaries is preserved.
- ``item_sources`` (parallel to ``items``) is the raw chat message each item
was converted from. Conversion can be lossy for summaries (a
merge-into-tail carrier becomes a typed ``function_call_output``, or an
assistant carrier is shadowed by a stale exact replay), so when a source
is itself a canonical summary carrier its content is read from the
SOURCE and retained as a synthesized ``role="assistant"`` message.
- ``enable_summary_retention`` is a function-level override for tests, not
a config surface.
- User messages are kept verbatim within ``retained_user_token_budget``; the boundary
message is head-truncated when it only partially fits (string content only). A
recognized image-only user message is retained whole at one-token cost.
- Summaries are retained whole within ``retained_summary_token_budget``, never sliced
(framing would corrupt) and never duplicated. Relative order is preserved.
- ``item_sources`` (parallel to ``items``) is the raw chat message each item came from.
Conversion can be lossy for summaries (merge-into-tail carrier → typed
``function_call_output``; assistant carrier shadowed by a stale replay), so a source
that is itself a canonical summary carrier is read from the SOURCE and retained as a
synthesized ``role="assistant"`` message.
- ``enable_summary_retention`` is a test override, not a config surface.
"""
if not isinstance(items, list) or not items:
return items
@@ -297,10 +250,8 @@ def prune_pre_checkpoint_items(
checkpoint_run = items[first_cp : last_cp + 1]
post = items[last_cp + 1 :]
if isinstance(item_sources, list) and len(item_sources) == len(items):
pre_sources: List[Any] = item_sources[:first_cp]
else:
pre_sources = [None] * len(pre)
has_sources = isinstance(item_sources, list) and len(item_sources) == len(items)
pre_sources: List[Any] = item_sources[:first_cp] if has_sources else [None] * len(pre)
retained_reversed: List[Dict[str, Any]] = []
user_remaining = max(0, int(retained_user_token_budget))
@@ -383,11 +334,10 @@ _REJECTION_MARKERS = (
def is_native_compaction_rejection(error: Any, status_code: Any = None) -> bool:
"""True when a provider error is a STRUCTURED rejection of ``context_management``.
Drives the loop's one-shot recovery (strip the field, disable for the
session, retry), so matching is narrow: a transient 5xx whose body merely
ECHOES the request must not permanently downgrade native compaction. Requires
``status_code`` 400 (or unknown — some transports surface only a message)
AND the field name alongside rejection language.
Drives one-shot recovery (strip, disable for the session, retry), so matching is
narrow: a transient 5xx that merely ECHOES the request must not downgrade native
compaction. Requires ``status_code`` 400 (or unknown) AND the field name with rejection
language.
"""
text = str(error or "").lower()
if "context_management" not in text and "compact_threshold" not in text:
@@ -404,22 +354,18 @@ def is_native_compaction_rejection(error: Any, status_code: Any = None) -> bool:
def has_compaction_checkpoint(items: Any) -> bool:
"""Does this ``codex_reasoning_items`` sidecar carry a compaction checkpoint?
A ``type: "compaction"`` item is cumulative context, not per-turn
reasoning, and exists in exactly one place: anything that rewrites or
discards the sidecar must ask this first or lose the compacted history.
A compaction item is cumulative context that exists in exactly one place: anything
that rewrites or discards the sidecar must ask this first or lose the history.
"""
return isinstance(items, list) and any(_is_compaction_item(item) for item in items)
def merge_interim_reasoning_items(
prior_items: Any,
new_items: Any,
) -> List[Dict[str, Any]]:
def merge_interim_reasoning_items(prior_items: Any, new_items: Any) -> List[Dict[str, Any]]:
"""Merge ``codex_reasoning_items`` across Codex incomplete-continuation dedup.
A checkpoint captured on the EARLIER response is not re-emitted by the
continuation, so a blind overwrite drops the only copy. Rule: newer items
win, but prior checkpoints are prepended unless the newer payload has its own.
A checkpoint on the EARLIER response is not re-emitted by the continuation, so a blind
overwrite drops the only copy: newer items win, prior checkpoints are prepended unless
the newer payload has its own.
"""
kept_checkpoints = [
item for item in (prior_items if isinstance(prior_items, list) else []) if _is_compaction_item(item)
+28 -58
View File
@@ -1,21 +1,12 @@
"""Rotation-stable logical cache scope for prompt_cache_key derivation.
Legacy compression rotation mints a new physical ``session_id`` mid-conversation,
which moved the conversation into a fresh cache bucket each time.
``resolve_prompt_cache_scope()`` maps the physical id to the ROOT of its compression
lineage via ``SessionDB.get_compression_lineage()`` — NOT ``get_conversation_root``
(the Portal-attribution walk), which follows ``parent_session_id`` blindly and would
collapse /branch children and delegate trees into one id.
Scope boundaries: rotation children walk back to the original segment; ``/new``
starts a fresh scope; ``/branch`` children, delegate subagents, and tool-tagged
children are explicit fork children with their own isolated scope; cron fires
keep their physical id. Hosts that mint one physical id per RESPONSE (Studio group
chat, ``/v1/responses`` with client-managed history) carry no lineage, so they
declare the conversation via ``gateway_session_key`` (``X-Hermes-Session-Key``),
consumed by ``declared_conversation_scope()``, which wins over the lineage walk and
is hashed to ``gwk_<sha256[:24]>`` because it embeds platform/chat/user identifiers.
Resolution is memoized per (agent, session_id, db-present).
Legacy compression rotation mints a new physical ``session_id`` mid-conversation, moving it
into a fresh cache bucket. ``resolve_prompt_cache_scope()`` maps the physical id to the ROOT
of its compression lineage — NOT ``get_conversation_root`` (the Portal-attribution walk),
which would collapse /branch children and delegate trees into one id. ``/new`` starts a
fresh scope; fork children (branch, delegate, tool-tagged) are isolated. Hosts minting one id
per RESPONSE declare the conversation via ``gateway_session_key``, which wins over the lineage
walk and is hashed to ``gwk_<sha256[:24]>`` (it embeds platform/chat/user identifiers).
"""
import hashlib
@@ -49,12 +40,10 @@ def _agent_source(
) -> str:
"""The ``sessions.source`` this agent's conversation is recorded under.
``row_source`` is the row's value when the caller already read it (``""``
= read, no source; ``None`` = not read yet, do the lookup). Before the row
lands, use the SAME resolver persistence uses
(``run_agent._session_source_for_agent``), not ``agent.platform``: they
diverge under ``HERMES_SESSION_SOURCE``, and the declared scope is memoized
immediately, so both sides of a ``/new`` would otherwise hash the same scope.
``row_source``: the row's value if already read (``""`` = read, none; ``None`` = look it
up). Before the row lands, use the SAME resolver persistence uses, not ``agent.platform``:
they diverge under ``HERMES_SESSION_SOURCE`` and the declared scope is memoized at once,
so both sides of a ``/new`` would otherwise hash the same scope.
"""
if row_source is None and session_id and session_db is not None:
try:
@@ -81,12 +70,9 @@ def _agent_source(
def _conversation_generation(session_key: str, source: str, session_db: Any) -> str:
"""Durable generation for *session_key*'s current conversation (``""`` if none).
The declared key names a chat and survives ``/new`` and policy resets, so
hashing it alone would reuse one scope across distinct conversations. The
``conversation_generations`` counter advances in the same transaction that
records a reset boundary and is independent of prunable rows and
wall-clock, so pruning or clock rollback cannot reissue a generation.
Compression does not advance it.
The declared key survives ``/new``, so hashing it alone would reuse one scope across
conversations. The counter advances with each reset boundary, independent of prunable
rows and wall-clock; compression does not advance it.
"""
reader = getattr(session_db, "latest_conversation_boundary", None)
if not callable(reader):
@@ -98,12 +84,10 @@ def _conversation_generation(session_key: str, source: str, session_db: Any) ->
def declared_conversation_scope(agent: Any) -> Optional[str]:
"""Host-declared logical conversation scope (``gwk_<sha256[:24]>``), or None.
Hashes ``(source, gateway_session_key, generation)`` so no platform/chat/
user identifier reaches a provider. None — fall back to the physical-id
scope — when no key is declared, when the agent is a background-review
fork (``_persist_disabled`` clones the live runtime incl. the key), when
the row is an explicit fork child, and on any DB error (fail closed rather
than merge a fork onto its parent's key).
Hashes ``(source, gateway_session_key, generation)``. None (fall back to the physical id)
when no key is declared, for a background-review fork (``_persist_disabled``), for an
explicit fork child, and on any DB error (fail closed rather than merge a fork onto its
parent's key).
"""
key = str(getattr(agent, "_gateway_session_key", "") or "").strip()
if not key or getattr(agent, "_persist_disabled", False):
@@ -114,9 +98,7 @@ def declared_conversation_scope(agent: Any) -> Optional[str]:
row_source: Optional[str] = None
if sid and db is not None:
try:
# One read for both halves of the row identity (fork verdict +
# source). A SessionDB without the combined view keeps the
# original call.
# One read for both halves of the row identity (fork verdict + source).
identity = getattr(db, "declared_scope_identity", None)
if callable(identity):
is_fork, row_source = identity(sid)
@@ -134,37 +116,29 @@ def declared_conversation_scope(agent: Any) -> Optional[str]:
except Exception:
logger.debug("declared-scope generation read failed", exc_info=True)
return None
# Same identity tuple the peer queries use: two hosts may declare the
# same key under different sources and must not collapse.
# Same identity tuple the peer queries use: same key under different sources must not collapse.
carrier = f"{source}|{key}|{generation}"
digest = hashlib.sha256(carrier.encode("utf-8", errors="replace")).hexdigest()[:24]
return f"{_DECLARED_SCOPE_PREFIX}{digest}"
def resolve_prompt_cache_scope(agent: Any) -> str:
"""Rotation-stable cache-scope id for *agent*'s conversation.
Declared scope when one applies, else the compression-lineage root of
``agent.session_id`` (the physical id when there is no ancestry, no DB, or
the walk fails). Memoized on the agent keyed by session id.
"""
"""Rotation-stable cache-scope id: declared scope, else the compression-lineage root of
``agent.session_id`` (the physical id without ancestry/DB). Memoized on the agent."""
sid = str(getattr(agent, "session_id", None) or "")
if not sid:
return ""
db = getattr(agent, "_session_db", None)
# DB presence is part of the key: an agent that gains a DB handle later
# must re-resolve instead of staying pinned to the physical id.
# DB presence is part of the key: an agent that gains a DB handle later must re-resolve.
key = (sid, db is not None)
memo = getattr(agent, _MEMO_ATTR, None)
if isinstance(memo, tuple) and len(memo) == 2 and memo[0] == key:
return memo[1]
root = declared_conversation_scope(agent) or _lineage_root(sid, db)
scope = root or sid
# Memoize on success, with no DB, or when the agent never persists a row
# (background-review forks hold a DB handle but set _persist_disabled).
# A failed/empty walk on a persisting agent is NOT memoized: the physical
# id is right for now (row not yet persisted, transient error) but would
# stay wrong for the whole segment once the row lands.
# Memoize on success, with no DB, or when the agent never persists a row. A failed/empty
# walk on a persisting agent is NOT memoized: the physical id is right for now (row not
# yet persisted) but would stay wrong for the whole segment once it lands.
if root is not None or db is None or getattr(agent, "_persist_disabled", False):
try:
setattr(agent, _MEMO_ATTR, (key, scope))
@@ -183,12 +157,8 @@ def declared_conversation_scope_safe(agent: Any) -> Optional[str]:
def resolve_prompt_cache_scope_safe(agent: Any) -> Optional[str]:
"""Never-raising variant of :func:`resolve_prompt_cache_scope` (None on failure/empty).
Consumers treat None as "use the physical session_id"; at turn_context's
call site an exception inside the ``set_runtime_main(...)`` argument list
would skip the whole runtime binding, not just the cache scope.
"""
"""Never-raising variant of :func:`resolve_prompt_cache_scope` (None = use the physical id).
At turn_context an exception inside ``set_runtime_main(...)`` would skip the whole binding."""
try:
return resolve_prompt_cache_scope(agent) or None
except Exception:
+80 -196
View File
@@ -1,10 +1,8 @@
"""Anthropic prompt caching strategy — pure functions, no AIAgent dependency.
Default layout: 4 cache_control breakpoints — the static system prefix, the end
of the system prompt, and the last 2 non-system messages. Without a static
prefix: one system breakpoint plus the last 3 messages. All markers share one
TTL (5m or 1h). This keeps intra-session caching while letting new sessions
reuse the stable system-prompt prefix.
Default layout: 4 cache_control breakpoints — the static system prefix, the end of the
system prompt, and the last 2 non-system messages (without a static prefix: one system
breakpoint plus the last 3 messages). All markers share one TTL (5m or 1h).
"""
import copy
@@ -22,17 +20,13 @@ class PromptCachePlan:
tools: List[Dict[str, Any]]
def envelope_tool_part_cache_markers_supported(
provider: str | None, base_url: str | None
) -> bool:
def envelope_tool_part_cache_markers_supported(provider: str | None, base_url: str | None) -> bool:
"""Whether the envelope-layout route honors part-level markers on role:tool.
OpenRouter (and Nous Portal, which proxies to it) relocate a part-level
``cache_control`` onto the ``tool_result`` block during OpenAI→Anthropic
translation. LiteLLM-style proxies copy parts verbatim, so the marker lands
at ``tool_result.content[0]`` — forbidden by the Anthropic schema, a
non-retryable 400. On those routes tool messages carry no part markers and
the breakpoint budget reallocates to the nearest eligible message.
OpenRouter (and Nous Portal) relocate a part-level ``cache_control`` onto the
``tool_result`` block; LiteLLM-style proxies copy parts verbatim, so the marker lands at
``tool_result.content[0]`` — a non-retryable 400. There, tool messages carry no part
markers and the breakpoint budget reallocates to the nearest eligible message.
"""
from agent.agent_runtime_helpers import _is_litellm_route
@@ -47,10 +41,7 @@ def _text_part(text: str, cache_marker: dict | None = None) -> dict:
def _apply_cache_marker(
msg: dict,
cache_marker: dict,
native_anthropic: bool = False,
tool_part_markers: bool = True,
msg: dict, cache_marker: dict, native_anthropic: bool = False, tool_part_markers: bool = True,
) -> None:
"""Add cache_control to a single message, handling all format variations."""
role = msg.get("role", "")
@@ -61,14 +52,12 @@ def _apply_cache_marker(
msg["cache_control"] = cache_marker
return
if role == "tool" and not tool_part_markers:
# LiteLLM-style envelope: a part marker becomes
# tool_result.content[0].cache_control → non-retryable 400.
# LiteLLM-style envelope: a part marker → tool_result.content[0] → non-retryable 400.
return
if content is None or content == "":
# Envelope layout: OpenRouter rejects top-level cache_control on
# role:tool (silent hang), and ignores it on empty assistant turns
# (pure tool_calls) — neither has a content part to carry it.
# Envelope layout: OpenRouter rejects top-level cache_control on role:tool (silent
# hang) and ignores it on empty assistant turns — no content part to carry it.
if role in ("tool", "assistant") and not native_anthropic:
return
msg["cache_control"] = cache_marker
@@ -77,10 +66,8 @@ def _apply_cache_marker(
if isinstance(content, str):
stable_prefix = find_stable_prefix(content) if role == "user" else None
if stable_prefix is not None and content[len(stable_prefix):].strip():
# Builder-declared boundary: the scaffold carries the breakpoint and
# the volatile tail rides unmarked, so a changed ticket ID/timestamp
# no longer invalidates the skill body. Request-local only — the
# stored message stays a plain string.
# Builder-declared boundary: the scaffold carries the breakpoint and the volatile
# tail rides unmarked. Request-local only — the stored message stays a string.
msg["content"] = [
_text_part(stable_prefix, cache_marker),
_text_part(content[len(stable_prefix):]),
@@ -93,17 +80,13 @@ def _apply_cache_marker(
content[-1]["cache_control"] = cache_marker
def _can_carry_marker(
msg: dict, native_anthropic: bool, tool_part_markers: bool = True
) -> bool:
def _can_carry_marker(msg: dict, native_anthropic: bool, tool_part_markers: bool = True) -> bool:
"""True if a marker on this message is actually honored by the provider.
Native Anthropic honors every message (the adapter relocates top-level
markers). The envelope layout only honors markers inside content parts, so
empty-content messages would waste one of the four breakpoints; with
``tool_part_markers=False`` (LiteLLM-style routes) every role:tool message
is excluded too, since its part marker would be rejected with a 400.
Must agree with :func:`_apply_cache_marker`, which marks only the LAST part.
Native Anthropic honors every message. The envelope layout only honors markers inside
content parts, so empty-content messages would waste a breakpoint; with
``tool_part_markers=False`` every role:tool message is excluded too (400). Must agree
with :func:`_apply_cache_marker`, which marks only the LAST part.
"""
if native_anthropic:
return True
@@ -111,8 +94,6 @@ def _can_carry_marker(
return False
content = msg.get("content")
if isinstance(content, list):
# Mirrors _apply_cache_marker (marks only the LAST part): a list whose
# last element isn't a dict cannot receive a marker.
return bool(content) and isinstance(content[-1], dict)
return isinstance(content, str) and content != ""
@@ -125,33 +106,19 @@ def _build_marker(ttl: str) -> Dict[str, str]:
return marker
# Alibaba-family providers (Qwen routes): documented five-minute context cache,
# Anthropic 1h tier rejected. Shared with
# agent_runtime_helpers.anthropic_prompt_cache_policy so the cache-policy
# opt-in and the TTL clamp never desync. Do NOT narrow this set to extend a
# TTL — it also drives the marker-layout opt-in, so narrowing DISABLES caching.
ALIBABA_FAMILY_PROVIDERS = frozenset({
"opencode",
"opencode-go",
"opencode-zen",
"alibaba",
})
# Alibaba-family providers (Qwen routes): five-minute context cache, 1h tier rejected. Shared
# with agent_runtime_helpers.anthropic_prompt_cache_policy so the opt-in and the TTL clamp
# never desync. Do NOT narrow this set to extend a TTL — narrowing DISABLES caching.
ALIBABA_FAMILY_PROVIDERS = frozenset({"opencode", "opencode-go", "opencode-zen", "alibaba"})
# 1h-tier ALLOW-list: only routes wire-measured to retain a 1h marker (delayed
# read past 5 minutes with no intervening call). Other opencode routes stay
# clamped because they are UNMEASURED, not known-bad. opencode-go labels every
# write `ephemeral_5m_input_tokens` regardless of requested ttl; that label is
# not evidence of the retention window.
MEASURED_1H_PROVIDERS = frozenset({
"opencode-go",
})
# 1h-tier ALLOW-list: only routes wire-measured to retain a 1h marker. Other opencode routes
# are UNMEASURED, not known-bad (opencode-go's `ephemeral_5m_input_tokens` label is not
# evidence of the retention window).
MEASURED_1H_PROVIDERS = frozenset({"opencode-go"})
# Models measured to ignore the 1h tier on a MEASURED_1H_PROVIDERS route.
# Consulted only there: the same model on its own Anthropic-compatible endpoint
# is a separate cache-eligible route and must not inherit this clamp.
NO_1H_TIER_MODELS = frozenset({
"minimax-m2.5",
})
# Models measured to ignore the 1h tier on a MEASURED_1H_PROVIDERS route; consulted only
# there (the same model on its own endpoint is a separate route).
NO_1H_TIER_MODELS = frozenset({"minimax-m2.5"})
def _flat_model(model: str) -> str:
@@ -160,34 +127,21 @@ def _flat_model(model: str) -> str:
def is_qwen_model(model: str) -> bool:
"""True when ``model`` names a Qwen-family model (case-insensitive).
Shared with ``agent_runtime_helpers.anthropic_prompt_cache_policy`` so the
cache-policy opt-in and the TTL clamp never desync.
"""
"""True when ``model`` names a Qwen-family model (shared with anthropic_prompt_cache_policy)."""
return "qwen" in (model or "").lower()
def effective_cache_ttl(
ttl: str | None,
*,
model: str = "",
provider: str = "",
) -> str:
"""Clamp a requested cache TTL to what the destination route supports.
def effective_cache_ttl(ttl: str | None, *, model: str = "", provider: str = "") -> str:
"""Clamp a requested cache TTL to what the destination route supports (``None`` → ``5m``).
Qwen/Alibaba routes document a five-minute window and drop the ``1h``
tier, so a configured ``1h`` regresses to ``5m`` there — except on
``MEASURED_1H_PROVIDERS``, which keep ``1h`` minus any ``NO_1H_TIER_MODELS``
model. The measured-route check runs BEFORE the generic Qwen clamp, which
would otherwise swallow every Qwen model on it. ``None`` resolves to ``5m``.
Qwen/Alibaba routes drop the ``1h`` tier, so ``1h`` regresses to ``5m`` there — except
on ``MEASURED_1H_PROVIDERS`` (minus ``NO_1H_TIER_MODELS``). The measured-route check
runs BEFORE the generic Qwen clamp, which would otherwise swallow every Qwen model on it.
"""
if ttl != "1h":
return ttl or "5m"
provider_lower = (provider or "").lower()
if provider_lower in MEASURED_1H_PROVIDERS:
# The per-model denial stays nested so an opencode-go observation
# cannot reclamp the same model on another route.
return "5m" if _flat_model(model) in NO_1H_TIER_MODELS else "1h"
if is_qwen_model(model) or provider_lower in ALIBABA_FAMILY_PROVIDERS:
return "5m"
@@ -195,25 +149,15 @@ def effective_cache_ttl(
def _apply_system_cache_markers(
message: dict,
cache_marker: dict,
static_system_prefix: str | None,
*,
native_anthropic: bool,
mark_suffix: bool = True,
fallback_to_whole: bool = True,
message: dict, cache_marker: dict, static_system_prefix: str | None, *,
native_anthropic: bool, mark_suffix: bool = True, fallback_to_whole: bool = True,
) -> int:
"""Mark the static system prefix (and optionally the full prompt).
"""Mark the static system prefix (and optionally the full prompt); returns markers applied.
The system prompt stays one stored string; it is split only in the
outgoing request so persistence and non-Anthropic transports are
unchanged. ``mark_suffix=False`` is the tool-cache-plan layout (suffix
unmarked, its budget spent on the tools array). ``fallback_to_whole=False``
marks nothing when the prefix split is impossible. When the prompt IS the
prefix (empty/whitespace suffix) the whole message is marked as one block —
never a split with an empty text block, which Anthropic rejects.
Returns the number of markers applied (0, 1, or 2).
The system prompt stays one stored string, split only in the outgoing request.
``mark_suffix=False`` is the tool-cache-plan layout (suffix budget spent on the tools
array). ``fallback_to_whole=False`` marks nothing when the split is impossible. When the
prompt IS the prefix the whole message is one block — never an empty text block (400).
"""
content = message.get("content")
if (
@@ -241,22 +185,15 @@ def _has_part_marker(content: Any) -> bool:
)
def strip_anthropic_cache_control(
api_messages: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""Remove ``cache_control`` markers and undo decoration-produced list shapes.
def strip_anthropic_cache_control(api_messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Remove ``cache_control`` markers and undo decoration-produced list shapes (in place).
Used before re-decorating after a mid-turn provider failover, so the
mutated undecorated shape is preserved while markers match the new
provider's policy. Flattening back to a plain string is restricted to the
exact shapes :func:`apply_anthropic_cache_control` produces from string
content — a single text part, the two-part ``[static, volatile]`` system
split, or the two-part skill split — so the ``""``-join is provably
byte-exact; organic multi-part text and parts with extra keys keep their
structure. Marker removal is copy-on-write on part dicts: parts can alias
caller-held lists and stripping must never rewrite the stored transcript.
Mutates the top-level message dicts in place and returns the same list.
Used before re-decorating after a mid-turn provider failover. Flattening back to a
string is restricted to the exact shapes :func:`apply_anthropic_cache_control` produces
from string content — a single text part, the two-part ``[static, volatile]`` system
split, or the two-part skill split — so the ``""``-join is provably byte-exact. Marker
removal is copy-on-write on part dicts: parts can alias caller-held lists and stripping
must never rewrite the stored transcript.
"""
for msg in api_messages:
if not isinstance(msg, dict):
@@ -265,11 +202,8 @@ def strip_anthropic_cache_control(
content = msg.get("content")
if not isinstance(content, list):
continue
# The builder-declared skill split is the only decoration that marks
# the FIRST part of a user message (list content is otherwise marked
# on the last part; the [static, volatile] split is system-only), so
# the shape alone identifies it even after the prefix registry has
# evicted the entry.
# The skill split is the only decoration marking the FIRST part of a user message,
# so the shape alone identifies it even after the prefix registry evicted the entry.
skill_split_shape = (
msg.get("role") == "user"
and len(content) == 2
@@ -322,14 +256,10 @@ def _count_cache_markers(messages: List[Dict[str, Any]], tools: List[Dict[str, A
count += sum(
1 for part in message["content"] if isinstance(part, dict) and "cache_control" in part
)
return count + sum(
1 for tool in tools if isinstance(tool, dict) and "cache_control" in tool
)
return count + sum(1 for tool in tools if isinstance(tool, dict) and "cache_control" in tool)
def _completed_transaction_endpoint_indexes(
messages: List[Dict[str, Any]], *, native_anthropic: bool,
) -> List[int]:
def _completed_transaction_endpoint_indexes(messages: List[Dict[str, Any]], *, native_anthropic: bool) -> List[int]:
"""Select legal ends of completed tool runs and ordinary turns."""
def _tool_run_end(start: int) -> int:
@@ -374,80 +304,49 @@ def _completed_transaction_endpoint_indexes(
def build_prompt_cache_plan(
api_messages: List[Dict[str, Any]],
tools: List[Dict[str, Any]] | None,
*,
cache_ttl: str = "5m",
native_anthropic: bool = False,
static_system_prefix: str | None = None,
direct_native_tool_cache: bool = False,
tool_part_markers: bool = True,
api_messages: List[Dict[str, Any]], tools: List[Dict[str, Any]] | None, *,
cache_ttl: str = "5m", native_anthropic: bool = False, static_system_prefix: str | None = None,
direct_native_tool_cache: bool = False, tool_part_markers: bool = True,
) -> PromptCachePlan:
"""Build isolated cache sections for one resolved request destination.
``tool_part_markers=False`` (LiteLLM-style envelope routes) keeps
``cache_control`` off role:tool content parts; breakpoints reallocate to
the nearest eligible non-tool message.
"""
"""Build isolated cache sections for one resolved request destination
(``tool_part_markers=False`` keeps markers off role:tool parts on LiteLLM-style routes)."""
messages = copy.deepcopy(api_messages or [])
strip_anthropic_cache_control(messages)
planned_tools = strip_anthropic_tool_cache_control(tools)
if not direct_native_tool_cache or not planned_tools:
planned_messages = apply_anthropic_cache_control(
messages,
cache_ttl=cache_ttl,
native_anthropic=native_anthropic,
static_system_prefix=static_system_prefix,
tool_part_markers=tool_part_markers,
messages, cache_ttl=cache_ttl, native_anthropic=native_anthropic,
static_system_prefix=static_system_prefix, tool_part_markers=tool_part_markers,
)
return PromptCachePlan(messages=planned_messages, tools=planned_tools)
marker = _build_marker(cache_ttl)
if (
messages
and isinstance(messages[0], dict)
and messages[0].get("role") == "system"
):
# Tool-cache layout: only the static prefix carries a system-side
# marker; the volatile suffix's budget is spent on the tools array.
if messages and isinstance(messages[0], dict) and messages[0].get("role") == "system":
# Tool-cache layout: only the static prefix carries a system-side marker; the
# volatile suffix's budget is spent on the tools array.
_apply_system_cache_markers(
messages[0],
marker,
static_system_prefix,
native_anthropic=True,
mark_suffix=False,
fallback_to_whole=False,
messages[0], marker, static_system_prefix,
native_anthropic=True, mark_suffix=False, fallback_to_whole=False,
)
planned_tools[-1]["cache_control"] = dict(marker)
for endpoint in _completed_transaction_endpoint_indexes(
messages,
native_anthropic=True,
)[-2:]:
for endpoint in _completed_transaction_endpoint_indexes(messages, native_anthropic=True)[-2:]:
_apply_cache_marker(messages[endpoint], marker, native_anthropic=True)
return PromptCachePlan(messages=messages, tools=planned_tools)
def apply_anthropic_cache_control(
api_messages: List[Dict[str, Any]],
cache_ttl: str = "5m",
native_anthropic: bool = False,
static_system_prefix: str | None = None,
tool_part_markers: bool = True,
api_messages: List[Dict[str, Any]], cache_ttl: str = "5m", native_anthropic: bool = False,
static_system_prefix: str | None = None, tool_part_markers: bool = True,
) -> List[Dict[str, Any]]:
"""Apply Anthropic cache-control markers to API messages.
With a matching ``static_system_prefix`` the prefix gets an early marker
and the full system prompt a trailing one; the remaining two markers go to
the latest cacheable non-system messages. Without it, the legacy
system-and-3 layout applies. Idempotent: pre-existing markers are stripped
from a per-message copy first, so repeated calls never accumulate past 4
markers; a shallow top-level copy suffices because
:func:`strip_anthropic_cache_control` is copy-on-write on content parts.
Returns:
Shallow copy of message list with selective deep copies of modified messages.
With a matching ``static_system_prefix`` the prefix and the full system prompt each get
a marker and the remaining two go to the latest cacheable non-system messages; without
it, the legacy system-and-3 layout applies. Idempotent: pre-existing markers are
stripped from a per-message copy first (shallow copy suffices — stripping is
copy-on-write on parts). Returns a shallow list copy with deep copies of modified messages.
"""
if not api_messages:
return api_messages
@@ -460,34 +359,19 @@ def apply_anthropic_cache_control(
messages[i] = strip_anthropic_cache_control([dict(msg)])[0]
breakpoints_used = 0
if messages[0].get("role") == "system":
messages[0] = copy.deepcopy(messages[0])
breakpoints_used = _apply_system_cache_markers(
messages[0],
marker,
static_system_prefix,
native_anthropic=native_anthropic,
messages[0], marker, static_system_prefix, native_anthropic=native_anthropic,
)
remaining = 4 - breakpoints_used
non_sys = [
i
for i in range(len(messages))
i for i in range(len(messages))
if messages[i].get("role") != "system"
and _can_carry_marker(
messages[i],
native_anthropic=native_anthropic,
tool_part_markers=tool_part_markers,
)
and _can_carry_marker(messages[i], native_anthropic=native_anthropic, tool_part_markers=tool_part_markers)
]
for idx in non_sys[-remaining:]:
for idx in non_sys[-(4 - breakpoints_used):]:
messages[idx] = copy.deepcopy(messages[idx])
_apply_cache_marker(
messages[idx],
marker,
native_anthropic=native_anthropic,
tool_part_markers=tool_part_markers,
)
_apply_cache_marker(messages[idx], marker, native_anthropic=native_anthropic, tool_part_markers=tool_part_markers)
return messages