diff --git a/agent/memory_manager.py b/agent/memory_manager.py index 07ebeb9ce1..eb6c589708 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -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'', 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 = "" @@ -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, ) diff --git a/agent/memory_provider.py b/agent/memory_provider.py index cbfd26033b..3c45d58169 100644 --- a/agent/memory_provider.py +++ b/agent/memory_provider.py @@ -1,9 +1,8 @@ """Abstract base class for pluggable memory providers. -Plugins ship in ``plugins/memory//`` 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//``, 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: diff --git a/agent/message_sanitization.py b/agent/message_sanitization.py index f7b22c0755..93c2c859cc 100644 --- a/agent/message_sanitization.py +++ b/agent/message_sanitization.py @@ -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 ``_d`` 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 ``_d`` 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``. diff --git a/agent/micro_compaction.py b/agent/micro_compaction.py index 12ec2b7da6..d043fb13e9 100644 --- a/agent/micro_compaction.py +++ b/agent/micro_compaction.py @@ -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) diff --git a/agent/native_compaction.py b/agent/native_compaction.py index c13967894f..b8b30fabc1 100644 --- a/agent/native_compaction.py +++ b/agent/native_compaction.py @@ -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) diff --git a/agent/prompt_cache_scope.py b/agent/prompt_cache_scope.py index 9be0d054b8..18437a6b53 100644 --- a/agent/prompt_cache_scope.py +++ b/agent/prompt_cache_scope.py @@ -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_`` 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_`` (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_``), 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: diff --git a/agent/prompt_caching.py b/agent/prompt_caching.py index 3c24628167..8398d6543d 100644 --- a/agent/prompt_caching.py +++ b/agent/prompt_caching.py @@ -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