diff --git a/tools/memory_tool.py b/tools/memory_tool.py index 260d2ef6e5..7a18ddb1e6 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,17 +29,15 @@ 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) 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" @@ -50,9 +47,9 @@ from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports) 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.""" + """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 @@ -72,9 +69,8 @@ def load_on_disk_store() -> "MemoryStore": # -- 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: @@ -126,11 +122,10 @@ def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Op # -- 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. A missing ``old_text`` - is recoverable (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, not a dead-end.""" + """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: @@ -162,10 +157,9 @@ def memory_tool( 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.""" + """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) @@ -191,8 +185,7 @@ def memory_tool( 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 @@ -200,9 +193,9 @@ def memory_tool( 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 diff --git a/tools/memory_tool_store.py b/tools/memory_tool_store.py index 242e205854..7999a7be3c 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,9 +14,8 @@ 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)"} @@ -29,15 +28,13 @@ def _memory_dir() -> Path: def _scan_memory_content(content: str) -> Optional[str]: - """Scan 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 externally added content.""" + """External drift: the file wouldn't round-trip, so flushing would discard content.""" return { "success": False, "error": ( @@ -55,8 +52,7 @@ def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]: 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, " @@ -80,10 +76,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, *, @@ -95,9 +90,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.""" @@ -108,10 +101,8 @@ 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 @@ -122,10 +113,9 @@ class MemoryStore: 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). @@ -139,9 +129,8 @@ class MemoryStore: @staticmethod def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]: - """*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] = [] @@ -158,8 +147,8 @@ class MemoryStore: @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 @@ -192,13 +181,11 @@ class MemoryStore: return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md") def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]: - """Re-read entries from disk (under file lock) before mutating; return the - abort error dict, or None to proceed. Aborts on external drift (flushing - would discard un-roundtrippable content) and when the file exists but can't - be read (rewriting from an assumed-empty view would wipe it — even + """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". - *skip_drift* skips the round-trip check (``add``).""" + 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: @@ -227,9 +214,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 showing 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)}) @@ -261,9 +253,8 @@ class MemoryStore: 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. + # Append-only: skip the drift guard (appending never clobbers foreign + # content) but still refuse a failed read — add rewrites the WHOLE file. err = self._reload_or_error(target, skip_drift=True) if err: return err @@ -282,23 +273,35 @@ class MemoryStore: 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} + 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*.""" 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, "replace") + idx, err = self._locate(target, old_text, "replace" if new_content else "remove") if err: return err entries = self._entries_for(target) + if new_content is None: + entries.pop(idx) + return self._commit(target, entries, "Entry removed.") limit = self._char_limit(target) new_total = len(ENTRY_DELIMITER.join(entries[:idx] + [new_content] + entries[idx + 1:])) if new_total > limit: @@ -309,24 +312,6 @@ class MemoryStore: entries[idx] = new_content return self._commit(target, entries, "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 -- - @staticmethod def _apply_batch_op(working: List[str], act: str, content: str, old_text: str, pos: str) -> Optional[str]: """Apply one batch op to *working* in place; return an error message or None.""" @@ -354,17 +339,14 @@ 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 one 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. + # Scan every add/replace content BEFORE touching disk -- one poisoned op rejects the batch. for i, op in enumerate(operations): op = op or {} scan_error = op.get("action") in {"add", "replace"} and op.get("content") and _scan_memory_content(op["content"]) @@ -386,7 +368,6 @@ class MemoryStore: pos = f"Operation {i + 1} ({act or 'unknown'})" msg = self._apply_batch_op(working, act, content, old_text, pos) if msg: - # Batch-abort error reporting live (uncommitted) state. return self._failure_with_entries( target, msg + " No operations were applied (batch is all-or-nothing).") # Budget check against the FINAL state only. @@ -399,38 +380,24 @@ class MemoryStore: return self._commit(target, working, f"Applied {len(operations)} operation(s).") def format_for_system_prompt(self, target: str) -> Optional[str]: - """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.""" @@ -443,12 +410,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: @@ -458,24 +423,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_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.""" + """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) @@ -491,8 +452,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 b2fbc73a5f..88b1ea9eb9 100644 --- a/tools/microsoft_graph_auth.py +++ b/tools/microsoft_graph_auth.py @@ -31,11 +31,8 @@ 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): @@ -114,18 +111,12 @@ 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``.""" @@ -145,11 +136,8 @@ class MicrosoftGraphTokenProvider: 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, @@ -163,7 +151,6 @@ class MicrosoftGraphTokenProvider: 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.") @@ -171,9 +158,7 @@ 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", @@ -181,16 +166,12 @@ class MicrosoftGraphTokenProvider: 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"] diff --git a/tools/microsoft_graph_client.py b/tools/microsoft_graph_client.py index 25a1aa40e5..bfdc3c011b 100644 --- a/tools/microsoft_graph_client.py +++ b/tools/microsoft_graph_client.py @@ -38,12 +38,9 @@ class MicrosoftGraphAPIError(MicrosoftGraphClientError): 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, *, @@ -71,15 +68,15 @@ 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}) + + 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]: @@ -102,9 +99,8 @@ class MicrosoftGraphClient: 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) @@ -145,12 +141,9 @@ class MicrosoftGraphClient: 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. - """ + """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 diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index f49e3889d7..e4f2159ac5 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -19,34 +19,29 @@ from hermes_state_common import _RESET_END_REASONS # 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. +# 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 handful of distinct -# sessions a query returns, so interactive matches buried under cron hits are -# still in hand for the demotion pass. +# 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 discovery-plan input; the response hydrates its own -# anchored window and bookends after lineage dedup. +# 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") -# Generated context-compaction handoff summaries (agent/context_compressor.py); -# excluded from bookends so huge compaction payloads aren't re-introduced. +# 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 -# this tool and the recovery fence cannot drift. +# /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): - """Call ``fn()``; on any exception debug-log *msg* (appending the exception - as a final ``%s`` arg when *with_exc*) and return *default*.""" + """``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: @@ -55,8 +50,8 @@ def _quiet(fn, default, msg, *log_args, with_exc: bool = False): 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.""" + """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" try: @@ -89,10 +84,9 @@ def _is_compaction_summary(content: str) -> bool: 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.""" + """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)``.""" if not session_id: return session_id, False visited: set[str] = set() @@ -127,12 +121,10 @@ def _session_end_reason(db, session_id: str) -> Optional[str]: 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.""" + """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.""" end_reason = _session_end_reason(db, session_id) return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS @@ -153,23 +145,20 @@ def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]: 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.""" + """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 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).""" + """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: - """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.""" + """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.""" status = _quiet(db.fts_rebuild_status, None, "fts_rebuild_status failed") if status is None: return @@ -180,21 +169,17 @@ def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None: 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.""" + """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 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.""" + """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: - # Recalled messages can carry ANSI escapes (archived terminal output). + if isinstance(content, str) and "\x1b" in content: # archived terminal output carries ANSI from tools.ansi_strip import strip_ansi content = strip_ansi(content) @@ -208,10 +193,9 @@ def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None, 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).""" + """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(): @@ -291,14 +275,10 @@ def _discover_payload(db, query: str, detail: str, results: list, **extra) -> st 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. - """ + 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 @@ -322,8 +302,8 @@ def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]: 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).""" + """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: @@ -402,8 +382,8 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort 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 @@ -417,8 +397,8 @@ 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: @@ -452,8 +432,7 @@ 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) @@ -477,10 +456,9 @@ 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. + # 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), @@ -493,9 +471,8 @@ def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_p 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. + # Compression continuation: the root was summarised into the live child, so + # hide it. /new-reset children carry no transcript — keep that root browsable. if has_compression_hop and current_root and sid == current_root: continue results.append({ @@ -522,10 +499,10 @@ def _clamp_int(value, default: int, lo: int, hi: int) -> int: 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.""" + """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. @@ -540,9 +517,8 @@ def _same_lineage(db, a: str, b: str) -> bool: 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)``.""" + """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)``.""" if not _same_lineage(db, session_id, owning): return None, None rebind_view = _quiet(lambda: db.get_messages_around(owning, around_message_id, window=window), @@ -554,8 +530,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() @@ -610,9 +586,8 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, 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 @@ -630,12 +605,10 @@ def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str: 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: @@ -643,9 +616,8 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_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. + # 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: @@ -676,8 +648,7 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_id, 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: