diff --git a/tests/tools/test_memory_tool.py b/tests/tools/test_memory_tool.py index 0582703abb..2af6107bfc 100644 --- a/tests/tools/test_memory_tool.py +++ b/tests/tools/test_memory_tool.py @@ -686,9 +686,7 @@ class TestBomToleranceInMemoryFiles: raw, read_ok = MemoryStore._read_raw_checked(path) assert read_ok is True assert not raw.startswith("\ufeff") - entries, ok = MemoryStore._read_entries_checked(path) - assert ok is True - assert entries == ["First fact."] + assert MemoryStore._read_file(path) == ["First fact."] def test_bom_file_add_keeps_existing_entry_intact(self, store): path = store._path_for("memory") diff --git a/tools/memory_tool.py b/tools/memory_tool.py index b4e0c7bd98..c8d42cb49a 100644 --- a/tools/memory_tool.py +++ b/tools/memory_tool.py @@ -1,9 +1,8 @@ #!/usr/bin/env python3 -"""Memory Tool - persistent curated memory (MEMORY.md = agent notes, USER.md = -user profile). Both enter the system prompt as a FROZEN snapshot at session -start; mid-session writes hit disk immediately but never change the prompt -(prefix cache stays intact). Single `memory` tool: add/replace/remove or a -batch `operations` list. The store lives in ``tools.memory_tool_store``.""" +"""Memory Tool - persistent curated memory (MEMORY.md = agent notes, USER.md = user +profile). Both enter the system prompt as a FROZEN snapshot at session start; +mid-session writes hit disk but never change the prompt (prefix cache intact). +Single `memory` tool: add/replace/remove or a batch `operations` list.""" import copy import json @@ -16,8 +15,8 @@ from typing import Dict, Any, List, Optional, Tuple from utils import is_truthy_value from tools.registry import no_cache_check_fn -# fcntl is Unix-only; on Windows use msvcrt for file locking. MemoryStore reads -# these lazily from this module (tests inspect ``memory_tool.fcntl``). +# fcntl is Unix-only; Windows uses msvcrt. MemoryStore reads both lazily from +# this module (tests patch ``memory_tool.fcntl``). msvcrt = None try: import fcntl @@ -30,63 +29,46 @@ except ImportError: logger = logging.getLogger(__name__) -# One tool-definition pass must use one config decision for both availability -# and the dynamic target schema. ContextVar keeps concurrent profile/session -# builds isolated while letting the check_fn result flow to the immediately -# following dynamic_schema_overrides call in ToolRegistry.get_definitions(). +# One tool-definition pass must use ONE config decision for availability and the +# dynamic target schema: the check_fn result flows to the immediately following +# dynamic_schema_overrides call; ContextVar isolates concurrent profile builds. _memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar( - "memory_surface_flags", default=None -) + "memory_surface_flags", default=None) def get_memory_dir() -> Path: - """Return the profile-scoped memories directory (resolved per call so - HERMES_HOME/profile switches after import are respected).""" + """Profile-scoped memories dir, resolved per call (HERMES_HOME may switch after import).""" return get_hermes_home() / "memories" from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports) - ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, _READ_FAILED, - _drift_error, _read_failed_error, _scan_memory_content, -) + ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, _scan_memory_content) def load_on_disk_store() -> "MemoryStore": - """Fresh on-disk MemoryStore with configured limits/flags, for contexts with - no live agent (gateway, Desktop, bare CLI ``/memory``) so approvals enforce - the SAME caps as ``agent_init``. Defaults if config can't load; never raises.""" - kwargs: Dict[str, Any] = {} + """Fresh on-disk MemoryStore with configured limits/flags for contexts with no + live agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps + as ``agent_init``. Falls back to defaults if config can't load; never raises.""" try: from hermes_cli.config import load_config - config = load_config() or {} mem_cfg = get_builtin_memory_config(config) memory_enabled, user_profile_enabled = get_builtin_memory_store_flags(config) - kwargs = { - "memory_char_limit": int(mem_cfg.get("memory_char_limit", 2200)), - "user_char_limit": int(mem_cfg.get("user_char_limit", 1375)), - "memory_enabled": memory_enabled, - "user_profile_enabled": user_profile_enabled, - } + kwargs = {"memory_char_limit": int(mem_cfg.get("memory_char_limit", 2200)), + "user_char_limit": int(mem_cfg.get("user_char_limit", 1375)), + "memory_enabled": memory_enabled, "user_profile_enabled": user_profile_enabled} except Exception: - kwargs = {} # config optional — fall back to defaults rather than break /memory + kwargs: Dict[str, Any] = {} # config optional — fall back to defaults rather than break /memory store = MemoryStore(**kwargs) store.load_from_disk() return store -# --------------------------------------------------------------------------- -# Write-approval gate -# --------------------------------------------------------------------------- - -def _target_label(target: str) -> str: - return "user profile" if target == "user" else "memory" - +# -- Write-approval gate -- def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]: - """Run the memory write gate. Returns a JSON tool-result string when the - write must NOT proceed (blocked, or staged for approval), None to proceed. - If the gate module can't load, fail open rather than block all writes.""" + """JSON tool-result string when the write must NOT proceed (blocked or staged + for approval), None to proceed. Fails open if the gate module can't load.""" try: from tools import write_approval as wa except Exception: @@ -101,103 +83,74 @@ def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Option ensure_ascii=False) +# action -> (store call, gate (summary, detail) text) for the live tool path and staged replay. +_STORE_ACTIONS = { + "add": (lambda store, target, content, old_text: store.add(target, content), + lambda label, content, old_text: (f"add to {label}", content or "")), + "replace": (lambda store, target, content, old_text: store.replace(target, old_text, content), + lambda label, content, old_text: (f"replace in {label}", f"old: {old_text}\nnew: {content}")), + "remove": (lambda store, target, content, old_text: store.remove(target, old_text), + lambda label, content, old_text: (f"remove from {label}", old_text or ""))} + + def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str]) -> Optional[str]: - """Gate a single mutating op (add/replace/remove); other actions pass.""" - if action not in _STORE_ACTIONS: - return None - label = _target_label(target) - if action == "add": - summary, detail = f"add to {label}", content or "" - elif action == "replace": - summary, detail = f"replace in {label}", f"old: {old_text}\nnew: {content}" - else: - summary, detail = f"remove from {label}", old_text or "" - payload = {"action": action, "target": target, "content": content, "old_text": old_text} - return _gate_or_stage(summary, detail, payload) + """Gate a single mutating op (add/replace/remove).""" + summary, detail = _STORE_ACTIONS[action][1]("user profile" if target == "user" else "memory", content, old_text) + return _gate_or_stage(summary, detail, + {"action": action, "target": target, "content": content, "old_text": old_text}) def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Optional[str]: """Gate a whole batch as a single unit.""" - summary = f"apply {len(operations)} op(s) to {_target_label(target)}" + summary = f"apply {len(operations)} op(s) to {'user profile' if target == 'user' else 'memory'}" detail_lines = [] for op in operations: op = op or {} act = op.get("action", "?") - _op_content = op.get("content") or op.get("new_text") or "" - if act == "remove": - detail_lines.append(f"- remove: {op.get('old_text', '')}") - elif act == "replace": - detail_lines.append(f"- replace: {op.get('old_text', '')} -> {_op_content}") - else: - detail_lines.append(f"- {act}: {_op_content}") - payload = {"action": "batch", "target": target, "operations": operations} - return _gate_or_stage(summary, "\n".join(detail_lines), payload) + content = op.get("content") or op.get("new_text") or "" + detail_lines.append(f"- remove: {op.get('old_text', '')}" if act == "remove" + else f"- replace: {op.get('old_text', '')} -> {content}" if act == "replace" + else f"- {act}: {content}") + return _gate_or_stage(summary, "\n".join(detail_lines), + {"action": "batch", "target": target, "operations": operations}) -# --------------------------------------------------------------------------- -# Tool entry point -# --------------------------------------------------------------------------- - -def _missing_old_text_error(store: "MemoryStore", target: str, action: str) -> str: - """Recoverable error for replace/remove without ``old_text``. It can't be - schema-required (needs a combinator the Codex backend rejects — see - test_memory_tool_schema.py) and some clients omit it, so return the current - inventory plus a retry instruction instead of a dead-end.""" - return json.dumps({ - "success": False, - "error": (f"'{action}' needs old_text -- a short unique substring of the entry " - f"to {action}. None was provided. Reissue the {action} with old_text " - f"set to part of one of the current_entries below."), - "current_entries": store._entries_for(target), - "usage": store._usage(target), - }, ensure_ascii=False) - +# -- Tool entry point -- def _validate_single_op(store, action, target, content, old_text) -> Optional[str]: - """Validate required params BEFORE the gate so an invalid write is rejected - now rather than staged and failing at approve time.""" + """Validate BEFORE the gate so an invalid write is rejected now, not at approve + time. Missing ``old_text`` is recoverable (it can't be schema-required — needs a + combinator the Codex backend rejects — and some clients omit it): return the + current inventory plus a retry instruction instead of a dead-end.""" if action == "add" and not content: return tool_error("Content is required for 'add' action.", success=False) if action in ("replace", "remove") and not old_text: - return _missing_old_text_error(store, target, action) + return json.dumps({ + "success": False, + "error": (f"'{action}' needs old_text -- a short unique substring of the entry " + f"to {action}. None was provided. Reissue the {action} with old_text " + f"set to part of one of the current_entries below."), + "current_entries": store._entries_for(target), "usage": store._usage(target)}, ensure_ascii=False) if action == "replace" and not content: return tool_error("content is required for 'replace' action.", success=False) return None -# action -> store call for both the live tool path and staged-write replay. -_STORE_ACTIONS = { - "add": lambda store, target, content, old_text: store.add(target, content), - "replace": lambda store, target, content, old_text: store.replace(target, old_text, content), - "remove": lambda store, target, content, old_text: store.remove(target, old_text), -} - - -def memory_tool( - action: str = None, - target: str = "memory", - content: str = None, - old_text: str = None, - new_text: str = None, - operations: Optional[List[Dict[str, Any]]] = None, - store: Optional[MemoryStore] = None, -) -> str: - """Tool entry point; returns a JSON string. Single op (action + content / - old_text) or batch (``operations`` applied atomically against the final - budget). ``new_text`` aliases ``content`` — callers mirror ``old_text`` - with it (patch-tool shape), which used to leave ``content`` empty.""" +def memory_tool(action: str = None, target: str = "memory", content: str = None, old_text: str = None, + new_text: str = None, operations: Optional[List[Dict[str, Any]]] = None, + store: Optional[MemoryStore] = None) -> str: + """Tool entry point; returns a JSON string. Single op (action + content/old_text) + or batch (``operations``, atomic against the final budget). ``new_text`` + aliases ``content`` — callers mirror ``old_text`` with it (patch-tool shape).""" if store is None: return tool_error("Memory is not available. It may be disabled in config or this environment.", success=False) - if content is None and new_text is not None: content = new_text # Strict providers send JSON null for optional fields; treat as omitted. - if target is None: - target = "memory" + target = "memory" if target is None else target target_error = _memory_target_error(store, target) if target_error is not None: return json.dumps(target_error) - if operations: if not isinstance(operations, list): return tool_error("operations must be a list of {action, content?, old_text?} objects.", success=False) @@ -205,29 +158,25 @@ def memory_tool( if gate_result is not None: return gate_result return json.dumps(store.apply_batch(target, operations), ensure_ascii=False) - - run = _STORE_ACTIONS.get(action) - if run is None: + if action not in _STORE_ACTIONS: return tool_error(f"Unknown action '{action}'. Use: add, replace, remove", success=False) invalid = _validate_single_op(store, action, target, content, old_text) if invalid is not None: return invalid - # Approval gate: when on, stages the write (background/gateway) or prompts - # inline (interactive CLI); when off (default) passes straight through. + # Approval gate: stages (background/gateway) or prompts inline (CLI); off by default. gate_result = _apply_write_gate(action, target, content, old_text) if gate_result is not None: return gate_result - return json.dumps(run(store, target, content, old_text), ensure_ascii=False) + return json.dumps(_STORE_ACTIONS[action][0](store, target, content, old_text), ensure_ascii=False) def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: - """Normalized ``memory`` config section ({} when missing/malformed → flags - default to enabled). ``agent_init`` consumes the same section so tool - availability and store construction cannot diverge.""" + """Normalized ``memory`` config section ({} when missing/malformed → flags default + to enabled). ``agent_init`` reads the same section so availability and store + construction cannot diverge.""" if config is None: try: from hermes_cli.config import load_config_readonly - config = load_config_readonly() except Exception: logger.debug("Could not read memory config for availability", exc_info=True) @@ -239,10 +188,7 @@ def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[s def get_builtin_memory_store_flags(config: Optional[Dict[str, Any]] = None) -> Tuple[bool, bool]: """Return ``(memory_enabled, user_profile_enabled)`` from resolved config.""" section = get_builtin_memory_config(config) - return ( - is_truthy_value(section.get("memory_enabled"), default=True), - is_truthy_value(section.get("user_profile_enabled"), default=True), - ) + return tuple(is_truthy_value(section.get(k), default=True) for k in ("memory_enabled", "user_profile_enabled")) @no_cache_check_fn @@ -258,7 +204,6 @@ def _memory_target_error(store: "MemoryStore", target: str) -> Optional[Dict[str """Return a shared validation error for an invalid or disabled target.""" if target not in {"memory", "user"}: from tools.registry import _bound_error_text - return {"success": False, "error": _bound_error_text(f"Invalid memory target '{target}'. Use 'memory' or 'user'.")} if store.target_enabled(target): @@ -276,15 +221,12 @@ def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[ return target_error if action == "batch": return store.apply_batch(target, payload.get("operations") or []) - run = _STORE_ACTIONS.get(action) - if run is None: + if action not in _STORE_ACTIONS: return {"success": False, "error": f"Unknown staged action '{action}'."} - return run(store, target, payload.get("content") or "", payload.get("old_text") or "") + return _STORE_ACTIONS[action][0](store, target, payload.get("content") or "", payload.get("old_text") or "") -# ============================================================================= -# OpenAI Function-Calling Schema -# ============================================================================= +# -- OpenAI Function-Calling Schema -- MEMORY_SCHEMA = { "name": "memory", @@ -361,24 +303,17 @@ MEMORY_SCHEMA = { # Schema text when only one built-in store is enabled: (target description, TARGETS replacement). _SINGLE_TARGET_TEXT = { - ("memory",): ( - "The enabled built-in store: 'memory' for personal notes.", - "TARGET: only 'memory' is enabled for personal notes (environment, conventions, " - "tool quirks, lessons).", - ), - ("user",): ( - "The enabled built-in store: 'user' for user profile.", - "TARGET: only 'user' is enabled for user profile facts (name, role, preferences, style).", - ), -} + ("memory",): ("The enabled built-in store: 'memory' for personal notes.", + "TARGET: only 'memory' is enabled for personal notes (environment, conventions, " + "tool quirks, lessons)."), + ("user",): ("The enabled built-in store: 'user' for user profile.", + "TARGET: only 'user' is enabled for user profile facts (name, role, preferences, style).")} def _build_memory_schema_overrides() -> Dict[str, Any]: """Narrow the advertised target surface using the availability snapshot.""" - flags = _memory_surface_flags.get() + flags = _memory_surface_flags.get() or get_builtin_memory_store_flags() _memory_surface_flags.set(None) - if flags is None: - flags = get_builtin_memory_store_flags() targets = [t for t, on in zip(("memory", "user"), flags) if on] parameters = copy.deepcopy(MEMORY_SCHEMA["parameters"]) target_schema = parameters["properties"]["target"] @@ -390,8 +325,7 @@ def _build_memory_schema_overrides() -> Dict[str, Any]: description = description.replace( "TARGETS: 'user' = who the user is (name, role, preferences, style). 'memory' = your " "notes (environment, conventions, tool quirks, lessons).", - replacement, - ) + replacement) return {"description": description, "parameters": parameters} @@ -403,14 +337,8 @@ registry.register( toolset="memory", schema=MEMORY_SCHEMA, handler=lambda args, **kw: memory_tool( - action=args.get("action", ""), - target=args.get("target", "memory"), - content=args.get("content"), - old_text=args.get("old_text"), - new_text=args.get("new_text"), - operations=args.get("operations"), - store=kw.get("store")), + action=args.get("action", ""), target=args.get("target", "memory"), store=kw.get("store"), + **{k: args.get(k) for k in ("content", "old_text", "new_text", "operations")}), check_fn=check_memory_requirements, emoji="🧠", - dynamic_schema_overrides=_build_memory_schema_overrides, -) + dynamic_schema_overrides=_build_memory_schema_overrides) diff --git a/tools/memory_tool_store.py b/tools/memory_tool_store.py index 2c9005cf71..1ee1408358 100644 --- a/tools/memory_tool_store.py +++ b/tools/memory_tool_store.py @@ -1,7 +1,7 @@ """MemoryStore — bounded, file-backed curated memory (MEMORY.md / USER.md). -Entries are joined by ``ENTRY_DELIMITER``; budgets are in chars (model- -independent). Module state that tests monkeypatch (``get_memory_dir``, -``fcntl``/``msvcrt``) stays in ``tools.memory_tool`` and is read lazily.""" +Entries are joined by ``ENTRY_DELIMITER``; budgets are in chars (model-independent). +Module state that tests monkeypatch (``get_memory_dir``, ``fcntl``/``msvcrt``) stays +in ``tools.memory_tool`` and is read lazily.""" import logging import time @@ -14,76 +14,54 @@ from tools.threat_patterns import first_threat_message as _first_threat_message logger = logging.getLogger("tools.memory_tool") -# System-prompt block header prefixes rendered by MemoryStore._render_block. -# agent/conversation_compression.py uses them to detect a leftover block for a -# target whose entries have since been emptied — keep in lockstep. +# Block header prefixes rendered by _render_block; agent/conversation_compression.py +# matches them to detect a leftover block for an emptied target — keep in lockstep. MEMORY_BLOCK_HEADERS = { - "memory": "MEMORY (your personal notes)", - "user": "USER PROFILE (who the user is)", -} + "memory": "MEMORY (your personal notes)", "user": "USER PROFILE (who the user is)"} ENTRY_DELIMITER = "\n§\n" -# Sentinel from ``_reload_target``: the file EXISTS but could not be read. -# Distinct from a drift-backup path (str) and a clean reload (None); the caller -# must abort rather than persist over an unreadable file. -_READ_FAILED = object() - def _memory_dir() -> Path: from tools import memory_tool - return memory_tool.get_memory_dir() def _scan_memory_content(content: str) -> Optional[str]: - """Scan memory content for injection/exfil patterns ("strict" scope: memory - enters the system prompt as a frozen snapshot, so a poisoned entry persists - across sessions). Returns the error string if blocked.""" + """Error string if *content* matches injection/exfil patterns. Strict scope: + memory enters the system prompt, so a poisoned entry persists across sessions.""" return _first_threat_message(content, scope="strict") def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]: - """Error dict for external drift: the on-disk file wouldn't round-trip - through the parser/serializer, so flushing would discard content added by a - patch tool, shell append, manual edit, or sister session.""" - return { - "success": False, - "error": ( - f"Refusing to write {path.name}: file on disk has content that wouldn't round-trip " - f"through the memory tool (likely added by the patch tool, a shell append, a manual edit, " - f"or a concurrent session). A snapshot was saved to {bak_path}. Resolve the drift first — " - f"either rewrite the file as a clean §-delimited list of entries, or move the extra " - f"content out — then retry. This guard exists to prevent silent data loss (issue #26045)." - ), - "drift_backup": bak_path, - "remediation": ( - "Open the .bak file, integrate the missing entries into the memory tool one at a time via " - "memory(action=add, content=...), then remove or rewrite the original file to a clean state." - ), - } + """External drift: the file wouldn't round-trip, so flushing would discard content.""" + return {"success": False, "error": ( + f"Refusing to write {path.name}: file on disk has content that wouldn't round-trip " + f"through the memory tool (likely added by the patch tool, a shell append, a manual edit, " + f"or a concurrent session). A snapshot was saved to {bak_path}. Resolve the drift first — " + f"either rewrite the file as a clean §-delimited list of entries, or move the extra " + f"content out — then retry. This guard exists to prevent silent data loss (issue #26045)." + ), "drift_backup": bak_path, "remediation": ( + "Open the .bak file, integrate the missing entries into the memory tool one at a time via " + "memory(action=add, content=...), then remove or rewrite the original file to a clean state.")} def _read_failed_error(path: "Path") -> Dict[str, Any]: - """Error dict for an unreadable (but existing) memory file. Treating it as - empty and saving would rewrite the whole file from ``[]`` — wiping memory.""" + """Existing-but-unreadable file: saving from an assumed-empty view would wipe it.""" return {"success": False, "error": ( f"Refusing to write {path.name}: the file exists on disk but could not be read right now " f"(temporarily locked by another program, a permission change, invalid/corrupt text encoding, " f"or a filesystem error). Treating an unreadable file as empty and saving would wipe existing " - f"memory, so the write is refused. Nothing was changed — retry in a moment." - )} + f"memory, so the write is refused. Nothing was changed — retry in a moment.")} def _find_unique_match(entries: List[str], old_text: str) -> Tuple[Optional[int], bool]: """``(index, ambiguous)`` for entries containing *old_text*. Exact-duplicate matches are safe (first wins); distinct matches → ``(None, True)``.""" matches = [i for i, e in enumerate(entries) if old_text in e] - if not matches: - return None, False if len({entries[i] for i in matches}) > 1: return None, True - return matches[0], False + return (matches[0] if matches else None), False class MemoryStore: @@ -91,10 +69,9 @@ class MemoryStore: ``_system_prompt_snapshot`` is frozen at load time (prefix-cache stable); ``memory_entries`` / ``user_entries`` are live state persisted to disk.""" - # After this many failed consolidation attempts (overflow / zero-match) in - # ONE turn, return a terminal "save skipped" result instead of "retry in - # this turn", so a fragile replace/add can't loop the turn to budget - # exhaustion and suppress the user's reply. + # Failed consolidation attempts (overflow / zero-match) allowed per turn before + # a TERMINAL "save skipped" result, so a fragile replace/add can't loop the turn + # to budget exhaustion and suppress the user's reply. _MAX_CONSOLIDATION_FAILURES_PER_TURN = 3 def __init__(self, memory_char_limit: int = 2200, user_char_limit: int = 1375, *, @@ -106,9 +83,7 @@ class MemoryStore: self.memory_enabled = memory_enabled self.user_profile_enabled = user_profile_enabled self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""} - # Per-turn counter of failed at-capacity consolidation attempts; reset - # at each turn boundary by reset_consolidation_failures(). - self._consolidation_failures = 0 + self._consolidation_failures = 0 # per turn; reset by reset_consolidation_failures() def target_enabled(self, target: str) -> bool: """Return whether this session's selected built-in store is writable.""" @@ -119,118 +94,91 @@ class MemoryStore: self._consolidation_failures = 0 def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]: - """Count an at-capacity consolidation failure. Under the per-turn cap - return ``response`` unchanged (it says how to retry); past it return a - TERMINAL result so the model stops looping — a failed memory side effect - must never block the turn's reply.""" + """Count a consolidation failure: under the per-turn cap return ``response`` + (it says how to retry); past it a TERMINAL result so the model stops looping.""" self._consolidation_failures += 1 if self._consolidation_failures <= self._MAX_CONSOLIDATION_FAILURES_PER_TURN: return response return {"success": False, "done": True, "error": ( f"Memory consolidation failed {self._consolidation_failures} times this turn. Stop retrying " "memory calls — leave memory unchanged for now and continue with your reply to the user. " - "The fact can be saved in a later turn." - )} + "The fact can be saved in a later turn.")} def load_from_disk(self): """Load MEMORY.md / USER.md and capture the frozen system-prompt snapshot. - Threat hits are replaced in the SNAPSHOT only by a ``[BLOCKED: …]`` - placeholder; live lists keep the raw text so the user can see and remove - poisoned entries (dropping them silently would hide the attack). - Scanning is deterministic from disk bytes, so the snapshot stays stable.""" + Threat hits are replaced by a ``[BLOCKED: …]`` placeholder in the SNAPSHOT only; + live lists keep the raw text so the user can see and remove poisoned entries + (dropping them silently would hide the attack).""" mem_dir = _memory_dir() mem_dir.mkdir(parents=True, exist_ok=True) # Deduplicate (order-preserving, first occurrence wins). self.memory_entries = list(dict.fromkeys(self._read_file(mem_dir / "MEMORY.md"))) self.user_entries = list(dict.fromkeys(self._read_file(mem_dir / "USER.md"))) self._system_prompt_snapshot = { - target: self._render_block(target, self._sanitize_entries_for_snapshot(entries, filename)) - for target, entries, filename in ( - ("memory", self.memory_entries, "MEMORY.md"), - ("user", self.user_entries, "USER.md"), - ) - } + "memory": self._render_block("memory", self._sanitize_entries_for_snapshot(self.memory_entries, "MEMORY.md")), + "user": self._render_block("user", self._sanitize_entries_for_snapshot(self.user_entries, "USER.md"))} @staticmethod def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]: - """Return *entries* with any threat-matching entry replaced by a - ``[BLOCKED: …]`` placeholder (strict scope, same as writes). Empty or - already-blocked entries pass through unchanged.""" + """*entries* with threat matches replaced by a ``[BLOCKED: …]`` placeholder + (strict scope, same as writes); empty / already-blocked entries pass through.""" from tools.threat_patterns import scan_for_threats - sanitized: List[str] = [] - for entry in entries: + def _one(entry): findings = scan_for_threats(entry, scope="strict") if entry and not entry.startswith("[BLOCKED:") else None if not findings: - sanitized.append(entry) - continue + return entry logger.warning("Memory entry from %s blocked at load time: %s", filename, ", ".join(findings)) - sanitized.append(f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. " - f"Removed from system prompt; use memory(action=remove) to delete the original.]") - return sanitized + return (f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. " + f"Removed from system prompt; use memory(action=remove) to delete the original.]") + + return [_one(e) for e in entries] @staticmethod @contextmanager def _file_lock(path: Path): - """Exclusive lock for read-modify-write safety, on a separate .lock file - so the memory file itself can still be atomically replaced.""" + """Exclusive lock on a separate .lock file so the memory file itself can + still be atomically replaced.""" from tools import memory_tool as _mt # fcntl/msvcrt live (and are patched) there - fcntl, msvcrt = _mt.fcntl, _mt.msvcrt lock_path = path.with_suffix(path.suffix + ".lock") lock_path.parent.mkdir(parents=True, exist_ok=True) if fcntl is None and msvcrt is None: yield return - fd = open(lock_path, "a+", encoding="utf-8") - try: - if fcntl: - fcntl.flock(fd, fcntl.LOCK_EX) - else: - fd.seek(0) - msvcrt.locking(fd.fileno(), msvcrt.LK_LOCK, 1) - yield - finally: - try: + with open(lock_path, "a+", encoding="utf-8") as fd: + def _flock(unlock: bool): if fcntl: - fcntl.flock(fd, fcntl.LOCK_UN) + fcntl.flock(fd, fcntl.LOCK_UN if unlock else fcntl.LOCK_EX) else: fd.seek(0) - msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK, 1) - except OSError: - pass - fd.close() + msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK if unlock else msvcrt.LK_LOCK, 1) + _flock(False) + try: + yield + finally: + try: + _flock(True) + except OSError: + pass @staticmethod def _path_for(target: str) -> Path: return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md") - def _reload_target(self, target: str, *, skip_drift: bool = False): - """Re-read entries from disk (under file lock) before mutating. - - Returns ``None`` on a clean reload; the backup path (str) on external - drift (caller must abort — flushing would discard un-roundtrippable - content); or ``_READ_FAILED`` when the file exists but could not be - read (caller MUST abort — rewriting from an assumed-empty view would - wipe it; even append-only ``add`` rewrites the whole file). Drift check - and parse use the SAME raw snapshot — a failed second read used to count - as "no drift". *skip_drift* skips the round-trip check (``add``). - """ - raw, read_ok = self._read_raw_checked(self._path_for(target)) + def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]: + """Re-read entries from disk (under lock) before mutating; return the abort + error dict or None. Aborts on external drift (flushing would discard + un-roundtrippable content) and on an existing-but-unreadable file (even + append-only ``add`` rewrites the whole file). Drift check and parse use the + SAME raw snapshot — a failed second read used to count as "no drift".""" + path = self._path_for(target) + raw, read_ok = self._read_raw_checked(path) if not read_ok: - return _READ_FAILED + return _read_failed_error(path) bak = None if skip_drift else self._detect_external_drift(target, raw) self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw)))) - return bak - - def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]: - """Reload under lock; return the abort error dict, or None to proceed.""" - bak = self._reload_target(target, skip_drift=skip_drift) - if bak is _READ_FAILED: - return _read_failed_error(self._path_for(target)) - if bak: - return _drift_error(self._path_for(target), bak) - return None + return _drift_error(path, bak) if bak else None def save_to_disk(self, target: str): """Persist entries to the appropriate file. Called after every mutation.""" @@ -241,10 +189,7 @@ class MemoryStore: return self.user_entries if target == "user" else self.memory_entries def _set_entries(self, target: str, entries: List[str]): - if target == "user": - self.user_entries = entries - else: - self.memory_entries = entries + setattr(self, "user_entries" if target == "user" else "memory_entries", entries) def _char_count(self, target: str) -> int: return len(ENTRY_DELIMITER.join(self._entries_for(target))) @@ -255,9 +200,14 @@ class MemoryStore: def _usage(self, target: str) -> str: return f"{self._char_count(target):,}/{self._char_limit(target):,}" + def _usage_pct(self, target: str, current: int) -> str: + """``"% — / chars"`` for the given target.""" + limit = self._char_limit(target) + pct = min(100, int((current / limit) * 100)) if limit > 0 else 0 + return f"{pct}% — {current:,}/{limit:,} chars" + def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]: - """Consolidation failure that shows the live entries so the model can - decide what to consolidate.""" + """Consolidation failure carrying the live entries so the model can consolidate.""" return self._consolidation_failure({"success": False, "error": message, "current_entries": self._entries_for(target), "usage": self._usage(target)}) @@ -267,19 +217,28 @@ class MemoryStore: idx, ambiguous = _find_unique_match(entries, old_text) if ambiguous: return None, {"success": False, "error": f"Multiple entries matched '{old_text}'. Be more specific.", - "matches": self._previews([e for e in entries if old_text in e])} + "matches": [e[:80] + ("..." if len(e) > 80 else "") for e in entries if old_text in e]} if idx is None: return None, self._consolidation_failure({ "success": False, "error": f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text of the entry you want to {verb}.", - "current_entries": entries, - }) + "current_entries": entries}) return idx, None - def _commit(self, target: str, entries: List[str], message: str) -> Dict[str, Any]: - self._set_entries(target, entries) - self.save_to_disk(target) - return self._success_response(target, message) + def _mutate(self, target: str, mutate, *, skip_drift: bool = False) -> Dict[str, Any]: + """Lock, reload, run ``mutate(entries, limit)`` -> ``(new_entries, message)`` or an + error dict, then persist and return the success response.""" + with self._file_lock(self._path_for(target)): + err = self._reload_or_error(target, skip_drift=skip_drift) + if err: + return err + result = mutate(self._entries_for(target), self._char_limit(target)) + if isinstance(result, dict): + return result + entries, message = result + self._set_entries(target, entries) + self.save_to_disk(target) + return self._success_response(target, message) def add(self, target: str, content: str) -> Dict[str, Any]: """Append a new entry. Returns error if it would exceed the char limit.""" @@ -289,15 +248,8 @@ class MemoryStore: scan_error = _scan_memory_content(content) if scan_error: return {"success": False, "error": scan_error} - with self._file_lock(self._path_for(target)): - # Append-only: skip the drift guard (appending never clobbers - # un-roundtrippable content), but still refuse on a failed read — - # add rewrites the WHOLE file from the parsed entries. - err = self._reload_or_error(target, skip_drift=True) - if err: - return err - entries = self._entries_for(target) - limit = self._char_limit(target) + + def _add(entries, limit): if content in entries: return self._success_response(target, "Entry already exists (no duplicate added).") if len(ENTRY_DELIMITER.join(entries + [content])) > limit: @@ -305,58 +257,49 @@ class MemoryStore: f"Memory at {self._char_count(target):,}/{limit:,} chars. Adding this entry " f"({len(content)} chars) would exceed the limit. Consolidate now: use 'replace' to merge " f"overlapping entries into shorter ones or 'remove' stale or less important entries (see " - f"current_entries below), then retry this add — all in this turn." - )) - entries.append(content) - return self._commit(target, entries, "Entry added.") + f"current_entries below), then retry this add — all in this turn.")) + return entries + [content], "Entry added." + + # Append-only: skip the drift guard (appending never clobbers foreign + # content) but still refuse a failed read — add rewrites the WHOLE file. + return self._mutate(target, _add, skip_drift=True) def replace(self, target: str, old_text: str, new_content: str) -> Dict[str, Any]: """Find entry containing old_text substring, replace it with new_content.""" - old_text = old_text.strip() new_content = new_content.strip() - if not old_text: + if not old_text.strip(): return {"success": False, "error": "old_text cannot be empty."} if not new_content: return {"success": False, "error": "new_content cannot be empty. Use 'remove' to delete entries."} scan_error = _scan_memory_content(new_content) if scan_error: return {"success": False, "error": scan_error} - with self._file_lock(self._path_for(target)): - err = self._reload_or_error(target) + return self._edit(target, old_text.strip(), new_content) + + def remove(self, target: str, old_text: str) -> Dict[str, Any]: + """Remove the entry containing old_text substring.""" + if not old_text.strip(): + return {"success": False, "error": "old_text cannot be empty."} + return self._edit(target, old_text.strip(), None) + + def _edit(self, target: str, old_text: str, new_content: Optional[str]) -> Dict[str, Any]: + """Locked replace (``new_content`` set) or remove (None) of the unique entry matching *old_text*.""" + def _apply(entries, limit): + idx, err = self._locate(target, old_text, "replace" if new_content else "remove") if err: return err - idx, err = self._locate(target, old_text, "replace") - if err: - return err - entries = self._entries_for(target) - limit = self._char_limit(target) - new_total = len(ENTRY_DELIMITER.join(entries[:idx] + [new_content] + entries[idx + 1:])) + if new_content is None: + return entries[:idx] + entries[idx + 1:], "Entry removed." + replaced = entries[:idx] + [new_content] + entries[idx + 1:] + new_total = len(ENTRY_DELIMITER.join(replaced)) if new_total > limit: return self._failure_with_entries(target, ( f"Replacement would put memory at {new_total:,}/{limit:,} chars. Shorten the new content, " f"or 'remove' other stale or less important entries to make room (see current_entries " - f"below), then retry — all in this turn." - )) - entries[idx] = new_content - return self._commit(target, entries, "Entry replaced.") + f"below), then retry — all in this turn.")) + return replaced, "Entry replaced." - def remove(self, target: str, old_text: str) -> Dict[str, Any]: - """Remove the entry containing old_text substring.""" - old_text = old_text.strip() - if not old_text: - return {"success": False, "error": "old_text cannot be empty."} - with self._file_lock(self._path_for(target)): - err = self._reload_or_error(target) - if err: - return err - idx, err = self._locate(target, old_text, "remove") - if err: - return err - entries = self._entries_for(target) - entries.pop(idx) - return self._commit(target, entries, "Entry removed.") - - # -- Batch -- + return self._mutate(target, _apply) @staticmethod def _apply_batch_op(working: List[str], act: str, content: str, old_text: str, pos: str) -> Optional[str]: @@ -385,90 +328,56 @@ class MemoryStore: return None def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]: - """Apply a sequence of add/replace/remove ops to one target atomically. - - Ops are validated and applied against the FINAL budget only — so a single - call can free space (remove/replace) and add new entries without the - multi-turn consolidate-then-retry dance. All-or-nothing: if any op is - malformed, doesn't match, or the net result exceeds the char limit, - NOTHING is written and the first failure plus live state is returned. - """ + """Apply add/replace/remove ops to one target atomically against the FINAL + budget, so one call can free space and add entries. All-or-nothing: any + malformed / unmatched op or an over-limit result writes NOTHING and returns + the first failure plus live state.""" if not operations: return {"success": False, "error": "operations list is empty."} - # Scan every add/replace content BEFORE touching disk -- one poisoned - # op rejects the whole batch. - for i, op in enumerate(operations): - op = op or {} + ops = [op or {} for op in operations] + # Scan every add/replace content BEFORE touching disk -- one poisoned op rejects the batch. + for i, op in enumerate(ops): scan_error = op.get("action") in {"add", "replace"} and op.get("content") and _scan_memory_content(op["content"]) if scan_error: return {"success": False, "error": f"Operation {i + 1}: {scan_error}"} - with self._file_lock(self._path_for(target)): - err = self._reload_or_error(target) - if err: - return err - # Work on a copy; only commit if the whole batch validates. - working: List[str] = list(self._entries_for(target)) - limit = self._char_limit(target) - for i, op in enumerate(operations): - op = op or {} + def _apply(entries, limit): + working = list(entries) # only committed if the whole batch validates + for i, op in enumerate(ops): act = op.get("action") content = (op.get("content") or op.get("new_text") or "").strip() old_text = (op.get("old_text") or "").strip() pos = f"Operation {i + 1} ({act or 'unknown'})" msg = self._apply_batch_op(working, act, content, old_text, pos) if msg: - return self._batch_error(target, msg) + return self._failure_with_entries( + target, msg + " No operations were applied (batch is all-or-nothing).") # Budget check against the FINAL state only. new_total = len(ENTRY_DELIMITER.join(working)) if new_total > limit: return self._failure_with_entries(target, ( f"After applying all {len(operations)} operations, memory would be at " f"{new_total:,}/{limit:,} chars -- over the limit. Remove or shorten more " - f"entries in the same batch (see current_entries below), then retry." - )) - return self._commit(target, working, f"Applied {len(operations)} operation(s).") + f"entries in the same batch (see current_entries below), then retry.")) + return working, f"Applied {len(operations)} operation(s)." - def _batch_error(self, target: str, message: str) -> Dict[str, Any]: - """Build a batch-abort error that reports live (uncommitted) state.""" - return self._failure_with_entries( - target, message + " No operations were applied (batch is all-or-nothing)." - ) + return self._mutate(target, _apply) def format_for_system_prompt(self, target: str) -> Optional[str]: - """Return the frozen load-time snapshot for system-prompt injection (NOT - live state — mid-session writes don't affect it, preserving the prefix - cache). None if the snapshot is empty.""" + """Frozen load-time snapshot for the system prompt (NOT live state — mid-session + writes don't touch it, preserving the prefix cache); None if empty.""" return self._system_prompt_snapshot.get(target, "") or None - # -- Internal helpers -- - - @staticmethod - def _previews(entries: List[str], width: int = 80) -> List[str]: - """Truncated one-line previews of entries for error feedback.""" - return [e[:width] + ("..." if len(e) > width else "") for e in entries] - - def _usage_pct(self, target: str, current: int) -> str: - """``"% — / chars"`` for the given target.""" - limit = self._char_limit(target) - pct = min(100, int((current / limit) * 100)) if limit > 0 else 0 - return f"{pct}% — {current:,}/{limit:,} chars" - def _success_response(self, target: str, message: str = None) -> Dict[str, Any]: - # A successful write means the consolidation loop made progress, so the - # per-turn failure budget resets (the cap counts consecutive failures). + # A successful write is progress: reset the per-turn (consecutive) failure budget. self._consolidation_failures = 0 - # Intentionally TERMINAL and without the entries list: echoing entries - # invites the model to "find more to fix" and re-issue the same ops. - # Entries are only shown on error/over-budget paths. - resp = {"success": True, "done": True, "target": target, + # TERMINAL and WITHOUT the entries list: echoing entries invites the model to + # "find more to fix" and re-issue the same ops. Entries only appear on errors. + return {"success": True, "done": True, "target": target, "usage": self._usage_pct(target, self._char_count(target)), - "entry_count": len(self._entries_for(target))} - if message: - resp["message"] = message - resp["note"] = "Write saved. This update is complete — do not repeat it." - return resp + "entry_count": len(self._entries_for(target)), **({"message": message} if message else {}), + "note": "Write saved. This update is complete — do not repeat it."} def _render_block(self, target: str, entries: List[str]) -> str: """Render a system prompt block with header and usage indicator.""" @@ -481,12 +390,10 @@ class MemoryStore: @staticmethod def _read_raw_checked(path: Path) -> Tuple[str, bool]: - """Read raw text as ``(raw, read_ok)``. ``read_ok`` is False ONLY when the - file EXISTS but can't be read (absent file → ``("", True)``). Invalid - UTF-8 counts as unreadable; decoding stays STRICT because - ``errors="replace"`` would hand callers a lossy view that a save then - persists over the real bytes. ``utf-8-sig`` strips a Notepad BOM that - otherwise glues U+FEFF onto the first entry forever.""" + """``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't + be read (absent → ``("", True)``). Decoding stays STRICT: ``errors="replace"`` + would hand callers a lossy view that a save then persists. ``utf-8-sig`` strips + a Notepad BOM that otherwise glues U+FEFF onto the first entry forever.""" if not path.exists(): return "", True try: @@ -496,30 +403,20 @@ class MemoryStore: @staticmethod def _parse_entries(raw: str) -> List[str]: - """Split raw memory-file text into stripped, non-empty entries. Splits on - the full ENTRY_DELIMITER so a bare "§" inside an entry is preserved.""" + """Stripped, non-empty entries; splits on the FULL delimiter so a bare "§" survives.""" return [e for e in (x.strip() for x in raw.split(ENTRY_DELIMITER)) if e] - @staticmethod - def _read_entries_checked(path: Path) -> Tuple[List[str], bool]: - """Read + parse as ``(entries, read_ok)`` — see ``_read_raw_checked``.""" - raw, read_ok = MemoryStore._read_raw_checked(path) - return MemoryStore._parse_entries(raw), read_ok - @staticmethod def _read_file(path: Path) -> List[str]: - """Read a memory file into entries (empty list on any error). Only for - read-only callers (``load_from_disk``, learning_mutations); mutation - paths must use ``_read_raw_checked`` so they can refuse to overwrite an - unreadable file.""" - return MemoryStore._read_entries_checked(path)[0] + """Entries of a memory file ([] on any error). Read-only callers only + (``load_from_disk``, learning_mutations); mutation paths must use + ``_read_raw_checked`` so they can refuse to overwrite an unreadable file.""" + return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0]) def _detect_external_drift(self, target: str, raw: str) -> Optional[str]: - """Backup-path string if *raw* (the caller's checked-read snapshot) shows - external drift, else None. Signals: round-trip mismatch, or one parsed - entry exceeding the whole-file char limit (no tool-written entry can — - an external writer appended free-form content a flush would truncate). - The file is snapshotted to ``.bak.`` so the operator can recover it.""" + """Backup path if *raw* shows external drift, else None. Signals: round-trip + mismatch, or one entry over the whole-file limit (no tool-written entry can — + an external writer appended free-form text). Snapshots to ``.bak.``.""" if not raw.strip(): return None parsed = self._parse_entries(raw) @@ -535,8 +432,7 @@ class MemoryStore: @staticmethod def _write_file(path: Path, entries: List[str]): - """Atomic temp-file + rename: readers see the old or the new complete - file, never a truncated one.""" + """Atomic temp-file + rename: readers never see a truncated file.""" try: atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_") except OSError as e: diff --git a/tools/microsoft_graph_auth.py b/tools/microsoft_graph_auth.py index f58eac5f16..5c8edc1160 100644 --- a/tools/microsoft_graph_auth.py +++ b/tools/microsoft_graph_auth.py @@ -31,19 +31,14 @@ class MicrosoftGraphTokenError(MicrosoftGraphAuthError): def format_graph_error(error: Any) -> str | None: - """Render Graph's ``{"error": {"code", "message"}}`` (or bare-string ``error``) body. - - Shared by the token endpoint and the REST client so both surface - ``code: message`` the same way. ``None`` means the shape was unusable. - """ + """Render Graph's ``{"error": {"code", "message"}}`` (or bare-string ``error``) as + ``code: message``; shared by token endpoint and REST client. None if unusable.""" if isinstance(error, str): return error if not isinstance(error, dict): return None code, message = error.get("code"), error.get("message") - if code and message: - return f"{code}: {message}" - return str(message) if message else None + return f"{code}: {message}" if code and message else (str(message) if message else None) @dataclass(frozen=True) @@ -68,15 +63,12 @@ class GraphCredentials: env = environ if environ is not None else os.environ values = [(env.get(name) or "").strip() for name in _REQUIRED_ENV] missing = [name for name, value in zip(_REQUIRED_ENV, values) if not value] + if missing and not required: + return None if missing: - if not required: - return None raise MicrosoftGraphConfigError(f"Missing Microsoft Graph configuration: {', '.join(missing)}") - return cls( - *values, - scope=(env.get("MSGRAPH_SCOPE") or DEFAULT_GRAPH_SCOPE).strip(), - authority_url=(env.get("MSGRAPH_AUTHORITY_URL") or DEFAULT_GRAPH_AUTHORITY_URL).strip(), - ) + return cls(*values, scope=(env.get("MSGRAPH_SCOPE") or DEFAULT_GRAPH_SCOPE).strip(), + authority_url=(env.get("MSGRAPH_AUTHORITY_URL") or DEFAULT_GRAPH_AUTHORITY_URL).strip()) @dataclass @@ -102,9 +94,7 @@ class MicrosoftGraphTokenProvider: self, credentials: GraphCredentials, *, timeout: float = 20.0, skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, transport: httpx.AsyncBaseTransport | None = None, ) -> None: - self.credentials = credentials - self.timeout = timeout - self.skew_seconds = max(0, int(skew_seconds)) + self.credentials, self.timeout, self.skew_seconds = credentials, timeout, max(0, int(skew_seconds)) self._transport = transport self._cached_token: CachedAccessToken | None = None self._lock = asyncio.Lock() @@ -117,26 +107,17 @@ class MicrosoftGraphTokenProvider: self._cached_token = None def inspect_token_health(self) -> dict[str, Any]: - cached = self._cached_token - return { - "configured": True, - "tenant_id": self.credentials.tenant_id, - "client_id": self.credentials.client_id, - "scope": self.credentials.scope, - "authority_url": self.credentials.authority_url, - "token_url": self.credentials.token_url, - "cached": bool(cached), - "expires_in_seconds": cached.expires_in_seconds if cached else None, - "is_expired": cached.is_expired(skew_seconds=0) if cached else None, - "refresh_skew_seconds": self.skew_seconds, - } + cached, creds = self._cached_token, self.credentials + return {"configured": True, "tenant_id": creds.tenant_id, "client_id": creds.client_id, + "scope": creds.scope, "authority_url": creds.authority_url, "token_url": creds.token_url, + "cached": bool(cached), "expires_in_seconds": cached.expires_in_seconds if cached else None, + "is_expired": cached.is_expired(skew_seconds=0) if cached else None, + "refresh_skew_seconds": self.skew_seconds} def _fresh_cached(self) -> CachedAccessToken | None: """The cached token unless it expires within ``skew_seconds``.""" cached = self._cached_token - if cached and not cached.is_expired(skew_seconds=self.skew_seconds): - return cached - return None + return cached if cached and not cached.is_expired(skew_seconds=self.skew_seconds) else None async def get_access_token(self, *, force_refresh: bool = False) -> str: # Double-checked under the lock so concurrent callers share one fetch. @@ -145,33 +126,22 @@ class MicrosoftGraphTokenProvider: async with self._lock: if not force_refresh and (cached := self._fresh_cached()): return cached.access_token - token = await self._fetch_access_token() - self._cached_token = token - return token.access_token + self._cached_token = await self._fetch_access_token() + return self._cached_token.access_token async def _fetch_access_token(self) -> CachedAccessToken: - data = { - "grant_type": "client_credentials", - "client_id": self.credentials.client_id, - "client_secret": self.credentials.client_secret, - "scope": self.credentials.scope, - } + data = {"grant_type": "client_credentials", "client_id": self.credentials.client_id, + "client_secret": self.credentials.client_secret, "scope": self.credentials.scope} async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client: - response = await client.post( - self.credentials.token_url, data=data, - headers={"Content-Type": "application/x-www-form-urlencoded"}, - ) - + response = await client.post(self.credentials.token_url, data=data, + headers={"Content-Type": "application/x-www-form-urlencoded"}) if response.status_code >= 400: - raise MicrosoftGraphTokenError( - "Microsoft Graph token request failed with HTTP " - f"{response.status_code}: {_extract_error_detail(response)}" - ) + raise MicrosoftGraphTokenError("Microsoft Graph token request failed with HTTP " + f"{response.status_code}: {_extract_error_detail(response)}") try: payload = response.json() except ValueError as exc: raise MicrosoftGraphTokenError("Microsoft Graph token response was not valid JSON.") from exc - access_token = str(payload.get("access_token") or "").strip() if not access_token: raise MicrosoftGraphTokenError("Microsoft Graph token response did not include access_token.") @@ -179,34 +149,26 @@ class MicrosoftGraphTokenProvider: expires_in_seconds = int(payload.get("expires_in")) except (TypeError, ValueError) as exc: raise MicrosoftGraphTokenError( - "Microsoft Graph token response did not include a valid expires_in." - ) from exc - + "Microsoft Graph token response did not include a valid expires_in.") from exc return CachedAccessToken( access_token=access_token, token_type=str(payload.get("token_type") or "Bearer").strip() or "Bearer", - expires_at=time.time() + max(0, expires_in_seconds), - ) + expires_at=time.time() + max(0, expires_in_seconds)) def _extract_error_detail(response: httpx.Response) -> str: - """Best human-readable detail from a token-endpoint error body. - - The OAuth endpoint prefers ``error_description``; fall back to the - Graph-style ``error`` object/string, then a bare ``code``, then raw text. - """ + """Best human-readable detail from a token-endpoint error body: ``error_description``, + then the Graph-style ``error`` object/string, then a bare ``code``, then raw text.""" try: payload = response.json() except ValueError: return response.text.strip() or "unknown error" - - if isinstance(payload, dict): - if isinstance(payload.get("error_description"), str): - return payload["error_description"] - error = payload.get("error") - detail = format_graph_error(error) - if detail is not None: - return detail - if isinstance(error, dict) and error.get("code"): - return str(error["code"]) - return str(payload) + if not isinstance(payload, dict): + return str(payload) + if isinstance(payload.get("error_description"), str): + return payload["error_description"] + error = payload.get("error") + detail = format_graph_error(error) + if detail is not None: + return detail + return str(error["code"]) if isinstance(error, dict) and error.get("code") else str(payload) diff --git a/tools/microsoft_graph_client.py b/tools/microsoft_graph_client.py index ab8842dcf4..9e3dd746c2 100644 --- a/tools/microsoft_graph_client.py +++ b/tools/microsoft_graph_client.py @@ -5,16 +5,12 @@ from __future__ import annotations import asyncio import os from pathlib import Path -from typing import Any, AsyncIterator, Awaitable, Callable +from typing import Any, Awaitable, Callable import httpx from agent.retry_utils import parse_retry_after_seconds -from tools.microsoft_graph_auth import ( - GraphCredentials, - MicrosoftGraphTokenProvider, - format_graph_error, -) +from tools.microsoft_graph_auth import GraphCredentials, MicrosoftGraphTokenProvider, format_graph_error DEFAULT_GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0" @@ -32,38 +28,26 @@ class MicrosoftGraphAPIError(MicrosoftGraphClientError): def __init__( self, status_code: int, method: str, url: str, message: str, *, - retry_after_seconds: float | None = None, payload: Any = None, - ) -> None: - self.status_code = status_code - self.method = method - self.url = url - self.retry_after_seconds = retry_after_seconds - self.payload = payload + retry_after_seconds: float | None = None, payload: Any = None) -> None: + self.status_code, self.method, self.url = status_code, method, url + self.retry_after_seconds, self.payload = retry_after_seconds, payload super().__init__(f"Microsoft Graph API error {status_code} for {method} {url}: {message}") class MicrosoftGraphClient: - """Minimal async Microsoft Graph client with retries and pagination. - - Retry policy (shared by JSON requests and streaming downloads): transport - errors back off exponentially; 401 clears the token cache and refetches; - 429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``. - """ + """Minimal async Graph client. Retry policy (JSON requests and streaming downloads + alike): transport errors back off exponentially; 401 clears the token cache and + refetches; 429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``.""" def __init__( self, token_provider: MicrosoftGraphTokenProvider, *, base_url: str = DEFAULT_GRAPH_BASE_URL, timeout: float = 60.0, max_retries: int = 3, transport: httpx.AsyncBaseTransport | None = None, sleep: Callable[[float], Awaitable[None]] | None = None, - user_agent: str = "Hermes-Agent/graph-client", - ) -> None: - self.token_provider = token_provider - self.base_url = base_url.rstrip("/") - self.timeout = timeout - self.max_retries = max(0, int(max_retries)) - self._transport = transport - self._sleep = sleep or asyncio.sleep - self.user_agent = user_agent + user_agent: str = "Hermes-Agent/graph-client") -> None: + self.token_provider, self.base_url, self.timeout = token_provider, base_url.rstrip("/"), timeout + self.max_retries, self.user_agent = max(0, int(max_retries)), user_agent + self._transport, self._sleep = transport, sleep or asyncio.sleep @classmethod def from_env(cls, **kwargs: Any) -> "MicrosoftGraphClient": @@ -77,19 +61,20 @@ class MicrosoftGraphClient: async def patch_json(self, path: str, *, json_body: Any | None = None, headers: Headers = None) -> Any: response = await self._request("PATCH", path, json_body=json_body, headers=headers) - if response.status_code == 204 or not response.content: - return {} - return self._decode_json(response) + return self._decode_json_or(response, {}) async def delete(self, path: str, *, headers: Headers = None) -> dict[str, Any]: response = await self._request("DELETE", path, headers=headers) - if response.status_code == 204 or not response.content: - return {"deleted": True, "status_code": response.status_code} - return self._decode_json(response) + return self._decode_json_or(response, {"deleted": True, "status_code": response.status_code}) - async def iterate_pages( - self, path: str, *, params: Params = None, headers: Headers = None - ) -> AsyncIterator[dict[str, Any]]: + def _decode_json_or(self, response: httpx.Response, empty: Any) -> Any: + """*empty* for a 204 / bodiless response, else the decoded JSON body.""" + return empty if response.status_code == 204 or not response.content else self._decode_json(response) + + async def collect_paginated( + self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]: + """Follow ``@odata.nextLink`` and concatenate every page's ``value`` list.""" + items: list[Any] = [] # Query params go on the first request only; @odata.nextLink already embeds them. next_url: str | None = self._resolve_url(path) next_params = dict(params or {}) @@ -98,28 +83,17 @@ class MicrosoftGraphClient: payload = self._decode_json(response) if not isinstance(payload, dict): raise MicrosoftGraphClientError( - f"Expected paginated Graph response dict, got {type(payload).__name__}." - ) - yield payload - next_url = payload.get("@odata.nextLink") - next_params = {} - - async def collect_paginated( - self, path: str, *, params: Params = None, headers: Headers = None - ) -> list[Any]: - items: list[Any] = [] - async for page in self.iterate_pages(path, params=params, headers=headers): - value = page.get("value") - if isinstance(value, list): - items.extend(value) + f"Expected paginated Graph response dict, got {type(payload).__name__}.") + if isinstance(payload.get("value"), list): + items.extend(payload["value"]) + next_url, next_params = payload.get("@odata.nextLink"), {} return items async def download_to_file( self, path: str, destination: str | Path, *, headers: Headers = None, chunk_size: int = 65536 ) -> dict[str, Any]: - """Download a Graph resource to disk, streaming the body chunk-by-chunk - (recordings and other large artifacts never need to fit in memory). - Written to a ``.part`` file and renamed into place only on success.""" + """Stream a Graph resource to disk chunk-by-chunk (large recordings never + fit in memory); written to ``.part`` and renamed into place only on success.""" url = self._resolve_url(path) target = Path(destination) target.parent.mkdir(parents=True, exist_ok=True) @@ -147,8 +121,7 @@ class MicrosoftGraphClient: async def _request( self, method: str, path_or_url: str, *, - params: Params = None, json_body: Any | None = None, headers: Headers = None, - ) -> httpx.Response: + params: Params = None, json_body: Any | None = None, headers: Headers = None) -> httpx.Response: url = self._resolve_url(path_or_url) async def perform(client: httpx.AsyncClient, request_headers: dict[str, str]): @@ -160,14 +133,10 @@ class MicrosoftGraphClient: async def _with_retries( self, method: str, url: str, accept: str, json_body: Any | None, headers: Headers, perform: Callable[[httpx.AsyncClient, dict[str, str]], Awaitable[tuple[httpx.Response, Any]]], - kind: str, - ) -> Any: - """Run ``perform`` (returning ``(response, result)``) under the retry policy. - - ``kind`` ("request"/"download") only labels the transport-failure messages. - A ``MicrosoftGraphAPIError`` for the failing status is raised once retries - are exhausted or the status is not retryable; only a 401 forces a token refresh. - """ + kind: str) -> Any: + """Run ``perform`` (-> ``(response, result)``) under the retry policy. ``kind`` + only labels transport-failure messages. Raises ``MicrosoftGraphAPIError`` once + retries are exhausted or the status is not retryable; only 401 forces a token refresh.""" attempt = 0 last_error: Exception | None = None @@ -175,8 +144,7 @@ class MicrosoftGraphClient: token = await self.token_provider.get_access_token( force_refresh=attempt > 0 and isinstance(last_error, MicrosoftGraphAPIError) - and last_error.status_code == 401 - ) + and last_error.status_code == 401) request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent} if json_body is not None: request_headers["Content-Type"] = "application/json" @@ -187,27 +155,21 @@ class MicrosoftGraphClient: async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client: response, result = await perform(client, request_headers) except httpx.HTTPError as exc: - last_error = exc + last_error, response = exc, None if attempt >= self.max_retries: raise MicrosoftGraphClientError( - f"Microsoft Graph {kind} failed for {method} {url}: {exc}" - ) from exc - await self._sleep(self._retry_delay(None, attempt)) - attempt += 1 - continue - - if response.status_code < 400: - return result - - api_error = last_error = self._build_api_error(method, url, response) - status = response.status_code - if attempt < self.max_retries and (status in (401, 429) or 500 <= status < 600): + f"Microsoft Graph {kind} failed for {method} {url}: {exc}") from exc + else: + if response.status_code < 400: + return result + last_error = self._build_api_error(method, url, response) + status = response.status_code + if attempt >= self.max_retries or not (status in (401, 429) or 500 <= status < 600): + raise last_error if status == 401: self.token_provider.clear_cache() - await self._sleep(self._retry_delay(response, attempt)) - attempt += 1 - continue - raise api_error + await self._sleep(self._retry_delay(response, attempt)) + attempt += 1 raise MicrosoftGraphClientError(f"Microsoft Graph {kind} exhausted retries for {method} {url}.") @@ -224,29 +186,21 @@ class MicrosoftGraphClient: except ValueError as exc: raise MicrosoftGraphClientError( "Microsoft Graph response was not valid JSON for " - f"{response.request.method} {response.request.url}" - ) from exc + f"{response.request.method} {response.request.url}") from exc @staticmethod def _retry_delay(response: httpx.Response | None, attempt: int) -> float: - if response is not None: - retry_after = parse_retry_after_seconds(response.headers) - if retry_after is not None: - return retry_after - return min(8.0, 0.5 * (2 ** attempt)) + retry_after = parse_retry_after_seconds(response.headers) if response is not None else None + return min(8.0, 0.5 * (2 ** attempt)) if retry_after is None else retry_after @staticmethod def _build_api_error(method: str, url: str, response: httpx.Response) -> MicrosoftGraphAPIError: - message = response.text.strip() or "unknown error" try: payload: Any = response.json() except ValueError: payload = None - if isinstance(payload, dict): - detail = format_graph_error(payload.get("error")) - if detail is not None: - message = detail + detail = format_graph_error(payload.get("error")) if isinstance(payload, dict) else None return MicrosoftGraphAPIError( - response.status_code, method, url, message, - retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload, - ) + response.status_code, method, url, + detail if detail is not None else (response.text.strip() or "unknown error"), + retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload) diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index 5a8ff6fd4d..5baeeda1ed 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -5,35 +5,361 @@ Single-shape tool; the mode is inferred from the args: DISCOVERY (``query``; FTS5 deduped by lineage, adaptive detail hydrates only the top result), SCROLL (``session_id`` + ``around_message_id``; ±window around the anchor), READ (``session_id`` alone; whole session or head/tail), BROWSE (no args). -No LLM calls — every shape returns actual DB messages. Helpers live in -``session_search_tool_common`` / ``_discover`` and are re-exported here. +No LLM calls — every shape returns actual DB messages. """ import json import logging -from typing import Any, List, Optional +from datetime import datetime +from typing import Any, Dict, List, Optional, Union -from tools.session_search_tool_common import ( # noqa: F401 (re-exports) - _COMPACTION_PREFIXES, _DEMOTED_SESSION_SOURCES, _DISCOVER_SCAN_LIMIT, - _DISCOVER_SEARCH_FIELDS, _FRESH_RESET_END_REASONS, _HIDDEN_SESSION_SOURCES, - _annotate_rebuild_status, _format_timestamp, _get_message_storage_state, - _is_compacted_message, _is_compacted_state, _is_compaction_summary, - _ok, _order_for_recall, _quiet, _resolve_lineage, _resolve_to_parent, _session_end_reason, - _session_left_live_context, _session_link, _session_meta_block, _shape_message, -) -from tools.session_search_tool_discover import ( # noqa: F401 (re-exports) - _discover, _normalize_title_query, _title_match_result, -) +from hermes_state_common import _RESET_END_REASONS + +# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool), +# delegate subagent runs, kanban workers — not the user's history. +_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool") + +# Searchable but DEMOTED below interactive sessions: cron vocabulary dominates bare +# BM25 and starves out the user's own sessions ("recall blindness"). +_DEMOTED_SESSION_SOURCES = ("cron",) + +# FTS rows scanned before dedup-by-lineage — well above the distinct sessions a query +# returns, so interactive matches buried under cron hits survive the demotion pass. +_DISCOVER_SCAN_LIMIT = 300 + +# Raw FTS rows are only a plan input; the response hydrates its own window/bookends. +_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started") + +# Compaction handoff summaries (agent/context_compressor.py); excluded from bookends. +_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:") + +# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the predecessor +# WITHOUT carrying its transcript forward — unlike compression continuations and +# live delegation children. Derived from the gateway set so the two cannot drift. +_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"} + + +def _quiet(fn, default, msg, *log_args, with_exc: bool = False): + """``fn()``, or *default* after debug-logging *msg* (exception appended as a + final ``%s`` arg when *with_exc*) on any exception.""" + try: + return fn() + except Exception as e: + logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True) + return default + + +def _loud(fn, log_msg, error_prefix, *log_args): + """``(value, None)`` from ``fn()``, or ``(None, tool_error_json)`` after an + error-level log — for DB calls whose failure the model must see.""" + try: + return fn(), None + except Exception as e: + logging.error(log_msg, *log_args, e, exc_info=True) + return None, tool_error(f"{error_prefix}: {e}", success=False) + + +def _format_timestamp(ts: Union[int, float, str, None]) -> str: + """Unix timestamp (number / numeric string) -> readable date; ISO strings pass + through; "unknown" for None; str(ts) if conversion fails.""" + if ts is None: + return "unknown" + if isinstance(ts, str) and not ts.replace(".", "").replace("-", "").isdigit(): + return ts + try: + return datetime.fromtimestamp(float(ts)).strftime("%B %d, %Y at %I:%M %p") + except Exception as e: + logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True) + return str(ts) + + +def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]: + """The ``session_meta`` sub-object shared by read/scroll responses.""" + return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"), + "model": meta.get("model"), "title": meta.get("title")} + + +def _ok(**payload) -> str: + """Serialize a successful tool result (``success`` first, then *payload* in order).""" + return json.dumps({"success": True, **payload}, ensure_ascii=False) + + +def _is_compaction_summary(content: str) -> bool: + """Return True if *content* looks like a generated compaction handoff.""" + return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES) + + +def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]: + """Walk parent_session_id to the root -> ``(root_id, has_compression_hop)``. The + flag separates a compression-split lineage (parent summarised away) from a + delegation lineage (child still visible to the parent). Errors -> ``(session_id, False)``.""" + visited: set[str] = set() + cur, has_compression = session_id, False + while cur and cur not in visited: + visited.add(cur) + s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True) + if not s: + break + has_compression = has_compression or s.get("end_reason") == "compression" + if not s.get("parent_session_id"): + break + cur = s["parent_session_id"] + return cur, has_compression + + +def _resolve_lineage(db, session_id: str) -> str: + """Return only the lineage root (ignores the compression hop).""" + return _resolve_to_parent(db, session_id)[0] + + +def _session_left_live_context(db, session_id: str) -> bool: + """True when the transcript left everyone's live context: ``compression`` + (summarised into the child) or a fresh reset (child starts empty). Live delegation + children (``end_reason is None``) and ``branched`` parents (copied verbatim into + the branch) ARE the current context, so they stay excluded from recall.""" + s = session_id and _quiet(lambda: db.get_session(session_id), None, "get_session failed for %s", session_id) + end_reason = (s.get("end_reason") or None) if s else None + return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS + + +def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]: + """Return the owning session and visibility flags for *message_id*.""" + if not message_id: + return None + + def _lookup(): + with db._lock: + return db._conn.execute( + "SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,)).fetchone() + row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id) + return dict(row) if row is not None else None + + +def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool: + """Compaction archives are ``active=0, compacted=1``; rewind/undo rows are + ``active=0, compacted=0`` and must stay hidden.""" + return state is not None and state["active"] == 0 and state["compacted"] == 1 + + +def _is_compacted_message(db, message_id) -> bool: + """True for a compaction-archived row — content no longer in live context, so + it stays discoverable even on the current session. False on any error.""" + return _is_compacted_state(_get_message_storage_state(db, message_id)) + + +def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None: + """Note rebuild progress while the deferred FTS backfill runs so the agent can + explain thin results instead of treating them as ground truth. Never raises.""" + try: + status = db.fts_rebuild_status() + except Exception: + return + if status is None: + return + payload["index_rebuild"] = {"percent": status["percent"], "note": ( + f"The search index is rebuilding in the background ({status['percent']}% done, " + f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages " + f"may be incomplete until it finishes.")} + + +def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """Stable-sort so interactive sessions rank above automation; BM25 order is + kept within each class, so a cron hit never displaces an interactive one.""" + return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0) + + +def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None, + max_content_len: Optional[int] = None) -> Dict[str, Any]: + """Slim a message row. Keeps ``content`` even when empty (tool-call-only + assistant turns); with *max_content_len* truncates and flags it.""" + content = m.get("content") + if isinstance(content, str) and "\x1b" in content: # archived terminal output carries ANSI + from tools.ansi_strip import strip_ansi + content = strip_ansi(content) + entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")} + entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)}) + if anchor_id is not None and m.get("id") == anchor_id: + entry["anchor"] = True + if max_content_len and content and len(content) > max_content_len: + entry.update(content=content[:max_content_len] + "…", content_truncated=True, original_content_chars=len(content)) + return {k: v for k, v in entry.items() if v is not None or k == "content"} + + +def _session_link(session_id: str, profile: str = None) -> str: + """The reference the agent writes for a session — same value the desktop composer + emits, so it renders as a titled link. The profile segment is omitted when it + can't be named confidently (a bare id still resolves, just not across profiles).""" + name = (profile or "").strip() + if not name: + def _active(): + from hermes_cli.profiles import get_active_profile_name + resolved = get_active_profile_name() + return "" if resolved == "custom" else resolved + name = _quiet(_active, "", "get_active_profile_name failed for session link") + return f"@session:{name}/{session_id}" if name else f"@session:{session_id}" + + +def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]: + """Return a discovery-shaped result when the query matches a session title.""" + title_query = query.strip().strip("`'\"") # models often quote a remembered title + if not title_query: + return None + session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None, + "resolve_session_by_title failed for %r", title_query) + if not session_id: + return None + lineage_root = _resolve_lineage(db, session_id) + # Same-lineage title hits are in-context only while the session is live; + # /new-reset and compression-ended parents are not. + if (current_lineage_root and lineage_root == current_lineage_root + and not _session_left_live_context(db, session_id)): + return None + session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None, + "get_session failed for title match %s", session_id) or {} + if session_meta.get("source") in _HIDDEN_SESSION_SOURCES: + return None + messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id) + anchor_id = messages[0].get("id") if messages else None + view = {} if anchor_id is None else _quiet( + lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {}, + "get_anchored_view failed for title match %s/%s", session_id, anchor_id) + title = session_meta.get("title") or title_query + entry = _discovery_entry( + lineage_root, session_id=session_id, when=_format_timestamp(session_meta.get("started_at")), + source=session_meta.get("source", "unknown"), model=session_meta.get("model") or "unknown", + title=title, matched_role="session_title", match_message_id=anchor_id, + snippet=f"Session title matched: {title}", + bookend_start=[_shape_message(m) for m in (view.get("bookend_start") or messages[:3])], + messages=[_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])], + bookend_end=[_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])], + messages_before=view.get("messages_before", 0), + messages_after=view.get("messages_after", max(len(messages) - 5, 0)), detail="full") + entry["_lineage_root"] = lineage_root + return entry + + +def _discovery_entry(lineage_root: Optional[str], **fields) -> Dict[str, Any]: + """One discovery result in canonical key order; ``parent_session_id`` is set + when the hit lives in a child of its lineage root.""" + entry = {k: fields[k] for k in ( + "session_id", "when", "source", "model", "title", "matched_role", "match_message_id", "snippet", + "bookend_start", "messages", "bookend_end", "messages_before", "messages_after", "detail")} + if lineage_root and lineage_root != entry["session_id"]: + entry["parent_session_id"] = lineage_root + return entry + + +def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str: + payload = {"success": True, "mode": "discover", "query": query, "detail": detail, + "results": results, "count": len(results), **extra} + _annotate_rebuild_status(db, payload) + return json.dumps(payload, ensure_ascii=False) + + +def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None: + """Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*. + The raw owning session_id stays on the row — only it pairs validly with the FTS + match id. Current-lineage hits are skipped UNLESS the transcript left live + context (compression-ended, /new-reset predecessor, or an in-place compacted row + on the SAME session); a live delegation child (end_reason=None) stays excluded.""" + for r in raw_results: + if len(seen_sessions) >= limit: + break + raw_sid = r["session_id"] + resolved_sid = _resolve_lineage(db, raw_sid) + is_compacted_hit = _is_compacted_message(db, r.get("id")) + in_current_lineage = bool(current_lineage_root) and resolved_sid == current_lineage_root + if in_current_lineage and not (_session_left_live_context(db, raw_sid) or is_compacted_hit): + continue + if current_session_id and raw_sid == current_session_id and not is_compacted_hit: + continue + seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid}) + + +def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]: + return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or []) + if not _is_compaction_summary(m.get("content", ""))] + + +def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]: + """One discovery result from a surviving FTS row; None (hit dropped) if the + anchored view can't be loaded.""" + hit_sid = match_info.get("session_id") or lineage_root + msg_id = match_info.get("id") + try: + view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3) + except Exception as e: + logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True) + return None + session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {} + full = result_detail == "full" + window_messages = [m for m in (view.get("window") or []) if full or m.get("id") == msg_id] + return _discovery_entry( + lineage_root, session_id=hit_sid, + when=_format_timestamp(session_meta.get("started_at") or match_info.get("session_started")), + source=session_meta.get("source") or match_info.get("source", "unknown"), + model=session_meta.get("model") or match_info.get("model") or "unknown", + title=session_meta.get("title") or None, matched_role=match_info.get("role"), + match_message_id=msg_id, snippet=match_info.get("snippet") or "", + bookend_start=_bookend(view, "bookend_start") if full else [], + messages=[_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages], + bookend_end=_bookend(view, "bookend_end") if full else [], + messages_before=view.get("messages_before", 0), messages_after=view.get("messages_after", 0), + detail=result_detail) + + +def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str], + detail: str, current_session_id: str = None, link_profile: str = None) -> str: + """Discovery shape: FTS5 plus adaptive or full result hydration.""" + current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None + title_result = _title_match_result(db, query, current_lineage_root) + raw_results, err = _loud(lambda: db.search_messages( + query=query, role_filter=role_filter or ["user", "assistant"], + exclude_sources=list(_HIDDEN_SESSION_SOURCES), limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, + fields=_DISCOVER_SEARCH_FIELDS), "FTS5 search failed: %s", "Search failed") + if err: + return err + # Demote cron rows below interactive ones BEFORE dedup so a high-volume cron + # corpus can't starve the user's own sessions out of the top `limit`. + raw_results = _order_for_recall(raw_results) + if not raw_results and not title_result: + return _discover_payload(db, query, detail, [], message=( + "No matching sessions found. FTS5 ANDs all terms by default — " + "broaden with OR (`alpha OR beta`), exact-match with quoted " + "phrases, exclude with NOT, or prefix-match with `deploy*`.")) + + seen_sessions: Dict[str, Dict[str, Any]] = {} + results = [] + if title_result: + title_lineage = title_result.pop("_lineage_root", None) + if title_lineage: + seen_sessions[title_lineage] = {"_title_only": True} + results.append(title_result) + _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) + for lineage_root, match_info in seen_sessions.items(): + if match_info.get("_title_only"): + continue + # Adaptive: only the top-ranked result is fully hydrated. + entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact") + if entry is not None: + results.append(entry) + for entry in results: + entry["link"] = _session_link(entry["session_id"], link_profile) + return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=( + "When referring the user to a session, write its `link` value " + "verbatim inline mid-sentence (it renders as a titled link) — never " + "as markdown, in backticks, on its own line, or next to the " + "title/id/date. To read more around a compact result, scroll: " + "session_search(session_id=..., around_message_id=match_message_id).")) def _resolve_profile_db(profile: str): - """Open another profile's ``state.db`` read-only (no write lock — safe on a - live DB), or None for the current profile.""" + """Another profile's ``state.db`` opened read-only (safe on a live DB), or None + for the current profile.""" if profile is None or not str(profile).strip(): return None from hermes_cli import profiles as profiles_mod from hermes_state import SessionDB - canon = profiles_mod.normalize_profile_name(profile) profiles_mod.validate_profile_name(canon) if not profiles_mod.profile_exists(canon): @@ -42,10 +368,9 @@ def _resolve_profile_db(profile: str): def _locate_session_db(session_id: str): - """Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)`` - or ``(None, None)``. Ids are globally unique, so the first hit is authoritative.""" + """Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)`` or + ``(None, None)``. Ids are globally unique, so the first hit is authoritative.""" from pathlib import Path - try: from hermes_cli import profiles as profiles_mod from hermes_state import SessionDB @@ -60,14 +385,12 @@ def _locate_session_db(session_id: str): if str(db_path) in seen or not db_path.exists(): continue seen.add(str(db_path)) - try: - pdb = SessionDB(db_path=db_path, read_only=True) - except Exception: - continue - if _quiet(lambda: pdb.get_session(session_id), None, - "get_session probe failed for %s in %s", session_id, name): + pdb = _quiet(lambda: SessionDB(db_path=db_path, read_only=True), None, "open %s failed", db_path) + if pdb and _quiet(lambda: pdb.get_session(session_id), None, + "get_session probe failed for %s in %s", session_id, name): return pdb, name - pdb.close() + if pdb: + pdb.close() return None, None @@ -78,18 +401,14 @@ def _get_session_meta(db, session_id: str) -> dict: def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_profile: str = None) -> str: - """Read shape: whole session, or first ``head`` + last ``tail`` messages - with a pointer to scroll the middle.""" + """Read shape: whole session, or ``head`` + ``tail`` messages with a scroll pointer.""" meta = _get_session_meta(db, session_id) if not meta: return tool_error(f"session_id not found: {session_id}", success=False) - - try: - rows = db.get_messages(session_id) - except Exception as e: - logging.error("get_messages failed for %s: %s", session_id, e, exc_info=True) - return tool_error(f"failed to load session: {e}", success=False) - + rows, err = _loud(lambda: db.get_messages(session_id), "get_messages failed for %s: %s", "failed to load session", + session_id) + if err: + return err shaped = [_shape_message(m) for m in rows] total = len(shaped) truncated = total > head + tail @@ -102,86 +421,67 @@ def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_prof def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_profile: str = None) -> str: """Browse shape: metadata for the most recent sessions (no LLM, no FTS5).""" - try: - # list_sessions_rich (include_children=False) already applies the - # canonical child classifier: roots, /branch children and /new-reset - # children are admitted, delegation/compression children hidden. - # Re-classifying here re-hid legacy reset children — trust the query. + def _browse(): + # list_sessions_rich already applies the canonical child classifier (roots, + # /branch and /new-reset children admitted; delegation/compression children + # hidden). Re-classifying here re-hid legacy reset children — trust the query. sessions = db.list_sessions_rich( limit=limit + 15, # extra so we can skip current / compression roots - exclude_sources=list(_HIDDEN_SESSION_SOURCES), - order_by_last_active=True, - ) - + exclude_sources=list(_HIDDEN_SESSION_SOURCES), order_by_last_active=True) current_root, has_compression_hop = ( - _resolve_to_parent(db, current_session_id) - if current_session_id else (None, False) - ) + _resolve_to_parent(db, current_session_id) if current_session_id else (None, False)) results = [] for s in sessions: sid = s.get("id", "") - if sid == current_session_id: - continue - # Compression continuation: the root's turns were summarised into - # the live child, so hide the root. /new-reset children share a - # root but carry no transcript — keep that root browsable. - if has_compression_hop and current_root and sid == current_root: + # Compression continuation: the root was summarised into the live child, so + # hide it. /new-reset children carry no transcript — keep that root browsable. + if sid == current_session_id or (has_compression_hop and current_root and sid == current_root): continue results.append({ "session_id": sid, "link": _session_link(sid, link_profile), "title": s.get("title") or None, - "source": s.get("source", ""), "started_at": s.get("started_at", ""), - "last_active": s.get("last_active", ""), "message_count": s.get("message_count", 0), - "preview": s.get("preview", ""), - }) + **{k: s.get(k, "") for k in ("source", "started_at", "last_active")}, + "message_count": s.get("message_count", 0), "preview": s.get("preview", "")}) if len(results) >= limit: break return _ok(mode="browse", results=results, count=len(results), message=( f"Showing {len(results)} most recent sessions. Pass a query= to search, " "or session_id+around_message_id to scroll.")) - except Exception as e: - logging.error("Error listing recent sessions: %s", e, exc_info=True) - return tool_error(f"Failed to list recent sessions: {e}", success=False) + + out, err = _loud(_browse, "Error listing recent sessions: %s", "Failed to list recent sessions") + return err or out def _clamp_int(value, default: int, lo: int, hi: int) -> int: - if not isinstance(value, int): - try: - value = int(value) - except (TypeError, ValueError): - value = default + try: + value = int(value) + except (TypeError, ValueError): + value = default return max(lo, min(value, hi)) def _anchor_in_live_context(db, anchor_state, anchor_session_id: str, current_session_id: str) -> bool: - """True when the scroll anchor is still in the caller's active context and - must be rejected. Same-lineage history that has LEFT live context (compacted - rows, compression-ended parents, /new-reset predecessors) is allowed, so - scroll never rejects a result discovery just returned.""" - a_root = _resolve_lineage(db, anchor_session_id) - c_root = _resolve_lineage(db, current_session_id) - if not (a_root and c_root and a_root == c_root): - return False - if _is_compacted_state(anchor_state): + """True when the scroll anchor is still in the caller's active context (reject). + Same-lineage history that LEFT live context (compacted rows, compression-ended + parents, /new-reset predecessors) is allowed, so scroll never rejects a result + discovery just returned.""" + if not _same_lineage(db, anchor_session_id, current_session_id) or _is_compacted_state(anchor_state): return False # Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history. - is_inactive_non_compacted = ( - anchor_state is not None - and anchor_state["active"] == 0 - and anchor_state["compacted"] != 1 - ) - return is_inactive_non_compacted or not _session_left_live_context(db, anchor_session_id) + inactive_non_compacted = anchor_state is not None and anchor_state["active"] == 0 and anchor_state["compacted"] != 1 + return inactive_non_compacted or not _session_left_live_context(db, anchor_session_id) + + +def _same_lineage(db, a: str, b: str) -> bool: + a_root, b_root = _resolve_lineage(db, a), _resolve_lineage(db, b) + return bool(a_root and b_root and a_root == b_root) def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, window: int): - """Lineage rebind: the caller paired a parent session_id with a message id - that lives in a descendant (compaction / delegation create child sessions). - Returns ``(view, warning)`` from the owning session, or ``(None, None)``.""" - a_root = _resolve_lineage(db, session_id) - o_root = _resolve_lineage(db, owning) - if not (a_root and o_root and a_root == o_root): - return None, None - rebind_view = _quiet(lambda: db.get_messages_around(owning, around_message_id, window=window), - None, "rebind get_messages_around failed: %s", with_exc=True) + """Lineage rebind when the caller paired a parent session_id with a message id + living in a descendant. ``(view, warning)`` from the owner, or ``(None, None)``.""" + rebind_view = _same_lineage(db, session_id, owning) and _quiet( + lambda: db.get_messages_around(owning, around_message_id, window=window), + None, "rebind get_messages_around failed: %s", with_exc=True) if not (rebind_view and rebind_view.get("window")): return None, None return rebind_view, f"around_message_id {around_message_id} lives in {owning} (child of {session_id}); rebound transparently" @@ -189,8 +489,8 @@ def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, w def _scroll(db, session_id: str, around_message_id: int, window: int = 5, current_session_id: str = None) -> str: - """Scroll shape: a window of messages centered on an anchor (no FTS5, no - bookends). Rebinds silently if the anchor lives in a same-lineage child.""" + """Scroll shape: a window centered on an anchor (no FTS5, no bookends); + rebinds silently if the anchor lives in a same-lineage child.""" if not isinstance(session_id, str) or not session_id.strip(): return tool_error("scroll requires session_id", success=False) session_id = session_id.strip() @@ -199,41 +499,29 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, except (TypeError, ValueError): return tool_error("scroll requires integer around_message_id", success=False) window = _clamp_int(window, 5, 1, 20) - # Locate the anchor BEFORE the current-lineage guard (see _anchor_in_live_context). anchor_state = _get_message_storage_state(db, around_message_id) owning_session_id = anchor_state.get("session_id") if anchor_state is not None else None - if current_session_id and _anchor_in_live_context( - db, anchor_state, owning_session_id or session_id, current_session_id - ): + db, anchor_state, owning_session_id or session_id, current_session_id): return tool_error("scroll rejected: anchor lives in the current session lineage (already in your active context)", success=False) - session_meta = _get_session_meta(db, session_id) if not session_meta: return tool_error(f"session_id not found: {session_id}", success=False) - - try: - view = db.get_messages_around(session_id, around_message_id, window=window) - except Exception as e: - logging.error("get_messages_around failed: %s", e, exc_info=True) - return tool_error(f"failed to load messages: {e}", success=False) - + view, err = _loud(lambda: db.get_messages_around(session_id, around_message_id, window=window), + "get_messages_around failed: %s", "failed to load messages") + if err: + return err messages = view.get("window") or [] rebind_warning = None if not messages and owning_session_id and owning_session_id != session_id: - rebind_view, rebind_warning = _rebind_to_owner( - db, session_id, owning_session_id, around_message_id, window - ) + rebind_view, rebind_warning = _rebind_to_owner(db, session_id, owning_session_id, around_message_id, window) if rebind_view is not None: - view = rebind_view - messages = view["window"] + view, messages = rebind_view, rebind_view["window"] session_meta = _get_session_meta(db, owning_session_id) or session_meta session_id = owning_session_id - if not messages: return tool_error(f"around_message_id {around_message_id} not in session_id {session_id}", success=False) - return _ok( mode="scroll", session_id=session_id, around_message_id=around_message_id, session_meta=_session_meta_block(session_meta), window=window, @@ -243,101 +531,84 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, "id; backward: the FIRST message's id (the boundary message repeats " "as an orientation marker). messages_before/messages_after < window " "means you've hit that end of the session."), - **({"warning": rebind_warning} if rebind_warning else {}), - ) + **({"warning": rebind_warning} if rebind_warning else {})) def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str: - """Read shape. On a miss in the target profile, scan every profile (the - model may have dropped the owning profile from the link) and tag the result - with the profile it was found in.""" + """Read shape; on a miss scan every profile (the model may have dropped the + owning profile from the link) and tag the result with where it was found.""" result = _read_session(db, sid, link_profile=profile) if json.loads(result).get("success"): return result located, owner = _locate_session_db(sid) - if located is not None: - try: - found = json.loads(_read_session(located, sid, link_profile=owner)) - finally: - located.close() - if found.get("success"): - found["profile"] = owner - return json.dumps(found, ensure_ascii=False) - return result + if located is None: + return result + try: + found = json.loads(_read_session(located, sid, link_profile=owner)) + finally: + located.close() + if not found.get("success"): + return result + found["profile"] = owner + return json.dumps(found, ensure_ascii=False) def _dispatch(query, role_filter, limit, db, current_session_id, session_id, around_message_id, window, sort, profile, detail, owned_dbs) -> str: - """Mode dispatch (see module docstring). Scroll wins over read/discovery when - an anchor is set — the agent asked for a specific slice. Profile DBs opened - here are appended to *owned_dbs* for the caller to close.""" - # A raw `@session:/` link passed as session_id: ids never - # contain "/", so a slash means profile/id — strip the prefix and adopt the - # embedded profile only when none was passed explicitly. + """Mode dispatch (see module docstring); scroll wins when an anchor is set. + Profile DBs opened here are appended to *owned_dbs* for the caller to close.""" + # A raw `@session:/` link as session_id: ids never contain "/", so + # split on it and adopt the embedded profile only when none was passed. if isinstance(session_id, str) and "/" in session_id: emb_profile, _, emb_id = session_id.partition("/") if emb_id: session_id = emb_id if emb_profile and (profile is None or not str(profile).strip()): profile = emb_profile - - # Cross-profile read: swap in the named profile's DB (read-only) for every - # shape. Current-lineage guards key off ids that won't collide, so they - # stay inert. - if profile is not None and str(profile).strip(): - try: - profile_db = _resolve_profile_db(profile) - except Exception as e: - return tool_error(f"profile '{profile}': {e}", success=False) - if profile_db is not None: - db, current_session_id = profile_db, None - owned_dbs.append(profile_db) - - has_session = isinstance(session_id, str) and bool(session_id.strip()) - if has_session and around_message_id is not None: - return _scroll(db, session_id, around_message_id, window, current_session_id) - if has_session: + # Cross-profile: swap in the named profile's DB (read-only) for every shape; + # current-lineage guards key off ids that won't collide, so they stay inert. + try: + profile_db = _resolve_profile_db(profile) + except Exception as e: + return tool_error(f"profile '{profile}': {e}", success=False) + if profile_db is not None: + db, current_session_id = profile_db, None + owned_dbs.append(profile_db) + if isinstance(session_id, str) and session_id.strip(): + if around_message_id is not None: + return _scroll(db, session_id, around_message_id, window, current_session_id) return _read_with_profile_fallback(db, session_id.strip(), profile) - limit = _clamp_int(limit, 3, 1, 10) if not query or not isinstance(query, str) or not query.strip(): return _list_recent_sessions(db, limit, current_session_id, link_profile=profile) - role_list = ([r.strip() for r in role_filter.split(",") if r.strip()] or None) if isinstance(role_filter, str) else None sort_norm = sort.strip().lower() if isinstance(sort, str) else None - if sort_norm not in ("newest", "oldest"): - sort_norm = None + sort_norm = sort_norm if sort_norm in ("newest", "oldest") else None detail_norm = "full" if isinstance(detail, str) and detail.strip().lower() == "full" else "adaptive" return _discover( db=db, query=query.strip(), role_filter=role_list, limit=limit, sort=sort_norm, - detail=detail_norm, current_session_id=current_session_id, link_profile=profile, - ) + detail=detail_norm, current_session_id=current_session_id, link_profile=profile) def session_search(query: str = "", role_filter: str = None, limit: int = 3, db=None, current_session_id: str = None, session_id: str = None, around_message_id: int = None, window: int = 5, sort: str = None, profile: str = None, detail: str = "adaptive") -> str: - """Run session search and close databases opened by this invocation. - Parameter order is positional-compatible with older callers.""" + """Run session search, closing DBs opened here. Positional order is frozen for old callers.""" owned_dbs: List[Any] = [] if db is None: try: from hermes_state import get_shared_session_db - db = get_shared_session_db() owned_dbs.append(db) except Exception: logging.debug("SessionDB unavailable for session_search", exc_info=True) from hermes_state import format_session_db_unavailable - return tool_error(format_session_db_unavailable(), success=False) - try: return _dispatch(query, role_filter, limit, db, current_session_id, session_id, around_message_id, window, sort, profile, detail, owned_dbs) finally: from hermes_state import release_or_close - for owned_db in reversed(owned_dbs): _quiet(lambda: release_or_close(owned_db), None, "Failed to close session_search SessionDB") @@ -465,18 +736,8 @@ registry.register( toolset="session_search", schema=SESSION_SEARCH_SCHEMA, handler=lambda args, **kw: session_search( - query=args.get("query") or "", - role_filter=args.get("role_filter"), - limit=args.get("limit", 3), - session_id=args.get("session_id"), - around_message_id=args.get("around_message_id"), - window=args.get("window", 5), - sort=args.get("sort"), - detail=args.get("detail", "adaptive"), - profile=args.get("profile"), - db=kw.get("db"), - current_session_id=kw.get("current_session_id"), - ), + query=args.get("query") or "", limit=args.get("limit", 3), window=args.get("window", 5), + detail=args.get("detail", "adaptive"), db=kw.get("db"), current_session_id=kw.get("current_session_id"), + **{k: args.get(k) for k in ("role_filter", "session_id", "around_message_id", "sort", "profile")}), check_fn=check_session_search_requirements, - emoji="🔍", -) + emoji="🔍") diff --git a/tools/session_search_tool_common.py b/tools/session_search_tool_common.py deleted file mode 100644 index 0a0d2c44f0..0000000000 --- a/tools/session_search_tool_common.py +++ /dev/null @@ -1,230 +0,0 @@ -"""Shared helpers for the session_search tool: source classification, lineage -resolution, message storage state, and response shaping. Imported by -``tools.session_search_tool`` (which re-exports the names) and -``tools.session_search_tool_discover``.""" - -import json -import logging -from datetime import datetime -from typing import Any, Dict, List, Optional, Union - -from hermes_state_common import _RESET_END_REASONS - -# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool), -# delegate subagent runs, kanban workers — not the user's history. -_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool") - -# Searchable but DEMOTED below interactive sessions: cron sessions' repetitive -# vocabulary dominates bare BM25 and starves out the user's own sessions -# ("recall blindness"). Demoting keeps them reachable when they're the only match. -_DEMOTED_SESSION_SOURCES = ("cron",) - -# FTS rows scanned before dedup-by-lineage — well above the handful of distinct -# sessions a query returns, so interactive matches buried under cron hits are -# still in hand for the demotion pass. -_DISCOVER_SCAN_LIMIT = 300 - -# Raw FTS rows are only a discovery-plan input; the response hydrates its own -# anchored window and bookends after lineage dedup. -_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started") - -# Generated context-compaction handoff summaries (agent/context_compressor.py); -# excluded from bookends so huge compaction payloads aren't re-introduced. -_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:") - -# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the -# predecessor WITHOUT carrying its transcript forward — unlike compression -# continuations and live delegation children. Derived from the gateway set so -# this tool and the recovery fence cannot drift. -_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"} - - -def _quiet(fn, default, msg, *log_args, with_exc: bool = False): - """Call ``fn()``; on any exception debug-log *msg* (appending the exception - as a final ``%s`` arg when *with_exc*) and return *default*.""" - try: - return fn() - except Exception as e: - logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True) - return default - - -def _format_timestamp(ts: Union[int, float, str, None]) -> str: - """Unix timestamp (number or numeric string) or ISO string -> readable date. - "unknown" for None; str(ts) if conversion fails.""" - if ts is None: - return "unknown" - try: - value = ts - if isinstance(ts, str): - if not ts.replace(".", "").replace("-", "").isdigit(): - return ts - value = float(ts) - if isinstance(value, (int, float)): - return datetime.fromtimestamp(value).strftime("%B %d, %Y at %I:%M %p") - except (ValueError, OSError, OverflowError) as e: - logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True) - except Exception as e: - logging.debug("Unexpected error formatting timestamp %s: %s", ts, e, exc_info=True) - return str(ts) - - -def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]: - """The ``session_meta`` sub-object shared by read/scroll responses.""" - return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"), - "model": meta.get("model"), "title": meta.get("title")} - - -def _ok(**payload) -> str: - """Serialize a successful tool result (``success`` first, then *payload* in order).""" - return json.dumps({"success": True, **payload}, ensure_ascii=False) - - -def _is_compaction_summary(content: str) -> bool: - """Return True if *content* looks like a generated compaction handoff.""" - return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES) - - -def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]: - """Walk parent_session_id to the lineage root -> ``(root_id, has_compression_hop)``. - The flag distinguishes a compression-split lineage (parent content summarised - away) from a delegation lineage (child content still visible to the parent). - Falls back to ``(session_id, False)`` on errors.""" - if not session_id: - return session_id, False - visited: set[str] = set() - cur, has_compression = session_id, False - while cur and cur not in visited: - visited.add(cur) - s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True) - if not s: - break - if s.get("end_reason") == "compression": - has_compression = True - if not s.get("parent_session_id"): - break - cur = s["parent_session_id"] - return cur, has_compression - - -def _resolve_lineage(db, session_id: str) -> str: - """Return only the lineage root (ignores the compression hop).""" - return _resolve_to_parent(db, session_id)[0] - - -def _session_end_reason(db, session_id: str) -> Optional[str]: - """Return the session's ``end_reason``, or None if missing/unended/error.""" - if not session_id: - return None - try: - s = db.get_session(session_id) - return (s.get("end_reason") or None) if s else None - except Exception: - return None - - -def _session_left_live_context(db, session_id: str) -> bool: - """True when *session_id*'s transcript is no longer in anyone's live context: - ``compression`` (summarised into the continuation child) or a fresh reset - (child starts empty). Everything else stays excluded from same-lineage - recall — live delegation children (``end_reason is None``) are visible to - the parent agent, and ``branched`` parents were copied verbatim into the - branch child, so their content IS the current context.""" - end_reason = _session_end_reason(db, session_id) - return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS - - -def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]: - """Return the owning session and visibility flags for *message_id*.""" - if not message_id: - return None - - def _lookup(): - with db._lock: - return db._conn.execute( - "SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,) - ).fetchone() - - row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id) - return dict(row) if row is not None else None - - -def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool: - """Compaction archives are ``active=0, compacted=1`` (content summarised - away by archive_and_compact). Rewind/undo rows are ``active=0, compacted=0`` - and must stay hidden.""" - return state is not None and state["active"] == 0 and state["compacted"] == 1 - - -def _is_compacted_message(db, message_id) -> bool: - """True if *message_id* is a compaction-archived row — pre-compaction content - no longer in live context, so it should stay discoverable even on the - current session. False on any error (caller falls back to skipping).""" - return _is_compacted_state(_get_message_storage_state(db, message_id)) - - -def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None: - """Add a rebuild-progress note while the deferred FTS backfill is running, - so the agent can explain thin/slow results instead of treating them as - ground truth. No-op (never raises) when no rebuild is pending.""" - try: - status = db.fts_rebuild_status() - except Exception: - return - if status is None: - return - payload["index_rebuild"] = {"percent": status["percent"], "note": ( - f"The search index is rebuilding in the background ({status['percent']}% done, " - f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages " - f"may be incomplete until it finishes." - )} - - -def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]: - """Stable-sort FTS rows so interactive sessions rank above automation. - BM25 order is preserved within each class; only cross-class order changes, - so a cron hit never displaces an interactive hit during lineage dedup.""" - return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0) - - -def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None, - max_content_len: Optional[int] = None) -> Dict[str, Any]: - """Slim a message row for the tool response. Keeps content even if empty - (absent content is meaningful — tool-call-only assistant turns). With - *max_content_len*, content is truncated and ``content_truncated`` / - ``original_content_chars`` added.""" - content = m.get("content") - if isinstance(content, str) and "\x1b" in content: - # Recalled messages can carry ANSI escapes (archived terminal output). - from tools.ansi_strip import strip_ansi - - content = strip_ansi(content) - original_chars = None - if max_content_len and content and len(content) > max_content_len: - original_chars = len(content) - content = content[:max_content_len] + "…" - entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")} - entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)}) - if anchor_id is not None and m.get("id") == anchor_id: - entry["anchor"] = True - if original_chars is not None: - entry["content_truncated"] = True - entry["original_content_chars"] = original_chars - return {k: v for k, v in entry.items() if v is not None or k == "content"} - - -def _session_link(session_id: str, profile: str = None) -> str: - """The reference the agent writes to point the user at a session — same - value the desktop composer emits, so it renders as a titled link. The - profile segment is omitted when it can't be named confidently (a bare id - still resolves, it just can't disambiguate across profiles).""" - name = (profile or "").strip() - if not name: - def _active(): - from hermes_cli.profiles import get_active_profile_name - - resolved = get_active_profile_name() - return "" if resolved == "custom" else resolved - - name = _quiet(_active, "", "get_active_profile_name failed for session link") - return f"@session:{name}/{session_id}" if name else f"@session:{session_id}" diff --git a/tools/session_search_tool_discover.py b/tools/session_search_tool_discover.py deleted file mode 100644 index 08ba526421..0000000000 --- a/tools/session_search_tool_discover.py +++ /dev/null @@ -1,192 +0,0 @@ -"""Discovery shape of session_search: FTS5 query, title match, lineage dedup -and adaptive/full hydration of the surviving results.""" - -import json -import logging -from typing import Any, Dict, List, Optional - -from tools.registry import tool_error -from tools.session_search_tool_common import ( - _DISCOVER_SCAN_LIMIT, _DISCOVER_SEARCH_FIELDS, _HIDDEN_SESSION_SOURCES, _annotate_rebuild_status, - _format_timestamp, _is_compacted_message, _is_compaction_summary, _order_for_recall, _quiet, - _resolve_lineage, _resolve_to_parent, _session_left_live_context, _session_link, _shape_message, -) - - -def _normalize_title_query(query: str) -> str: - """Strip common quoting the model may include around a remembered title.""" - return query.strip().strip("`'\"") - - -def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]: - """Return a discovery-shaped result when the query matches a session title.""" - title_query = _normalize_title_query(query) - if not title_query: - return None - session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None, - "resolve_session_by_title failed for %r", title_query) - if not session_id: - return None - lineage_root = _resolve_lineage(db, session_id) - # Same-lineage title hits are in-context only while the session is live; - # /new-reset and compression-ended parents are not. - if ( - current_lineage_root - and lineage_root == current_lineage_root - and not _session_left_live_context(db, session_id) - ): - return None - - session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None, - "get_session failed for title match %s", session_id) or {} - if session_meta.get("source") in _HIDDEN_SESSION_SOURCES: - return None - messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id) - anchor_id = messages[0].get("id") if messages else None - view = {} - if anchor_id is not None: - view = _quiet(lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {}, - "get_anchored_view failed for title match %s/%s", session_id, anchor_id) - entry = { - "session_id": session_id, "when": _format_timestamp(session_meta.get("started_at")), - "source": session_meta.get("source", "unknown"), "model": session_meta.get("model") or "unknown", - "title": session_meta.get("title") or title_query, "matched_role": "session_title", - "match_message_id": anchor_id, - "snippet": f"Session title matched: {session_meta.get('title') or title_query}", - "bookend_start": [_shape_message(m) for m in (view.get("bookend_start") or messages[:3])], - "messages": [_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])], - "bookend_end": [_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])], - "messages_before": view.get("messages_before", 0), - "messages_after": view.get("messages_after", max(len(messages) - 5, 0)), - "detail": "full", "_lineage_root": lineage_root, - } - if lineage_root and lineage_root != session_id: - entry["parent_session_id"] = lineage_root - return entry - - -def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str: - payload = {"success": True, "mode": "discover", "query": query, "detail": detail, - "results": results, "count": len(results), **extra} - _annotate_rebuild_status(db, payload) - return json.dumps(payload, ensure_ascii=False) - - -def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None: - """Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*. - - The raw owning session_id stays on the row — only it pairs validly with the - FTS match id for the anchored window. Current-lineage hits are skipped - UNLESS the transcript left live context: compression-ended session, /new- - reset predecessor (hiding it made gateway recall blind after every /new), - or an in-place compacted row on the SAME session_id. A live delegation - child has end_reason=None, so it stays excluded. - """ - for r in raw_results: - if len(seen_sessions) >= limit: - break - raw_sid = r["session_id"] - resolved_sid, _ = _resolve_to_parent(db, raw_sid) - is_compacted_hit = _is_compacted_message(db, r.get("id")) - is_ended_session = _session_left_live_context(db, raw_sid) - if ( - current_lineage_root - and resolved_sid == current_lineage_root - and not (is_ended_session or is_compacted_hit) - ): - continue - if current_session_id and raw_sid == current_session_id and not is_compacted_hit: - continue - seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid}) - - -def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]: - return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or []) - if not _is_compaction_summary(m.get("content", ""))] - - -def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]: - """Build one discovery result from a surviving FTS row; None if the anchored - view can't be loaded (the hit is dropped).""" - hit_sid = match_info.get("session_id") or lineage_root - msg_id = match_info.get("id") - try: - view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3) - except Exception as e: - logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True) - return None - session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {} - full = result_detail == "full" - window_messages = view.get("window") or [] - if not full: - window_messages = [m for m in window_messages if m.get("id") == msg_id] - entry = { - "session_id": hit_sid, - "when": _format_timestamp(session_meta.get("started_at") or match_info.get("session_started")), - "source": session_meta.get("source") or match_info.get("source", "unknown"), - "model": session_meta.get("model") or match_info.get("model") or "unknown", - "title": session_meta.get("title") or None, "matched_role": match_info.get("role"), - "match_message_id": msg_id, "snippet": match_info.get("snippet") or "", - "bookend_start": _bookend(view, "bookend_start") if full else [], - "messages": [_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages], - "bookend_end": _bookend(view, "bookend_end") if full else [], - "messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", 0), - "detail": result_detail, - } - if lineage_root and lineage_root != hit_sid: - entry["parent_session_id"] = lineage_root - return entry - - -def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str], - detail: str, current_session_id: str = None, link_profile: str = None) -> str: - """Discovery shape: FTS5 plus adaptive or full result hydration.""" - role_list = role_filter if role_filter else ["user", "assistant"] - current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None - title_result = _title_match_result(db, query, current_lineage_root) - - try: - raw_results = db.search_messages( - query=query, role_filter=role_list, exclude_sources=list(_HIDDEN_SESSION_SOURCES), - limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS, - ) - except Exception as e: - logging.error("FTS5 search failed: %s", e, exc_info=True) - return tool_error(f"Search failed: {e}", success=False) - - # Demote cron rows below interactive ones BEFORE dedup so a high-volume - # cron corpus can't starve the user's own sessions out of the top `limit`. - raw_results = _order_for_recall(raw_results) - - if not raw_results and not title_result: - return _discover_payload(db, query, detail, [], message=( - "No matching sessions found. FTS5 ANDs all terms by default — " - "broaden with OR (`alpha OR beta`), exact-match with quoted " - "phrases, exclude with NOT, or prefix-match with `deploy*`." - )) - - seen_sessions: Dict[str, Dict[str, Any]] = {} - results = [] - if title_result: - title_lineage = title_result.pop("_lineage_root", None) - if title_lineage: - seen_sessions[title_lineage] = {"_title_only": True} - results.append(title_result) - _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) - - for lineage_root, match_info in seen_sessions.items(): - if match_info.get("_title_only"): - continue - # Adaptive: only the top-ranked result is fully hydrated. - entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact") - if entry is not None: - results.append(entry) - for entry in results: - entry["link"] = _session_link(entry["session_id"], link_profile) - return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=( - "When referring the user to a session, write its `link` value " - "verbatim inline mid-sentence (it renders as a titled link) — never " - "as markdown, in backticks, on its own line, or next to the " - "title/id/date. To read more around a compact result, scroll: " - "session_search(session_id=..., around_message_id=match_message_id)." - ))