"""MemoryManager — orchestrates memory providers for the agent.
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.
Usage in run_agent.py:
self._memory_manager = MemoryManager()
self._memory_manager.add_provider(plugin_provider) # at most one external
prompt_parts.append(self._memory_manager.build_system_prompt())
context = self._memory_manager.prefetch_all(user_message) # pre-turn
self._memory_manager.sync_all(user_msg, assistant_response) # post-turn
self._memory_manager.queue_prefetch_all(user_msg)
"""
from __future__ import annotations
import contextvars
import inspect
import json
import logging
import re
import threading
from concurrent.futures import Future, ThreadPoolExecutor, wait
from functools import partial
from typing import Any, Callable, Dict, List, Optional
from agent.memory_provider import MemoryProvider, PRE_COMPRESS_CHECKPOINT_API_VERSION
from agent.skill_commands import extract_user_instruction_from_skill_message
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).
_LEGACY_PRE_COMPRESS_API_VERSION = 1
# How long shutdown_all() waits for in-flight background sync/prefetch work to
# drain before abandoning it. Worker threads are daemon, so a wedged provider
# never blocks interpreter exit — it dies with the process past this window.
_SYNC_DRAIN_TIMEOUT_S = 5.0
_EXTERNAL_PREFETCH_TIMEOUT_S = 8.0
_VAR_KEYWORD = inspect.Parameter.VAR_KEYWORD
# ---------------------------------------------------------------------------
# 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)."""
try:
return inspect.signature(fn).parameters
except (TypeError, ValueError):
return None
def _has_var_kwargs(params) -> bool:
return any(p.kind is _VAR_KEYWORD for p in params.values())
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
re-raise as a checkpoint failure even though the durable write succeeded.
Unreadable signatures conservatively report False.
"""
params = _signature_params(fn)
if params is None:
return False
if _has_var_kwargs(params):
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,
)
def _ctx_bound(fn: Callable[[], Any]) -> Callable[[], Any]:
"""Bind ``fn`` to the CALLER's contextvars for execution on another thread.
Profile isolation in multi-profile processes (gateway multiplexer, dashboard,
cron) 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.
"""
return partial(contextvars.copy_context().run, fn)
# ---------------------------------------------------------------------------
# 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"}`` which callers
wrap as ``{"type": "function", "function": schema}``. 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 (HTTP 400),
disabling every tool. Both shapes are normalized here so callers can skip
nameless entries with a warning instead of poisoning the request.
"""
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
def memory_provider_tools_enabled(
enabled_toolsets: Optional[List[str]],
disabled_toolsets: Optional[List[str]] = None,
*,
memory_tool_present: bool = False,
) -> bool:
"""Return whether external memory-provider tools should be exposed."""
if disabled_toolsets and "memory" in disabled_toolsets:
return False
if memory_tool_present or enabled_toolsets is None:
return True
if not enabled_toolsets:
return False
if "memory" in enabled_toolsets:
return True
try:
from toolsets import resolve_toolset
return any("memory" in resolve_toolset(name) for name in enabled_toolsets)
except Exception:
logger.debug("Failed to resolve enabled toolsets for memory-provider tools", exc_info=True)
return False
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.
"""
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,
)
def inject_memory_provider_tools(agent: Any) -> int:
"""Append external memory-provider tool schemas to an agent tool surface; return count added."""
memory_manager = getattr(agent, "_memory_manager", None)
tools = getattr(agent, "tools", None)
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)
}
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.
_providers = [
p for p in (getattr(memory_manager, "providers", None) or [])
if getattr(p, "name", "") != "builtin"
]
if _providers:
logger.info(
"Memory provider(s) %s configured but the 'memory' toolset is "
"gated off for this session (platform_toolsets / "
"agent.disabled_toolsets) — provider tools and system-prompt "
"block are both withheld.",
[getattr(p, "name", type(p).__name__) for p in _providers],
)
return 0
get_schemas = getattr(memory_manager, "get_all_tool_schemas", None)
if not callable(get_schemas):
return 0
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
added = 0
for raw_schema in get_schemas():
schema = normalize_tool_schema(raw_schema)
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,
)
continue
tool_name = schema["name"]
if tool_name in existing_tool_names:
continue
tools.append({"type": "function", "function": schema})
valid_tool_names.add(tool_name)
existing_tool_names.add(tool_name)
added += 1
return added
# ---------------------------------------------------------------------------
# Context fencing helpers
# ---------------------------------------------------------------------------
_FENCE_TAG_RE = re.compile(r'?\s*memory-context\s*>', re.IGNORECASE)
_INTERNAL_CONTEXT_RE = re.compile(
r'<\s*memory-context\s*>[\s\S]*?\s*memory-context\s*>',
re.IGNORECASE,
)
_INTERNAL_NOTE_RE = re.compile(
r'\[System note:\s*The following is recalled memory context,\s*NOT new user input\.\s*Treat as (?:informational background data|authoritative reference data[^\]]*)\.\]\s*',
re.IGNORECASE,
)
def sanitize_context(text: str) -> str:
"""Strip fence tags, injected context blocks, and system notes from provider output."""
text = _INTERNAL_CONTEXT_RE.sub('', text)
text = _INTERNAL_NOTE_RE.sub('', text)
return _FENCE_TAG_RE.sub('', text)
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 (including the system-note line). Create a
fresh scrubber (or ``reset()``) per top-level response; call ``flush()`` at
end of stream.
"""
_OPEN_TAG = ""
_CLOSE_TAG = ""
def __init__(self) -> None:
self.reset()
def reset(self) -> None:
self._in_span: bool = False
self._buf: str = ""
self._at_block_boundary: bool = True
def feed(self, text: str) -> str:
"""Return the visible portion of ``text``; a possible partial tag tail is held for the next call."""
if not text:
return ""
buf = self._buf + text
self._buf = ""
out: list[str] = []
while buf:
if self._in_span:
idx = buf.lower().find(self._CLOSE_TAG)
if idx == -1:
# Hold back a potential partial close tag; drop the rest.
held = self._max_partial_suffix(buf, self._CLOSE_TAG)
self._buf = buf[-held:] if held else ""
return "".join(out)
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)
)
self._append_visible(out, buf[:-held] if held else buf)
if held:
self._buf = buf[-held:]
return "".join(out)
if idx > 0:
self._append_visible(out, buf[:idx])
buf = buf[idx + len(self._OPEN_TAG):]
self._in_span = True
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.
"""
if self._in_span:
self._buf = ""
self._in_span = False
return ""
tail = self._buf
self._buf = ""
return tail
@staticmethod
def _max_partial_suffix(buf: str, tag: str) -> int:
"""Length of the longest buf-suffix that is a (case-insensitive) prefix of ``tag``, else 0."""
tag_lower = tag.lower()
buf_lower = buf.lower()
for i in range(min(len(buf_lower), len(tag_lower) - 1), 0, -1):
if tag_lower.startswith(buf_lower[-i:]):
return i
return 0
def _find_boundary_open_tag(self, buf: str) -> int:
"""Find an opening fence only when it starts a block-like span (own line, newline after)."""
buf_lower = buf.lower()
search_start = 0
while True:
idx = buf_lower.find(self._OPEN_TAG, search_start)
if idx == -1:
return -1
if self._is_block_boundary(buf, idx) and self._has_block_opener_suffix(buf, idx):
return idx
search_start = idx + 1
def _max_pending_open_suffix(self, buf: str) -> int:
"""Hold a complete boundary tag at the buffer end until the following char confirms it."""
if not buf.lower().endswith(self._OPEN_TAG):
return 0
if not self._is_block_boundary(buf, len(buf) - len(self._OPEN_TAG)):
return 0
return len(self._OPEN_TAG)
def _has_block_opener_suffix(self, buf: str, idx: int) -> bool:
after_idx = idx + len(self._OPEN_TAG)
return after_idx < len(buf) and buf[after_idx] in "\r\n"
def _is_block_boundary(self, buf: str, idx: int) -> bool:
if idx == 0:
return self._at_block_boundary
preceding = buf[:idx]
last_newline = preceding.rfind("\n")
if last_newline == -1:
return self._at_block_boundary and preceding.strip() == ""
return preceding[last_newline + 1:].strip() == ""
def _append_visible(self, out: list[str], text: str) -> None:
if not text:
return
out.append(text)
last_newline = text.rfind("\n")
if last_newline != -1:
self._at_block_boundary = text[last_newline + 1:].strip() == ""
else:
self._at_block_boundary = self._at_block_boundary and text.strip() == ""
def build_memory_context_block(raw_context: str) -> str:
"""Wrap prefetched memory in a fenced block with system note."""
if not raw_context or not raw_context.strip():
return ""
clean = sanitize_context(raw_context)
if clean != raw_context:
logger.warning("memory provider returned pre-wrapped context; stripped")
return (
"\n"
"[System note: The following is recalled memory context, "
"NOT new user input. Treat as authoritative reference data — "
"this is the agent's persistent memory and should inform all responses.]\n\n"
f"{clean}\n"
""
)
def _nonblank(text: Any) -> Any:
"""Return ``text`` when it has non-whitespace content, else None."""
return text if text and text.strip() else None
class MemoryManager:
"""Orchestrates the built-in provider 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.
"""
def __init__(self, *, external_prefetch_timeout: Optional[float] = None) -> None:
self._providers: List[MemoryProvider] = []
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)
)
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).
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.
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,
}
# -- 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,
) -> List[Any]:
"""Call ``call(provider)`` for each provider, logging and swallowing failures.
``label`` completes the log line ``Memory provider ''