Files
hermes-agent/agent/memory_manager.py
Teknium 8dcb2b6ada refactor(agent/providers): shared ProviderBase/CatalogProviderBase and provider_media; compact contract docs
- provider_base.py: ProviderBase (name/display_name/get_setup_schema) and
  CatalogProviderBase (default_model/list_models/is_available) replace the
  identical default-method bodies duplicated across 7 provider ABCs
- provider_media.py: one save_b64/save_bytes/save_url/cache_dir implementation
  behind image_gen_provider and video_gen_provider
- memory_manager.py: _each_provider fan-out helper replaces per-hook
  try/except loops; _signature_params/_has_var_kwargs unify signature probes
- image_routing.py: _resolve_inference_value shared by base_url/api_key
  resolution; _dict_or_empty/_clean_str/_custom_provider_entries helpers
- MemoryProvider/ContextEngine/TTS/browser/web/terminal-env ABC docstrings
  compacted to their invariants; method names and signatures unchanged
2026-09-02 13:53:28 -07:00

1200 lines
48 KiB
Python

"""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 = "<memory-context>"
_CLOSE_TAG = "</memory-context>"
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 (
"<memory-context>\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"
"</memory-context>"
)
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 '<name>' <label>: <exc>``.
Returns the successful results in provider order.
"""
results: List[Any] = []
for provider in self._providers if providers is None else providers:
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,
)
return results
# -- Registration --------------------------------------------------------
def add_provider(self, provider: MemoryProvider) -> None:
"""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"
)
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,
)
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, like the TTS/browser/search
# provider registries do.
from toolsets import _HERMES_CORE_TOOLS
_core_tool_names = set(_HERMES_CORE_TOOLS)
for raw_schema in provider.get_tool_schemas():
schema = normalize_tool_schema(raw_schema)
if schema is None:
continue
tool_name = schema["name"]
if tool_name in _core_tool_names:
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,
)
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,
)
else:
self._tool_to_provider[tool_name] = provider
logger.info(
"Memory provider '%s' registered (%d tools)",
provider.name,
len(provider.get_tool_schemas()),
)
@property
def providers(self) -> List[MemoryProvider]:
"""All registered providers in order."""
return list(self._providers)
def get_provider(self, name: str) -> Optional[MemoryProvider]:
"""Get a provider by name, or None if not registered."""
return next((p for p in self._providers if p.name == name), None)
# -- System prompt -------------------------------------------------------
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,
)
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. Recover just the user's instruction once, for the whole
fan-out. Non-skill text passes through; a bare invocation (no instruction)
yields None since there is nothing worth remembering.
"""
return extract_user_instruction_from_skill_message(text)
def prefetch_all(self, query: str, *, session_id: str = "") -> str:
"""Merge non-empty prefetch context from all providers (failures are non-fatal)."""
clean_query = self._strip_skill_scaffolding(query)
if not clean_query:
return ""
parts = self._each_provider(
"prefetch failed (non-fatal)",
lambda p: _nonblank(self._prefetch_provider(p, clean_query, session_id=session_id)),
)
return "\n\n".join(p for p in parts if p)
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.
"""
if provider.name == "builtin":
return provider.prefetch(query, session_id=session_id)
result_box: Dict[str, str] = {}
error_box: Dict[str, Exception] = {}
def _run() -> None:
try:
result_box["value"] = provider.prefetch(query, session_id=session_id) or ""
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}",
)
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,
)
return ""
self._external_prefetch_threads.pop(provider.name, None)
self._external_prefetch_threads[provider.name] = thread
thread.start()
thread.join(self._external_prefetch_timeout)
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,
)
return ""
with self._external_prefetch_lock:
if self._external_prefetch_threads.get(provider.name) is thread:
self._external_prefetch_threads.pop(provider.name, None)
if error_box:
raise error_box["value"]
return result_box.get("value", "")
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.
"""
segments: List[str] = []
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).
detail = "recalled relevant memory"
segments.append(f"{status.glyph} {status.provider_label} — {detail}")
return " ".join(segments)
def queue_prefetch_all(self, query: str, *, session_id: str = "") -> None:
"""Queue background prefetch on all providers for the next turn (see ``sync_all``)."""
providers = list(self._providers)
if not providers:
return
clean_query = self._strip_skill_scaffolding(query)
if not clean_query:
return
self._submit_background(
lambda: self._each_provider(
"queue_prefetch failed (non-fatal)",
lambda p: p.queue_prefetch(clean_query, session_id=session_id),
providers=providers,
),
kind="prefetch",
)
# -- Sync ----------------------------------------------------------------
@staticmethod
def _provider_sync_accepts_messages(provider: MemoryProvider) -> bool:
"""Whether ``sync_turn`` accepts a ``messages`` keyword (uninspectable → assume yes)."""
params = _signature_params(provider.sync_turn)
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,
) -> 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, so every interface showed the agent "running" and follow-up
messages triggered interrupts. The single worker also serializes writes so
turn N lands before turn N+1 without provider-side ordering logic.
"""
providers = list(self._providers)
if not providers:
return
clean_user_content = self._strip_skill_scaffolding(user_content)
if not clean_user_content:
return
def _sync(provider: MemoryProvider) -> None:
kwargs: Dict[str, Any] = {"session_id": session_id}
if messages is not None and self._provider_sync_accepts_messages(provider):
kwargs["messages"] = messages
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
)
)
# -- Background dispatch -------------------------------------------------
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).
"""
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
try:
# 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)
except RuntimeError:
if self._shutting_down:
logger.warning("Memory manager shut down during %s submission; task rejected", kind)
return
_run_inline()
def _forget_background_future(self, future: Future) -> None:
with self._sync_executor_lock:
self._background_futures.pop(future, None)
def _get_sync_executor(self) -> Optional[ThreadPoolExecutor]:
"""Lazily create the single-worker background executor (None once shutting down)."""
if self._shutting_down:
return None
if self._sync_executor is not None:
return self._sync_executor
with self._sync_executor_lock:
if self._shutting_down:
return None
if self._sync_executor is None:
try:
# 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",
)
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.
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.
"""
executor = self._sync_executor
if executor is None:
return True
try:
fut = executor.submit(lambda: None)
except RuntimeError:
# Executor already shut down — nothing pending.
return True
try:
fut.result(timeout=timeout)
return True
except Exception:
return False
# -- Tools ---------------------------------------------------------------
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.
"""
from toolsets import _HERMES_CORE_TOOLS
_core_tool_names = set(_HERMES_CORE_TOOLS)
schemas: List[Dict[str, Any]] = []
seen = set()
def _collect(provider: MemoryProvider) -> None:
for raw_schema in provider.get_tool_schemas():
schema = normalize_tool_schema(raw_schema)
if schema is None:
logger.warning(
"Memory provider '%s' returned a tool schema with "
"no resolvable name; skipping (%r)",
provider.name, raw_schema,
)
continue
name = schema["name"]
if name not in _core_tool_names and name not in seen:
schemas.append(schema)
seen.add(name)
self._each_provider("get_tool_schemas() failed", _collect, level=logging.WARNING)
return schemas
def get_all_tool_names(self) -> set:
"""Return set of all tool names across all providers."""
return set(self._tool_to_provider.keys())
def has_tool(self, tool_name: str) -> bool:
"""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:
"""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:
return tool_error(f"No memory provider handles tool '{tool_name}'")
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,
)
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),
)
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,
)
def commit_session_boundary_async(
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
to the new session. One task on the single FIFO worker gives both an
immediate return and ordering against every other provider write. If
the executor is unavailable, ``_submit_background`` runs it inline.
"""
if not self._providers:
return
snapshot = list(messages or [])
def _run() -> None:
try:
self.on_session_end(snapshot)
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,
)
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,
) -> None:
"""Notify providers that ``AIAgent.session_id`` rotated without teardown.
Fires on ``/resume``, ``/branch``, ``/reset``, ``/new`` and compression;
providers refresh cached per-session state so later writes land in the
right record. ``rewound=True`` (``/undo``) means the id is unchanged but
the transcript was truncated.
"""
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.
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
),
)
@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,
)
)
except (TypeError, ValueError):
return None
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
for p in self._providers
)
def on_pre_compress(
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.
"""
parts = []
checkpoint_succeeded = False
for provider in self._providers:
provider_version = self._checkpoint_api_version(provider)
if provider_version is None:
provider_version = _LEGACY_PRE_COMPRESS_API_VERSION
is_checkpoint_provider = provider_version >= checkpoint_api_version
provider_messages = messages
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.
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,
)
if require_checkpoint and is_checkpoint_provider:
raise
else:
if is_checkpoint_provider:
checkpoint_succeeded = True
if require_checkpoint and not checkpoint_succeeded:
raise RuntimeError(
"No active memory provider completed pre-compress checkpoint "
f"API v{checkpoint_api_version}"
)
return "\n\n".join(parts)
@staticmethod
def _provider_memory_write_metadata_mode(provider: MemoryProvider) -> str:
"""How to pass metadata to ``on_memory_write``: "keyword", "positional", or "legacy" (none)."""
params = _signature_params(provider.on_memory_write)
if params is None or _has_var_kwargs(params) or "metadata" in params:
return "keyword"
accepted = sum(p.kind is not inspect.Parameter.VAR_POSITIONAL for p in params.values())
return "positional" if accepted >= 4 else "legacy"
def on_memory_write(
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)."""
def _notify(provider: MemoryProvider) -> None:
mode = self._provider_memory_write_metadata_mode(provider)
if mode == "keyword":
provider.on_memory_write(action, target, content, metadata=dict(metadata or {}))
elif mode == "positional":
provider.on_memory_write(action, target, content, dict(metadata or {}))
else:
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"],
)
# 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.
_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.
"""
if isinstance(result, str):
try:
result = json.loads(result)
except Exception:
return False
if not isinstance(result, dict):
return False
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],
*,
build_metadata: Optional[Callable[[], Dict[str, Any]]] = None,
) -> None:
"""Mirror a built-in memory tool call to external providers.
Single entry point the agent loop calls after the ``memory`` tool runs:
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).
"""
if not self._memory_tool_result_succeeded(tool_result):
return
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"),
}]
for op in operations:
if not isinstance(op, dict):
continue
action = str(op.get("action") or "")
if action not in self._MIRRORED_MEMORY_ACTIONS:
continue
try:
metadata = dict(build_metadata() if build_metadata else {})
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,
)
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:
"""Notify all providers that a subagent completed."""
self._each_provider(
"on_delegation failed",
lambda p: p.on_delegation(task, result, child_session_id=child_session_id, **kwargs),
)
def shutdown_all(self) -> None:
"""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)),
)
@property
def shutdown_drain_state(self) -> Dict[str, Any]:
"""Snapshot of the most recent bounded shutdown drain outcome."""
with self._sync_executor_lock:
return dict(self._shutdown_drain_state)
def _drain_sync_executor(self) -> None:
"""Give queued FIFO work a bounded chance, then abandon explicitly."""
with self._sync_executor_lock:
self._shutting_down = True
executor = self._sync_executor
self._sync_executor = None
tracked = dict(self._background_futures)
self._shutdown_drain_state = {
"status": "draining" if executor is not None else "drained",
"abandoned_writes": 0,
"abandoned_prefetches": 0,
"active_tasks": sum(not future.done() for future in tracked),
}
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.
executor.shutdown(wait=False, cancel_futures=False)
_, pending = wait(tuple(tracked), timeout=_SYNC_DRAIN_TIMEOUT_S)
if not pending:
with self._sync_executor_lock:
self._shutdown_drain_state.update(status="drained", active_tasks=0)
return
abandoned_writes = abandoned_prefetches = active_tasks = 0
for future in pending:
if not future.cancel():
active_tasks += 1
elif tracked[future] == "prefetch":
abandoned_prefetches += 1
else:
abandoned_writes += 1
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,
)
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,
)
def initialize_all(self, session_id: str, **kwargs) -> None:
"""Initialize all providers, injecting ``hermes_home`` so they resolve profile-scoped paths."""
if "hermes_home" not in kwargs:
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,
)