"""MemoryManager — fans the agent's memory hooks out to registered providers. The builtin provider is always allowed; only ONE external plugin provider may be registered at a time (tool-schema bloat, conflicting backends). """ 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, ctx_bound, spawn_context_thread from agent.skill_commands import extract_user_instruction_from_skill_message from tools.hook_output_spill import get_spill_config, spill_if_oversized from tools.registry import tool_error logger = logging.getLogger(__name__) # Providers that predate the checkpoint-API attribute are on the best-effort v1 contract. _LEGACY_PRE_COMPRESS_API_VERSION = 1 # 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) ----- def _signature_params(fn: Callable[..., Any]): """``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 inspect.Parameter.VAR_KEYWORD for p in params.values()) def _accepts_require_checkpoint(fn: Callable[..., Any]) -> bool: """True if ``fn`` can receive the ``require_checkpoint`` keyword (unreadable signatures -> False). Bare-shape v2 providers (``on_pre_compress(self, messages)``) would raise TypeError on the keyword, which the host would re-raise as a checkpoint failure despite a successful write. """ params = _signature_params(fn) if params is None: return False kind = getattr(params.get("require_checkpoint"), "kind", None) return _has_var_kwargs(params) or kind in (inspect.Parameter.KEYWORD_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD) # -- 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"}`` 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", "") return schema if name and isinstance(name, str) else None 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 _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()`` never advertises tools absent from the tool surface. """ tools = getattr(agent, "tools", None) present = isinstance(tools, (list, tuple)) and any(_tool_name(t) == "memory" for t in tools) enabled, disabled = getattr(agent, "enabled_toolsets", None), getattr(agent, "disabled_toolsets", None) return memory_provider_tools_enabled(enabled, disabled, memory_tool_present=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 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. # See #81014. _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 if getattr(agent, "valid_tool_names", None) is None: agent.valid_tool_names = set() existing_tool_names = {_tool_name(tool) for tool in tools if isinstance(tool, dict)} 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, ) elif schema["name"] not in existing_tool_names: tools.append({"type": "function", "function": schema}) agent.valid_tool_names.add(schema["name"]) existing_tool_names.add(schema["name"]) added += 1 return added # -- Context fencing helpers -------------------------------------------------- _FENCE_TAG_RE = re.compile(r'', re.IGNORECASE) _INTERNAL_CONTEXT_RE = re.compile(r'<\s*memory-context\s*>[\s\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.""" for pattern in (_INTERNAL_CONTEXT_RE, _INTERNAL_NOTE_RE, _FENCE_TAG_RE): text = pattern.sub('', text) return text class StreamingContextScrubber: """Stateful scrubber for streaming text whose memory-context spans may straddle deltas. ``sanitize_context`` needs both tags in one string, so a split span would leak to the UI; this holds back partial-tag tails between ``feed()`` calls and drops span interiors. One 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: tag = self._CLOSE_TAG idx = buf.lower().find(tag) held = self._max_partial_suffix(buf, tag) # potential partial close tag else: tag = self._OPEN_TAG idx = self._find_boundary_open_tag(buf) # A complete boundary tag at the buffer end is held until the next char confirms it. n = len(tag) pending = n if buf.lower().endswith(tag) and self._ends_at_block_boundary(buf[:-n]) else 0 held = pending or self._max_partial_suffix(buf, tag) if idx == -1: # Hold back the possible partial tag; inside a span the rest is dropped. if not self._in_span: self._append_visible(out, buf[:-held] if held else buf) self._buf = buf[-held:] if held else "" break if not self._in_span: self._append_visible(out, buf[:idx]) buf = buf[idx + len(tag):] 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 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 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, buf_lower = tag.lower(), buf.lower() span = range(min(len(buf_lower), len(tag_lower) - 1), 0, -1) return next((i for i in span if tag_lower.startswith(buf_lower[-i:])), 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, tag_len = buf.lower(), len(self._OPEN_TAG) idx = buf_lower.find(self._OPEN_TAG) while idx != -1: after_idx = idx + tag_len if self._ends_at_block_boundary(buf[:idx]) and after_idx < len(buf) and buf[after_idx] in "\r\n": return idx idx = buf_lower.find(self._OPEN_TAG, idx + 1) return -1 def _ends_at_block_boundary(self, text: str) -> bool: """Whether emitting ``text`` leaves the stream at a line start (blank tail after the last newline; no newline at all -> only whitespace and already at a boundary).""" head, sep, tail = text.rpartition("\n") return tail.strip() == "" and (bool(sep) or self._at_block_boundary) def _append_visible(self, out: list[str], text: str) -> None: if text: out.append(text) self._at_block_boundary = self._ends_at_block_boundary(text) # A markdown bullet: a marker, whitespace, then content. The whitespace matters — it is what keeps # ``**Preferences**`` (a bold heading) and ``*emphasis*`` out of the rule. _RECALL_BULLET_RE = re.compile(r"[-*+]\s+\S") def _drop_repeated_recall_lines(text: str) -> str: """Drop a recalled bullet that an EARLIER line of this same block already states. Providers merge several stores (and this merges several providers), so one prefetch routinely surfaces the same fact two or three times. A byte-identical repeat inside one block tells the model nothing the block has not already said, and it is not free: the composed block is stamped into the user row's ``api_content`` sidecar and replayed verbatim on every later request for as long as that row is in context, so each duplicate is paid once per turn, forever. The ``seen`` set is scoped per section — every non-bullet line at column 0 (a heading of any style, a ``---`` rule, prose) starts a new one — so a repeat is only dropped when the SAME section already states it. Only a SELF-CONTAINED bullet is considered — a marker, whitespace, content, and no continuation line indented beneath it. A bullet that carries continuation lines is never dropped and never suppresses a later one, because two entries can share a headline and differ underneath it (``- prefers draft PRs`` / `` (logged 12 Jan, builtin)`` vs the same headline logged elsewhere): dropping one would re-parent its provenance under the other and invent a record neither provider reported. Headings — including ``**bold**`` ones — prose, blank lines, separators and numbered items are left exactly as written. """ lines = text.split("\n") seen: set[str] = set() kept: list[str] = [] for index, line in enumerate(lines): stripped = line.strip() # An indented line is a continuation of the bullet above it (nested child, provenance, # wrapped prose). It never participates in dedupe and is never dropped. if stripped and line[0].isspace(): kept.append(line) continue is_bullet = bool(_RECALL_BULLET_RE.match(stripped)) # Any column-0 non-bullet line (heading, rule, paragraph) opens a fresh dedupe scope. if stripped and not is_bullet: seen.clear() if is_bullet: following = lines[index + 1] if index + 1 < len(lines) else "" carries_continuation = bool(following.strip()) and following[0].isspace() if not carries_continuation: if stripped in seen: continue seen.add(stripped) kept.append(line) return "\n".join(kept) 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 "" sanitized = sanitize_context(raw_context) if sanitized != raw_context: # Stays keyed on sanitization alone: a deduped bullet is routine, not a provider fault. logger.warning("memory provider returned pre-wrapped context; stripped") clean = _drop_repeated_recall_lines(sanitized) 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" "" ) class MemoryManager: """Builtin provider (always first) plus at most one external provider. 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._external_prefetch_spill_config: Optional[Dict[str, Any]] = None self._has_external: bool = False timeout = external_prefetch_timeout timeout = _EXTERNAL_PREFETCH_TIMEOUT_S if timeout is None else float(timeout) if timeout <= 0: raise ValueError("external_prefetch_timeout must be positive") self._external_prefetch_timeout = timeout 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. self._sync_executor: Optional[ThreadPoolExecutor] = None self._sync_executor_lock = threading.Lock() # 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, } 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)`` per provider, logging+swallowing failures; returns successes in order. ``label`` completes the log line ``Memory provider ''