From 17f95967f46e9211b8bd6bad1cbac3e1f3d3a475 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:43:42 -0700 Subject: [PATCH] refactor(tools): wave-2 compaction of memory/session_search/graph/notification modules memory_tool_store: _error() helper for every failure dict, reload folded into _mutate, load_from_disk loops over targets with an inline snapshot sanitizer, _locate inlined into _edit, batch-op splice, drift check compacted; on-disk format and every result string unchanged (golden corpus). memory_tool: single _apply_write_gate handles op + batch, _batch_op_line helper, gate/validation returns folded, on-disk store built directly. session_search_tool: _get_session_meta unifies 5 get_session lookups, rebuild note folded into _discover_payload via _ok, lineage dedupe inlined into _discover, scroll rebind inlined, title-match shaping via one closure, browse comprehension, unreachable scroll guard dropped, single hermes_state import. microsoft_graph_client/auth: signatures hugged, bodiless-response handling inlined, header dict built in one expression, dead from_env removed. process_registry_notifications: _preamble merges header/task-source/role lines, _notice_lines shared, table-driven completion status, comprehension headers. Schemas byte-identical (SCHEMA-OK); golden corpus old-vs-new identical. --- tools/memory_tool.py | 98 +++--- tools/memory_tool_store.py | 259 ++++++--------- tools/microsoft_graph_auth.py | 20 +- tools/microsoft_graph_client.py | 92 ++---- tools/process_registry_notifications.py | 178 ++++------- tools/session_search_tool.py | 407 ++++++++++-------------- 6 files changed, 411 insertions(+), 643 deletions(-) diff --git a/tools/memory_tool.py b/tools/memory_tool.py index c8d42cb49a..7bdad746f4 100644 --- a/tools/memory_tool.py +++ b/tools/memory_tool.py @@ -23,7 +23,7 @@ try: except ImportError: fcntl = None try: - import msvcrt + import msvcrt # noqa: F401 except ImportError: pass @@ -32,8 +32,7 @@ logger = logging.getLogger(__name__) # 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: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar("memory_surface_flags", default=None) def get_memory_dir() -> Path: @@ -46,26 +45,22 @@ 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, ``/memory``) so approvals enforce the SAME caps - as ``agent_init``. Falls back to 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 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} + store = MemoryStore(int(mem_cfg.get("memory_char_limit", 2200)), int(mem_cfg.get("user_char_limit", 1375)), + memory_enabled=memory_enabled, user_profile_enabled=user_profile_enabled) except Exception: - kwargs: Dict[str, Any] = {} # config optional — fall back to defaults rather than break /memory - store = MemoryStore(**kwargs) + store = MemoryStore() # config optional — fall back to defaults rather than break /memory store.load_from_disk() return store -# -- Write-approval gate -- - def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]: """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.""" @@ -93,35 +88,30 @@ _STORE_ACTIONS = { 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).""" - summary, detail = _STORE_ACTIONS[action][1]("user profile" if target == "user" else "memory", content, old_text) - return _gate_or_stage(summary, detail, +def _batch_op_line(op: Dict[str, Any]) -> str: + op = op or {} + act, content, old = op.get("action", "?"), op.get("content") or op.get("new_text") or "", op.get("old_text", "") + if act == "remove": + return f"- remove: {old}" + return f"- replace: {old} -> {content}" if act == "replace" else f"- {act}: {content}" + + +def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str], + operations: Optional[List[Dict[str, Any]]] = None) -> Optional[str]: + """Gate one mutating op, or (``operations`` set) a whole batch as a single unit.""" + label = "user profile" if target == "user" else "memory" + if operations is not None: + return _gate_or_stage(f"apply {len(operations)} op(s) to {label}", + "\n".join(_batch_op_line(op) for op in operations), + {"action": "batch", "target": target, "operations": operations}) + return _gate_or_stage(*_STORE_ACTIONS[action][1](label, content, old_text), {"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 {'user profile' if target == 'user' else 'memory'}" - detail_lines = [] - for op in operations: - op = op or {} - act = op.get("action", "?") - 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 _validate_single_op(store, action, target, content, old_text) -> Optional[str]: - """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.""" + """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): return the inventory plus a retry instruction.""" 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: @@ -154,26 +144,23 @@ def memory_tool(action: str = None, target: str = "memory", content: str = None, if operations: if not isinstance(operations, list): return tool_error("operations must be a list of {action, content?, old_text?} objects.", success=False) - gate_result = _apply_batch_write_gate(target, operations) + # Approval gate: stages (background/gateway) or prompts inline (CLI); off by default. + gate_result = _apply_write_gate("batch", target, None, None, operations) if gate_result is not None: return gate_result return json.dumps(store.apply_batch(target, operations), ensure_ascii=False) 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) + invalid = (_validate_single_op(store, action, target, content, old_text) + or _apply_write_gate(action, target, content, old_text)) if invalid is not None: return invalid - # 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(_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`` reads the same section so 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 cannot diverge.""" if config is None: try: from hermes_cli.config import load_config_readonly @@ -214,8 +201,7 @@ def _memory_target_error(store: "MemoryStore", target: str) -> Optional[Dict[str def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[str, Any]: """Replay a staged write against the store, bypassing the gate (/memory approve).""" - action = payload.get("action") - target = payload.get("target", "memory") + action, target = payload.get("action"), payload.get("target", "memory") target_error = _memory_target_error(store, target) if target_error is not None: return target_error @@ -226,8 +212,6 @@ def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[ return _STORE_ACTIONS[action][0](store, target, payload.get("content") or "", payload.get("old_text") or "") -# -- OpenAI Function-Calling Schema -- - MEMORY_SCHEMA = { "name": "memory", "description": ( @@ -316,21 +300,17 @@ def _build_memory_schema_overrides() -> Dict[str, Any]: _memory_surface_flags.set(None) targets = [t for t, on in zip(("memory", "user"), flags) if on] parameters = copy.deepcopy(MEMORY_SCHEMA["parameters"]) - target_schema = parameters["properties"]["target"] + target_schema, description = parameters["properties"]["target"], MEMORY_SCHEMA["description"] target_schema["enum"] = targets - description = MEMORY_SCHEMA["description"] - narrowed = _SINGLE_TARGET_TEXT.get(tuple(targets)) - if narrowed: + if narrowed := _SINGLE_TARGET_TEXT.get(tuple(targets)): target_schema["description"], replacement = narrowed description = description.replace( "TARGETS: 'user' = who the user is (name, role, preferences, style). 'memory' = your " - "notes (environment, conventions, tool quirks, lessons).", - replacement) + "notes (environment, conventions, tool quirks, lessons).", replacement) return {"description": description, "parameters": parameters} -# --- Registry --- -from tools.registry import registry, tool_error +from tools.registry import registry, tool_error # noqa: E402 (registration at import time) registry.register( name="memory", diff --git a/tools/memory_tool_store.py b/tools/memory_tool_store.py index 1ee1408358..6abd008dd0 100644 --- a/tools/memory_tool_store.py +++ b/tools/memory_tool_store.py @@ -5,7 +5,7 @@ in ``tools.memory_tool`` and is read lazily.""" import logging import time -from contextlib import contextmanager +from contextlib import contextmanager, suppress from pathlib import Path from typing import Any, Dict, List, Optional, Tuple @@ -22,37 +22,36 @@ MEMORY_BLOCK_HEADERS = { ENTRY_DELIMITER = "\n§\n" -def _memory_dir() -> Path: - from tools import memory_tool - return memory_tool.get_memory_dir() - - def _scan_memory_content(content: str) -> Optional[str]: """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]: +def _error(message: str, **extra) -> Dict[str, Any]: + return {"success": False, "error": message, **extra} + + +def _drift_error(path: Path, bak_path: str) -> Dict[str, Any]: """External drift: the file wouldn't round-trip, so flushing would discard content.""" - return {"success": False, "error": ( + return _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": ( + ), 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.")} + "memory(action=add, content=...), then remove or rewrite the original file to a clean state.")) -def _read_failed_error(path: "Path") -> Dict[str, Any]: +def _read_failed_error(path: Path) -> Dict[str, Any]: """Existing-but-unreadable file: saving from an assumed-empty view would wipe it.""" - return {"success": False, "error": ( + return _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]: @@ -78,19 +77,16 @@ class MemoryStore: memory_enabled: bool = True, user_profile_enabled: bool = True): self.memory_entries: List[str] = [] self.user_entries: List[str] = [] - self.memory_char_limit = memory_char_limit - self.user_char_limit = user_char_limit - self.memory_enabled = memory_enabled - self.user_profile_enabled = user_profile_enabled + self.memory_char_limit, self.user_char_limit = memory_char_limit, user_char_limit + self.memory_enabled, self.user_profile_enabled = memory_enabled, user_profile_enabled self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""} 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.""" return self.user_profile_enabled if target == "user" else self.memory_enabled def reset_consolidation_failures(self) -> None: - """Reset the per-turn consolidation-failure counter (call at turn start).""" + """Call at turn start.""" self._consolidation_failures = 0 def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]: @@ -109,22 +105,10 @@ class MemoryStore: 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 = { - "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]: - """*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 - def _one(entry): + def _sanitize(entry, filename): + # Strict scope, same as writes; empty / already-blocked entries pass through. findings = scan_for_threats(entry, scope="strict") if entry and not entry.startswith("[BLOCKED:") else None if not findings: return entry @@ -132,7 +116,13 @@ class MemoryStore: 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] + for target in ("memory", "user"): + path = self._path_for(target) + path.parent.mkdir(parents=True, exist_ok=True) + # Deduplicate (order-preserving, first occurrence wins). + entries = list(dict.fromkeys(self._read_file(path))) + self._set_entries(target, entries) + self._system_prompt_snapshot[target] = self._render_block(target, [_sanitize(e, path.name) for e in entries]) @staticmethod @contextmanager @@ -157,33 +147,13 @@ class MemoryStore: try: yield finally: - try: + with suppress(OSError): _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_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_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 _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.""" - _memory_dir().mkdir(parents=True, exist_ok=True) - self._write_file(self._path_for(target), self._entries_for(target)) + from tools import memory_tool # get_memory_dir is monkeypatched there + return memory_tool.get_memory_dir() / ("USER.md" if target == "user" else "MEMORY.md") def _entries_for(self, target: str) -> List[str]: return self.user_entries if target == "user" else self.memory_entries @@ -201,53 +171,45 @@ class MemoryStore: 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" + return f"{min(100, int((current / limit) * 100)) if limit > 0 else 0}% — {current:,}/{limit:,} chars" def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]: """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)}) - - def _locate(self, target: str, old_text: str, verb: str): - """Resolve *old_text* to a unique entry index, or an error dict.""" - entries = self._entries_for(target) - 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": [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}) - return idx, None + return self._consolidation_failure( + _error(message, current_entries=self._entries_for(target), usage=self._usage(target))) 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 + """Lock, re-read from disk, run ``mutate(entries, limit)`` -> ``(new_entries, message)`` + or an error dict, then persist and return the success response. The reload aborts + on an existing-but-unreadable file (even append-only ``add`` rewrites the whole + file) and, unless *skip_drift*, on external drift (flushing would discard + un-roundtrippable content). Drift check and parse use the SAME raw snapshot — + a failed second read used to count as "no drift".""" + path = self._path_for(target) + with self._file_lock(path): + raw, read_ok = self._read_raw_checked(path) + if not read_ok: + 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)))) + if bak: + return _drift_error(path, bak) 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) + self._set_entries(target, result[0]) + path.parent.mkdir(parents=True, exist_ok=True) + self._write_file(path, result[0]) + return self._success_response(target, result[1]) def add(self, target: str, content: str) -> Dict[str, Any]: """Append a new entry. Returns error if it would exceed the char limit.""" content = content.strip() if not content: - return {"success": False, "error": "Content cannot be empty."} - scan_error = _scan_memory_content(content) - if scan_error: - return {"success": False, "error": scan_error} + return _error("Content cannot be empty.") + if scan_error := _scan_memory_content(content): + return _error(scan_error) def _add(entries, limit): if content in entries: @@ -259,7 +221,6 @@ class MemoryStore: 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.")) 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) @@ -268,29 +229,33 @@ class MemoryStore: """Find entry containing old_text substring, replace it with new_content.""" new_content = new_content.strip() if not old_text.strip(): - return {"success": False, "error": "old_text cannot be empty."} + return _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 _error("new_content cannot be empty. Use 'remove' to delete entries.") + if scan_error := _scan_memory_content(new_content): + return _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 _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*.""" + """Locked replace (``new_content`` set) or remove (None) of the 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, ambiguous = _find_unique_match(entries, old_text) + if ambiguous: + return _error(f"Multiple entries matched '{old_text}'. Be more specific.", + matches=[e[:80] + ("..." if len(e) > 80 else "") for e in entries if old_text in e]) + if idx is None: + return self._consolidation_failure(_error( + f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text " + f"of the entry you want to {'replace' if new_content else 'remove'}.", current_entries=entries)) + replaced = entries[:idx] + ([] if new_content is None else [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:] + return replaced, "Entry removed." new_total = len(ENTRY_DELIMITER.join(replaced)) if new_total > limit: return self._failure_with_entries(target, ( @@ -298,7 +263,6 @@ class MemoryStore: f"or 'remove' other stale or less important entries to make room (see current_entries " f"below), then retry — all in this turn.")) return replaced, "Entry replaced." - return self._mutate(target, _apply) @staticmethod @@ -321,79 +285,67 @@ class MemoryStore: return f"{pos}: '{old_text}' matched multiple distinct entries -- be more specific." if idx is None: return f"{pos}: no entry matched '{old_text}'." - if act == "replace": - working[idx] = content - else: - working.pop(idx) + working[idx:idx + 1] = [content] if act == "replace" else [] return None def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]: - """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.""" + """Apply add/replace/remove ops 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."} - + return _error("operations list is empty.") 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}"} + return _error(f"Operation {i + 1}: {scan_error}") 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) + msg = self._apply_batch_op(working, act, (op.get("content") or op.get("new_text") or "").strip(), + (op.get("old_text") or "").strip(), f"Operation {i + 1} ({act or 'unknown'})") if 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)) + return self._failure_with_entries(target, msg + " No operations were applied (batch is all-or-nothing).") + new_total = len(ENTRY_DELIMITER.join(working)) # budget check against the FINAL state only 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 working, f"Applied {len(operations)} operation(s)." - return self._mutate(target, _apply) def format_for_system_prompt(self, target: str) -> Optional[str]: - """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.""" + """Frozen load-time snapshot (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 def _success_response(self, target: str, message: str = None) -> Dict[str, Any]: - # A successful write is progress: reset the per-turn (consecutive) failure budget. + """TERMINAL and WITHOUT the entries list: echoing entries invites the model to + "find more to fix" and re-issue the same ops. A successful write resets the + per-turn failure budget.""" self._consolidation_failures = 0 - # 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)), **({"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.""" + """System prompt block: header + usage indicator + entries ("" when empty).""" if not entries: return "" - content = ENTRY_DELIMITER.join(entries) + content, sep = ENTRY_DELIMITER.join(entries), "═" * 46 title = MEMORY_BLOCK_HEADERS["user" if target == "user" else "memory"] - separator = "═" * 46 - return f"{separator}\n{title} [{self._usage_pct(target, len(content))}]\n{separator}\n{content}" + return f"{sep}\n{title} [{self._usage_pct(target, len(content))}]\n{sep}\n{content}" @staticmethod def _read_raw_checked(path: Path) -> Tuple[str, bool]: - """``(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.""" + """``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't be + read. Decoding stays STRICT (``errors="replace"`` would hand callers a lossy view + a save then persists); ``utf-8-sig`` strips a Notepad BOM off the first entry.""" if not path.exists(): return "", True try: @@ -408,19 +360,26 @@ class MemoryStore: @staticmethod def _read_file(path: Path) -> List[str]: - """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.""" + """Entries of a memory file ([] on any error). Read-only callers only; mutation + paths use ``_read_raw_checked`` so they can refuse to overwrite an unreadable file.""" return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0]) + @staticmethod + def _write_file(path: Path, entries: List[str]): + """Atomic temp-file + rename: readers never see a truncated file. Also used by + agent/learning_mutations.py.""" + try: + atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_") + except OSError as e: + raise RuntimeError(f"Failed to write memory file {path}: {e}") + def _detect_external_drift(self, target: str, raw: str) -> Optional[str]: - """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 + """``.bak.`` snapshot 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 be — an external writer appended free-form text).""" parsed = self._parse_entries(raw) - if raw.strip() == ENTRY_DELIMITER.join(parsed) and max(map(len, parsed), default=0) <= self._char_limit(target): + if not raw.strip() or (raw.strip() == ENTRY_DELIMITER.join(parsed) + and max(map(len, parsed), default=0) <= self._char_limit(target)): return None path = self._path_for(target) bak_path = path.with_suffix(path.suffix + f".bak.{int(time.time())}") @@ -429,11 +388,3 @@ class MemoryStore: except OSError: return str(bak_path) + " (BACKUP FAILED — file unchanged on disk)" return str(bak_path) - - @staticmethod - def _write_file(path: Path, entries: List[str]): - """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: - raise RuntimeError(f"Failed to write memory file {path}: {e}") diff --git a/tools/microsoft_graph_auth.py b/tools/microsoft_graph_auth.py index 5c8edc1160..538597dcf4 100644 --- a/tools/microsoft_graph_auth.py +++ b/tools/microsoft_graph_auth.py @@ -53,13 +53,10 @@ class GraphCredentials: @property def token_url(self) -> str: - tenant = self.tenant_id.strip().strip("/") - return f"{self.authority_url.rstrip('/')}/{tenant}/oauth2/v2.0/token" + return f"{self.authority_url.rstrip('/')}/{self.tenant_id.strip().strip('/')}/oauth2/v2.0/token" @classmethod - def from_env( - cls, environ: dict[str, str] | None = None, *, required: bool = True - ) -> "GraphCredentials | None": + def from_env(cls, environ: dict[str, str] | None = None, *, required: bool = True) -> "GraphCredentials | None": 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] @@ -90,10 +87,9 @@ class CachedAccessToken: class MicrosoftGraphTokenProvider: """Acquire and cache Microsoft Graph app-only access tokens.""" - def __init__( - self, credentials: GraphCredentials, *, timeout: float = 20.0, - skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, transport: httpx.AsyncBaseTransport | None = None, - ) -> None: + def __init__(self, credentials: GraphCredentials, *, timeout: float = 20.0, + skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, + transport: httpx.AsyncBaseTransport | None = None) -> None: self.credentials, self.timeout, self.skew_seconds = credentials, timeout, max(0, int(skew_seconds)) self._transport = transport self._cached_token: CachedAccessToken | None = None @@ -150,10 +146,8 @@ class MicrosoftGraphTokenProvider: except (TypeError, ValueError) as exc: raise MicrosoftGraphTokenError( "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)) + return CachedAccessToken(access_token, time.time() + max(0, expires_in_seconds), + str(payload.get("token_type") or "Bearer").strip() or "Bearer") def _extract_error_detail(response: httpx.Response) -> str: diff --git a/tools/microsoft_graph_client.py b/tools/microsoft_graph_client.py index 9e3dd746c2..faf29d7acb 100644 --- a/tools/microsoft_graph_client.py +++ b/tools/microsoft_graph_client.py @@ -10,7 +10,7 @@ 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 MicrosoftGraphTokenProvider, format_graph_error DEFAULT_GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0" @@ -26,9 +26,8 @@ class MicrosoftGraphClientError(RuntimeError): class MicrosoftGraphAPIError(MicrosoftGraphClientError): """Raised when a Graph API request fails.""" - def __init__( - self, status_code: int, method: str, url: str, message: str, *, - retry_after_seconds: float | None = None, payload: Any = None) -> None: + 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, 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}") @@ -39,20 +38,15 @@ class MicrosoftGraphClient: 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: + 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, 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": - return cls(MicrosoftGraphTokenProvider(GraphCredentials.from_env()), **kwargs) - async def get_json(self, path: str, *, params: Params = None, headers: Headers = None) -> Any: return self._decode_json(await self._request("GET", path, params=params, headers=headers)) @@ -60,27 +54,24 @@ class MicrosoftGraphClient: return self._decode_json(await self._request("POST", path, json_body=json_body, headers=headers)) async def patch_json(self, path: str, *, json_body: Any | None = None, headers: Headers = None) -> Any: + """Decoded body, or ``{}`` for a 204 / bodiless response.""" response = await self._request("PATCH", path, json_body=json_body, headers=headers) - return self._decode_json_or(response, {}) + return self._decode_json(response) if response.status_code != 204 and response.content else {} async def delete(self, path: str, *, headers: Headers = None) -> dict[str, Any]: + """Decoded body, or ``{"deleted": True, "status_code"}`` for a 204 / bodiless response.""" response = await self._request("DELETE", path, headers=headers) - return self._decode_json_or(response, {"deleted": True, "status_code": response.status_code}) + if response.status_code != 204 and response.content: + return self._decode_json(response) + return {"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]: + 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 {}) + next_url, next_params = self._resolve_url(path), dict(params or {}) while next_url: - response = await self._request("GET", next_url, params=next_params or None, headers=headers) - payload = self._decode_json(response) + payload = self._decode_json(await self._request("GET", next_url, params=next_params or None, headers=headers)) if not isinstance(payload, dict): raise MicrosoftGraphClientError( f"Expected paginated Graph response dict, got {type(payload).__name__}.") @@ -89,13 +80,11 @@ class MicrosoftGraphClient: 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]: + async def download_to_file(self, path: str, destination: str | Path, *, headers: Headers = None, + chunk_size: int = 65536) -> dict[str, Any]: """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) + url, target = self._resolve_url(path), Path(destination) target.parent.mkdir(parents=True, exist_ok=True) tmp_target = target.with_suffix(target.suffix + ".part") @@ -103,8 +92,7 @@ class MicrosoftGraphClient: try: async with client.stream("GET", url, headers=request_headers) as response: if response.status_code >= 400: - # Materialize the (small) error body so the message is meaningful. - await response.aread() + await response.aread() # small error body -> meaningful message return response, None with tmp_target.open("wb") as handle: async for chunk in response.aiter_bytes(chunk_size=chunk_size): @@ -119,9 +107,8 @@ class MicrosoftGraphClient: os.replace(tmp_target, target) return {"path": str(target), "size_bytes": target.stat().st_size, "content_type": content_type} - async def _request( - self, method: str, path_or_url: str, *, - params: Params = None, json_body: Any | None = None, headers: Headers = None) -> httpx.Response: + async def _request(self, method: str, path_or_url: str, *, 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]): @@ -137,47 +124,37 @@ class MicrosoftGraphClient: """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 - - while attempt <= self.max_retries: + for attempt in range(self.max_retries + 1): token = await self.token_provider.get_access_token( - force_refresh=attempt > 0 - and isinstance(last_error, MicrosoftGraphAPIError) - 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" - if headers: - request_headers.update(headers) - + force_refresh=isinstance(last_error, MicrosoftGraphAPIError) and last_error.status_code == 401) + request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent, + **({"Content-Type": "application/json"} if json_body is not None else {}), + **(headers or {})} + exhausted = attempt >= self.max_retries try: 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, response = exc, None - if attempt >= self.max_retries: + if exhausted: raise MicrosoftGraphClientError( 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): + last_error, status = self._build_api_error(method, url, response), response.status_code + if exhausted 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 - raise MicrosoftGraphClientError(f"Microsoft Graph {kind} exhausted retries for {method} {url}.") def _resolve_url(self, path_or_url: str) -> str: if path_or_url.startswith(("http://", "https://")): return path_or_url - path = path_or_url if path_or_url.startswith("/") else f"/{path_or_url}" - return f"{self.base_url}{path}" + return f"{self.base_url}{path_or_url if path_or_url.startswith('/') else '/' + path_or_url}" @staticmethod def _decode_json(response: httpx.Response) -> Any: @@ -201,6 +178,5 @@ class MicrosoftGraphClient: payload = None detail = format_graph_error(payload.get("error")) if isinstance(payload, dict) else None return MicrosoftGraphAPIError( - response.status_code, method, url, - detail if detail is not None else (response.text.strip() or "unknown error"), + response.status_code, method, url, response.text.strip() or "unknown error" if detail is None else detail, retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload) diff --git a/tools/process_registry_notifications.py b/tools/process_registry_notifications.py index b470be43a3..0f70f371fd 100644 --- a/tools/process_registry_notifications.py +++ b/tools/process_registry_notifications.py @@ -1,14 +1,13 @@ -"""Human-readable rendering of background-process notification events. - -Events come off ``ProcessRegistry.completion_queue`` (completion, watch_match, -watch_disabled, watch_overflow_*, async_delegation) and are turned into the -``[IMPORTANT: ...]`` / ``[ASYNC DELEGATION ...]`` text injected into the agent -conversation by the CLI drain loop, the gateway, and the TUI. -""" +"""Human-readable rendering of ``ProcessRegistry.completion_queue`` events (completion, +watch_match, watch_disabled, watch_overflow_*, async_delegation) into the +``[IMPORTANT: ...]`` / ``[ASYNC DELEGATION ...]`` text the CLI drain loop, gateway and +TUI inject into the agent conversation.""" import time from contextlib import suppress +_DONE = ("completed", "success") + def _format_age(seconds: float) -> str: """Human-friendly elapsed string ('18m', '2h3m', '45s').""" @@ -26,34 +25,28 @@ def _format_age(seconds: float) -> str: def _model_not_found_patterns() -> "list[str]": - """Model-not-found phrases from ``agent.error_classifier`` (same classification - the failover path uses, no hand-copied list to drift); a minimal built-in set - if the import fails so per-task blocks are never hidden.""" + """Model-not-found phrases from ``agent.error_classifier`` (the failover path's + own list, so nothing drifts); a minimal built-in set if the import fails.""" try: from agent.error_classifier import _MODEL_NOT_FOUND_PATTERNS - return list(_MODEL_NOT_FOUND_PATTERNS) except Exception: return ["is not a valid model", "model not found", "model_not_found"] def _delegation_config() -> dict: - """Active delegation config (model/provider/fallbacks); ``{}`` on any error. - Lazy ``tools.delegate_tool._load_config`` so the renderer sees the dispatcher's - model/provider without importing the heavy delegation module at import time.""" + """Active delegation config; ``{}`` on any error. Lazy: delegate_tool is heavy.""" try: from tools.delegate_tool import _load_config as _cfg - return _cfg() or {} except Exception: return {} def _delegation_model_not_found(results, config) -> bool: - """True when a result reflects a config-level model_not_found rejection. - Requires both a model-not-found phrase AND the currently-configured model - name in the same error/summary text, so a stale task failing on a - different (removed) model is not mis-attributed to the config.""" + """True when a result reflects a config-level model_not_found rejection: needs a + model-not-found phrase AND the currently-configured model name in the same text, + so a stale task failing on a removed model is not mis-attributed to the config.""" model = str((config or {}).get("model") or "").lower() if not model: return False @@ -74,11 +67,9 @@ def _delegation_model_not_found_notice(results) -> "list[str] | None": f'"{model}" was rejected by provider "{provider}" ' "(HTTP 400: not a valid model ID).", "Every task in this batch failed for this reason before doing any work.", - "Check Settings → Advanced → Subagent Model (or: hermes config get delegation.model).", - ] + "Check Settings → Advanced → Subagent Model (or: hermes config get delegation.model)."] with suppress(Exception): from hermes_cli.fallback_config import get_fallback_chain - if not get_fallback_chain(config): lines.append("No fallback chain is configured, so no failover was attempted.") return lines @@ -94,71 +85,59 @@ def _is_truncated(entry: dict) -> bool: return bool(entry.get("truncated") or entry.get("exit_reason") == "max_iterations") -def _header_lines(evt: dict, title: str, intro: str, completed_at: float) -> "list[str]": - """Shared preamble: title, intro, blank, dispatch time and task-source lines.""" +def _notice_lines(results) -> "list[str]": + """Blank + model_not_found notice block, or [] when the notice does not apply.""" + notice = _delegation_model_not_found_notice(results) + return ["", *notice] if notice else [] + + +def _preamble(evt: dict, title: str, intro: str, completed_at: float, *, with_goal: bool) -> "list[str]": + """Shared preamble: title, intro, blank, dispatch time, [goal], context/toolsets, role+model.""" lines = [title, intro, ""] dispatched_at = evt.get("dispatched_at") if isinstance(dispatched_at, (int, float)): ts = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(dispatched_at)) lines.append(f"Dispatched: {ts} ({_format_age(completed_at - dispatched_at)} ago)") - return lines - - -def _task_source_lines(evt: dict) -> "list[str]": - lines = [] + if with_goal: + lines.append(f"Original goal: {evt.get('goal', '') or ''}") if evt.get("context"): lines.append(f"Context you provided: {evt['context']}") if evt.get("toolsets"): lines.append(f"Toolsets: {', '.join(evt['toolsets'])}") + lines.append(f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}") return lines -def _role_model(evt: dict) -> str: - return f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}" - - def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> str: """Consolidated block for a delegate_task fan-out that finished as one unit.""" - results = evt.get("results") or [] - goals = evt.get("goals") or [] + results, goals = evt.get("results") or [], evt.get("goals") or [] n = len(results) if results else len(goals) - total_dur = evt.get("total_duration_seconds", evt.get("duration_seconds", "?")) - error = evt.get("error") - lines = _header_lines( + lines = _preamble( evt, f"[ASYNC DELEGATION BATCH COMPLETE — {deleg_id}]", f"A background fan-out of {n} subagent(s) you dispatched earlier " "has finished. All ran in parallel and waited on each other; their " "consolidated results are below. You may have moved on since " "dispatching — act on these or re-dispatch if things have changed.", - completed_at) - lines.extend(_task_source_lines(evt)) - lines.append(f"{_role_model(evt)} Total duration: {total_dur}s") - if error and not results: - lines += ["--- ERROR ---", f"The batch did not complete successfully: {error}"] + completed_at, with_goal=False) + lines[-1] += f" Total duration: {evt.get('total_duration_seconds', evt.get('duration_seconds', '?'))}s" + if evt.get("error") and not results: + lines += ["--- ERROR ---", f"The batch did not complete successfully: {evt['error']}"] return "\n".join(lines) # Config-level rejection notice BEFORE the per-task wall — a rejected # delegation model fails every task identically and must not stay buried. - _notice = _delegation_model_not_found_notice(results) - if _notice: - lines += ["", *_notice] + lines += _notice_lines(results) for r in sorted(results, key=lambda x: x.get("task_index", 0)): - idx = r.get("task_index", 0) - r_status = r.get("status", "?") - r_summary = r.get("summary") - r_error = r.get("error") + idx, r_truncated = r.get("task_index", 0), _is_truncated(r) + r_status, r_summary, r_error = r.get("status", "?"), r.get("summary"), r.get("error") r_goal = goals[idx] if idx < len(goals) else r.get("goal", "") - r_truncated = _is_truncated(r) - icon = "⚠" if r_truncated else ("✓" if r_status in ("completed", "success") else "✗") - header = f"--- {icon} TASK {idx + 1}/{n}" + (f": {r_goal}" if r_goal else "") + f" (status={r_status}" - if r.get("api_calls"): - header += f", api_calls={r['api_calls']}" - if r.get("duration_seconds") is not None: - header += f", {r['duration_seconds']}s" - if r_truncated: - header += ", TRUNCATED: hit max_iterations — work may be incomplete" + icon = "⚠" if r_truncated else ("✓" if r_status in _DONE else "✗") + header = (f"--- {icon} TASK {idx + 1}/{n}" + (f": {r_goal}" if r_goal else "") + f" (status={r_status}" + + (f", api_calls={r['api_calls']}" if r.get("api_calls") else "") + + (f", {r['duration_seconds']}s" if r.get("duration_seconds") is not None else "") + + (", TRUNCATED: hit max_iterations — work may be incomplete" if r_truncated else "")) lines += ["", header + ") ---"] - if r_status in ("completed", "success") and r_summary: + if r_status in _DONE and r_summary: if r_truncated: lines.append(_TRUNCATED_SUMMARY_NOTE) lines.append(r_summary) @@ -174,39 +153,27 @@ def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> s def _format_async_delegation(evt: dict) -> str: - """Format an async-delegation completion into a self-contained re-injection. - Carries the FULL original task source (goal, context, toolsets, role, model) plus - dispatch time, status, and the complete result summary: when this re-enters the - conversation the agent may be deep in unrelated context and must be able to use - the result OR re-dispatch without remembering why the subagent existed.""" + """Self-contained re-injection for an async-delegation completion: the FULL + original task source (goal, context, toolsets, role, model), dispatch time, status + and result, so an agent deep in unrelated context can act on it or re-dispatch.""" deleg_id = evt.get("delegation_id", "unknown") completed_at = evt.get("completed_at") or time.time() if evt.get("is_batch") or isinstance(evt.get("results"), list): return _format_batch_delegation(evt, deleg_id, completed_at) - status = evt.get("status") or "completed" - summary = evt.get("summary") - error = evt.get("error") + status, summary, error = evt.get("status") or "completed", evt.get("summary"), evt.get("error") truncated = _is_truncated(evt) - lines = _header_lines( + lines = _preamble( evt, f"[ASYNC DELEGATION COMPLETE — {deleg_id}]", "A background subagent you dispatched earlier has finished. You may " "have moved on since dispatching it; the full task source is below so " "you can act on the result or re-dispatch if things have changed.", - completed_at) - lines.append(f"Original goal: {evt.get('goal', '') or ''}") - lines.extend(_task_source_lines(evt)) - lines.append(_role_model(evt)) - _notice = _delegation_model_not_found_notice([evt]) - if _notice: - lines += ["", *_notice] - _trunc = " [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else "" - lines += [ - f"Status: {status} API calls: {evt.get('api_calls', 0)} " - f"Duration: {evt.get('duration_seconds', '?')}s{_trunc}", - "--- RESULT ---", - ] - if status in ("completed", "success") and summary: + completed_at, with_goal=True) + lines += _notice_lines([evt]) + [ + f"Status: {status} API calls: {evt.get('api_calls', 0)} Duration: {evt.get('duration_seconds', '?')}s" + + (" [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else ""), + "--- RESULT ---"] + if status in _DONE and summary: if truncated: lines.append(_TRUNCATED_SUMMARY_NOTE) lines.append(summary) @@ -215,36 +182,30 @@ def _format_async_delegation(evt: dict) -> str: lines.append("The subagent was interrupted before completing" + (f": {error}" if error else ".")) else: # error / timeout / failed lines.append( - f"The subagent did not complete successfully (status={status})." + (f"\n{error}" if error else "") - ) + f"The subagent did not complete successfully (status={status})." + (f"\n{error}" if error else "")) if summary: lines += ["Partial output:", summary] return "\n".join(lines) def _delegation_attribution_line(evt: dict) -> "str | None": - """One-line provenance for a subagent-owned process event, else None. - A background process a subagent started outlives the child and is routed to - the PARENT conversation, which otherwise sees an anonymous raw output wall. - Judged on ``owner_task_id`` (the raw spawning id) — ``task_id`` is the - container key and may be collapsed to the session key.""" + """One-line provenance for a subagent-owned process event, else None. Such a process + outlives the child and lands in the PARENT conversation, which would otherwise see an + anonymous output wall. Keyed on ``owner_task_id`` — ``task_id`` may be the session key.""" task_id = str(evt.get("owner_task_id") or evt.get("task_id") or "") if not task_id.startswith("sa-"): return None info = None with suppress(Exception): from tools.delegate_tool import get_subagent_attribution - info = get_subagent_attribution(task_id) if not info: # Registry entry aged out — still attribute generically, not anonymously. return f"Started by subagent {task_id} (delegate_task)." - goal = str(info.get("goal") or "").strip() - if len(goal) > 120: - goal = goal[:117] + "..." - deleg = info.get("delegation_id") - line = f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "." - return line + (f' Task: "{goal}"' if goal else "") + goal, deleg = str(info.get("goal") or "").strip(), info.get("delegation_id") + goal = goal[:117] + "..." if len(goal) > 120 else goal + return (f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "." + + (f' Task: "{goal}"' if goal else "")) _REASON_STATUS = {"lost": "marked lost because the process backend disappeared", "failed_start": "failed to start"} @@ -254,35 +215,28 @@ def _completion_status(evt: dict) -> str: reason = evt.get("completion_reason") or "exited" if reason == "killed": return f"terminated by {evt.get('termination_source') or 'Hermes'}" - if reason in _REASON_STATUS: - return _REASON_STATUS[reason] - return "completed normally" if evt.get("exit_code", "?") == 0 else "exited" + return _REASON_STATUS.get(reason) or ("completed normally" if evt.get("exit_code", "?") == 0 else "exited") def format_process_notification(evt: dict) -> "str | None": """Format a completion_queue event into an ``[IMPORTANT: ...]`` message.""" evt_type = evt.get("type", "completion") - _sid = evt.get("session_id", "unknown") - _cmd = evt.get("command", "unknown") - _attribution = _delegation_attribution_line(evt) - # watch_disabled and overflow events carry their own human-readable `message`; # otherwise overflow events would fall through to the completion formatter as a # phantom "process exited (exit code ?)". if evt_type in ("watch_disabled", "watch_overflow_tripped", "watch_overflow_released"): return f"[IMPORTANT: {evt.get('message', '')}]" + if evt_type == "async_delegation": + return _format_async_delegation(evt) + _sid, _cmd = evt.get("session_id", "unknown"), evt.get("command", "unknown") + _attribution = _delegation_attribution_line(evt) attribution = f"{_attribution}\n" if _attribution else "" if evt_type == "watch_match": _sup = evt.get("suppressed", 0) - text = ( + return ( f"[IMPORTANT: Background process {_sid} matched watch pattern \"{evt.get('pattern', '?')}\".\n" - f"{attribution}Command: {_cmd}\nMatched output:\n{evt.get('output', '')}") - if _sup: - text += f"\n({_sup} earlier matches were suppressed by rate limit)" - return text + "]" - if evt_type == "async_delegation": - return _format_async_delegation(evt) - + f"{attribution}Command: {_cmd}\nMatched output:\n{evt.get('output', '')}" + + (f"\n({_sup} earlier matches were suppressed by rate limit)" if _sup else "") + "]") _exit = evt.get("exit_code", "?") _out = evt.get("output", "") _signal = ", SIGTERM" if _exit in {-15, 143, "-15", "143"} else "" diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index 5baeeda1ed..03ab134f36 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -15,33 +15,27 @@ 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 from browsing/searching — integrations (HERMES_SESSION_SOURCE=tool), delegate +# subagent runs, kanban workers are 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. +# /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.""" + """``fn()``, or *default* after debug-logging *msg* (+ the exception when *with_exc*).""" try: return fn() except Exception as e: @@ -50,8 +44,8 @@ def _quiet(fn, default, msg, *log_args, with_exc: bool = False): 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.""" + """``(fn(), None)``, 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: @@ -60,46 +54,43 @@ def _loud(fn, log_msg, error_prefix, *log_args): 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.""" + """Unix timestamp -> readable date; ISO strings pass through; "unknown" for None.""" 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) + return _quiet(lambda: datetime.fromtimestamp(float(ts)).strftime("%B %d, %Y at %I:%M %p"), str(ts), + "Failed to format timestamp %s: %s", ts, with_exc=True) + + +def _get_session_meta(db, session_id: str) -> dict: + """``db.get_session`` that degrades to ``{}`` on error.""" + return _quiet(lambda: db.get_session(session_id), None, + "get_session failed for %s: %s", session_id, with_exc=True) or {} 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)``.""" + """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).""" 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 + s = _get_session_meta(db, cur) has_compression = has_compression or s.get("end_reason") == "compression" if not s.get("parent_session_id"): break @@ -108,31 +99,31 @@ def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]: 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 _same_lineage(db, a: str, b: str) -> bool: + a_root = _resolve_lineage(db, a) + return bool(a_root and a_root == _resolve_lineage(db, b)) + + 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 + end_reason = (session_id and _get_session_meta(db, session_id).get("end_reason")) or 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 - + """Owning session and visibility flags for *message_id* (None if missing/error).""" 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 + row = message_id and _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id) + return dict(row) if row else None def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool: @@ -142,36 +133,14 @@ def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool: 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.""" + """True for a compaction-archived row: no longer in live context, so 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.""" + """Slim a message row; keeps ``content`` even when empty (tool-call-only turns).""" content = m.get("content") if isinstance(content, str) and "\x1b" in content: # archived terminal output carries ANSI from tools.ansi_strip import strip_ansi @@ -181,7 +150,8 @@ def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None, 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)) + 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"} @@ -189,23 +159,29 @@ 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") + def _active(): + from hermes_cli.profiles import get_active_profile_name + resolved = get_active_profile_name() + return "" if resolved == "custom" else resolved + name = (profile or "").strip() or _quiet(_active, "", "get_active_profile_name failed for session link") return f"@session:{name}/{session_id}" if name else f"@session:{session_id}" +def _discovery_entry(lineage_root: Optional[str], **fields) -> Dict[str, Any]: + """Canonical key order; ``parent_session_id`` set when the hit lives in a child.""" + 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 _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.""" + """Discovery-shaped result when the query matches a session title, else None.""" 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) + session_id = title_query and _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) @@ -224,56 +200,28 @@ def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> 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( + def shape(key, fallback, anchor=None): + return [_shape_message(m, anchor_id=anchor) for m in (view.get(key) or fallback)] + return {**_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 + bookend_start=shape("bookend_start", messages[:3]), messages=shape("window", messages[:5], anchor_id), + bookend_end=shape("bookend_end", 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} 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}) + """Discovery response; notes FTS backfill progress so the agent can explain thin + results instead of treating them as ground truth.""" + status = _quiet(db.fts_rebuild_status, None, "fts_rebuild_status failed") + rebuild = {} if status is None else {"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.")}} + return _ok(mode="discover", query=query, detail=detail, results=results, count=len(results), **extra, **rebuild) def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]: @@ -282,18 +230,14 @@ 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]]: - """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") + """Discovery result from a surviving FTS row; None (dropped) if the view can't load.""" + hit_sid, msg_id = match_info.get("session_id") or lineage_root, 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] + session_meta, full = _get_session_meta(db, lineage_root), result_detail == "full" return _discovery_entry( lineage_root, session_id=hit_sid, when=_format_timestamp(session_meta.get("started_at") or match_info.get("session_started")), @@ -302,7 +246,8 @@ def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detai 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], + messages=[_shape_message(m, anchor_id=msg_id, max_content_len=4000) + for m in (view.get("window") or []) if full or m.get("id") == msg_id], 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) @@ -319,23 +264,35 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, 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) + # 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`; stable sort keeps BM25 + # order within each class. + raw_results = sorted(raw_results, key=lambda r: (r.get("source") or "") in _DEMOTED_SESSION_SOURCES) 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) + results = [title_result] if title_result else [] + if title_result and (title_lineage := title_result.pop("_lineage_root", None)): + seen_sessions[title_lineage] = {"_title_only": True} + # Dedupe by lineage (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, resolved_sid = r["session_id"], _resolve_lineage(db, r["session_id"]) + is_compacted_hit = _is_compacted_message(db, r.get("id")) + if current_lineage_root and resolved_sid == current_lineage_root 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}) for lineage_root, match_info in seen_sessions.items(): if match_info.get("_title_only"): continue @@ -354,8 +311,7 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort def _resolve_profile_db(profile: str): - """Another profile's ``state.db`` opened read-only (safe on a live DB), or None - for the current profile.""" + """Another profile's ``state.db`` opened read-only (safe on a live DB); None = current.""" if profile is None or not str(profile).strip(): return None from hermes_cli import profiles as profiles_mod @@ -368,17 +324,17 @@ 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`` -> ``(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 except Exception: return None, None - targets = [("default", profiles_mod.get_profile_dir("default"))] - targets += _quiet(lambda: [(info.name, info.path) for info in profiles_mod.list_profiles()], - [], "list_profiles failed during session locate") + targets = [("default", profiles_mod.get_profile_dir("default"))] + _quiet( + lambda: [(info.name, info.path) for info in profiles_mod.list_profiles()], [], + "list_profiles failed during session locate") seen: set = set() for name, home in targets: db_path = Path(home) / "state.db" @@ -386,20 +342,13 @@ def _locate_session_db(session_id: str): continue seen.add(str(db_path)) 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): + if pdb and _get_session_meta(pdb, session_id): return pdb, name if pdb: pdb.close() return None, None -def _get_session_meta(db, session_id: str) -> dict: - """``db.get_session`` that degrades to ``{}`` on error.""" - return _quiet(lambda: db.get_session(session_id), None, - "get_session failed for %s: %s", session_id, with_exc=True) or {} - - def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_profile: str = None) -> str: """Read shape: whole session, or ``head`` + ``tail`` messages with a scroll pointer.""" meta = _get_session_meta(db, session_id) @@ -410,13 +359,26 @@ def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_prof if err: return err shaped = [_shape_message(m) for m in rows] - total = len(shaped) - truncated = total > head + tail - extra = {"message": (f"Session has {total} messages; showing first {head} + last {tail}. " - "Pass around_message_id (any id above) to scroll the middle.")} if truncated else {} + total, truncated = len(shaped), len(shaped) > head + tail return _ok(mode="read", session_id=session_id, link=_session_link(session_id, link_profile), session_meta=_session_meta_block(meta), message_count=total, truncated=truncated, - messages=shaped[:head] + shaped[-tail:] if truncated else shaped, **extra) + messages=shaped[:head] + shaped[-tail:] if truncated else shaped, + **({"message": (f"Session has {total} messages; showing first {head} + last {tail}. " + "Pass around_message_id (any id above) to scroll the middle.")} if truncated else {})) + + +def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str: + """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) + located, owner = (None, None) if json.loads(result).get("success") else _locate_session_db(sid) + if located is None: + return result + try: + found = json.loads(_read_session(located, sid, link_profile=owner)) + finally: + located.close() + return json.dumps({**found, "profile": owner}, ensure_ascii=False) if found.get("success") else result def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_profile: str = None) -> str: @@ -430,19 +392,14 @@ def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_p 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)) - results = [] - for s in sessions: - sid = s.get("id", "") - # 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, - **{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 + # Compression continuation: the root was summarised into the live child, so hide + # it. /new-reset children carry no transcript — keep that root browsable. + hidden = {current_session_id, current_root if has_compression_hop and current_root else None} + results = [{ + "session_id": s.get("id", ""), "link": _session_link(s.get("id", ""), link_profile), + "title": s.get("title") or None, **{k: s.get(k, "") for k in ("source", "started_at", "last_active")}, + "message_count": s.get("message_count", 0), "preview": s.get("preview", "")} + for s in [x for x in sessions if x.get("id", "") not in hidden][:limit]] 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.")) @@ -459,41 +416,19 @@ def _clamp_int(value, default: int, lo: int, hi: int) -> int: return max(lo, min(value, hi)) -def _anchor_in_live_context(db, anchor_state, anchor_session_id: str, current_session_id: str) -> bool: +def _anchor_in_live_context(db, anchor_state, anchor_sid: str, current_session_id: str) -> bool: """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): + parents, /new-reset predecessors) passes, so scroll never rejects a discovery result. + Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history.""" + if not _same_lineage(db, anchor_sid, 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. - 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 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" + return (anchor_state is not None and anchor_state["active"] == 0) or not _session_left_live_context(db, anchor_sid) def _scroll(db, session_id: str, around_message_id: int, window: int = 5, current_session_id: str = None) -> str: - """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() + """Scroll shape: a window centered on an anchor (no FTS5, no bookends).""" try: around_message_id = int(around_message_id) except (TypeError, ValueError): @@ -501,9 +436,8 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, 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): + owning = (anchor_state or {}).get("session_id") + if current_session_id and _anchor_in_live_context(db, anchor_state, owning 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: @@ -513,13 +447,18 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, 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) - if rebind_view is not None: - view, messages = rebind_view, rebind_view["window"] - session_meta = _get_session_meta(db, owning_session_id) or session_meta - session_id = owning_session_id + extra = {} + if not messages and owning and owning != session_id: + # Lineage rebind: the caller paired a parent session_id with a message id + # living in a descendant — serve the owner's window transparently. + 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 rebind_view and rebind_view.get("window"): + extra["warning"] = (f"around_message_id {around_message_id} lives in {owning} " + f"(child of {session_id}); rebound transparently") + view, messages, session_id = rebind_view, rebind_view["window"], owning + session_meta = _get_session_meta(db, owning) or session_meta if not messages: return tool_error(f"around_message_id {around_message_id} not in session_id {session_id}", success=False) return _ok( @@ -530,27 +469,7 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, hint=("Scroll forward: re-call with around_message_id = the LAST message's " "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 {})) - - -def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str: - """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 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) + "means you've hit that end of the session."), **extra) def _dispatch(query, role_filter, limit, db, current_session_id, session_id, @@ -576,39 +495,34 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_id, 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 _scroll(db, session_id.strip(), 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 - 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) + db=db, query=query.strip(), limit=limit, sort=sort_norm if sort_norm in ("newest", "oldest") else None, + role_filter=([r.strip() for r in role_filter.split(",") if r.strip()] or None) if isinstance(role_filter, str) else None, + detail="full" if isinstance(detail, str) and detail.strip().lower() == "full" else "adaptive", + 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, closing DBs opened here. Positional order is frozen for old callers.""" + from hermes_state import format_session_db_unavailable, get_shared_session_db, release_or_close 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 + db = _quiet(get_shared_session_db, None, "SessionDB unavailable for session_search") + if db is None: return tool_error(format_session_db_unavailable(), success=False) + owned_dbs.append(db) 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") @@ -728,8 +642,7 @@ SESSION_SEARCH_SCHEMA = { } -# --- Registry --- -from tools.registry import registry, tool_error +from tools.registry import registry, tool_error # noqa: E402 (registration at import time) registry.register( name="session_search",