From 0c9b0256d4923d8bacd27e254568bf123a6936f1 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:14:49 -0700 Subject: [PATCH 1/3] =?UTF-8?q?refactor(tools):=20memory=20store=20?= =?UTF-8?q?=E2=80=94=20inline=20reload/read=20helpers,=20dict-dispatch=20w?= =?UTF-8?q?rite-gate=20text?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/tools/test_memory_tool.py | 4 +- tools/memory_tool.py | 67 +++++++--------- tools/memory_tool_store.py | 132 ++++++++++++-------------------- 3 files changed, 78 insertions(+), 125 deletions(-) diff --git a/tests/tools/test_memory_tool.py b/tests/tools/test_memory_tool.py index 0582703abb..2af6107bfc 100644 --- a/tests/tools/test_memory_tool.py +++ b/tests/tools/test_memory_tool.py @@ -686,9 +686,7 @@ class TestBomToleranceInMemoryFiles: raw, read_ok = MemoryStore._read_raw_checked(path) assert read_ok is True assert not raw.startswith("\ufeff") - entries, ok = MemoryStore._read_entries_checked(path) - assert ok is True - assert entries == ["First fact."] + assert MemoryStore._read_file(path) == ["First fact."] def test_bom_file_add_keeps_existing_entry_intact(self, store): path = store._path_for("memory") diff --git a/tools/memory_tool.py b/tools/memory_tool.py index b4e0c7bd98..f3d5ecc037 100644 --- a/tools/memory_tool.py +++ b/tools/memory_tool.py @@ -46,7 +46,7 @@ def get_memory_dir() -> Path: from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports) - ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, _READ_FAILED, + ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, _drift_error, _read_failed_error, _scan_memory_content, ) @@ -75,13 +75,7 @@ def load_on_disk_store() -> "MemoryStore": return store -# --------------------------------------------------------------------------- -# Write-approval gate -# --------------------------------------------------------------------------- - -def _target_label(target: str) -> str: - return "user profile" if target == "user" else "memory" - +# -- Write-approval gate -- def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]: """Run the memory write gate. Returns a JSON tool-result string when the @@ -101,24 +95,26 @@ def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Option ensure_ascii=False) +# action -> (gate summary verb, gate detail) for a single mutating op. +_GATE_TEXT = { + "add": lambda label, content, old_text: (f"add to {label}", content or ""), + "replace": lambda label, content, old_text: (f"replace in {label}", f"old: {old_text}\nnew: {content}"), + "remove": lambda label, content, old_text: (f"remove from {label}", old_text or ""), +} + + def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str]) -> Optional[str]: """Gate a single mutating op (add/replace/remove); other actions pass.""" - if action not in _STORE_ACTIONS: + if action not in _GATE_TEXT: return None - label = _target_label(target) - if action == "add": - summary, detail = f"add to {label}", content or "" - elif action == "replace": - summary, detail = f"replace in {label}", f"old: {old_text}\nnew: {content}" - else: - summary, detail = f"remove from {label}", old_text or "" + summary, detail = _GATE_TEXT[action]("user profile" if target == "user" else "memory", content, old_text) payload = {"action": action, "target": target, "content": content, "old_text": old_text} return _gate_or_stage(summary, detail, payload) def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Optional[str]: """Gate a whole batch as a single unit.""" - summary = f"apply {len(operations)} op(s) to {_target_label(target)}" + summary = f"apply {len(operations)} op(s) to {'user profile' if target == 'user' else 'memory'}" detail_lines = [] for op in operations: op = op or {} @@ -134,32 +130,25 @@ def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Op return _gate_or_stage(summary, "\n".join(detail_lines), payload) -# --------------------------------------------------------------------------- -# Tool entry point -# --------------------------------------------------------------------------- - -def _missing_old_text_error(store: "MemoryStore", target: str, action: str) -> str: - """Recoverable error for replace/remove without ``old_text``. It can't be - schema-required (needs a combinator the Codex backend rejects — see - test_memory_tool_schema.py) and some clients omit it, so return the current - inventory plus a retry instruction instead of a dead-end.""" - return json.dumps({ - "success": False, - "error": (f"'{action}' needs old_text -- a short unique substring of the entry " - f"to {action}. None was provided. Reissue the {action} with old_text " - f"set to part of one of the current_entries below."), - "current_entries": store._entries_for(target), - "usage": store._usage(target), - }, ensure_ascii=False) - +# -- Tool entry point -- def _validate_single_op(store, action, target, content, old_text) -> Optional[str]: """Validate required params BEFORE the gate so an invalid write is rejected - now rather than staged and failing at approve time.""" + now rather than staged and failing at approve time. A missing ``old_text`` + is recoverable (it can't be schema-required — needs a combinator the Codex + backend rejects, see test_memory_tool_schema.py — and some clients omit it), + so return the current inventory plus a retry instruction, not a dead-end.""" if action == "add" and not content: return tool_error("Content is required for 'add' action.", success=False) if action in ("replace", "remove") and not old_text: - return _missing_old_text_error(store, target, action) + return json.dumps({ + "success": False, + "error": (f"'{action}' needs old_text -- a short unique substring of the entry " + f"to {action}. None was provided. Reissue the {action} with old_text " + f"set to part of one of the current_entries below."), + "current_entries": store._entries_for(target), + "usage": store._usage(target), + }, ensure_ascii=False) if action == "replace" and not content: return tool_error("content is required for 'replace' action.", success=False) return None @@ -282,9 +271,7 @@ def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[ return run(store, target, payload.get("content") or "", payload.get("old_text") or "") -# ============================================================================= -# OpenAI Function-Calling Schema -# ============================================================================= +# -- OpenAI Function-Calling Schema -- MEMORY_SCHEMA = { "name": "memory", diff --git a/tools/memory_tool_store.py b/tools/memory_tool_store.py index 2c9005cf71..b5cffcdf0a 100644 --- a/tools/memory_tool_store.py +++ b/tools/memory_tool_store.py @@ -24,12 +24,6 @@ MEMORY_BLOCK_HEADERS = { ENTRY_DELIMITER = "\n§\n" -# Sentinel from ``_reload_target``: the file EXISTS but could not be read. -# Distinct from a drift-backup path (str) and a clean reload (None); the caller -# must abort rather than persist over an unreadable file. -_READ_FAILED = object() - - def _memory_dir() -> Path: from tools import memory_tool @@ -37,16 +31,15 @@ def _memory_dir() -> Path: def _scan_memory_content(content: str) -> Optional[str]: - """Scan memory content for injection/exfil patterns ("strict" scope: memory - enters the system prompt as a frozen snapshot, so a poisoned entry persists - across sessions). Returns the error string if blocked.""" + """Scan for injection/exfil patterns ("strict" scope: memory enters the system + prompt as a frozen snapshot, so a poisoned entry persists across sessions). + Returns the error string if blocked.""" return _first_threat_message(content, scope="strict") def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]: - """Error dict for external drift: the on-disk file wouldn't round-trip - through the parser/serializer, so flushing would discard content added by a - patch tool, shell append, manual edit, or sister session.""" + """Error dict for external drift: the on-disk file wouldn't round-trip through + the parser/serializer, so flushing would discard externally added content.""" return { "success": False, "error": ( @@ -65,7 +58,7 @@ def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]: def _read_failed_error(path: "Path") -> Dict[str, Any]: - """Error dict for an unreadable (but existing) memory file. Treating it as + """Error dict for an unreadable (but existing) memory file: treating it as empty and saving would rewrite the whole file from ``[]`` — wiping memory.""" return {"success": False, "error": ( f"Refusing to write {path.name}: the file exists on disk but could not be read right now " @@ -119,10 +112,10 @@ class MemoryStore: self._consolidation_failures = 0 def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]: - """Count an at-capacity consolidation failure. Under the per-turn cap - return ``response`` unchanged (it says how to retry); past it return a - TERMINAL result so the model stops looping — a failed memory side effect - must never block the turn's reply.""" + """Count an at-capacity consolidation failure. Under the per-turn cap return + ``response`` unchanged (it says how to retry); past it return a TERMINAL + result so the model stops looping — a failed memory side effect must never + block the turn's reply.""" self._consolidation_failures += 1 if self._consolidation_failures <= self._MAX_CONSOLIDATION_FAILURES_PER_TURN: return response @@ -146,16 +139,14 @@ class MemoryStore: self._system_prompt_snapshot = { target: self._render_block(target, self._sanitize_entries_for_snapshot(entries, filename)) for target, entries, filename in ( - ("memory", self.memory_entries, "MEMORY.md"), - ("user", self.user_entries, "USER.md"), - ) + ("memory", self.memory_entries, "MEMORY.md"), ("user", self.user_entries, "USER.md")) } @staticmethod def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]: - """Return *entries* with any threat-matching entry replaced by a - ``[BLOCKED: …]`` placeholder (strict scope, same as writes). Empty or - already-blocked entries pass through unchanged.""" + """*entries* with any threat-matching entry replaced by a ``[BLOCKED: …]`` + placeholder (strict scope, same as writes). Empty or already-blocked + entries pass through unchanged.""" from tools.threat_patterns import scan_for_threats sanitized: List[str] = [] @@ -205,32 +196,21 @@ class MemoryStore: def _path_for(target: str) -> Path: return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md") - def _reload_target(self, target: str, *, skip_drift: bool = False): - """Re-read entries from disk (under file lock) before mutating. - - Returns ``None`` on a clean reload; the backup path (str) on external - drift (caller must abort — flushing would discard un-roundtrippable - content); or ``_READ_FAILED`` when the file exists but could not be - read (caller MUST abort — rewriting from an assumed-empty view would - wipe it; even append-only ``add`` rewrites the whole file). Drift check - and parse use the SAME raw snapshot — a failed second read used to count - as "no drift". *skip_drift* skips the round-trip check (``add``). - """ - raw, read_ok = self._read_raw_checked(self._path_for(target)) + def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]: + """Re-read entries from disk (under file lock) before mutating; return the + abort error dict, or None to proceed. Aborts on external drift (flushing + would discard un-roundtrippable content) and when the file exists but can't + be read (rewriting from an assumed-empty view would wipe it — even + append-only ``add`` rewrites the whole file). Drift check and parse use the + SAME raw snapshot — a failed second read used to count as "no drift". + *skip_drift* skips the round-trip check (``add``).""" + path = self._path_for(target) + raw, read_ok = self._read_raw_checked(path) if not read_ok: - return _READ_FAILED + return _read_failed_error(path) bak = None if skip_drift else self._detect_external_drift(target, raw) self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw)))) - return bak - - def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]: - """Reload under lock; return the abort error dict, or None to proceed.""" - bak = self._reload_target(target, skip_drift=skip_drift) - if bak is _READ_FAILED: - return _read_failed_error(self._path_for(target)) - if bak: - return _drift_error(self._path_for(target), bak) - return None + return _drift_error(path, bak) if bak else None def save_to_disk(self, target: str): """Persist entries to the appropriate file. Called after every mutation.""" @@ -256,8 +236,8 @@ class MemoryStore: return f"{self._char_count(target):,}/{self._char_limit(target):,}" def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]: - """Consolidation failure that shows the live entries so the model can - decide what to consolidate.""" + """Consolidation failure showing the live entries so the model can decide + what to consolidate.""" return self._consolidation_failure({"success": False, "error": message, "current_entries": self._entries_for(target), "usage": self._usage(target)}) @@ -386,13 +366,11 @@ class MemoryStore: def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]: """Apply a sequence of add/replace/remove ops to one target atomically. - - Ops are validated and applied against the FINAL budget only — so a single - call can free space (remove/replace) and add new entries without the - multi-turn consolidate-then-retry dance. All-or-nothing: if any op is - malformed, doesn't match, or the net result exceeds the char limit, - NOTHING is written and the first failure plus live state is returned. - """ + Ops are validated and applied against the FINAL budget only — so one call + can free space (remove/replace) and add new entries without the multi-turn + consolidate-then-retry dance. All-or-nothing: if any op is malformed, + doesn't match, or the net result exceeds the char limit, NOTHING is written + and the first failure plus live state is returned.""" if not operations: return {"success": False, "error": "operations list is empty."} @@ -419,7 +397,9 @@ class MemoryStore: pos = f"Operation {i + 1} ({act or 'unknown'})" msg = self._apply_batch_op(working, act, content, old_text, pos) if msg: - return self._batch_error(target, msg) + # Batch-abort error reporting live (uncommitted) state. + return self._failure_with_entries( + target, msg + " No operations were applied (batch is all-or-nothing).") # Budget check against the FINAL state only. new_total = len(ENTRY_DELIMITER.join(working)) if new_total > limit: @@ -430,16 +410,10 @@ class MemoryStore: )) return self._commit(target, working, f"Applied {len(operations)} operation(s).") - def _batch_error(self, target: str, message: str) -> Dict[str, Any]: - """Build a batch-abort error that reports live (uncommitted) state.""" - return self._failure_with_entries( - target, message + " No operations were applied (batch is all-or-nothing)." - ) - def format_for_system_prompt(self, target: str) -> Optional[str]: - """Return the frozen load-time snapshot for system-prompt injection (NOT - live state — mid-session writes don't affect it, preserving the prefix - cache). None if the snapshot is empty.""" + """The frozen load-time snapshot for system-prompt injection (NOT live + state — mid-session writes don't affect it, preserving the prefix cache). + None if the snapshot is empty.""" return self._system_prompt_snapshot.get(target, "") or None # -- Internal helpers -- @@ -482,11 +456,11 @@ class MemoryStore: @staticmethod def _read_raw_checked(path: Path) -> Tuple[str, bool]: """Read raw text as ``(raw, read_ok)``. ``read_ok`` is False ONLY when the - file EXISTS but can't be read (absent file → ``("", True)``). Invalid - UTF-8 counts as unreadable; decoding stays STRICT because - ``errors="replace"`` would hand callers a lossy view that a save then - persists over the real bytes. ``utf-8-sig`` strips a Notepad BOM that - otherwise glues U+FEFF onto the first entry forever.""" + file EXISTS but can't be read (absent file → ``("", True)``). Invalid UTF-8 + counts as unreadable; decoding stays STRICT because ``errors="replace"`` + would hand callers a lossy view that a save then persists over the real + bytes. ``utf-8-sig`` strips a Notepad BOM that otherwise glues U+FEFF onto + the first entry forever.""" if not path.exists(): return "", True try: @@ -500,26 +474,20 @@ class MemoryStore: the full ENTRY_DELIMITER so a bare "§" inside an entry is preserved.""" return [e for e in (x.strip() for x in raw.split(ENTRY_DELIMITER)) if e] - @staticmethod - def _read_entries_checked(path: Path) -> Tuple[List[str], bool]: - """Read + parse as ``(entries, read_ok)`` — see ``_read_raw_checked``.""" - raw, read_ok = MemoryStore._read_raw_checked(path) - return MemoryStore._parse_entries(raw), read_ok - @staticmethod def _read_file(path: Path) -> List[str]: """Read a memory file into entries (empty list on any error). Only for - read-only callers (``load_from_disk``, learning_mutations); mutation - paths must use ``_read_raw_checked`` so they can refuse to overwrite an + read-only callers (``load_from_disk``, learning_mutations); mutation paths + must use ``_read_raw_checked`` so they can refuse to overwrite an unreadable file.""" - return MemoryStore._read_entries_checked(path)[0] + return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0]) def _detect_external_drift(self, target: str, raw: str) -> Optional[str]: """Backup-path string if *raw* (the caller's checked-read snapshot) shows external drift, else None. Signals: round-trip mismatch, or one parsed - entry exceeding the whole-file char limit (no tool-written entry can — - an external writer appended free-form content a flush would truncate). - The file is snapshotted to ``.bak.`` so the operator can recover it.""" + entry exceeding the whole-file char limit (no tool-written entry can — an + external writer appended free-form content a flush would truncate). The + file is snapshotted to ``.bak.`` so the operator can recover it.""" if not raw.strip(): return None parsed = self._parse_entries(raw) From 2e0ba89317479e11b44b359a21690a2b9accb6e4 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:28:46 -0700 Subject: [PATCH 2/3] =?UTF-8?q?refactor(tools):=20graph=20client/auth,=20s?= =?UTF-8?q?ession=5Fsearch,=20memory=20=E2=80=94=20collapse=20defensive=20?= =?UTF-8?q?layers,=20fold=20iterate=5Fpages,=20hug=20brackets?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/memory_tool.py | 46 ++++++----------- tools/memory_tool_store.py | 46 +++++++---------- tools/microsoft_graph_auth.py | 27 ++++------ tools/microsoft_graph_client.py | 72 ++++++++------------------- tools/session_search_tool.py | 51 ++++++------------- tools/session_search_tool_common.py | 5 +- tools/session_search_tool_discover.py | 24 +++------ 7 files changed, 86 insertions(+), 185 deletions(-) diff --git a/tools/memory_tool.py b/tools/memory_tool.py index f3d5ecc037..f7ff4e7223 100644 --- a/tools/memory_tool.py +++ b/tools/memory_tool.py @@ -35,8 +35,7 @@ logger = logging.getLogger(__name__) # builds isolated while letting the check_fn result flow to the immediately # following dynamic_schema_overrides call in ToolRegistry.get_definitions(). _memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar( - "memory_surface_flags", default=None -) + "memory_surface_flags", default=None) def get_memory_dir() -> Path: @@ -47,15 +46,13 @@ def get_memory_dir() -> Path: from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports) ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, - _drift_error, _read_failed_error, _scan_memory_content, -) + _drift_error, _read_failed_error, _scan_memory_content) def load_on_disk_store() -> "MemoryStore": """Fresh on-disk MemoryStore with configured limits/flags, for contexts with no live agent (gateway, Desktop, bare CLI ``/memory``) so approvals enforce the SAME caps as ``agent_init``. Defaults if config can't load; never raises.""" - kwargs: Dict[str, Any] = {} try: from hermes_cli.config import load_config @@ -66,10 +63,9 @@ def load_on_disk_store() -> "MemoryStore": "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, - } + "user_profile_enabled": user_profile_enabled} except Exception: - kwargs = {} # config optional — fall back to defaults rather than break /memory + kwargs: Dict[str, Any] = {} # config optional — fall back to defaults rather than break /memory store = MemoryStore(**kwargs) store.load_from_disk() return store @@ -99,8 +95,7 @@ def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Option _GATE_TEXT = { "add": lambda label, content, old_text: (f"add to {label}", content or ""), "replace": lambda label, content, old_text: (f"replace in {label}", f"old: {old_text}\nnew: {content}"), - "remove": lambda label, content, old_text: (f"remove from {label}", old_text or ""), -} + "remove": 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]: @@ -158,8 +153,7 @@ def _validate_single_op(store, action, target, content, old_text) -> Optional[st _STORE_ACTIONS = { "add": lambda store, target, content, old_text: store.add(target, content), "replace": lambda store, target, content, old_text: store.replace(target, old_text, content), - "remove": lambda store, target, content, old_text: store.remove(target, old_text), -} + "remove": lambda store, target, content, old_text: store.remove(target, old_text)} def memory_tool( @@ -169,8 +163,7 @@ def memory_tool( old_text: str = None, new_text: str = None, operations: Optional[List[Dict[str, Any]]] = None, - store: Optional[MemoryStore] = None, -) -> str: + store: Optional[MemoryStore] = None) -> str: """Tool entry point; returns a JSON string. Single op (action + content / old_text) or batch (``operations`` applied atomically against the final budget). ``new_text`` aliases ``content`` — callers mirror ``old_text`` @@ -230,8 +223,7 @@ def get_builtin_memory_store_flags(config: Optional[Dict[str, Any]] = None) -> T section = get_builtin_memory_config(config) return ( is_truthy_value(section.get("memory_enabled"), default=True), - is_truthy_value(section.get("user_profile_enabled"), default=True), - ) + is_truthy_value(section.get("user_profile_enabled"), default=True)) @no_cache_check_fn @@ -351,13 +343,10 @@ _SINGLE_TARGET_TEXT = { ("memory",): ( "The enabled built-in store: 'memory' for personal notes.", "TARGET: only 'memory' is enabled for personal notes (environment, conventions, " - "tool quirks, lessons).", - ), + "tool quirks, lessons)."), ("user",): ( "The enabled built-in store: 'user' for user profile.", - "TARGET: only 'user' is enabled for user profile facts (name, role, preferences, style).", - ), -} + "TARGET: only 'user' is enabled for user profile facts (name, role, preferences, style).")} def _build_memory_schema_overrides() -> Dict[str, Any]: @@ -377,8 +366,7 @@ def _build_memory_schema_overrides() -> Dict[str, Any]: description = description.replace( "TARGETS: 'user' = who the user is (name, role, preferences, style). 'memory' = your " "notes (environment, conventions, tool quirks, lessons).", - replacement, - ) + replacement) return {"description": description, "parameters": parameters} @@ -390,14 +378,8 @@ registry.register( toolset="memory", schema=MEMORY_SCHEMA, handler=lambda args, **kw: memory_tool( - action=args.get("action", ""), - target=args.get("target", "memory"), - content=args.get("content"), - old_text=args.get("old_text"), - new_text=args.get("new_text"), - operations=args.get("operations"), - store=kw.get("store")), + action=args.get("action", ""), target=args.get("target", "memory"), store=kw.get("store"), + **{k: args.get(k) for k in ("content", "old_text", "new_text", "operations")}), check_fn=check_memory_requirements, emoji="🧠", - dynamic_schema_overrides=_build_memory_schema_overrides, -) + dynamic_schema_overrides=_build_memory_schema_overrides) diff --git a/tools/memory_tool_store.py b/tools/memory_tool_store.py index b5cffcdf0a..242e205854 100644 --- a/tools/memory_tool_store.py +++ b/tools/memory_tool_store.py @@ -18,9 +18,7 @@ logger = logging.getLogger("tools.memory_tool") # agent/conversation_compression.py uses them to detect a leftover block for a # target whose entries have since been emptied — keep in lockstep. MEMORY_BLOCK_HEADERS = { - "memory": "MEMORY (your personal notes)", - "user": "USER PROFILE (who the user is)", -} + "memory": "MEMORY (your personal notes)", "user": "USER PROFILE (who the user is)"} ENTRY_DELIMITER = "\n§\n" @@ -53,8 +51,7 @@ def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]: "remediation": ( "Open the .bak file, integrate the missing entries into the memory tool one at a time via " "memory(action=add, content=...), then remove or rewrite the original file to a clean state." - ), - } + )} def _read_failed_error(path: "Path") -> Dict[str, Any]: @@ -64,8 +61,7 @@ def _read_failed_error(path: "Path") -> Dict[str, Any]: 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]: @@ -122,8 +118,7 @@ class MemoryStore: return {"success": False, "done": True, "error": ( f"Memory consolidation failed {self._consolidation_failures} times this turn. Stop retrying " "memory calls — leave memory unchanged for now and continue with your reply to the user. " - "The fact can be saved in a later turn." - )} + "The fact can be saved in a later turn.")} def load_from_disk(self): """Load MEMORY.md / USER.md and capture the frozen system-prompt snapshot. @@ -174,20 +169,20 @@ class MemoryStore: yield return fd = open(lock_path, "a+", encoding="utf-8") - try: + + def _flock(unlock: bool): if fcntl: - fcntl.flock(fd, fcntl.LOCK_EX) + fcntl.flock(fd, fcntl.LOCK_UN if unlock else fcntl.LOCK_EX) else: fd.seek(0) - msvcrt.locking(fd.fileno(), msvcrt.LK_LOCK, 1) + msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK if unlock else msvcrt.LK_LOCK, 1) + + try: + _flock(False) yield finally: try: - if fcntl: - fcntl.flock(fd, fcntl.LOCK_UN) - else: - fd.seek(0) - msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK, 1) + _flock(True) except OSError: pass fd.close() @@ -221,10 +216,7 @@ class MemoryStore: return self.user_entries if target == "user" else self.memory_entries def _set_entries(self, target: str, entries: List[str]): - if target == "user": - self.user_entries = entries - else: - self.memory_entries = entries + setattr(self, "user_entries" if target == "user" else "memory_entries", entries) def _char_count(self, target: str) -> int: return len(ENTRY_DELIMITER.join(self._entries_for(target))) @@ -252,8 +244,7 @@ class MemoryStore: return None, self._consolidation_failure({ "success": False, "error": f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text of the entry you want to {verb}.", - "current_entries": entries, - }) + "current_entries": entries}) return idx, None def _commit(self, target: str, entries: List[str], message: str) -> Dict[str, Any]: @@ -285,8 +276,7 @@ class MemoryStore: f"Memory at {self._char_count(target):,}/{limit:,} chars. Adding this entry " f"({len(content)} chars) would exceed the limit. Consolidate now: use 'replace' to merge " f"overlapping entries into shorter ones or 'remove' stale or less important entries (see " - f"current_entries below), then retry this add — all in this turn." - )) + f"current_entries below), then retry this add — all in this turn.")) entries.append(content) return self._commit(target, entries, "Entry added.") @@ -315,8 +305,7 @@ class MemoryStore: return self._failure_with_entries(target, ( f"Replacement would put memory at {new_total:,}/{limit:,} chars. Shorten the new content, " f"or 'remove' other stale or less important entries to make room (see current_entries " - f"below), then retry — all in this turn." - )) + f"below), then retry — all in this turn.")) entries[idx] = new_content return self._commit(target, entries, "Entry replaced.") @@ -406,8 +395,7 @@ class MemoryStore: 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." - )) + f"entries in the same batch (see current_entries below), then retry.")) return self._commit(target, working, f"Applied {len(operations)} operation(s).") def format_for_system_prompt(self, target: str) -> Optional[str]: diff --git a/tools/microsoft_graph_auth.py b/tools/microsoft_graph_auth.py index f58eac5f16..b2fbc73a5f 100644 --- a/tools/microsoft_graph_auth.py +++ b/tools/microsoft_graph_auth.py @@ -41,9 +41,7 @@ def format_graph_error(error: Any) -> str | None: if not isinstance(error, dict): return None code, message = error.get("code"), error.get("message") - if code and message: - return f"{code}: {message}" - return str(message) if message else None + return f"{code}: {message}" if code and message else (str(message) if message else None) @dataclass(frozen=True) @@ -75,8 +73,7 @@ class GraphCredentials: return cls( *values, scope=(env.get("MSGRAPH_SCOPE") or DEFAULT_GRAPH_SCOPE).strip(), - authority_url=(env.get("MSGRAPH_AUTHORITY_URL") or DEFAULT_GRAPH_AUTHORITY_URL).strip(), - ) + authority_url=(env.get("MSGRAPH_AUTHORITY_URL") or DEFAULT_GRAPH_AUTHORITY_URL).strip()) @dataclass @@ -128,8 +125,7 @@ class MicrosoftGraphTokenProvider: "cached": bool(cached), "expires_in_seconds": cached.expires_in_seconds if cached else None, "is_expired": cached.is_expired(skew_seconds=0) if cached else None, - "refresh_skew_seconds": self.skew_seconds, - } + "refresh_skew_seconds": self.skew_seconds} def _fresh_cached(self) -> CachedAccessToken | None: """The cached token unless it expires within ``skew_seconds``.""" @@ -145,28 +141,24 @@ class MicrosoftGraphTokenProvider: async with self._lock: if not force_refresh and (cached := self._fresh_cached()): return cached.access_token - token = await self._fetch_access_token() - self._cached_token = token - return token.access_token + self._cached_token = await self._fetch_access_token() + return self._cached_token.access_token async def _fetch_access_token(self) -> CachedAccessToken: data = { "grant_type": "client_credentials", "client_id": self.credentials.client_id, "client_secret": self.credentials.client_secret, - "scope": self.credentials.scope, - } + "scope": self.credentials.scope} async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client: response = await client.post( self.credentials.token_url, data=data, - headers={"Content-Type": "application/x-www-form-urlencoded"}, - ) + headers={"Content-Type": "application/x-www-form-urlencoded"}) if response.status_code >= 400: raise MicrosoftGraphTokenError( "Microsoft Graph token request failed with HTTP " - f"{response.status_code}: {_extract_error_detail(response)}" - ) + f"{response.status_code}: {_extract_error_detail(response)}") try: payload = response.json() except ValueError as exc: @@ -185,8 +177,7 @@ class MicrosoftGraphTokenProvider: return CachedAccessToken( access_token=access_token, token_type=str(payload.get("token_type") or "Bearer").strip() or "Bearer", - expires_at=time.time() + max(0, expires_in_seconds), - ) + expires_at=time.time() + max(0, expires_in_seconds)) def _extract_error_detail(response: httpx.Response) -> str: diff --git a/tools/microsoft_graph_client.py b/tools/microsoft_graph_client.py index ab8842dcf4..25a1aa40e5 100644 --- a/tools/microsoft_graph_client.py +++ b/tools/microsoft_graph_client.py @@ -5,16 +5,12 @@ from __future__ import annotations import asyncio import os from pathlib import Path -from typing import Any, AsyncIterator, Awaitable, Callable +from typing import Any, Awaitable, Callable import httpx from agent.retry_utils import parse_retry_after_seconds -from tools.microsoft_graph_auth import ( - GraphCredentials, - MicrosoftGraphTokenProvider, - format_graph_error, -) +from tools.microsoft_graph_auth import GraphCredentials, MicrosoftGraphTokenProvider, format_graph_error DEFAULT_GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0" @@ -32,8 +28,7 @@ class MicrosoftGraphAPIError(MicrosoftGraphClientError): def __init__( self, status_code: int, method: str, url: str, message: str, *, - retry_after_seconds: float | None = None, payload: Any = None, - ) -> None: + retry_after_seconds: float | None = None, payload: Any = None) -> None: self.status_code = status_code self.method = method self.url = url @@ -55,8 +50,7 @@ class MicrosoftGraphClient: 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: + user_agent: str = "Hermes-Agent/graph-client") -> None: self.token_provider = token_provider self.base_url = base_url.rstrip("/") self.timeout = timeout @@ -87,9 +81,10 @@ class MicrosoftGraphClient: return {"deleted": True, "status_code": response.status_code} return self._decode_json(response) - async def iterate_pages( - self, path: str, *, params: Params = None, headers: Headers = None - ) -> AsyncIterator[dict[str, 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 {}) @@ -98,20 +93,10 @@ class MicrosoftGraphClient: payload = self._decode_json(response) if not isinstance(payload, dict): raise MicrosoftGraphClientError( - f"Expected paginated Graph response dict, got {type(payload).__name__}." - ) - yield payload - next_url = payload.get("@odata.nextLink") - next_params = {} - - async def collect_paginated( - self, path: str, *, params: Params = None, headers: Headers = None - ) -> list[Any]: - items: list[Any] = [] - async for page in self.iterate_pages(path, params=params, headers=headers): - value = page.get("value") - if isinstance(value, list): - items.extend(value) + f"Expected paginated Graph response dict, got {type(payload).__name__}.") + if isinstance(payload.get("value"), list): + items.extend(payload["value"]) + next_url, next_params = payload.get("@odata.nextLink"), {} return items async def download_to_file( @@ -147,8 +132,7 @@ class MicrosoftGraphClient: async def _request( self, method: str, path_or_url: str, *, - params: Params = None, json_body: Any | None = None, headers: Headers = None, - ) -> httpx.Response: + params: Params = None, json_body: Any | None = None, headers: Headers = None) -> httpx.Response: url = self._resolve_url(path_or_url) async def perform(client: httpx.AsyncClient, request_headers: dict[str, str]): @@ -160,8 +144,7 @@ class MicrosoftGraphClient: async def _with_retries( self, method: str, url: str, accept: str, json_body: Any | None, headers: Headers, perform: Callable[[httpx.AsyncClient, dict[str, str]], Awaitable[tuple[httpx.Response, Any]]], - kind: str, - ) -> Any: + kind: str) -> Any: """Run ``perform`` (returning ``(response, result)``) under the retry policy. ``kind`` ("request"/"download") only labels the transport-failure messages. @@ -175,8 +158,7 @@ class MicrosoftGraphClient: token = await self.token_provider.get_access_token( force_refresh=attempt > 0 and isinstance(last_error, MicrosoftGraphAPIError) - and last_error.status_code == 401 - ) + and last_error.status_code == 401) request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent} if json_body is not None: request_headers["Content-Type"] = "application/json" @@ -190,8 +172,7 @@ class MicrosoftGraphClient: last_error = exc if attempt >= self.max_retries: raise MicrosoftGraphClientError( - f"Microsoft Graph {kind} failed for {method} {url}: {exc}" - ) from exc + f"Microsoft Graph {kind} failed for {method} {url}: {exc}") from exc await self._sleep(self._retry_delay(None, attempt)) attempt += 1 continue @@ -224,29 +205,20 @@ class MicrosoftGraphClient: except ValueError as exc: raise MicrosoftGraphClientError( "Microsoft Graph response was not valid JSON for " - f"{response.request.method} {response.request.url}" - ) from exc + f"{response.request.method} {response.request.url}") from exc @staticmethod def _retry_delay(response: httpx.Response | None, attempt: int) -> float: - if response is not None: - retry_after = parse_retry_after_seconds(response.headers) - if retry_after is not None: - return retry_after - return min(8.0, 0.5 * (2 ** attempt)) + retry_after = parse_retry_after_seconds(response.headers) if response is not None else None + return min(8.0, 0.5 * (2 ** attempt)) if retry_after is None else retry_after @staticmethod def _build_api_error(method: str, url: str, response: httpx.Response) -> MicrosoftGraphAPIError: - message = response.text.strip() or "unknown error" try: payload: Any = response.json() except ValueError: payload = None - if isinstance(payload, dict): - detail = format_graph_error(payload.get("error")) - if detail is not None: - message = detail + detail = format_graph_error(payload.get("error")) if isinstance(payload, dict) else None return MicrosoftGraphAPIError( - response.status_code, method, url, message, - retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload, - ) + response.status_code, method, url, detail if detail is not None else (response.text.strip() or "unknown error"), + retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload) diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index 5a8ff6fd4d..d8d25c8c41 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -19,11 +19,9 @@ from tools.session_search_tool_common import ( # noqa: F401 (re-exports) _annotate_rebuild_status, _format_timestamp, _get_message_storage_state, _is_compacted_message, _is_compacted_state, _is_compaction_summary, _ok, _order_for_recall, _quiet, _resolve_lineage, _resolve_to_parent, _session_end_reason, - _session_left_live_context, _session_link, _session_meta_block, _shape_message, -) + _session_left_live_context, _session_link, _session_meta_block, _shape_message) from tools.session_search_tool_discover import ( # noqa: F401 (re-exports) - _discover, _normalize_title_query, _title_match_result, -) + _discover, _normalize_title_query, _title_match_result) def _resolve_profile_db(profile: str): @@ -110,13 +108,10 @@ def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_p sessions = db.list_sessions_rich( limit=limit + 15, # extra so we can skip current / compression roots exclude_sources=list(_HIDDEN_SESSION_SOURCES), - order_by_last_active=True, - ) + order_by_last_active=True) current_root, has_compression_hop = ( - _resolve_to_parent(db, current_session_id) - if current_session_id else (None, False) - ) + _resolve_to_parent(db, current_session_id) if current_session_id else (None, False)) results = [] for s in sessions: sid = s.get("id", "") @@ -131,8 +126,7 @@ def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_p "session_id": sid, "link": _session_link(sid, link_profile), "title": s.get("title") or None, "source": s.get("source", ""), "started_at": s.get("started_at", ""), "last_active": s.get("last_active", ""), "message_count": s.get("message_count", 0), - "preview": s.get("preview", ""), - }) + "preview": s.get("preview", "")}) if len(results) >= limit: break return _ok(mode="browse", results=results, count=len(results), message=( @@ -165,10 +159,7 @@ def _anchor_in_live_context(db, anchor_state, anchor_session_id: str, current_se return False # Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history. is_inactive_non_compacted = ( - anchor_state is not None - and anchor_state["active"] == 0 - and anchor_state["compacted"] != 1 - ) + anchor_state is not None and anchor_state["active"] == 0 and anchor_state["compacted"] != 1) return is_inactive_non_compacted or not _session_left_live_context(db, anchor_session_id) @@ -205,8 +196,7 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, owning_session_id = anchor_state.get("session_id") if anchor_state is not None else None if current_session_id and _anchor_in_live_context( - db, anchor_state, owning_session_id or session_id, current_session_id - ): + db, anchor_state, owning_session_id or session_id, current_session_id): return tool_error("scroll rejected: anchor lives in the current session lineage (already in your active context)", success=False) session_meta = _get_session_meta(db, session_id) @@ -223,8 +213,7 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, 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 - ) + db, session_id, owning_session_id, around_message_id, window) if rebind_view is not None: view = rebind_view messages = view["window"] @@ -243,8 +232,7 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5, "id; backward: the FIRST message's id (the boundary message repeats " "as an orientation marker). messages_before/messages_after < window " "means you've hit that end of the session."), - **({"warning": rebind_warning} if rebind_warning else {}), - ) + **({"warning": rebind_warning} if rebind_warning else {})) def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str: @@ -310,8 +298,7 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_id, detail_norm = "full" if isinstance(detail, str) and detail.strip().lower() == "full" else "adaptive" return _discover( db=db, query=query.strip(), role_filter=role_list, limit=limit, sort=sort_norm, - detail=detail_norm, current_session_id=current_session_id, link_profile=profile, - ) + detail=detail_norm, current_session_id=current_session_id, link_profile=profile) def session_search(query: str = "", role_filter: str = None, limit: int = 3, db=None, @@ -465,18 +452,8 @@ registry.register( toolset="session_search", schema=SESSION_SEARCH_SCHEMA, handler=lambda args, **kw: session_search( - query=args.get("query") or "", - role_filter=args.get("role_filter"), - limit=args.get("limit", 3), - session_id=args.get("session_id"), - around_message_id=args.get("around_message_id"), - window=args.get("window", 5), - sort=args.get("sort"), - detail=args.get("detail", "adaptive"), - profile=args.get("profile"), - db=kw.get("db"), - current_session_id=kw.get("current_session_id"), - ), + query=args.get("query") or "", limit=args.get("limit", 3), window=args.get("window", 5), + detail=args.get("detail", "adaptive"), db=kw.get("db"), current_session_id=kw.get("current_session_id"), + **{k: args.get(k) for k in ("role_filter", "session_id", "around_message_id", "sort", "profile")}), check_fn=check_session_search_requirements, - emoji="🔍", -) + emoji="🔍") diff --git a/tools/session_search_tool_common.py b/tools/session_search_tool_common.py index 0a0d2c44f0..f9613796ea 100644 --- a/tools/session_search_tool_common.py +++ b/tools/session_search_tool_common.py @@ -118,9 +118,9 @@ def _session_end_reason(db, session_id: str) -> Optional[str]: return None try: s = db.get_session(session_id) - return (s.get("end_reason") or None) if s else None except Exception: return None + return (s.get("end_reason") or None) if s else None def _session_left_live_context(db, session_id: str) -> bool: @@ -176,8 +176,7 @@ def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None: 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." - )} + f"may be incomplete until it finishes.")} def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]: diff --git a/tools/session_search_tool_discover.py b/tools/session_search_tool_discover.py index 08ba526421..4867f3d421 100644 --- a/tools/session_search_tool_discover.py +++ b/tools/session_search_tool_discover.py @@ -9,8 +9,7 @@ from tools.registry import tool_error from tools.session_search_tool_common import ( _DISCOVER_SCAN_LIMIT, _DISCOVER_SEARCH_FIELDS, _HIDDEN_SESSION_SOURCES, _annotate_rebuild_status, _format_timestamp, _is_compacted_message, _is_compaction_summary, _order_for_recall, _quiet, - _resolve_lineage, _resolve_to_parent, _session_left_live_context, _session_link, _shape_message, -) + _resolve_lineage, _resolve_to_parent, _session_left_live_context, _session_link, _shape_message) def _normalize_title_query(query: str) -> str: @@ -33,8 +32,7 @@ def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> if ( current_lineage_root and lineage_root == current_lineage_root - and not _session_left_live_context(db, session_id) - ): + and not _session_left_live_context(db, session_id)): return None session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None, @@ -58,8 +56,7 @@ def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> "bookend_end": [_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])], "messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", max(len(messages) - 5, 0)), - "detail": "full", "_lineage_root": lineage_root, - } + "detail": "full", "_lineage_root": lineage_root} if lineage_root and lineage_root != session_id: entry["parent_session_id"] = lineage_root return entry @@ -92,8 +89,7 @@ def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id if ( current_lineage_root and resolved_sid == current_lineage_root - and not (is_ended_session or is_compacted_hit) - ): + and not (is_ended_session or is_compacted_hit)): continue if current_session_id and raw_sid == current_session_id and not is_compacted_hit: continue @@ -131,8 +127,7 @@ def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detai "messages": [_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages], "bookend_end": _bookend(view, "bookend_end") if full else [], "messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", 0), - "detail": result_detail, - } + "detail": result_detail} if lineage_root and lineage_root != hit_sid: entry["parent_session_id"] = lineage_root return entry @@ -148,8 +143,7 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort try: raw_results = db.search_messages( query=query, role_filter=role_list, exclude_sources=list(_HIDDEN_SESSION_SOURCES), - limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS, - ) + limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS) except Exception as e: logging.error("FTS5 search failed: %s", e, exc_info=True) return tool_error(f"Search failed: {e}", success=False) @@ -162,8 +156,7 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort 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*`." - )) + "phrases, exclude with NOT, or prefix-match with `deploy*`.")) seen_sessions: Dict[str, Dict[str, Any]] = {} results = [] @@ -188,5 +181,4 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort "verbatim inline mid-sentence (it renders as a titled link) — never " "as markdown, in backticks, on its own line, or next to the " "title/id/date. To read more around a compact result, scroll: " - "session_search(session_id=..., around_message_id=match_message_id)." - )) + "session_search(session_id=..., around_message_id=match_message_id).")) From 25038334c1050159516a0d09a955335888fa2054 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:34:48 -0700 Subject: [PATCH 3/3] refactor(tools): fold session_search_tool_common/_discover back into session_search_tool (single module, no re-export shim) --- tools/session_search_tool.py | 406 +++++++++++++++++++++++++- tools/session_search_tool_common.py | 229 --------------- tools/session_search_tool_discover.py | 184 ------------ 3 files changed, 394 insertions(+), 425 deletions(-) delete mode 100644 tools/session_search_tool_common.py delete mode 100644 tools/session_search_tool_discover.py diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index d8d25c8c41..3fe3f7c285 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -5,23 +5,405 @@ Single-shape tool; the mode is inferred from the args: DISCOVERY (``query``; FTS5 deduped by lineage, adaptive detail hydrates only the top result), SCROLL (``session_id`` + ``around_message_id``; ±window around the anchor), READ (``session_id`` alone; whole session or head/tail), BROWSE (no args). -No LLM calls — every shape returns actual DB messages. Helpers live in -``session_search_tool_common`` / ``_discover`` and are re-exported here. +No LLM calls — every shape returns actual DB messages. """ import json import logging -from typing import Any, List, Optional +from datetime import datetime +from typing import Any, Dict, List, Optional, Union -from tools.session_search_tool_common import ( # noqa: F401 (re-exports) - _COMPACTION_PREFIXES, _DEMOTED_SESSION_SOURCES, _DISCOVER_SCAN_LIMIT, - _DISCOVER_SEARCH_FIELDS, _FRESH_RESET_END_REASONS, _HIDDEN_SESSION_SOURCES, - _annotate_rebuild_status, _format_timestamp, _get_message_storage_state, - _is_compacted_message, _is_compacted_state, _is_compaction_summary, - _ok, _order_for_recall, _quiet, _resolve_lineage, _resolve_to_parent, _session_end_reason, - _session_left_live_context, _session_link, _session_meta_block, _shape_message) -from tools.session_search_tool_discover import ( # noqa: F401 (re-exports) - _discover, _normalize_title_query, _title_match_result) +from hermes_state_common import _RESET_END_REASONS + +# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool), +# delegate subagent runs, kanban workers — not the user's history. +_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool") + +# Searchable but DEMOTED below interactive sessions: cron sessions' repetitive +# vocabulary dominates bare BM25 and starves out the user's own sessions +# ("recall blindness"). Demoting keeps them reachable when they're the only match. +_DEMOTED_SESSION_SOURCES = ("cron",) + +# FTS rows scanned before dedup-by-lineage — well above the handful of distinct +# sessions a query returns, so interactive matches buried under cron hits are +# still in hand for the demotion pass. +_DISCOVER_SCAN_LIMIT = 300 + +# Raw FTS rows are only a discovery-plan input; the response hydrates its own +# anchored window and bookends after lineage dedup. +_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started") + +# Generated context-compaction handoff summaries (agent/context_compressor.py); +# excluded from bookends so huge compaction payloads aren't re-introduced. +_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:") + +# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the +# predecessor WITHOUT carrying its transcript forward — unlike compression +# continuations and live delegation children. Derived from the gateway set so +# this tool and the recovery fence cannot drift. +_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"} + + +def _quiet(fn, default, msg, *log_args, with_exc: bool = False): + """Call ``fn()``; on any exception debug-log *msg* (appending the exception + as a final ``%s`` arg when *with_exc*) and return *default*.""" + try: + return fn() + except Exception as e: + logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True) + return default + + +def _format_timestamp(ts: Union[int, float, str, None]) -> str: + """Unix timestamp (number or numeric string) or ISO string -> readable date. + "unknown" for None; str(ts) if conversion fails.""" + if ts is None: + return "unknown" + try: + value = ts + if isinstance(ts, str): + if not ts.replace(".", "").replace("-", "").isdigit(): + return ts + value = float(ts) + if isinstance(value, (int, float)): + return datetime.fromtimestamp(value).strftime("%B %d, %Y at %I:%M %p") + except (ValueError, OSError, OverflowError) as e: + logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True) + except Exception as e: + logging.debug("Unexpected error formatting timestamp %s: %s", ts, e, exc_info=True) + return str(ts) + + +def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]: + """The ``session_meta`` sub-object shared by read/scroll responses.""" + return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"), + "model": meta.get("model"), "title": meta.get("title")} + + +def _ok(**payload) -> str: + """Serialize a successful tool result (``success`` first, then *payload* in order).""" + return json.dumps({"success": True, **payload}, ensure_ascii=False) + + +def _is_compaction_summary(content: str) -> bool: + """Return True if *content* looks like a generated compaction handoff.""" + return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES) + + +def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]: + """Walk parent_session_id to the lineage root -> ``(root_id, has_compression_hop)``. + The flag distinguishes a compression-split lineage (parent content summarised + away) from a delegation lineage (child content still visible to the parent). + Falls back to ``(session_id, False)`` on errors.""" + if not session_id: + return session_id, False + visited: set[str] = set() + cur, has_compression = session_id, False + while cur and cur not in visited: + visited.add(cur) + s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True) + if not s: + break + if s.get("end_reason") == "compression": + has_compression = True + if not s.get("parent_session_id"): + break + cur = s["parent_session_id"] + return cur, has_compression + + +def _resolve_lineage(db, session_id: str) -> str: + """Return only the lineage root (ignores the compression hop).""" + return _resolve_to_parent(db, session_id)[0] + + +def _session_end_reason(db, session_id: str) -> Optional[str]: + """Return the session's ``end_reason``, or None if missing/unended/error.""" + if not session_id: + return None + try: + s = db.get_session(session_id) + except Exception: + return None + return (s.get("end_reason") or None) if s else None + + +def _session_left_live_context(db, session_id: str) -> bool: + """True when *session_id*'s transcript is no longer in anyone's live context: + ``compression`` (summarised into the continuation child) or a fresh reset + (child starts empty). Everything else stays excluded from same-lineage + recall — live delegation children (``end_reason is None``) are visible to + the parent agent, and ``branched`` parents were copied verbatim into the + branch child, so their content IS the current context.""" + end_reason = _session_end_reason(db, session_id) + return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS + + +def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]: + """Return the owning session and visibility flags for *message_id*.""" + if not message_id: + return None + + def _lookup(): + with db._lock: + return db._conn.execute( + "SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,) + ).fetchone() + + row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id) + return dict(row) if row is not None else None + + +def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool: + """Compaction archives are ``active=0, compacted=1`` (content summarised + away by archive_and_compact). Rewind/undo rows are ``active=0, compacted=0`` + and must stay hidden.""" + return state is not None and state["active"] == 0 and state["compacted"] == 1 + + +def _is_compacted_message(db, message_id) -> bool: + """True if *message_id* is a compaction-archived row — pre-compaction content + no longer in live context, so it should stay discoverable even on the + current session. False on any error (caller falls back to skipping).""" + return _is_compacted_state(_get_message_storage_state(db, message_id)) + + +def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None: + """Add a rebuild-progress note while the deferred FTS backfill is running, + so the agent can explain thin/slow results instead of treating them as + ground truth. No-op (never raises) when no rebuild is pending.""" + try: + status = db.fts_rebuild_status() + except Exception: + return + if status is None: + return + payload["index_rebuild"] = {"percent": status["percent"], "note": ( + f"The search index is rebuilding in the background ({status['percent']}% done, " + f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages " + f"may be incomplete until it finishes.")} + + +def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """Stable-sort FTS rows so interactive sessions rank above automation. + BM25 order is preserved within each class; only cross-class order changes, + so a cron hit never displaces an interactive hit during lineage dedup.""" + return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0) + + +def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None, + max_content_len: Optional[int] = None) -> Dict[str, Any]: + """Slim a message row for the tool response. Keeps content even if empty + (absent content is meaningful — tool-call-only assistant turns). With + *max_content_len*, content is truncated and ``content_truncated`` / + ``original_content_chars`` added.""" + content = m.get("content") + if isinstance(content, str) and "\x1b" in content: + # Recalled messages can carry ANSI escapes (archived terminal output). + from tools.ansi_strip import strip_ansi + + content = strip_ansi(content) + original_chars = None + if max_content_len and content and len(content) > max_content_len: + original_chars = len(content) + content = content[:max_content_len] + "…" + entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")} + entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)}) + if anchor_id is not None and m.get("id") == anchor_id: + entry["anchor"] = True + if original_chars is not None: + entry["content_truncated"] = True + entry["original_content_chars"] = original_chars + return {k: v for k, v in entry.items() if v is not None or k == "content"} + + +def _session_link(session_id: str, profile: str = None) -> str: + """The reference the agent writes to point the user at a session — same + value the desktop composer emits, so it renders as a titled link. The + profile segment is omitted when it can't be named confidently (a bare id + still resolves, it just can't disambiguate across profiles).""" + name = (profile or "").strip() + if not name: + def _active(): + from hermes_cli.profiles import get_active_profile_name + + resolved = get_active_profile_name() + return "" if resolved == "custom" else resolved + + name = _quiet(_active, "", "get_active_profile_name failed for session link") + return f"@session:{name}/{session_id}" if name else f"@session:{session_id}" + + +def _normalize_title_query(query: str) -> str: + """Strip common quoting the model may include around a remembered title.""" + return query.strip().strip("`'\"") + + +def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]: + """Return a discovery-shaped result when the query matches a session title.""" + title_query = _normalize_title_query(query) + if not title_query: + return None + session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None, + "resolve_session_by_title failed for %r", title_query) + if not session_id: + return None + lineage_root = _resolve_lineage(db, session_id) + # Same-lineage title hits are in-context only while the session is live; + # /new-reset and compression-ended parents are not. + if ( + current_lineage_root + and lineage_root == current_lineage_root + and not _session_left_live_context(db, session_id)): + return None + + session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None, + "get_session failed for title match %s", session_id) or {} + if session_meta.get("source") in _HIDDEN_SESSION_SOURCES: + return None + messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id) + anchor_id = messages[0].get("id") if messages else None + view = {} + if anchor_id is not None: + view = _quiet(lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {}, + "get_anchored_view failed for title match %s/%s", session_id, anchor_id) + entry = { + "session_id": session_id, "when": _format_timestamp(session_meta.get("started_at")), + "source": session_meta.get("source", "unknown"), "model": session_meta.get("model") or "unknown", + "title": session_meta.get("title") or title_query, "matched_role": "session_title", + "match_message_id": anchor_id, + "snippet": f"Session title matched: {session_meta.get('title') or title_query}", + "bookend_start": [_shape_message(m) for m in (view.get("bookend_start") or messages[:3])], + "messages": [_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])], + "bookend_end": [_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])], + "messages_before": view.get("messages_before", 0), + "messages_after": view.get("messages_after", max(len(messages) - 5, 0)), + "detail": "full", "_lineage_root": lineage_root} + if lineage_root and lineage_root != session_id: + entry["parent_session_id"] = lineage_root + return entry + + +def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str: + payload = {"success": True, "mode": "discover", "query": query, "detail": detail, + "results": results, "count": len(results), **extra} + _annotate_rebuild_status(db, payload) + return json.dumps(payload, ensure_ascii=False) + + +def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None: + """Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*. + + The raw owning session_id stays on the row — only it pairs validly with the + FTS match id for the anchored window. Current-lineage hits are skipped + UNLESS the transcript left live context: compression-ended session, /new- + reset predecessor (hiding it made gateway recall blind after every /new), + or an in-place compacted row on the SAME session_id. A live delegation + child has end_reason=None, so it stays excluded. + """ + for r in raw_results: + if len(seen_sessions) >= limit: + break + raw_sid = r["session_id"] + resolved_sid, _ = _resolve_to_parent(db, raw_sid) + is_compacted_hit = _is_compacted_message(db, r.get("id")) + is_ended_session = _session_left_live_context(db, raw_sid) + if ( + current_lineage_root + and resolved_sid == current_lineage_root + and not (is_ended_session or is_compacted_hit)): + continue + if current_session_id and raw_sid == current_session_id and not is_compacted_hit: + continue + seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid}) + + +def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]: + return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or []) + if not _is_compaction_summary(m.get("content", ""))] + + +def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]: + """Build one discovery result from a surviving FTS row; None if the anchored + view can't be loaded (the hit is dropped).""" + hit_sid = match_info.get("session_id") or lineage_root + msg_id = match_info.get("id") + try: + view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3) + except Exception as e: + logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True) + return None + session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {} + full = result_detail == "full" + window_messages = view.get("window") or [] + if not full: + window_messages = [m for m in window_messages if m.get("id") == msg_id] + entry = { + "session_id": hit_sid, + "when": _format_timestamp(session_meta.get("started_at") or match_info.get("session_started")), + "source": session_meta.get("source") or match_info.get("source", "unknown"), + "model": session_meta.get("model") or match_info.get("model") or "unknown", + "title": session_meta.get("title") or None, "matched_role": match_info.get("role"), + "match_message_id": msg_id, "snippet": match_info.get("snippet") or "", + "bookend_start": _bookend(view, "bookend_start") if full else [], + "messages": [_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages], + "bookend_end": _bookend(view, "bookend_end") if full else [], + "messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", 0), + "detail": result_detail} + if lineage_root and lineage_root != hit_sid: + entry["parent_session_id"] = lineage_root + return entry + + +def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str], + detail: str, current_session_id: str = None, link_profile: str = None) -> str: + """Discovery shape: FTS5 plus adaptive or full result hydration.""" + role_list = role_filter if role_filter else ["user", "assistant"] + current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None + title_result = _title_match_result(db, query, current_lineage_root) + + try: + raw_results = db.search_messages( + query=query, role_filter=role_list, exclude_sources=list(_HIDDEN_SESSION_SOURCES), + limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS) + except Exception as e: + logging.error("FTS5 search failed: %s", e, exc_info=True) + return tool_error(f"Search failed: {e}", success=False) + + # Demote cron rows below interactive ones BEFORE dedup so a high-volume + # cron corpus can't starve the user's own sessions out of the top `limit`. + raw_results = _order_for_recall(raw_results) + + if not raw_results and not title_result: + return _discover_payload(db, query, detail, [], message=( + "No matching sessions found. FTS5 ANDs all terms by default — " + "broaden with OR (`alpha OR beta`), exact-match with quoted " + "phrases, exclude with NOT, or prefix-match with `deploy*`.")) + + seen_sessions: Dict[str, Dict[str, Any]] = {} + results = [] + if title_result: + title_lineage = title_result.pop("_lineage_root", None) + if title_lineage: + seen_sessions[title_lineage] = {"_title_only": True} + results.append(title_result) + _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) + + for lineage_root, match_info in seen_sessions.items(): + if match_info.get("_title_only"): + continue + # Adaptive: only the top-ranked result is fully hydrated. + entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact") + if entry is not None: + results.append(entry) + for entry in results: + entry["link"] = _session_link(entry["session_id"], link_profile) + return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=( + "When referring the user to a session, write its `link` value " + "verbatim inline mid-sentence (it renders as a titled link) — never " + "as markdown, in backticks, on its own line, or next to the " + "title/id/date. To read more around a compact result, scroll: " + "session_search(session_id=..., around_message_id=match_message_id).")) def _resolve_profile_db(profile: str): diff --git a/tools/session_search_tool_common.py b/tools/session_search_tool_common.py deleted file mode 100644 index f9613796ea..0000000000 --- a/tools/session_search_tool_common.py +++ /dev/null @@ -1,229 +0,0 @@ -"""Shared helpers for the session_search tool: source classification, lineage -resolution, message storage state, and response shaping. Imported by -``tools.session_search_tool`` (which re-exports the names) and -``tools.session_search_tool_discover``.""" - -import json -import logging -from datetime import datetime -from typing import Any, Dict, List, Optional, Union - -from hermes_state_common import _RESET_END_REASONS - -# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool), -# delegate subagent runs, kanban workers — not the user's history. -_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool") - -# Searchable but DEMOTED below interactive sessions: cron sessions' repetitive -# vocabulary dominates bare BM25 and starves out the user's own sessions -# ("recall blindness"). Demoting keeps them reachable when they're the only match. -_DEMOTED_SESSION_SOURCES = ("cron",) - -# FTS rows scanned before dedup-by-lineage — well above the handful of distinct -# sessions a query returns, so interactive matches buried under cron hits are -# still in hand for the demotion pass. -_DISCOVER_SCAN_LIMIT = 300 - -# Raw FTS rows are only a discovery-plan input; the response hydrates its own -# anchored window and bookends after lineage dedup. -_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started") - -# Generated context-compaction handoff summaries (agent/context_compressor.py); -# excluded from bookends so huge compaction payloads aren't re-introduced. -_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:") - -# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the -# predecessor WITHOUT carrying its transcript forward — unlike compression -# continuations and live delegation children. Derived from the gateway set so -# this tool and the recovery fence cannot drift. -_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"} - - -def _quiet(fn, default, msg, *log_args, with_exc: bool = False): - """Call ``fn()``; on any exception debug-log *msg* (appending the exception - as a final ``%s`` arg when *with_exc*) and return *default*.""" - try: - return fn() - except Exception as e: - logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True) - return default - - -def _format_timestamp(ts: Union[int, float, str, None]) -> str: - """Unix timestamp (number or numeric string) or ISO string -> readable date. - "unknown" for None; str(ts) if conversion fails.""" - if ts is None: - return "unknown" - try: - value = ts - if isinstance(ts, str): - if not ts.replace(".", "").replace("-", "").isdigit(): - return ts - value = float(ts) - if isinstance(value, (int, float)): - return datetime.fromtimestamp(value).strftime("%B %d, %Y at %I:%M %p") - except (ValueError, OSError, OverflowError) as e: - logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True) - except Exception as e: - logging.debug("Unexpected error formatting timestamp %s: %s", ts, e, exc_info=True) - return str(ts) - - -def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]: - """The ``session_meta`` sub-object shared by read/scroll responses.""" - return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"), - "model": meta.get("model"), "title": meta.get("title")} - - -def _ok(**payload) -> str: - """Serialize a successful tool result (``success`` first, then *payload* in order).""" - return json.dumps({"success": True, **payload}, ensure_ascii=False) - - -def _is_compaction_summary(content: str) -> bool: - """Return True if *content* looks like a generated compaction handoff.""" - return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES) - - -def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]: - """Walk parent_session_id to the lineage root -> ``(root_id, has_compression_hop)``. - The flag distinguishes a compression-split lineage (parent content summarised - away) from a delegation lineage (child content still visible to the parent). - Falls back to ``(session_id, False)`` on errors.""" - if not session_id: - return session_id, False - visited: set[str] = set() - cur, has_compression = session_id, False - while cur and cur not in visited: - visited.add(cur) - s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True) - if not s: - break - if s.get("end_reason") == "compression": - has_compression = True - if not s.get("parent_session_id"): - break - cur = s["parent_session_id"] - return cur, has_compression - - -def _resolve_lineage(db, session_id: str) -> str: - """Return only the lineage root (ignores the compression hop).""" - return _resolve_to_parent(db, session_id)[0] - - -def _session_end_reason(db, session_id: str) -> Optional[str]: - """Return the session's ``end_reason``, or None if missing/unended/error.""" - if not session_id: - return None - try: - s = db.get_session(session_id) - except Exception: - return None - return (s.get("end_reason") or None) if s else None - - -def _session_left_live_context(db, session_id: str) -> bool: - """True when *session_id*'s transcript is no longer in anyone's live context: - ``compression`` (summarised into the continuation child) or a fresh reset - (child starts empty). Everything else stays excluded from same-lineage - recall — live delegation children (``end_reason is None``) are visible to - the parent agent, and ``branched`` parents were copied verbatim into the - branch child, so their content IS the current context.""" - end_reason = _session_end_reason(db, session_id) - return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS - - -def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]: - """Return the owning session and visibility flags for *message_id*.""" - if not message_id: - return None - - def _lookup(): - with db._lock: - return db._conn.execute( - "SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,) - ).fetchone() - - row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id) - return dict(row) if row is not None else None - - -def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool: - """Compaction archives are ``active=0, compacted=1`` (content summarised - away by archive_and_compact). Rewind/undo rows are ``active=0, compacted=0`` - and must stay hidden.""" - return state is not None and state["active"] == 0 and state["compacted"] == 1 - - -def _is_compacted_message(db, message_id) -> bool: - """True if *message_id* is a compaction-archived row — pre-compaction content - no longer in live context, so it should stay discoverable even on the - current session. False on any error (caller falls back to skipping).""" - return _is_compacted_state(_get_message_storage_state(db, message_id)) - - -def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None: - """Add a rebuild-progress note while the deferred FTS backfill is running, - so the agent can explain thin/slow results instead of treating them as - ground truth. No-op (never raises) when no rebuild is pending.""" - try: - status = db.fts_rebuild_status() - except Exception: - return - if status is None: - return - payload["index_rebuild"] = {"percent": status["percent"], "note": ( - f"The search index is rebuilding in the background ({status['percent']}% done, " - f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages " - f"may be incomplete until it finishes.")} - - -def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]: - """Stable-sort FTS rows so interactive sessions rank above automation. - BM25 order is preserved within each class; only cross-class order changes, - so a cron hit never displaces an interactive hit during lineage dedup.""" - return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0) - - -def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None, - max_content_len: Optional[int] = None) -> Dict[str, Any]: - """Slim a message row for the tool response. Keeps content even if empty - (absent content is meaningful — tool-call-only assistant turns). With - *max_content_len*, content is truncated and ``content_truncated`` / - ``original_content_chars`` added.""" - content = m.get("content") - if isinstance(content, str) and "\x1b" in content: - # Recalled messages can carry ANSI escapes (archived terminal output). - from tools.ansi_strip import strip_ansi - - content = strip_ansi(content) - original_chars = None - if max_content_len and content and len(content) > max_content_len: - original_chars = len(content) - content = content[:max_content_len] + "…" - entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")} - entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)}) - if anchor_id is not None and m.get("id") == anchor_id: - entry["anchor"] = True - if original_chars is not None: - entry["content_truncated"] = True - entry["original_content_chars"] = original_chars - return {k: v for k, v in entry.items() if v is not None or k == "content"} - - -def _session_link(session_id: str, profile: str = None) -> str: - """The reference the agent writes to point the user at a session — same - value the desktop composer emits, so it renders as a titled link. The - profile segment is omitted when it can't be named confidently (a bare id - still resolves, it just can't disambiguate across profiles).""" - name = (profile or "").strip() - if not name: - def _active(): - from hermes_cli.profiles import get_active_profile_name - - resolved = get_active_profile_name() - return "" if resolved == "custom" else resolved - - name = _quiet(_active, "", "get_active_profile_name failed for session link") - return f"@session:{name}/{session_id}" if name else f"@session:{session_id}" diff --git a/tools/session_search_tool_discover.py b/tools/session_search_tool_discover.py deleted file mode 100644 index 4867f3d421..0000000000 --- a/tools/session_search_tool_discover.py +++ /dev/null @@ -1,184 +0,0 @@ -"""Discovery shape of session_search: FTS5 query, title match, lineage dedup -and adaptive/full hydration of the surviving results.""" - -import json -import logging -from typing import Any, Dict, List, Optional - -from tools.registry import tool_error -from tools.session_search_tool_common import ( - _DISCOVER_SCAN_LIMIT, _DISCOVER_SEARCH_FIELDS, _HIDDEN_SESSION_SOURCES, _annotate_rebuild_status, - _format_timestamp, _is_compacted_message, _is_compaction_summary, _order_for_recall, _quiet, - _resolve_lineage, _resolve_to_parent, _session_left_live_context, _session_link, _shape_message) - - -def _normalize_title_query(query: str) -> str: - """Strip common quoting the model may include around a remembered title.""" - return query.strip().strip("`'\"") - - -def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]: - """Return a discovery-shaped result when the query matches a session title.""" - title_query = _normalize_title_query(query) - if not title_query: - return None - session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None, - "resolve_session_by_title failed for %r", title_query) - if not session_id: - return None - lineage_root = _resolve_lineage(db, session_id) - # Same-lineage title hits are in-context only while the session is live; - # /new-reset and compression-ended parents are not. - if ( - current_lineage_root - and lineage_root == current_lineage_root - and not _session_left_live_context(db, session_id)): - return None - - session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None, - "get_session failed for title match %s", session_id) or {} - if session_meta.get("source") in _HIDDEN_SESSION_SOURCES: - return None - messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id) - anchor_id = messages[0].get("id") if messages else None - view = {} - if anchor_id is not None: - view = _quiet(lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {}, - "get_anchored_view failed for title match %s/%s", session_id, anchor_id) - entry = { - "session_id": session_id, "when": _format_timestamp(session_meta.get("started_at")), - "source": session_meta.get("source", "unknown"), "model": session_meta.get("model") or "unknown", - "title": session_meta.get("title") or title_query, "matched_role": "session_title", - "match_message_id": anchor_id, - "snippet": f"Session title matched: {session_meta.get('title') or title_query}", - "bookend_start": [_shape_message(m) for m in (view.get("bookend_start") or messages[:3])], - "messages": [_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])], - "bookend_end": [_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])], - "messages_before": view.get("messages_before", 0), - "messages_after": view.get("messages_after", max(len(messages) - 5, 0)), - "detail": "full", "_lineage_root": lineage_root} - if lineage_root and lineage_root != session_id: - entry["parent_session_id"] = lineage_root - return entry - - -def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str: - payload = {"success": True, "mode": "discover", "query": query, "detail": detail, - "results": results, "count": len(results), **extra} - _annotate_rebuild_status(db, payload) - return json.dumps(payload, ensure_ascii=False) - - -def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None: - """Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*. - - The raw owning session_id stays on the row — only it pairs validly with the - FTS match id for the anchored window. Current-lineage hits are skipped - UNLESS the transcript left live context: compression-ended session, /new- - reset predecessor (hiding it made gateway recall blind after every /new), - or an in-place compacted row on the SAME session_id. A live delegation - child has end_reason=None, so it stays excluded. - """ - for r in raw_results: - if len(seen_sessions) >= limit: - break - raw_sid = r["session_id"] - resolved_sid, _ = _resolve_to_parent(db, raw_sid) - is_compacted_hit = _is_compacted_message(db, r.get("id")) - is_ended_session = _session_left_live_context(db, raw_sid) - if ( - current_lineage_root - and resolved_sid == current_lineage_root - and not (is_ended_session or is_compacted_hit)): - continue - if current_session_id and raw_sid == current_session_id and not is_compacted_hit: - continue - seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid}) - - -def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]: - return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or []) - if not _is_compaction_summary(m.get("content", ""))] - - -def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]: - """Build one discovery result from a surviving FTS row; None if the anchored - view can't be loaded (the hit is dropped).""" - hit_sid = match_info.get("session_id") or lineage_root - msg_id = match_info.get("id") - try: - view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3) - except Exception as e: - logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True) - return None - session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {} - full = result_detail == "full" - window_messages = view.get("window") or [] - if not full: - window_messages = [m for m in window_messages if m.get("id") == msg_id] - entry = { - "session_id": hit_sid, - "when": _format_timestamp(session_meta.get("started_at") or match_info.get("session_started")), - "source": session_meta.get("source") or match_info.get("source", "unknown"), - "model": session_meta.get("model") or match_info.get("model") or "unknown", - "title": session_meta.get("title") or None, "matched_role": match_info.get("role"), - "match_message_id": msg_id, "snippet": match_info.get("snippet") or "", - "bookend_start": _bookend(view, "bookend_start") if full else [], - "messages": [_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages], - "bookend_end": _bookend(view, "bookend_end") if full else [], - "messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", 0), - "detail": result_detail} - if lineage_root and lineage_root != hit_sid: - entry["parent_session_id"] = lineage_root - return entry - - -def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str], - detail: str, current_session_id: str = None, link_profile: str = None) -> str: - """Discovery shape: FTS5 plus adaptive or full result hydration.""" - role_list = role_filter if role_filter else ["user", "assistant"] - current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None - title_result = _title_match_result(db, query, current_lineage_root) - - try: - raw_results = db.search_messages( - query=query, role_filter=role_list, exclude_sources=list(_HIDDEN_SESSION_SOURCES), - limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS) - except Exception as e: - logging.error("FTS5 search failed: %s", e, exc_info=True) - return tool_error(f"Search failed: {e}", success=False) - - # Demote cron rows below interactive ones BEFORE dedup so a high-volume - # cron corpus can't starve the user's own sessions out of the top `limit`. - raw_results = _order_for_recall(raw_results) - - if not raw_results and not title_result: - return _discover_payload(db, query, detail, [], message=( - "No matching sessions found. FTS5 ANDs all terms by default — " - "broaden with OR (`alpha OR beta`), exact-match with quoted " - "phrases, exclude with NOT, or prefix-match with `deploy*`.")) - - seen_sessions: Dict[str, Dict[str, Any]] = {} - results = [] - if title_result: - title_lineage = title_result.pop("_lineage_root", None) - if title_lineage: - seen_sessions[title_lineage] = {"_title_only": True} - results.append(title_result) - _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) - - for lineage_root, match_info in seen_sessions.items(): - if match_info.get("_title_only"): - continue - # Adaptive: only the top-ranked result is fully hydrated. - entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact") - if entry is not None: - results.append(entry) - for entry in results: - entry["link"] = _session_link(entry["session_id"], link_profile) - return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=( - "When referring the user to a session, write its `link` value " - "verbatim inline mid-sentence (it renders as a titled link) — never " - "as markdown, in backticks, on its own line, or next to the " - "title/id/date. To read more around a compact result, scroll: " - "session_search(session_id=..., around_message_id=match_message_id)."))