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