refactor(tools): group I — collapse replace/remove into _edit, decode_json_or, compact docstrings keeping every WHY

This commit is contained in:
Teknium
2026-09-02 23:05:32 -07:00
parent 12412693e4
commit 634ed5f4fc
5 changed files with 195 additions and 297 deletions
+26 -33
View File
@@ -1,9 +1,8 @@
#!/usr/bin/env python3
"""Memory Tool - persistent curated memory (MEMORY.md = agent notes, USER.md =
user profile). Both enter the system prompt as a FROZEN snapshot at session
start; mid-session writes hit disk immediately but never change the prompt
(prefix cache stays intact). Single `memory` tool: add/replace/remove or a
batch `operations` list. The store lives in ``tools.memory_tool_store``."""
"""Memory Tool - persistent curated memory (MEMORY.md = agent notes, USER.md = user
profile). Both enter the system prompt as a FROZEN snapshot at session start;
mid-session writes hit disk but never change the prompt (prefix cache intact).
Single `memory` tool: add/replace/remove or a batch `operations` list."""
import copy
import json
@@ -16,8 +15,8 @@ from typing import Dict, Any, List, Optional, Tuple
from utils import is_truthy_value
from tools.registry import no_cache_check_fn
# fcntl is Unix-only; on Windows use msvcrt for file locking. MemoryStore reads
# these lazily from this module (tests inspect ``memory_tool.fcntl``).
# fcntl is Unix-only; Windows uses msvcrt. MemoryStore reads both lazily from
# this module (tests patch ``memory_tool.fcntl``).
msvcrt = None
try:
import fcntl
@@ -30,17 +29,15 @@ except ImportError:
logger = logging.getLogger(__name__)
# One tool-definition pass must use one config decision for both availability
# and the dynamic target schema. ContextVar keeps concurrent profile/session
# builds isolated while letting the check_fn result flow to the immediately
# following dynamic_schema_overrides call in ToolRegistry.get_definitions().
# One tool-definition pass must use ONE config decision for availability and the
# dynamic target schema: the check_fn result flows to the immediately following
# dynamic_schema_overrides call; ContextVar isolates concurrent profile builds.
_memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar(
"memory_surface_flags", default=None)
def get_memory_dir() -> Path:
"""Return the profile-scoped memories directory (resolved per call so
HERMES_HOME/profile switches after import are respected)."""
"""Profile-scoped memories dir, resolved per call (HERMES_HOME may switch after import)."""
return get_hermes_home() / "memories"
@@ -50,9 +47,9 @@ from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports)
def load_on_disk_store() -> "MemoryStore":
"""Fresh on-disk MemoryStore with configured limits/flags, for contexts with
no live agent (gateway, Desktop, bare CLI ``/memory``) so approvals enforce
the SAME caps as ``agent_init``. Defaults if config can't load; never raises."""
"""Fresh on-disk MemoryStore with configured limits/flags for contexts with no
live agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps
as ``agent_init``. Falls back to defaults if config can't load; never raises."""
try:
from hermes_cli.config import load_config
@@ -72,9 +69,8 @@ def load_on_disk_store() -> "MemoryStore":
# -- Write-approval gate --
def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]:
"""Run the memory write gate. Returns a JSON tool-result string when the
write must NOT proceed (blocked, or staged for approval), None to proceed.
If the gate module can't load, fail open rather than block all writes."""
"""JSON tool-result string when the write must NOT proceed (blocked or staged
for approval), None to proceed. Fails open if the gate module can't load."""
try:
from tools import write_approval as wa
except Exception:
@@ -126,11 +122,10 @@ def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Op
# -- Tool entry point --
def _validate_single_op(store, action, target, content, old_text) -> Optional[str]:
"""Validate required params BEFORE the gate so an invalid write is rejected
now rather than staged and failing at approve time. A missing ``old_text``
is recoverable (it can't be schema-required — needs a combinator the Codex
backend rejects, see test_memory_tool_schema.py — and some clients omit it),
so return the current inventory plus a retry instruction, not a dead-end."""
"""Validate BEFORE the gate so an invalid write is rejected now, not at approve
time. Missing ``old_text`` is recoverable (it can't be schema-required — needs a
combinator the Codex backend rejects — and some clients omit it): return the
current inventory plus a retry instruction instead of a dead-end."""
if action == "add" and not content:
return tool_error("Content is required for 'add' action.", success=False)
if action in ("replace", "remove") and not old_text:
@@ -162,10 +157,9 @@ def memory_tool(
new_text: str = None,
operations: Optional[List[Dict[str, Any]]] = None,
store: Optional[MemoryStore] = None) -> str:
"""Tool entry point; returns a JSON string. Single op (action + content /
old_text) or batch (``operations`` applied atomically against the final
budget). ``new_text`` aliases ``content`` — callers mirror ``old_text``
with it (patch-tool shape), which used to leave ``content`` empty."""
"""Tool entry point; returns a JSON string. Single op (action + content/old_text)
or batch (``operations``, atomic against the final budget). ``new_text``
aliases ``content`` — callers mirror ``old_text`` with it (patch-tool shape)."""
if store is None:
return tool_error("Memory is not available. It may be disabled in config or this environment.", success=False)
@@ -191,8 +185,7 @@ def memory_tool(
invalid = _validate_single_op(store, action, target, content, old_text)
if invalid is not None:
return invalid
# Approval gate: when on, stages the write (background/gateway) or prompts
# inline (interactive CLI); when off (default) passes straight through.
# Approval gate: stages (background/gateway) or prompts inline (CLI); off by default.
gate_result = _apply_write_gate(action, target, content, old_text)
if gate_result is not None:
return gate_result
@@ -200,9 +193,9 @@ def memory_tool(
def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Normalized ``memory`` config section ({} when missing/malformed → flags
default to enabled). ``agent_init`` consumes the same section so tool
availability and store construction cannot diverge."""
"""Normalized ``memory`` config section ({} when missing/malformed → flags default
to enabled). ``agent_init`` reads the same section so availability and store
construction cannot diverge."""
if config is None:
try:
from hermes_cli.config import load_config_readonly
+75 -115
View File
@@ -1,7 +1,7 @@
"""MemoryStore — bounded, file-backed curated memory (MEMORY.md / USER.md).
Entries are joined by ``ENTRY_DELIMITER``; budgets are in chars (model-
independent). Module state that tests monkeypatch (``get_memory_dir``,
``fcntl``/``msvcrt``) stays in ``tools.memory_tool`` and is read lazily."""
Entries are joined by ``ENTRY_DELIMITER``; budgets are in chars (model-independent).
Module state that tests monkeypatch (``get_memory_dir``, ``fcntl``/``msvcrt``) stays
in ``tools.memory_tool`` and is read lazily."""
import logging
import time
@@ -14,9 +14,8 @@ from tools.threat_patterns import first_threat_message as _first_threat_message
logger = logging.getLogger("tools.memory_tool")
# System-prompt block header prefixes rendered by MemoryStore._render_block.
# agent/conversation_compression.py uses them to detect a leftover block for a
# target whose entries have since been emptied — keep in lockstep.
# Block header prefixes rendered by _render_block; agent/conversation_compression.py
# matches them to detect a leftover block for an emptied target — keep in lockstep.
MEMORY_BLOCK_HEADERS = {
"memory": "MEMORY (your personal notes)", "user": "USER PROFILE (who the user is)"}
@@ -29,15 +28,13 @@ def _memory_dir() -> Path:
def _scan_memory_content(content: str) -> Optional[str]:
"""Scan for injection/exfil patterns ("strict" scope: memory enters the system
prompt as a frozen snapshot, so a poisoned entry persists across sessions).
Returns the error string if blocked."""
"""Error string if *content* matches injection/exfil patterns. Strict scope:
memory enters the system prompt, so a poisoned entry persists across sessions."""
return _first_threat_message(content, scope="strict")
def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]:
"""Error dict for external drift: the on-disk file wouldn't round-trip through
the parser/serializer, so flushing would discard externally added content."""
"""External drift: the file wouldn't round-trip, so flushing would discard content."""
return {
"success": False,
"error": (
@@ -55,8 +52,7 @@ def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]:
def _read_failed_error(path: "Path") -> Dict[str, Any]:
"""Error dict for an unreadable (but existing) memory file: treating it as
empty and saving would rewrite the whole file from ``[]`` — wiping memory."""
"""Existing-but-unreadable file: saving from an assumed-empty view would wipe it."""
return {"success": False, "error": (
f"Refusing to write {path.name}: the file exists on disk but could not be read right now "
f"(temporarily locked by another program, a permission change, invalid/corrupt text encoding, "
@@ -80,10 +76,9 @@ class MemoryStore:
``_system_prompt_snapshot`` is frozen at load time (prefix-cache stable);
``memory_entries`` / ``user_entries`` are live state persisted to disk."""
# After this many failed consolidation attempts (overflow / zero-match) in
# ONE turn, return a terminal "save skipped" result instead of "retry in
# this turn", so a fragile replace/add can't loop the turn to budget
# exhaustion and suppress the user's reply.
# Failed consolidation attempts (overflow / zero-match) allowed per turn before
# a TERMINAL "save skipped" result, so a fragile replace/add can't loop the turn
# to budget exhaustion and suppress the user's reply.
_MAX_CONSOLIDATION_FAILURES_PER_TURN = 3
def __init__(self, memory_char_limit: int = 2200, user_char_limit: int = 1375, *,
@@ -95,9 +90,7 @@ class MemoryStore:
self.memory_enabled = memory_enabled
self.user_profile_enabled = user_profile_enabled
self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""}
# Per-turn counter of failed at-capacity consolidation attempts; reset
# at each turn boundary by reset_consolidation_failures().
self._consolidation_failures = 0
self._consolidation_failures = 0 # per turn; reset by reset_consolidation_failures()
def target_enabled(self, target: str) -> bool:
"""Return whether this session's selected built-in store is writable."""
@@ -108,10 +101,8 @@ class MemoryStore:
self._consolidation_failures = 0
def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]:
"""Count an at-capacity consolidation failure. Under the per-turn cap return
``response`` unchanged (it says how to retry); past it return a TERMINAL
result so the model stops looping — a failed memory side effect must never
block the turn's reply."""
"""Count a consolidation failure: under the per-turn cap return ``response``
(it says how to retry); past it a TERMINAL result so the model stops looping."""
self._consolidation_failures += 1
if self._consolidation_failures <= self._MAX_CONSOLIDATION_FAILURES_PER_TURN:
return response
@@ -122,10 +113,9 @@ class MemoryStore:
def load_from_disk(self):
"""Load MEMORY.md / USER.md and capture the frozen system-prompt snapshot.
Threat hits are replaced in the SNAPSHOT only by a ``[BLOCKED: …]``
placeholder; live lists keep the raw text so the user can see and remove
poisoned entries (dropping them silently would hide the attack).
Scanning is deterministic from disk bytes, so the snapshot stays stable."""
Threat hits are replaced by a ``[BLOCKED: …]`` placeholder in the SNAPSHOT only;
live lists keep the raw text so the user can see and remove poisoned entries
(dropping them silently would hide the attack)."""
mem_dir = _memory_dir()
mem_dir.mkdir(parents=True, exist_ok=True)
# Deduplicate (order-preserving, first occurrence wins).
@@ -139,9 +129,8 @@ class MemoryStore:
@staticmethod
def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]:
"""*entries* with any threat-matching entry replaced by a ``[BLOCKED: …]``
placeholder (strict scope, same as writes). Empty or already-blocked
entries pass through unchanged."""
"""*entries* with threat matches replaced by a ``[BLOCKED: …]`` placeholder
(strict scope, same as writes); empty / already-blocked entries pass through."""
from tools.threat_patterns import scan_for_threats
sanitized: List[str] = []
@@ -158,8 +147,8 @@ class MemoryStore:
@staticmethod
@contextmanager
def _file_lock(path: Path):
"""Exclusive lock for read-modify-write safety, on a separate .lock file
so the memory file itself can still be atomically replaced."""
"""Exclusive lock on a separate .lock file so the memory file itself can
still be atomically replaced."""
from tools import memory_tool as _mt # fcntl/msvcrt live (and are patched) there
fcntl, msvcrt = _mt.fcntl, _mt.msvcrt
@@ -192,13 +181,11 @@ class MemoryStore:
return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md")
def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]:
"""Re-read entries from disk (under file lock) before mutating; return the
abort error dict, or None to proceed. Aborts on external drift (flushing
would discard un-roundtrippable content) and when the file exists but can't
be read (rewriting from an assumed-empty view would wipe it — even
"""Re-read entries from disk (under lock) before mutating; return the abort
error dict or None. Aborts on external drift (flushing would discard
un-roundtrippable content) and on an existing-but-unreadable file (even
append-only ``add`` rewrites the whole file). Drift check and parse use the
SAME raw snapshot — a failed second read used to count as "no drift".
*skip_drift* skips the round-trip check (``add``)."""
SAME raw snapshot — a failed second read used to count as "no drift"."""
path = self._path_for(target)
raw, read_ok = self._read_raw_checked(path)
if not read_ok:
@@ -227,9 +214,14 @@ class MemoryStore:
def _usage(self, target: str) -> str:
return f"{self._char_count(target):,}/{self._char_limit(target):,}"
def _usage_pct(self, target: str, current: int) -> str:
"""``"<pct>% — <current>/<limit> chars"`` for the given target."""
limit = self._char_limit(target)
pct = min(100, int((current / limit) * 100)) if limit > 0 else 0
return f"{pct}% — {current:,}/{limit:,} chars"
def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]:
"""Consolidation failure showing the live entries so the model can decide
what to consolidate."""
"""Consolidation failure carrying the live entries so the model can consolidate."""
return self._consolidation_failure({"success": False, "error": message,
"current_entries": self._entries_for(target), "usage": self._usage(target)})
@@ -261,9 +253,8 @@ class MemoryStore:
if scan_error:
return {"success": False, "error": scan_error}
with self._file_lock(self._path_for(target)):
# Append-only: skip the drift guard (appending never clobbers
# un-roundtrippable content), but still refuse on a failed read —
# add rewrites the WHOLE file from the parsed entries.
# Append-only: skip the drift guard (appending never clobbers foreign
# content) but still refuse a failed read — add rewrites the WHOLE file.
err = self._reload_or_error(target, skip_drift=True)
if err:
return err
@@ -282,23 +273,35 @@ class MemoryStore:
def replace(self, target: str, old_text: str, new_content: str) -> Dict[str, Any]:
"""Find entry containing old_text substring, replace it with new_content."""
old_text = old_text.strip()
new_content = new_content.strip()
if not old_text:
if not old_text.strip():
return {"success": False, "error": "old_text cannot be empty."}
if not new_content:
return {"success": False, "error": "new_content cannot be empty. Use 'remove' to delete entries."}
scan_error = _scan_memory_content(new_content)
if scan_error:
return {"success": False, "error": scan_error}
return self._edit(target, old_text.strip(), new_content)
def remove(self, target: str, old_text: str) -> Dict[str, Any]:
"""Remove the entry containing old_text substring."""
if not old_text.strip():
return {"success": False, "error": "old_text cannot be empty."}
return self._edit(target, old_text.strip(), None)
def _edit(self, target: str, old_text: str, new_content: Optional[str]) -> Dict[str, Any]:
"""Locked replace (``new_content`` set) or remove (None) of the unique entry matching *old_text*."""
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target)
if err:
return err
idx, err = self._locate(target, old_text, "replace")
idx, err = self._locate(target, old_text, "replace" if new_content else "remove")
if err:
return err
entries = self._entries_for(target)
if new_content is None:
entries.pop(idx)
return self._commit(target, entries, "Entry removed.")
limit = self._char_limit(target)
new_total = len(ENTRY_DELIMITER.join(entries[:idx] + [new_content] + entries[idx + 1:]))
if new_total > limit:
@@ -309,24 +312,6 @@ class MemoryStore:
entries[idx] = new_content
return self._commit(target, entries, "Entry replaced.")
def remove(self, target: str, old_text: str) -> Dict[str, Any]:
"""Remove the entry containing old_text substring."""
old_text = old_text.strip()
if not old_text:
return {"success": False, "error": "old_text cannot be empty."}
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target)
if err:
return err
idx, err = self._locate(target, old_text, "remove")
if err:
return err
entries = self._entries_for(target)
entries.pop(idx)
return self._commit(target, entries, "Entry removed.")
# -- Batch --
@staticmethod
def _apply_batch_op(working: List[str], act: str, content: str, old_text: str, pos: str) -> Optional[str]:
"""Apply one batch op to *working* in place; return an error message or None."""
@@ -354,17 +339,14 @@ class MemoryStore:
return None
def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]:
"""Apply a sequence of add/replace/remove ops to one target atomically.
Ops are validated and applied against the FINAL budget only — so one call
can free space (remove/replace) and add new entries without the multi-turn
consolidate-then-retry dance. All-or-nothing: if any op is malformed,
doesn't match, or the net result exceeds the char limit, NOTHING is written
and the first failure plus live state is returned."""
"""Apply add/replace/remove ops to one target atomically against the FINAL
budget, so one call can free space and add entries. All-or-nothing: any
malformed / unmatched op or an over-limit result writes NOTHING and returns
the first failure plus live state."""
if not operations:
return {"success": False, "error": "operations list is empty."}
# Scan every add/replace content BEFORE touching disk -- one poisoned
# op rejects the whole batch.
# Scan every add/replace content BEFORE touching disk -- one poisoned op rejects the batch.
for i, op in enumerate(operations):
op = op or {}
scan_error = op.get("action") in {"add", "replace"} and op.get("content") and _scan_memory_content(op["content"])
@@ -386,7 +368,6 @@ class MemoryStore:
pos = f"Operation {i + 1} ({act or 'unknown'})"
msg = self._apply_batch_op(working, act, content, old_text, pos)
if msg:
# Batch-abort error reporting live (uncommitted) state.
return self._failure_with_entries(
target, msg + " No operations were applied (batch is all-or-nothing).")
# Budget check against the FINAL state only.
@@ -399,38 +380,24 @@ class MemoryStore:
return self._commit(target, working, f"Applied {len(operations)} operation(s).")
def format_for_system_prompt(self, target: str) -> Optional[str]:
"""The frozen load-time snapshot for system-prompt injection (NOT live
state — mid-session writes don't affect it, preserving the prefix cache).
None if the snapshot is empty."""
"""Frozen load-time snapshot for the system prompt (NOT live state — mid-session
writes don't touch it, preserving the prefix cache); None if empty."""
return self._system_prompt_snapshot.get(target, "") or None
# -- Internal helpers --
@staticmethod
def _previews(entries: List[str], width: int = 80) -> List[str]:
"""Truncated one-line previews of entries for error feedback."""
return [e[:width] + ("..." if len(e) > width else "") for e in entries]
def _usage_pct(self, target: str, current: int) -> str:
"""``"<pct>% — <current>/<limit> chars"`` for the given target."""
limit = self._char_limit(target)
pct = min(100, int((current / limit) * 100)) if limit > 0 else 0
return f"{pct}% — {current:,}/{limit:,} chars"
def _success_response(self, target: str, message: str = None) -> Dict[str, Any]:
# A successful write means the consolidation loop made progress, so the
# per-turn failure budget resets (the cap counts consecutive failures).
# A successful write is progress: reset the per-turn (consecutive) failure budget.
self._consolidation_failures = 0
# Intentionally TERMINAL and without the entries list: echoing entries
# invites the model to "find more to fix" and re-issue the same ops.
# Entries are only shown on error/over-budget paths.
resp = {"success": True, "done": True, "target": target,
# TERMINAL and WITHOUT the entries list: echoing entries invites the model to
# "find more to fix" and re-issue the same ops. Entries only appear on errors.
return {"success": True, "done": True, "target": target,
"usage": self._usage_pct(target, self._char_count(target)),
"entry_count": len(self._entries_for(target))}
if message:
resp["message"] = message
resp["note"] = "Write saved. This update is complete — do not repeat it."
return resp
"entry_count": len(self._entries_for(target)), **({"message": message} if message else {}),
"note": "Write saved. This update is complete — do not repeat it."}
def _render_block(self, target: str, entries: List[str]) -> str:
"""Render a system prompt block with header and usage indicator."""
@@ -443,12 +410,10 @@ class MemoryStore:
@staticmethod
def _read_raw_checked(path: Path) -> Tuple[str, bool]:
"""Read raw text as ``(raw, read_ok)``. ``read_ok`` is False ONLY when the
file EXISTS but can't be read (absent file → ``("", True)``). Invalid UTF-8
counts as unreadable; decoding stays STRICT because ``errors="replace"``
would hand callers a lossy view that a save then persists over the real
bytes. ``utf-8-sig`` strips a Notepad BOM that otherwise glues U+FEFF onto
the first entry forever."""
"""``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't
be read (absent → ``("", True)``). Decoding stays STRICT: ``errors="replace"``
would hand callers a lossy view that a save then persists. ``utf-8-sig`` strips
a Notepad BOM that otherwise glues U+FEFF onto the first entry forever."""
if not path.exists():
return "", True
try:
@@ -458,24 +423,20 @@ class MemoryStore:
@staticmethod
def _parse_entries(raw: str) -> List[str]:
"""Split raw memory-file text into stripped, non-empty entries. Splits on
the full ENTRY_DELIMITER so a bare "§" inside an entry is preserved."""
"""Stripped, non-empty entries; splits on the FULL delimiter so a bare "§" survives."""
return [e for e in (x.strip() for x in raw.split(ENTRY_DELIMITER)) if e]
@staticmethod
def _read_file(path: Path) -> List[str]:
"""Read a memory file into entries (empty list on any error). Only for
read-only callers (``load_from_disk``, learning_mutations); mutation paths
must use ``_read_raw_checked`` so they can refuse to overwrite an
unreadable file."""
"""Entries of a memory file ([] on any error). Read-only callers only
(``load_from_disk``, learning_mutations); mutation paths must use
``_read_raw_checked`` so they can refuse to overwrite an unreadable file."""
return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0])
def _detect_external_drift(self, target: str, raw: str) -> Optional[str]:
"""Backup-path string if *raw* (the caller's checked-read snapshot) shows
external drift, else None. Signals: round-trip mismatch, or one parsed
entry exceeding the whole-file char limit (no tool-written entry can — an
external writer appended free-form content a flush would truncate). The
file is snapshotted to ``.bak.<ts>`` so the operator can recover it."""
"""Backup path if *raw* shows external drift, else None. Signals: round-trip
mismatch, or one entry over the whole-file limit (no tool-written entry can —
an external writer appended free-form text). Snapshots to ``.bak.<ts>``."""
if not raw.strip():
return None
parsed = self._parse_entries(raw)
@@ -491,8 +452,7 @@ class MemoryStore:
@staticmethod
def _write_file(path: Path, entries: List[str]):
"""Atomic temp-file + rename: readers see the old or the new complete
file, never a truncated one."""
"""Atomic temp-file + rename: readers never see a truncated file."""
try:
atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_")
except OSError as e:
+13 -32
View File
@@ -31,11 +31,8 @@ class MicrosoftGraphTokenError(MicrosoftGraphAuthError):
def format_graph_error(error: Any) -> str | None:
"""Render Graph's ``{"error": {"code", "message"}}`` (or bare-string ``error``) body.
Shared by the token endpoint and the REST client so both surface
``code: message`` the same way. ``None`` means the shape was unusable.
"""
"""Render Graph's ``{"error": {"code", "message"}}`` (or bare-string ``error``) as
``code: message``; shared by token endpoint and REST client. None if unusable."""
if isinstance(error, str):
return error
if not isinstance(error, dict):
@@ -114,18 +111,12 @@ class MicrosoftGraphTokenProvider:
self._cached_token = None
def inspect_token_health(self) -> dict[str, Any]:
cached = self._cached_token
return {
"configured": True,
"tenant_id": self.credentials.tenant_id,
"client_id": self.credentials.client_id,
"scope": self.credentials.scope,
"authority_url": self.credentials.authority_url,
"token_url": self.credentials.token_url,
"cached": bool(cached),
"expires_in_seconds": cached.expires_in_seconds if cached else None,
"is_expired": cached.is_expired(skew_seconds=0) if cached else None,
"refresh_skew_seconds": self.skew_seconds}
cached, creds = self._cached_token, self.credentials
return {"configured": True, "tenant_id": creds.tenant_id, "client_id": creds.client_id,
"scope": creds.scope, "authority_url": creds.authority_url, "token_url": creds.token_url,
"cached": bool(cached), "expires_in_seconds": cached.expires_in_seconds if cached else None,
"is_expired": cached.is_expired(skew_seconds=0) if cached else None,
"refresh_skew_seconds": self.skew_seconds}
def _fresh_cached(self) -> CachedAccessToken | None:
"""The cached token unless it expires within ``skew_seconds``."""
@@ -145,11 +136,8 @@ class MicrosoftGraphTokenProvider:
return self._cached_token.access_token
async def _fetch_access_token(self) -> CachedAccessToken:
data = {
"grant_type": "client_credentials",
"client_id": self.credentials.client_id,
"client_secret": self.credentials.client_secret,
"scope": self.credentials.scope}
data = {"grant_type": "client_credentials", "client_id": self.credentials.client_id,
"client_secret": self.credentials.client_secret, "scope": self.credentials.scope}
async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client:
response = await client.post(
self.credentials.token_url, data=data,
@@ -163,7 +151,6 @@ class MicrosoftGraphTokenProvider:
payload = response.json()
except ValueError as exc:
raise MicrosoftGraphTokenError("Microsoft Graph token response was not valid JSON.") from exc
access_token = str(payload.get("access_token") or "").strip()
if not access_token:
raise MicrosoftGraphTokenError("Microsoft Graph token response did not include access_token.")
@@ -171,9 +158,7 @@ class MicrosoftGraphTokenProvider:
expires_in_seconds = int(payload.get("expires_in"))
except (TypeError, ValueError) as exc:
raise MicrosoftGraphTokenError(
"Microsoft Graph token response did not include a valid expires_in."
) from exc
"Microsoft Graph token response did not include a valid expires_in.") from exc
return CachedAccessToken(
access_token=access_token,
token_type=str(payload.get("token_type") or "Bearer").strip() or "Bearer",
@@ -181,16 +166,12 @@ class MicrosoftGraphTokenProvider:
def _extract_error_detail(response: httpx.Response) -> str:
"""Best human-readable detail from a token-endpoint error body.
The OAuth endpoint prefers ``error_description``; fall back to the
Graph-style ``error`` object/string, then a bare ``code``, then raw text.
"""
"""Best human-readable detail from a token-endpoint error body: ``error_description``,
then the Graph-style ``error`` object/string, then a bare ``code``, then raw text."""
try:
payload = response.json()
except ValueError:
return response.text.strip() or "unknown error"
if isinstance(payload, dict):
if isinstance(payload.get("error_description"), str):
return payload["error_description"]
+14 -21
View File
@@ -38,12 +38,9 @@ class MicrosoftGraphAPIError(MicrosoftGraphClientError):
class MicrosoftGraphClient:
"""Minimal async Microsoft Graph client with retries and pagination.
Retry policy (shared by JSON requests and streaming downloads): transport
errors back off exponentially; 401 clears the token cache and refetches;
429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``.
"""
"""Minimal async Graph client. Retry policy (JSON requests and streaming downloads
alike): transport errors back off exponentially; 401 clears the token cache and
refetches; 429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``."""
def __init__(
self, token_provider: MicrosoftGraphTokenProvider, *,
@@ -71,15 +68,15 @@ class MicrosoftGraphClient:
async def patch_json(self, path: str, *, json_body: Any | None = None, headers: Headers = None) -> Any:
response = await self._request("PATCH", path, json_body=json_body, headers=headers)
if response.status_code == 204 or not response.content:
return {}
return self._decode_json(response)
return self._decode_json_or(response, {})
async def delete(self, path: str, *, headers: Headers = None) -> dict[str, Any]:
response = await self._request("DELETE", path, headers=headers)
if response.status_code == 204 or not response.content:
return {"deleted": True, "status_code": response.status_code}
return self._decode_json(response)
return self._decode_json_or(response, {"deleted": True, "status_code": response.status_code})
def _decode_json_or(self, response: httpx.Response, empty: Any) -> Any:
"""*empty* for a 204 / bodiless response, else the decoded JSON body."""
return empty if response.status_code == 204 or not response.content else self._decode_json(response)
async def collect_paginated(
self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]:
@@ -102,9 +99,8 @@ class MicrosoftGraphClient:
async def download_to_file(
self, path: str, destination: str | Path, *, headers: Headers = None, chunk_size: int = 65536
) -> dict[str, Any]:
"""Download a Graph resource to disk, streaming the body chunk-by-chunk
(recordings and other large artifacts never need to fit in memory).
Written to a ``.part`` file and renamed into place only on success."""
"""Stream a Graph resource to disk chunk-by-chunk (large recordings never
fit in memory); written to ``.part`` and renamed into place only on success."""
url = self._resolve_url(path)
target = Path(destination)
target.parent.mkdir(parents=True, exist_ok=True)
@@ -145,12 +141,9 @@ class MicrosoftGraphClient:
self, method: str, url: str, accept: str, json_body: Any | None, headers: Headers,
perform: Callable[[httpx.AsyncClient, dict[str, str]], Awaitable[tuple[httpx.Response, Any]]],
kind: str) -> Any:
"""Run ``perform`` (returning ``(response, result)``) under the retry policy.
``kind`` ("request"/"download") only labels the transport-failure messages.
A ``MicrosoftGraphAPIError`` for the failing status is raised once retries
are exhausted or the status is not retryable; only a 401 forces a token refresh.
"""
"""Run ``perform`` (-> ``(response, result)``) under the retry policy. ``kind``
only labels transport-failure messages. Raises ``MicrosoftGraphAPIError`` once
retries are exhausted or the status is not retryable; only 401 forces a token refresh."""
attempt = 0
last_error: Exception | None = None
+67 -96
View File
@@ -19,34 +19,29 @@ from hermes_state_common import _RESET_END_REASONS
# delegate subagent runs, kanban workers — not the user's history.
_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool")
# Searchable but DEMOTED below interactive sessions: cron sessions' repetitive
# vocabulary dominates bare BM25 and starves out the user's own sessions
# ("recall blindness"). Demoting keeps them reachable when they're the only match.
# Searchable but DEMOTED below interactive sessions: cron vocabulary dominates bare
# BM25 and starves out the user's own sessions ("recall blindness").
_DEMOTED_SESSION_SOURCES = ("cron",)
# FTS rows scanned before dedup-by-lineage — well above the handful of distinct
# sessions a query returns, so interactive matches buried under cron hits are
# still in hand for the demotion pass.
# FTS rows scanned before dedup-by-lineage — well above the distinct sessions a query
# returns, so interactive matches buried under cron hits survive the demotion pass.
_DISCOVER_SCAN_LIMIT = 300
# Raw FTS rows are only a discovery-plan input; the response hydrates its own
# anchored window and bookends after lineage dedup.
# Raw FTS rows are only a plan input; the response hydrates its own window/bookends.
_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started")
# Generated context-compaction handoff summaries (agent/context_compressor.py);
# excluded from bookends so huge compaction payloads aren't re-introduced.
# Compaction handoff summaries (agent/context_compressor.py); excluded from bookends.
_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:")
# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the
# predecessor WITHOUT carrying its transcript forward — unlike compression
# continuations and live delegation children. Derived from the gateway set so
# this tool and the recovery fence cannot drift.
# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the predecessor
# WITHOUT carrying its transcript forward — unlike compression continuations and
# live delegation children. Derived from the gateway set so the two cannot drift.
_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"}
def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
"""Call ``fn()``; on any exception debug-log *msg* (appending the exception
as a final ``%s`` arg when *with_exc*) and return *default*."""
"""``fn()``, or *default* after debug-logging *msg* (exception appended as a
final ``%s`` arg when *with_exc*) on any exception."""
try:
return fn()
except Exception as e:
@@ -55,8 +50,8 @@ def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
def _format_timestamp(ts: Union[int, float, str, None]) -> str:
"""Unix timestamp (number or numeric string) or ISO string -> readable date.
"unknown" for None; str(ts) if conversion fails."""
"""Unix timestamp (number / numeric string) -> readable date; ISO strings pass
through; "unknown" for None; str(ts) if conversion fails."""
if ts is None:
return "unknown"
try:
@@ -89,10 +84,9 @@ def _is_compaction_summary(content: str) -> bool:
def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]:
"""Walk parent_session_id to the lineage root -> ``(root_id, has_compression_hop)``.
The flag distinguishes a compression-split lineage (parent content summarised
away) from a delegation lineage (child content still visible to the parent).
Falls back to ``(session_id, False)`` on errors."""
"""Walk parent_session_id to the root -> ``(root_id, has_compression_hop)``. The
flag separates a compression-split lineage (parent summarised away) from a
delegation lineage (child still visible to the parent). Errors -> ``(session_id, False)``."""
if not session_id:
return session_id, False
visited: set[str] = set()
@@ -127,12 +121,10 @@ def _session_end_reason(db, session_id: str) -> Optional[str]:
def _session_left_live_context(db, session_id: str) -> bool:
"""True when *session_id*'s transcript is no longer in anyone's live context:
``compression`` (summarised into the continuation child) or a fresh reset
(child starts empty). Everything else stays excluded from same-lineage
recall — live delegation children (``end_reason is None``) are visible to
the parent agent, and ``branched`` parents were copied verbatim into the
branch child, so their content IS the current context."""
"""True when the transcript left everyone's live context: ``compression``
(summarised into the child) or a fresh reset (child starts empty). Live delegation
children (``end_reason is None``) and ``branched`` parents (copied verbatim into
the branch) ARE the current context, so they stay excluded from recall."""
end_reason = _session_end_reason(db, session_id)
return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS
@@ -153,23 +145,20 @@ def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]:
def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool:
"""Compaction archives are ``active=0, compacted=1`` (content summarised
away by archive_and_compact). Rewind/undo rows are ``active=0, compacted=0``
and must stay hidden."""
"""Compaction archives are ``active=0, compacted=1``; rewind/undo rows are
``active=0, compacted=0`` and must stay hidden."""
return state is not None and state["active"] == 0 and state["compacted"] == 1
def _is_compacted_message(db, message_id) -> bool:
"""True if *message_id* is a compaction-archived row — pre-compaction content
no longer in live context, so it should stay discoverable even on the
current session. False on any error (caller falls back to skipping)."""
"""True for a compaction-archived row — content no longer in live context, so
it stays discoverable even on the current session. False on any error."""
return _is_compacted_state(_get_message_storage_state(db, message_id))
def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None:
"""Add a rebuild-progress note while the deferred FTS backfill is running,
so the agent can explain thin/slow results instead of treating them as
ground truth. No-op (never raises) when no rebuild is pending."""
"""Note rebuild progress while the deferred FTS backfill runs so the agent can
explain thin results instead of treating them as ground truth. Never raises."""
status = _quiet(db.fts_rebuild_status, None, "fts_rebuild_status failed")
if status is None:
return
@@ -180,21 +169,17 @@ def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None:
def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Stable-sort FTS rows so interactive sessions rank above automation.
BM25 order is preserved within each class; only cross-class order changes,
so a cron hit never displaces an interactive hit during lineage dedup."""
"""Stable-sort so interactive sessions rank above automation; BM25 order is
kept within each class, so a cron hit never displaces an interactive one."""
return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0)
def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None,
max_content_len: Optional[int] = None) -> Dict[str, Any]:
"""Slim a message row for the tool response. Keeps content even if empty
(absent content is meaningful — tool-call-only assistant turns). With
*max_content_len*, content is truncated and ``content_truncated`` /
``original_content_chars`` added."""
"""Slim a message row. Keeps ``content`` even when empty (tool-call-only
assistant turns); with *max_content_len* truncates and flags it."""
content = m.get("content")
if isinstance(content, str) and "\x1b" in content:
# Recalled messages can carry ANSI escapes (archived terminal output).
if isinstance(content, str) and "\x1b" in content: # archived terminal output carries ANSI
from tools.ansi_strip import strip_ansi
content = strip_ansi(content)
@@ -208,10 +193,9 @@ def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None,
def _session_link(session_id: str, profile: str = None) -> str:
"""The reference the agent writes to point the user at a session — same
value the desktop composer emits, so it renders as a titled link. The
profile segment is omitted when it can't be named confidently (a bare id
still resolves, it just can't disambiguate across profiles)."""
"""The reference the agent writes for a session — same value the desktop composer
emits, so it renders as a titled link. The profile segment is omitted when it
can't be named confidently (a bare id still resolves, just not across profiles)."""
name = (profile or "").strip()
if not name:
def _active():
@@ -291,14 +275,10 @@ def _discover_payload(db, query: str, detail: str, results: list, **extra) -> st
def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None:
"""Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*.
The raw owning session_id stays on the row — only it pairs validly with the
FTS match id for the anchored window. Current-lineage hits are skipped
UNLESS the transcript left live context: compression-ended session, /new-
reset predecessor (hiding it made gateway recall blind after every /new),
or an in-place compacted row on the SAME session_id. A live delegation
child has end_reason=None, so it stays excluded.
"""
The raw owning session_id stays on the row — only it pairs validly with the FTS
match id. Current-lineage hits are skipped UNLESS the transcript left live
context (compression-ended, /new-reset predecessor, or an in-place compacted row
on the SAME session); a live delegation child (end_reason=None) stays excluded."""
for r in raw_results:
if len(seen_sessions) >= limit:
break
@@ -322,8 +302,8 @@ def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]:
def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]:
"""Build one discovery result from a surviving FTS row; None if the anchored
view can't be loaded (the hit is dropped)."""
"""One discovery result from a surviving FTS row; None (hit dropped) if the
anchored view can't be loaded."""
hit_sid = match_info.get("session_id") or lineage_root
msg_id = match_info.get("id")
try:
@@ -402,8 +382,8 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort
def _resolve_profile_db(profile: str):
"""Open another profile's ``state.db`` read-only (no write lock — safe on a
live DB), or None for the current profile."""
"""Another profile's ``state.db`` opened read-only (safe on a live DB), or None
for the current profile."""
if profile is None or not str(profile).strip():
return None
from hermes_cli import profiles as profiles_mod
@@ -417,8 +397,8 @@ def _resolve_profile_db(profile: str):
def _locate_session_db(session_id: str):
"""Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)``
or ``(None, None)``. Ids are globally unique, so the first hit is authoritative."""
"""Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)`` or
``(None, None)``. Ids are globally unique, so the first hit is authoritative."""
from pathlib import Path
try:
@@ -452,8 +432,7 @@ def _get_session_meta(db, session_id: str) -> dict:
def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_profile: str = None) -> str:
"""Read shape: whole session, or first ``head`` + last ``tail`` messages
with a pointer to scroll the middle."""
"""Read shape: whole session, or ``head`` + ``tail`` messages with a scroll pointer."""
meta = _get_session_meta(db, session_id)
if not meta:
return tool_error(f"session_id not found: {session_id}", success=False)
@@ -477,10 +456,9 @@ def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_prof
def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_profile: str = None) -> str:
"""Browse shape: metadata for the most recent sessions (no LLM, no FTS5)."""
try:
# list_sessions_rich (include_children=False) already applies the
# canonical child classifier: roots, /branch children and /new-reset
# children are admitted, delegation/compression children hidden.
# Re-classifying here re-hid legacy reset children — trust the query.
# list_sessions_rich already applies the canonical child classifier (roots,
# /branch and /new-reset children admitted; delegation/compression children
# hidden). Re-classifying here re-hid legacy reset children — trust the query.
sessions = db.list_sessions_rich(
limit=limit + 15, # extra so we can skip current / compression roots
exclude_sources=list(_HIDDEN_SESSION_SOURCES),
@@ -493,9 +471,8 @@ def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_p
sid = s.get("id", "")
if sid == current_session_id:
continue
# Compression continuation: the root's turns were summarised into
# the live child, so hide the root. /new-reset children share a
# root but carry no transcript — keep that root browsable.
# Compression continuation: the root was summarised into the live child, so
# hide it. /new-reset children carry no transcript — keep that root browsable.
if has_compression_hop and current_root and sid == current_root:
continue
results.append({
@@ -522,10 +499,10 @@ def _clamp_int(value, default: int, lo: int, hi: int) -> int:
def _anchor_in_live_context(db, anchor_state, anchor_session_id: str, current_session_id: str) -> bool:
"""True when the scroll anchor is still in the caller's active context and
must be rejected. Same-lineage history that has LEFT live context (compacted
rows, compression-ended parents, /new-reset predecessors) is allowed, so
scroll never rejects a result discovery just returned."""
"""True when the scroll anchor is still in the caller's active context (reject).
Same-lineage history that LEFT live context (compacted rows, compression-ended
parents, /new-reset predecessors) is allowed, so scroll never rejects a result
discovery just returned."""
if not _same_lineage(db, anchor_session_id, current_session_id) or _is_compacted_state(anchor_state):
return False
# Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history.
@@ -540,9 +517,8 @@ def _same_lineage(db, a: str, b: str) -> bool:
def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, window: int):
"""Lineage rebind: the caller paired a parent session_id with a message id
that lives in a descendant (compaction / delegation create child sessions).
Returns ``(view, warning)`` from the owning session, or ``(None, None)``."""
"""Lineage rebind when the caller paired a parent session_id with a message id
living in a descendant. ``(view, warning)`` from the owner, or ``(None, None)``."""
if not _same_lineage(db, session_id, owning):
return None, None
rebind_view = _quiet(lambda: db.get_messages_around(owning, around_message_id, window=window),
@@ -554,8 +530,8 @@ def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, w
def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
current_session_id: str = None) -> str:
"""Scroll shape: a window of messages centered on an anchor (no FTS5, no
bookends). Rebinds silently if the anchor lives in a same-lineage child."""
"""Scroll shape: a window centered on an anchor (no FTS5, no bookends);
rebinds silently if the anchor lives in a same-lineage child."""
if not isinstance(session_id, str) or not session_id.strip():
return tool_error("scroll requires session_id", success=False)
session_id = session_id.strip()
@@ -610,9 +586,8 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str:
"""Read shape. On a miss in the target profile, scan every profile (the
model may have dropped the owning profile from the link) and tag the result
with the profile it was found in."""
"""Read shape; on a miss scan every profile (the model may have dropped the
owning profile from the link) and tag the result with where it was found."""
result = _read_session(db, sid, link_profile=profile)
if json.loads(result).get("success"):
return result
@@ -630,12 +605,10 @@ def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str:
def _dispatch(query, role_filter, limit, db, current_session_id, session_id,
around_message_id, window, sort, profile, detail, owned_dbs) -> str:
"""Mode dispatch (see module docstring). Scroll wins over read/discovery when
an anchor is set — the agent asked for a specific slice. Profile DBs opened
here are appended to *owned_dbs* for the caller to close."""
# A raw `@session:<profile>/<id>` link passed as session_id: ids never
# contain "/", so a slash means profile/id — strip the prefix and adopt the
# embedded profile only when none was passed explicitly.
"""Mode dispatch (see module docstring); scroll wins when an anchor is set.
Profile DBs opened here are appended to *owned_dbs* for the caller to close."""
# A raw `@session:<profile>/<id>` link as session_id: ids never contain "/", so
# split on it and adopt the embedded profile only when none was passed.
if isinstance(session_id, str) and "/" in session_id:
emb_profile, _, emb_id = session_id.partition("/")
if emb_id:
@@ -643,9 +616,8 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_id,
if emb_profile and (profile is None or not str(profile).strip()):
profile = emb_profile
# Cross-profile read: swap in the named profile's DB (read-only) for every
# shape. Current-lineage guards key off ids that won't collide, so they
# stay inert.
# Cross-profile: swap in the named profile's DB (read-only) for every shape;
# current-lineage guards key off ids that won't collide, so they stay inert.
try:
profile_db = _resolve_profile_db(profile)
except Exception as e:
@@ -676,8 +648,7 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_id,
def session_search(query: str = "", role_filter: str = None, limit: int = 3, db=None,
current_session_id: str = None, session_id: str = None, around_message_id: int = None,
window: int = 5, sort: str = None, profile: str = None, detail: str = "adaptive") -> str:
"""Run session search and close databases opened by this invocation.
Parameter order is positional-compatible with older callers."""
"""Run session search, closing DBs opened here. Positional order is frozen for old callers."""
owned_dbs: List[Any] = []
if db is None:
try: