Merge branch 'simp/r3-33-I' into simp/r3-33

This commit is contained in:
Teknium
2026-09-03 00:00:26 -07:00
8 changed files with 758 additions and 1181 deletions

View File

@@ -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")

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,63 +29,46 @@ except ImportError:
logger = logging.getLogger(__name__)
# One tool-definition pass must use one config decision for both availability
# and the dynamic target schema. ContextVar keeps concurrent profile/session
# builds isolated while letting the check_fn result flow to the immediately
# following dynamic_schema_overrides call in ToolRegistry.get_definitions().
# One tool-definition pass must use ONE config decision for availability and the
# dynamic target schema: the check_fn result flows to the immediately following
# dynamic_schema_overrides call; ContextVar isolates concurrent profile builds.
_memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar(
"memory_surface_flags", default=None
)
"memory_surface_flags", default=None)
def get_memory_dir() -> Path:
"""Return the profile-scoped memories directory (resolved per call so
HERMES_HOME/profile switches after import are respected)."""
"""Profile-scoped memories dir, resolved per call (HERMES_HOME may switch after import)."""
return get_hermes_home() / "memories"
from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports)
ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, _READ_FAILED,
_drift_error, _read_failed_error, _scan_memory_content,
)
ENTRY_DELIMITER, MEMORY_BLOCK_HEADERS, MemoryStore, _scan_memory_content)
def load_on_disk_store() -> "MemoryStore":
"""Fresh on-disk MemoryStore with configured limits/flags, for contexts with
no live agent (gateway, Desktop, bare CLI ``/memory``) so approvals enforce
the SAME caps as ``agent_init``. Defaults if config can't load; never raises."""
kwargs: Dict[str, Any] = {}
"""Fresh on-disk MemoryStore with configured limits/flags for contexts with no
live agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps
as ``agent_init``. Falls back to defaults if config can't load; never raises."""
try:
from hermes_cli.config import load_config
config = load_config() or {}
mem_cfg = get_builtin_memory_config(config)
memory_enabled, user_profile_enabled = get_builtin_memory_store_flags(config)
kwargs = {
"memory_char_limit": int(mem_cfg.get("memory_char_limit", 2200)),
"user_char_limit": int(mem_cfg.get("user_char_limit", 1375)),
"memory_enabled": memory_enabled,
"user_profile_enabled": user_profile_enabled,
}
kwargs = {"memory_char_limit": int(mem_cfg.get("memory_char_limit", 2200)),
"user_char_limit": int(mem_cfg.get("user_char_limit", 1375)),
"memory_enabled": memory_enabled, "user_profile_enabled": user_profile_enabled}
except Exception:
kwargs = {} # config optional — fall back to defaults rather than break /memory
kwargs: Dict[str, Any] = {} # config optional — fall back to defaults rather than break /memory
store = MemoryStore(**kwargs)
store.load_from_disk()
return store
# ---------------------------------------------------------------------------
# Write-approval gate
# ---------------------------------------------------------------------------
def _target_label(target: str) -> str:
return "user profile" if target == "user" else "memory"
# -- Write-approval gate --
def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]:
"""Run the memory write gate. Returns a JSON tool-result string when the
write must NOT proceed (blocked, or staged for approval), None to proceed.
If the gate module can't load, fail open rather than block all writes."""
"""JSON tool-result string when the write must NOT proceed (blocked or staged
for approval), None to proceed. Fails open if the gate module can't load."""
try:
from tools import write_approval as wa
except Exception:
@@ -101,103 +83,74 @@ def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Option
ensure_ascii=False)
# action -> (store call, gate (summary, detail) text) for the live tool path and staged replay.
_STORE_ACTIONS = {
"add": (lambda store, target, content, old_text: store.add(target, content),
lambda label, content, old_text: (f"add to {label}", content or "")),
"replace": (lambda store, target, content, old_text: store.replace(target, old_text, content),
lambda label, content, old_text: (f"replace in {label}", f"old: {old_text}\nnew: {content}")),
"remove": (lambda store, target, content, old_text: store.remove(target, old_text),
lambda label, content, old_text: (f"remove from {label}", old_text or ""))}
def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str]) -> Optional[str]:
"""Gate a single mutating op (add/replace/remove); other actions pass."""
if action not in _STORE_ACTIONS:
return None
label = _target_label(target)
if action == "add":
summary, detail = f"add to {label}", content or ""
elif action == "replace":
summary, detail = f"replace in {label}", f"old: {old_text}\nnew: {content}"
else:
summary, detail = f"remove from {label}", old_text or ""
payload = {"action": action, "target": target, "content": content, "old_text": old_text}
return _gate_or_stage(summary, detail, payload)
"""Gate a single mutating op (add/replace/remove)."""
summary, detail = _STORE_ACTIONS[action][1]("user profile" if target == "user" else "memory", content, old_text)
return _gate_or_stage(summary, detail,
{"action": action, "target": target, "content": content, "old_text": old_text})
def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Optional[str]:
"""Gate a whole batch as a single unit."""
summary = f"apply {len(operations)} op(s) to {_target_label(target)}"
summary = f"apply {len(operations)} op(s) to {'user profile' if target == 'user' else 'memory'}"
detail_lines = []
for op in operations:
op = op or {}
act = op.get("action", "?")
_op_content = op.get("content") or op.get("new_text") or ""
if act == "remove":
detail_lines.append(f"- remove: {op.get('old_text', '')}")
elif act == "replace":
detail_lines.append(f"- replace: {op.get('old_text', '')} -> {_op_content}")
else:
detail_lines.append(f"- {act}: {_op_content}")
payload = {"action": "batch", "target": target, "operations": operations}
return _gate_or_stage(summary, "\n".join(detail_lines), payload)
content = op.get("content") or op.get("new_text") or ""
detail_lines.append(f"- remove: {op.get('old_text', '')}" if act == "remove"
else f"- replace: {op.get('old_text', '')} -> {content}" if act == "replace"
else f"- {act}: {content}")
return _gate_or_stage(summary, "\n".join(detail_lines),
{"action": "batch", "target": target, "operations": operations})
# ---------------------------------------------------------------------------
# Tool entry point
# ---------------------------------------------------------------------------
def _missing_old_text_error(store: "MemoryStore", target: str, action: str) -> str:
"""Recoverable error for replace/remove without ``old_text``. It can't be
schema-required (needs a combinator the Codex backend rejects — see
test_memory_tool_schema.py) and some clients omit it, so return the current
inventory plus a retry instruction instead of a dead-end."""
return json.dumps({
"success": False,
"error": (f"'{action}' needs old_text -- a short unique substring of the entry "
f"to {action}. None was provided. Reissue the {action} with old_text "
f"set to part of one of the current_entries below."),
"current_entries": store._entries_for(target),
"usage": store._usage(target),
}, ensure_ascii=False)
# -- Tool entry point --
def _validate_single_op(store, action, target, content, old_text) -> Optional[str]:
"""Validate required params BEFORE the gate so an invalid write is rejected
now rather than staged and failing at approve time."""
"""Validate BEFORE the gate so an invalid write is rejected now, not at approve
time. Missing ``old_text`` is recoverable (it can't be schema-required — needs a
combinator the Codex backend rejects — and some clients omit it): return the
current inventory plus a retry instruction instead of a dead-end."""
if action == "add" and not content:
return tool_error("Content is required for 'add' action.", success=False)
if action in ("replace", "remove") and not old_text:
return _missing_old_text_error(store, target, action)
return json.dumps({
"success": False,
"error": (f"'{action}' needs old_text -- a short unique substring of the entry "
f"to {action}. None was provided. Reissue the {action} with old_text "
f"set to part of one of the current_entries below."),
"current_entries": store._entries_for(target), "usage": store._usage(target)}, ensure_ascii=False)
if action == "replace" and not content:
return tool_error("content is required for 'replace' action.", success=False)
return None
# action -> store call for both the live tool path and staged-write replay.
_STORE_ACTIONS = {
"add": lambda store, target, content, old_text: store.add(target, content),
"replace": lambda store, target, content, old_text: store.replace(target, old_text, content),
"remove": lambda store, target, content, old_text: store.remove(target, old_text),
}
def memory_tool(
action: str = None,
target: str = "memory",
content: str = None,
old_text: str = None,
new_text: str = None,
operations: Optional[List[Dict[str, Any]]] = None,
store: Optional[MemoryStore] = None,
) -> str:
"""Tool entry point; returns a JSON string. Single op (action + content /
old_text) or batch (``operations`` applied atomically against the final
budget). ``new_text`` aliases ``content`` — callers mirror ``old_text``
with it (patch-tool shape), which used to leave ``content`` empty."""
def memory_tool(action: str = None, target: str = "memory", content: str = None, old_text: str = None,
new_text: str = None, operations: Optional[List[Dict[str, Any]]] = None,
store: Optional[MemoryStore] = None) -> str:
"""Tool entry point; returns a JSON string. Single op (action + content/old_text)
or batch (``operations``, atomic against the final budget). ``new_text``
aliases ``content`` — callers mirror ``old_text`` with it (patch-tool shape)."""
if store is None:
return tool_error("Memory is not available. It may be disabled in config or this environment.", success=False)
if content is None and new_text is not None:
content = new_text
# Strict providers send JSON null for optional fields; treat as omitted.
if target is None:
target = "memory"
target = "memory" if target is None else target
target_error = _memory_target_error(store, target)
if target_error is not None:
return json.dumps(target_error)
if operations:
if not isinstance(operations, list):
return tool_error("operations must be a list of {action, content?, old_text?} objects.", success=False)
@@ -205,29 +158,25 @@ def memory_tool(
if gate_result is not None:
return gate_result
return json.dumps(store.apply_batch(target, operations), ensure_ascii=False)
run = _STORE_ACTIONS.get(action)
if run is None:
if action not in _STORE_ACTIONS:
return tool_error(f"Unknown action '{action}'. Use: add, replace, remove", success=False)
invalid = _validate_single_op(store, action, target, content, old_text)
if invalid is not None:
return invalid
# Approval gate: when on, stages the write (background/gateway) or prompts
# inline (interactive CLI); when off (default) passes straight through.
# Approval gate: stages (background/gateway) or prompts inline (CLI); off by default.
gate_result = _apply_write_gate(action, target, content, old_text)
if gate_result is not None:
return gate_result
return json.dumps(run(store, target, content, old_text), ensure_ascii=False)
return json.dumps(_STORE_ACTIONS[action][0](store, target, content, old_text), ensure_ascii=False)
def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Normalized ``memory`` config section ({} when missing/malformed → flags
default to enabled). ``agent_init`` consumes the same section so tool
availability and store construction cannot diverge."""
"""Normalized ``memory`` config section ({} when missing/malformed → flags default
to enabled). ``agent_init`` reads the same section so availability and store
construction cannot diverge."""
if config is None:
try:
from hermes_cli.config import load_config_readonly
config = load_config_readonly()
except Exception:
logger.debug("Could not read memory config for availability", exc_info=True)
@@ -239,10 +188,7 @@ def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[s
def get_builtin_memory_store_flags(config: Optional[Dict[str, Any]] = None) -> Tuple[bool, bool]:
"""Return ``(memory_enabled, user_profile_enabled)`` from resolved config."""
section = get_builtin_memory_config(config)
return (
is_truthy_value(section.get("memory_enabled"), default=True),
is_truthy_value(section.get("user_profile_enabled"), default=True),
)
return tuple(is_truthy_value(section.get(k), default=True) for k in ("memory_enabled", "user_profile_enabled"))
@no_cache_check_fn
@@ -258,7 +204,6 @@ def _memory_target_error(store: "MemoryStore", target: str) -> Optional[Dict[str
"""Return a shared validation error for an invalid or disabled target."""
if target not in {"memory", "user"}:
from tools.registry import _bound_error_text
return {"success": False,
"error": _bound_error_text(f"Invalid memory target '{target}'. Use 'memory' or 'user'.")}
if store.target_enabled(target):
@@ -276,15 +221,12 @@ def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[
return target_error
if action == "batch":
return store.apply_batch(target, payload.get("operations") or [])
run = _STORE_ACTIONS.get(action)
if run is None:
if action not in _STORE_ACTIONS:
return {"success": False, "error": f"Unknown staged action '{action}'."}
return run(store, target, payload.get("content") or "", payload.get("old_text") or "")
return _STORE_ACTIONS[action][0](store, target, payload.get("content") or "", payload.get("old_text") or "")
# =============================================================================
# OpenAI Function-Calling Schema
# =============================================================================
# -- OpenAI Function-Calling Schema --
MEMORY_SCHEMA = {
"name": "memory",
@@ -361,24 +303,17 @@ MEMORY_SCHEMA = {
# Schema text when only one built-in store is enabled: (target description, TARGETS replacement).
_SINGLE_TARGET_TEXT = {
("memory",): (
"The enabled built-in store: 'memory' for personal notes.",
"TARGET: only 'memory' is enabled for personal notes (environment, conventions, "
"tool quirks, lessons).",
),
("user",): (
"The enabled built-in store: 'user' for user profile.",
"TARGET: only 'user' is enabled for user profile facts (name, role, preferences, style).",
),
}
("memory",): ("The enabled built-in store: 'memory' for personal notes.",
"TARGET: only 'memory' is enabled for personal notes (environment, conventions, "
"tool quirks, lessons)."),
("user",): ("The enabled built-in store: 'user' for user profile.",
"TARGET: only 'user' is enabled for user profile facts (name, role, preferences, style).")}
def _build_memory_schema_overrides() -> Dict[str, Any]:
"""Narrow the advertised target surface using the availability snapshot."""
flags = _memory_surface_flags.get()
flags = _memory_surface_flags.get() or get_builtin_memory_store_flags()
_memory_surface_flags.set(None)
if flags is None:
flags = get_builtin_memory_store_flags()
targets = [t for t, on in zip(("memory", "user"), flags) if on]
parameters = copy.deepcopy(MEMORY_SCHEMA["parameters"])
target_schema = parameters["properties"]["target"]
@@ -390,8 +325,7 @@ def _build_memory_schema_overrides() -> Dict[str, Any]:
description = description.replace(
"TARGETS: 'user' = who the user is (name, role, preferences, style). 'memory' = your "
"notes (environment, conventions, tool quirks, lessons).",
replacement,
)
replacement)
return {"description": description, "parameters": parameters}
@@ -403,14 +337,8 @@ registry.register(
toolset="memory",
schema=MEMORY_SCHEMA,
handler=lambda args, **kw: memory_tool(
action=args.get("action", ""),
target=args.get("target", "memory"),
content=args.get("content"),
old_text=args.get("old_text"),
new_text=args.get("new_text"),
operations=args.get("operations"),
store=kw.get("store")),
action=args.get("action", ""), target=args.get("target", "memory"), store=kw.get("store"),
**{k: args.get(k) for k in ("content", "old_text", "new_text", "operations")}),
check_fn=check_memory_requirements,
emoji="🧠",
dynamic_schema_overrides=_build_memory_schema_overrides,
)
dynamic_schema_overrides=_build_memory_schema_overrides)

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,76 +14,54 @@ from tools.threat_patterns import first_threat_message as _first_threat_message
logger = logging.getLogger("tools.memory_tool")
# System-prompt block header prefixes rendered by MemoryStore._render_block.
# agent/conversation_compression.py uses them to detect a leftover block for a
# target whose entries have since been emptied — keep in lockstep.
# Block header prefixes rendered by _render_block; agent/conversation_compression.py
# matches them to detect a leftover block for an emptied target — keep in lockstep.
MEMORY_BLOCK_HEADERS = {
"memory": "MEMORY (your personal notes)",
"user": "USER PROFILE (who the user is)",
}
"memory": "MEMORY (your personal notes)", "user": "USER PROFILE (who the user is)"}
ENTRY_DELIMITER = "\n§\n"
# Sentinel from ``_reload_target``: the file EXISTS but could not be read.
# Distinct from a drift-backup path (str) and a clean reload (None); the caller
# must abort rather than persist over an unreadable file.
_READ_FAILED = object()
def _memory_dir() -> Path:
from tools import memory_tool
return memory_tool.get_memory_dir()
def _scan_memory_content(content: str) -> Optional[str]:
"""Scan memory content for injection/exfil patterns ("strict" scope: memory
enters the system prompt as a frozen snapshot, so a poisoned entry persists
across sessions). Returns the error string if blocked."""
"""Error string if *content* matches injection/exfil patterns. Strict scope:
memory enters the system prompt, so a poisoned entry persists across sessions."""
return _first_threat_message(content, scope="strict")
def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]:
"""Error dict for external drift: the on-disk file wouldn't round-trip
through the parser/serializer, so flushing would discard content added by a
patch tool, shell append, manual edit, or sister session."""
return {
"success": False,
"error": (
f"Refusing to write {path.name}: file on disk has content that wouldn't round-trip "
f"through the memory tool (likely added by the patch tool, a shell append, a manual edit, "
f"or a concurrent session). A snapshot was saved to {bak_path}. Resolve the drift first — "
f"either rewrite the file as a clean §-delimited list of entries, or move the extra "
f"content out — then retry. This guard exists to prevent silent data loss (issue #26045)."
),
"drift_backup": bak_path,
"remediation": (
"Open the .bak file, integrate the missing entries into the memory tool one at a time via "
"memory(action=add, content=...), then remove or rewrite the original file to a clean state."
),
}
"""External drift: the file wouldn't round-trip, so flushing would discard content."""
return {"success": False, "error": (
f"Refusing to write {path.name}: file on disk has content that wouldn't round-trip "
f"through the memory tool (likely added by the patch tool, a shell append, a manual edit, "
f"or a concurrent session). A snapshot was saved to {bak_path}. Resolve the drift first — "
f"either rewrite the file as a clean §-delimited list of entries, or move the extra "
f"content out — then retry. This guard exists to prevent silent data loss (issue #26045)."
), "drift_backup": bak_path, "remediation": (
"Open the .bak file, integrate the missing entries into the memory tool one at a time via "
"memory(action=add, content=...), then remove or rewrite the original file to a clean state.")}
def _read_failed_error(path: "Path") -> Dict[str, Any]:
"""Error dict for an unreadable (but existing) memory file. Treating it as
empty and saving would rewrite the whole file from ``[]`` — wiping memory."""
"""Existing-but-unreadable file: saving from an assumed-empty view would wipe it."""
return {"success": False, "error": (
f"Refusing to write {path.name}: the file exists on disk but could not be read right now "
f"(temporarily locked by another program, a permission change, invalid/corrupt text encoding, "
f"or a filesystem error). Treating an unreadable file as empty and saving would wipe existing "
f"memory, so the write is refused. Nothing was changed — retry in a moment."
)}
f"memory, so the write is refused. Nothing was changed — retry in a moment.")}
def _find_unique_match(entries: List[str], old_text: str) -> Tuple[Optional[int], bool]:
"""``(index, ambiguous)`` for entries containing *old_text*. Exact-duplicate
matches are safe (first wins); distinct matches → ``(None, True)``."""
matches = [i for i, e in enumerate(entries) if old_text in e]
if not matches:
return None, False
if len({entries[i] for i in matches}) > 1:
return None, True
return matches[0], False
return (matches[0] if matches else None), False
class MemoryStore:
@@ -91,10 +69,9 @@ class MemoryStore:
``_system_prompt_snapshot`` is frozen at load time (prefix-cache stable);
``memory_entries`` / ``user_entries`` are live state persisted to disk."""
# After this many failed consolidation attempts (overflow / zero-match) in
# ONE turn, return a terminal "save skipped" result instead of "retry in
# this turn", so a fragile replace/add can't loop the turn to budget
# exhaustion and suppress the user's reply.
# Failed consolidation attempts (overflow / zero-match) allowed per turn before
# a TERMINAL "save skipped" result, so a fragile replace/add can't loop the turn
# to budget exhaustion and suppress the user's reply.
_MAX_CONSOLIDATION_FAILURES_PER_TURN = 3
def __init__(self, memory_char_limit: int = 2200, user_char_limit: int = 1375, *,
@@ -106,9 +83,7 @@ class MemoryStore:
self.memory_enabled = memory_enabled
self.user_profile_enabled = user_profile_enabled
self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""}
# Per-turn counter of failed at-capacity consolidation attempts; reset
# at each turn boundary by reset_consolidation_failures().
self._consolidation_failures = 0
self._consolidation_failures = 0 # per turn; reset by reset_consolidation_failures()
def target_enabled(self, target: str) -> bool:
"""Return whether this session's selected built-in store is writable."""
@@ -119,118 +94,91 @@ class MemoryStore:
self._consolidation_failures = 0
def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]:
"""Count an at-capacity consolidation failure. Under the per-turn cap
return ``response`` unchanged (it says how to retry); past it return a
TERMINAL result so the model stops looping — a failed memory side effect
must never block the turn's reply."""
"""Count a consolidation failure: under the per-turn cap return ``response``
(it says how to retry); past it a TERMINAL result so the model stops looping."""
self._consolidation_failures += 1
if self._consolidation_failures <= self._MAX_CONSOLIDATION_FAILURES_PER_TURN:
return response
return {"success": False, "done": True, "error": (
f"Memory consolidation failed {self._consolidation_failures} times this turn. Stop retrying "
"memory calls — leave memory unchanged for now and continue with your reply to the user. "
"The fact can be saved in a later turn."
)}
"The fact can be saved in a later turn.")}
def load_from_disk(self):
"""Load MEMORY.md / USER.md and capture the frozen system-prompt snapshot.
Threat hits are replaced in the SNAPSHOT only by a ``[BLOCKED: …]``
placeholder; live lists keep the raw text so the user can see and remove
poisoned entries (dropping them silently would hide the attack).
Scanning is deterministic from disk bytes, so the snapshot stays stable."""
Threat hits are replaced by a ``[BLOCKED: …]`` placeholder in the SNAPSHOT only;
live lists keep the raw text so the user can see and remove poisoned entries
(dropping them silently would hide the attack)."""
mem_dir = _memory_dir()
mem_dir.mkdir(parents=True, exist_ok=True)
# Deduplicate (order-preserving, first occurrence wins).
self.memory_entries = list(dict.fromkeys(self._read_file(mem_dir / "MEMORY.md")))
self.user_entries = list(dict.fromkeys(self._read_file(mem_dir / "USER.md")))
self._system_prompt_snapshot = {
target: self._render_block(target, self._sanitize_entries_for_snapshot(entries, filename))
for target, entries, filename in (
("memory", self.memory_entries, "MEMORY.md"),
("user", self.user_entries, "USER.md"),
)
}
"memory": self._render_block("memory", self._sanitize_entries_for_snapshot(self.memory_entries, "MEMORY.md")),
"user": self._render_block("user", self._sanitize_entries_for_snapshot(self.user_entries, "USER.md"))}
@staticmethod
def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]:
"""Return *entries* with any threat-matching entry replaced by a
``[BLOCKED: …]`` placeholder (strict scope, same as writes). Empty or
already-blocked entries pass through unchanged."""
"""*entries* with threat matches replaced by a ``[BLOCKED: …]`` placeholder
(strict scope, same as writes); empty / already-blocked entries pass through."""
from tools.threat_patterns import scan_for_threats
sanitized: List[str] = []
for entry in entries:
def _one(entry):
findings = scan_for_threats(entry, scope="strict") if entry and not entry.startswith("[BLOCKED:") else None
if not findings:
sanitized.append(entry)
continue
return entry
logger.warning("Memory entry from %s blocked at load time: %s", filename, ", ".join(findings))
sanitized.append(f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. "
f"Removed from system prompt; use memory(action=remove) to delete the original.]")
return sanitized
return (f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. "
f"Removed from system prompt; use memory(action=remove) to delete the original.]")
return [_one(e) for e in entries]
@staticmethod
@contextmanager
def _file_lock(path: Path):
"""Exclusive lock for read-modify-write safety, on a separate .lock file
so the memory file itself can still be atomically replaced."""
"""Exclusive lock on a separate .lock file so the memory file itself can
still be atomically replaced."""
from tools import memory_tool as _mt # fcntl/msvcrt live (and are patched) there
fcntl, msvcrt = _mt.fcntl, _mt.msvcrt
lock_path = path.with_suffix(path.suffix + ".lock")
lock_path.parent.mkdir(parents=True, exist_ok=True)
if fcntl is None and msvcrt is None:
yield
return
fd = open(lock_path, "a+", encoding="utf-8")
try:
if fcntl:
fcntl.flock(fd, fcntl.LOCK_EX)
else:
fd.seek(0)
msvcrt.locking(fd.fileno(), msvcrt.LK_LOCK, 1)
yield
finally:
try:
with open(lock_path, "a+", encoding="utf-8") as fd:
def _flock(unlock: bool):
if fcntl:
fcntl.flock(fd, fcntl.LOCK_UN)
fcntl.flock(fd, fcntl.LOCK_UN if unlock else fcntl.LOCK_EX)
else:
fd.seek(0)
msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK, 1)
except OSError:
pass
fd.close()
msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK if unlock else msvcrt.LK_LOCK, 1)
_flock(False)
try:
yield
finally:
try:
_flock(True)
except OSError:
pass
@staticmethod
def _path_for(target: str) -> Path:
return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md")
def _reload_target(self, target: str, *, skip_drift: bool = False):
"""Re-read entries from disk (under file lock) before mutating.
Returns ``None`` on a clean reload; the backup path (str) on external
drift (caller must abort — flushing would discard un-roundtrippable
content); or ``_READ_FAILED`` when the file exists but could not be
read (caller MUST abort — rewriting from an assumed-empty view would
wipe it; even append-only ``add`` rewrites the whole file). Drift check
and parse use the SAME raw snapshot — a failed second read used to count
as "no drift". *skip_drift* skips the round-trip check (``add``).
"""
raw, read_ok = self._read_raw_checked(self._path_for(target))
def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]:
"""Re-read entries from disk (under lock) before mutating; return the abort
error dict or None. Aborts on external drift (flushing would discard
un-roundtrippable content) and on an existing-but-unreadable file (even
append-only ``add`` rewrites the whole file). Drift check and parse use the
SAME raw snapshot — a failed second read used to count as "no drift"."""
path = self._path_for(target)
raw, read_ok = self._read_raw_checked(path)
if not read_ok:
return _READ_FAILED
return _read_failed_error(path)
bak = None if skip_drift else self._detect_external_drift(target, raw)
self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw))))
return bak
def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]:
"""Reload under lock; return the abort error dict, or None to proceed."""
bak = self._reload_target(target, skip_drift=skip_drift)
if bak is _READ_FAILED:
return _read_failed_error(self._path_for(target))
if bak:
return _drift_error(self._path_for(target), bak)
return None
return _drift_error(path, bak) if bak else None
def save_to_disk(self, target: str):
"""Persist entries to the appropriate file. Called after every mutation."""
@@ -241,10 +189,7 @@ class MemoryStore:
return self.user_entries if target == "user" else self.memory_entries
def _set_entries(self, target: str, entries: List[str]):
if target == "user":
self.user_entries = entries
else:
self.memory_entries = entries
setattr(self, "user_entries" if target == "user" else "memory_entries", entries)
def _char_count(self, target: str) -> int:
return len(ENTRY_DELIMITER.join(self._entries_for(target)))
@@ -255,9 +200,14 @@ class MemoryStore:
def _usage(self, target: str) -> str:
return f"{self._char_count(target):,}/{self._char_limit(target):,}"
def _usage_pct(self, target: str, current: int) -> str:
"""``"<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 that shows the live entries so the model can
decide what to consolidate."""
"""Consolidation failure carrying the live entries so the model can consolidate."""
return self._consolidation_failure({"success": False, "error": message,
"current_entries": self._entries_for(target), "usage": self._usage(target)})
@@ -267,19 +217,28 @@ class MemoryStore:
idx, ambiguous = _find_unique_match(entries, old_text)
if ambiguous:
return None, {"success": False, "error": f"Multiple entries matched '{old_text}'. Be more specific.",
"matches": self._previews([e for e in entries if old_text in e])}
"matches": [e[:80] + ("..." if len(e) > 80 else "") for e in entries if old_text in e]}
if idx is None:
return None, self._consolidation_failure({
"success": False,
"error": f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text of the entry you want to {verb}.",
"current_entries": entries,
})
"current_entries": entries})
return idx, None
def _commit(self, target: str, entries: List[str], message: str) -> Dict[str, Any]:
self._set_entries(target, entries)
self.save_to_disk(target)
return self._success_response(target, message)
def _mutate(self, target: str, mutate, *, skip_drift: bool = False) -> Dict[str, Any]:
"""Lock, reload, run ``mutate(entries, limit)`` -> ``(new_entries, message)`` or an
error dict, then persist and return the success response."""
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target, skip_drift=skip_drift)
if err:
return err
result = mutate(self._entries_for(target), self._char_limit(target))
if isinstance(result, dict):
return result
entries, message = result
self._set_entries(target, entries)
self.save_to_disk(target)
return self._success_response(target, message)
def add(self, target: str, content: str) -> Dict[str, Any]:
"""Append a new entry. Returns error if it would exceed the char limit."""
@@ -289,15 +248,8 @@ class MemoryStore:
scan_error = _scan_memory_content(content)
if scan_error:
return {"success": False, "error": scan_error}
with self._file_lock(self._path_for(target)):
# Append-only: skip the drift guard (appending never clobbers
# un-roundtrippable content), but still refuse on a failed read —
# add rewrites the WHOLE file from the parsed entries.
err = self._reload_or_error(target, skip_drift=True)
if err:
return err
entries = self._entries_for(target)
limit = self._char_limit(target)
def _add(entries, limit):
if content in entries:
return self._success_response(target, "Entry already exists (no duplicate added).")
if len(ENTRY_DELIMITER.join(entries + [content])) > limit:
@@ -305,58 +257,49 @@ class MemoryStore:
f"Memory at {self._char_count(target):,}/{limit:,} chars. Adding this entry "
f"({len(content)} chars) would exceed the limit. Consolidate now: use 'replace' to merge "
f"overlapping entries into shorter ones or 'remove' stale or less important entries (see "
f"current_entries below), then retry this add — all in this turn."
))
entries.append(content)
return self._commit(target, entries, "Entry added.")
f"current_entries below), then retry this add — all in this turn."))
return entries + [content], "Entry added."
# Append-only: skip the drift guard (appending never clobbers foreign
# content) but still refuse a failed read — add rewrites the WHOLE file.
return self._mutate(target, _add, skip_drift=True)
def replace(self, target: str, old_text: str, new_content: str) -> Dict[str, Any]:
"""Find entry containing old_text substring, replace it with new_content."""
old_text = old_text.strip()
new_content = new_content.strip()
if not old_text:
if not old_text.strip():
return {"success": False, "error": "old_text cannot be empty."}
if not new_content:
return {"success": False, "error": "new_content cannot be empty. Use 'remove' to delete entries."}
scan_error = _scan_memory_content(new_content)
if scan_error:
return {"success": False, "error": scan_error}
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target)
return self._edit(target, old_text.strip(), new_content)
def remove(self, target: str, old_text: str) -> Dict[str, Any]:
"""Remove the entry containing old_text substring."""
if not old_text.strip():
return {"success": False, "error": "old_text cannot be empty."}
return self._edit(target, old_text.strip(), None)
def _edit(self, target: str, old_text: str, new_content: Optional[str]) -> Dict[str, Any]:
"""Locked replace (``new_content`` set) or remove (None) of the unique entry matching *old_text*."""
def _apply(entries, limit):
idx, err = self._locate(target, old_text, "replace" if new_content else "remove")
if err:
return err
idx, err = self._locate(target, old_text, "replace")
if err:
return err
entries = self._entries_for(target)
limit = self._char_limit(target)
new_total = len(ENTRY_DELIMITER.join(entries[:idx] + [new_content] + entries[idx + 1:]))
if new_content is None:
return entries[:idx] + entries[idx + 1:], "Entry removed."
replaced = entries[:idx] + [new_content] + entries[idx + 1:]
new_total = len(ENTRY_DELIMITER.join(replaced))
if new_total > limit:
return self._failure_with_entries(target, (
f"Replacement would put memory at {new_total:,}/{limit:,} chars. Shorten the new content, "
f"or 'remove' other stale or less important entries to make room (see current_entries "
f"below), then retry — all in this turn."
))
entries[idx] = new_content
return self._commit(target, entries, "Entry replaced.")
f"below), then retry — all in this turn."))
return replaced, "Entry replaced."
def remove(self, target: str, old_text: str) -> Dict[str, Any]:
"""Remove the entry containing old_text substring."""
old_text = old_text.strip()
if not old_text:
return {"success": False, "error": "old_text cannot be empty."}
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target)
if err:
return err
idx, err = self._locate(target, old_text, "remove")
if err:
return err
entries = self._entries_for(target)
entries.pop(idx)
return self._commit(target, entries, "Entry removed.")
# -- Batch --
return self._mutate(target, _apply)
@staticmethod
def _apply_batch_op(working: List[str], act: str, content: str, old_text: str, pos: str) -> Optional[str]:
@@ -385,90 +328,56 @@ class MemoryStore:
return None
def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]:
"""Apply a sequence of add/replace/remove ops to one target atomically.
Ops are validated and applied against the FINAL budget only — so a single
call can free space (remove/replace) and add new entries without the
multi-turn consolidate-then-retry dance. All-or-nothing: if any op is
malformed, doesn't match, or the net result exceeds the char limit,
NOTHING is written and the first failure plus live state is returned.
"""
"""Apply add/replace/remove ops to one target atomically against the FINAL
budget, so one call can free space and add entries. All-or-nothing: any
malformed / unmatched op or an over-limit result writes NOTHING and returns
the first failure plus live state."""
if not operations:
return {"success": False, "error": "operations list is empty."}
# Scan every add/replace content BEFORE touching disk -- one poisoned
# op rejects the whole batch.
for i, op in enumerate(operations):
op = op or {}
ops = [op or {} for op in operations]
# Scan every add/replace content BEFORE touching disk -- one poisoned op rejects the batch.
for i, op in enumerate(ops):
scan_error = op.get("action") in {"add", "replace"} and op.get("content") and _scan_memory_content(op["content"])
if scan_error:
return {"success": False, "error": f"Operation {i + 1}: {scan_error}"}
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target)
if err:
return err
# Work on a copy; only commit if the whole batch validates.
working: List[str] = list(self._entries_for(target))
limit = self._char_limit(target)
for i, op in enumerate(operations):
op = op or {}
def _apply(entries, limit):
working = list(entries) # only committed if the whole batch validates
for i, op in enumerate(ops):
act = op.get("action")
content = (op.get("content") or op.get("new_text") or "").strip()
old_text = (op.get("old_text") or "").strip()
pos = f"Operation {i + 1} ({act or 'unknown'})"
msg = self._apply_batch_op(working, act, content, old_text, pos)
if msg:
return self._batch_error(target, msg)
return self._failure_with_entries(
target, msg + " No operations were applied (batch is all-or-nothing).")
# Budget check against the FINAL state only.
new_total = len(ENTRY_DELIMITER.join(working))
if new_total > limit:
return self._failure_with_entries(target, (
f"After applying all {len(operations)} operations, memory would be at "
f"{new_total:,}/{limit:,} chars -- over the limit. Remove or shorten more "
f"entries in the same batch (see current_entries below), then retry."
))
return self._commit(target, working, f"Applied {len(operations)} operation(s).")
f"entries in the same batch (see current_entries below), then retry."))
return working, f"Applied {len(operations)} operation(s)."
def _batch_error(self, target: str, message: str) -> Dict[str, Any]:
"""Build a batch-abort error that reports live (uncommitted) state."""
return self._failure_with_entries(
target, message + " No operations were applied (batch is all-or-nothing)."
)
return self._mutate(target, _apply)
def format_for_system_prompt(self, target: str) -> Optional[str]:
"""Return the frozen load-time snapshot for system-prompt injection (NOT
live state — mid-session writes don't affect it, preserving the prefix
cache). None if the snapshot is empty."""
"""Frozen load-time snapshot for the system prompt (NOT live state — mid-session
writes don't touch it, preserving the prefix cache); None if empty."""
return self._system_prompt_snapshot.get(target, "") or None
# -- Internal helpers --
@staticmethod
def _previews(entries: List[str], width: int = 80) -> List[str]:
"""Truncated one-line previews of entries for error feedback."""
return [e[:width] + ("..." if len(e) > width else "") for e in entries]
def _usage_pct(self, target: str, current: int) -> str:
"""``"<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."""
@@ -481,12 +390,10 @@ class MemoryStore:
@staticmethod
def _read_raw_checked(path: Path) -> Tuple[str, bool]:
"""Read raw text as ``(raw, read_ok)``. ``read_ok`` is False ONLY when the
file EXISTS but can't be read (absent file → ``("", True)``). Invalid
UTF-8 counts as unreadable; decoding stays STRICT because
``errors="replace"`` would hand callers a lossy view that a save then
persists over the real bytes. ``utf-8-sig`` strips a Notepad BOM that
otherwise glues U+FEFF onto the first entry forever."""
"""``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't
be read (absent → ``("", True)``). Decoding stays STRICT: ``errors="replace"``
would hand callers a lossy view that a save then persists. ``utf-8-sig`` strips
a Notepad BOM that otherwise glues U+FEFF onto the first entry forever."""
if not path.exists():
return "", True
try:
@@ -496,30 +403,20 @@ class MemoryStore:
@staticmethod
def _parse_entries(raw: str) -> List[str]:
"""Split raw memory-file text into stripped, non-empty entries. Splits on
the full ENTRY_DELIMITER so a bare "§" inside an entry is preserved."""
"""Stripped, non-empty entries; splits on the FULL delimiter so a bare "§" survives."""
return [e for e in (x.strip() for x in raw.split(ENTRY_DELIMITER)) if e]
@staticmethod
def _read_entries_checked(path: Path) -> Tuple[List[str], bool]:
"""Read + parse as ``(entries, read_ok)`` — see ``_read_raw_checked``."""
raw, read_ok = MemoryStore._read_raw_checked(path)
return MemoryStore._parse_entries(raw), read_ok
@staticmethod
def _read_file(path: Path) -> List[str]:
"""Read a memory file into entries (empty list on any error). Only for
read-only callers (``load_from_disk``, learning_mutations); mutation
paths must use ``_read_raw_checked`` so they can refuse to overwrite an
unreadable file."""
return MemoryStore._read_entries_checked(path)[0]
"""Entries of a memory file ([] on any error). Read-only callers only
(``load_from_disk``, learning_mutations); mutation paths must use
``_read_raw_checked`` so they can refuse to overwrite an unreadable file."""
return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0])
def _detect_external_drift(self, target: str, raw: str) -> Optional[str]:
"""Backup-path string if *raw* (the caller's checked-read snapshot) shows
external drift, else None. Signals: round-trip mismatch, or one parsed
entry exceeding the whole-file char limit (no tool-written entry can —
an external writer appended free-form content a flush would truncate).
The file is snapshotted to ``.bak.<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)
@@ -535,8 +432,7 @@ class MemoryStore:
@staticmethod
def _write_file(path: Path, entries: List[str]):
"""Atomic temp-file + rename: readers see the old or the new complete
file, never a truncated one."""
"""Atomic temp-file + rename: readers never see a truncated file."""
try:
atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_")
except OSError as e:

View File

@@ -31,19 +31,14 @@ class MicrosoftGraphTokenError(MicrosoftGraphAuthError):
def format_graph_error(error: Any) -> str | None:
"""Render Graph's ``{"error": {"code", "message"}}`` (or bare-string ``error``) body.
Shared by the token endpoint and the REST client so both surface
``code: message`` the same way. ``None`` means the shape was unusable.
"""
"""Render Graph's ``{"error": {"code", "message"}}`` (or bare-string ``error``) as
``code: message``; shared by token endpoint and REST client. None if unusable."""
if isinstance(error, str):
return error
if not isinstance(error, dict):
return None
code, message = error.get("code"), error.get("message")
if code and message:
return f"{code}: {message}"
return str(message) if message else None
return f"{code}: {message}" if code and message else (str(message) if message else None)
@dataclass(frozen=True)
@@ -68,15 +63,12 @@ class GraphCredentials:
env = environ if environ is not None else os.environ
values = [(env.get(name) or "").strip() for name in _REQUIRED_ENV]
missing = [name for name, value in zip(_REQUIRED_ENV, values) if not value]
if missing and not required:
return None
if missing:
if not required:
return None
raise MicrosoftGraphConfigError(f"Missing Microsoft Graph configuration: {', '.join(missing)}")
return cls(
*values,
scope=(env.get("MSGRAPH_SCOPE") or DEFAULT_GRAPH_SCOPE).strip(),
authority_url=(env.get("MSGRAPH_AUTHORITY_URL") or DEFAULT_GRAPH_AUTHORITY_URL).strip(),
)
return cls(*values, scope=(env.get("MSGRAPH_SCOPE") or DEFAULT_GRAPH_SCOPE).strip(),
authority_url=(env.get("MSGRAPH_AUTHORITY_URL") or DEFAULT_GRAPH_AUTHORITY_URL).strip())
@dataclass
@@ -102,9 +94,7 @@ class MicrosoftGraphTokenProvider:
self, credentials: GraphCredentials, *, timeout: float = 20.0,
skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, transport: httpx.AsyncBaseTransport | None = None,
) -> None:
self.credentials = credentials
self.timeout = timeout
self.skew_seconds = max(0, int(skew_seconds))
self.credentials, self.timeout, self.skew_seconds = credentials, timeout, max(0, int(skew_seconds))
self._transport = transport
self._cached_token: CachedAccessToken | None = None
self._lock = asyncio.Lock()
@@ -117,26 +107,17 @@ class MicrosoftGraphTokenProvider:
self._cached_token = None
def inspect_token_health(self) -> dict[str, Any]:
cached = self._cached_token
return {
"configured": True,
"tenant_id": self.credentials.tenant_id,
"client_id": self.credentials.client_id,
"scope": self.credentials.scope,
"authority_url": self.credentials.authority_url,
"token_url": self.credentials.token_url,
"cached": bool(cached),
"expires_in_seconds": cached.expires_in_seconds if cached else None,
"is_expired": cached.is_expired(skew_seconds=0) if cached else None,
"refresh_skew_seconds": self.skew_seconds,
}
cached, creds = self._cached_token, self.credentials
return {"configured": True, "tenant_id": creds.tenant_id, "client_id": creds.client_id,
"scope": creds.scope, "authority_url": creds.authority_url, "token_url": creds.token_url,
"cached": bool(cached), "expires_in_seconds": cached.expires_in_seconds if cached else None,
"is_expired": cached.is_expired(skew_seconds=0) if cached else None,
"refresh_skew_seconds": self.skew_seconds}
def _fresh_cached(self) -> CachedAccessToken | None:
"""The cached token unless it expires within ``skew_seconds``."""
cached = self._cached_token
if cached and not cached.is_expired(skew_seconds=self.skew_seconds):
return cached
return None
return cached if cached and not cached.is_expired(skew_seconds=self.skew_seconds) else None
async def get_access_token(self, *, force_refresh: bool = False) -> str:
# Double-checked under the lock so concurrent callers share one fetch.
@@ -145,33 +126,22 @@ class MicrosoftGraphTokenProvider:
async with self._lock:
if not force_refresh and (cached := self._fresh_cached()):
return cached.access_token
token = await self._fetch_access_token()
self._cached_token = token
return token.access_token
self._cached_token = await self._fetch_access_token()
return self._cached_token.access_token
async def _fetch_access_token(self) -> CachedAccessToken:
data = {
"grant_type": "client_credentials",
"client_id": self.credentials.client_id,
"client_secret": self.credentials.client_secret,
"scope": self.credentials.scope,
}
data = {"grant_type": "client_credentials", "client_id": self.credentials.client_id,
"client_secret": self.credentials.client_secret, "scope": self.credentials.scope}
async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client:
response = await client.post(
self.credentials.token_url, data=data,
headers={"Content-Type": "application/x-www-form-urlencoded"},
)
response = await client.post(self.credentials.token_url, data=data,
headers={"Content-Type": "application/x-www-form-urlencoded"})
if response.status_code >= 400:
raise MicrosoftGraphTokenError(
"Microsoft Graph token request failed with HTTP "
f"{response.status_code}: {_extract_error_detail(response)}"
)
raise MicrosoftGraphTokenError("Microsoft Graph token request failed with HTTP "
f"{response.status_code}: {_extract_error_detail(response)}")
try:
payload = response.json()
except ValueError as exc:
raise MicrosoftGraphTokenError("Microsoft Graph token response was not valid JSON.") from exc
access_token = str(payload.get("access_token") or "").strip()
if not access_token:
raise MicrosoftGraphTokenError("Microsoft Graph token response did not include access_token.")
@@ -179,34 +149,26 @@ class MicrosoftGraphTokenProvider:
expires_in_seconds = int(payload.get("expires_in"))
except (TypeError, ValueError) as exc:
raise MicrosoftGraphTokenError(
"Microsoft Graph token response did not include a valid expires_in."
) from exc
"Microsoft Graph token response did not include a valid expires_in.") from exc
return CachedAccessToken(
access_token=access_token,
token_type=str(payload.get("token_type") or "Bearer").strip() or "Bearer",
expires_at=time.time() + max(0, expires_in_seconds),
)
expires_at=time.time() + max(0, expires_in_seconds))
def _extract_error_detail(response: httpx.Response) -> str:
"""Best human-readable detail from a token-endpoint error body.
The OAuth endpoint prefers ``error_description``; fall back to the
Graph-style ``error`` object/string, then a bare ``code``, then raw text.
"""
"""Best human-readable detail from a token-endpoint error body: ``error_description``,
then the Graph-style ``error`` object/string, then a bare ``code``, then raw text."""
try:
payload = response.json()
except ValueError:
return response.text.strip() or "unknown error"
if isinstance(payload, dict):
if isinstance(payload.get("error_description"), str):
return payload["error_description"]
error = payload.get("error")
detail = format_graph_error(error)
if detail is not None:
return detail
if isinstance(error, dict) and error.get("code"):
return str(error["code"])
return str(payload)
if not isinstance(payload, dict):
return str(payload)
if isinstance(payload.get("error_description"), str):
return payload["error_description"]
error = payload.get("error")
detail = format_graph_error(error)
if detail is not None:
return detail
return str(error["code"]) if isinstance(error, dict) and error.get("code") else str(payload)

View File

@@ -5,16 +5,12 @@ from __future__ import annotations
import asyncio
import os
from pathlib import Path
from typing import Any, AsyncIterator, Awaitable, Callable
from typing import Any, Awaitable, Callable
import httpx
from agent.retry_utils import parse_retry_after_seconds
from tools.microsoft_graph_auth import (
GraphCredentials,
MicrosoftGraphTokenProvider,
format_graph_error,
)
from tools.microsoft_graph_auth import GraphCredentials, MicrosoftGraphTokenProvider, format_graph_error
DEFAULT_GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0"
@@ -32,38 +28,26 @@ class MicrosoftGraphAPIError(MicrosoftGraphClientError):
def __init__(
self, status_code: int, method: str, url: str, message: str, *,
retry_after_seconds: float | None = None, payload: Any = None,
) -> None:
self.status_code = status_code
self.method = method
self.url = url
self.retry_after_seconds = retry_after_seconds
self.payload = payload
retry_after_seconds: float | None = None, payload: Any = None) -> None:
self.status_code, self.method, self.url = status_code, method, url
self.retry_after_seconds, self.payload = retry_after_seconds, payload
super().__init__(f"Microsoft Graph API error {status_code} for {method} {url}: {message}")
class MicrosoftGraphClient:
"""Minimal async Microsoft Graph client with retries and pagination.
Retry policy (shared by JSON requests and streaming downloads): transport
errors back off exponentially; 401 clears the token cache and refetches;
429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``.
"""
"""Minimal async Graph client. Retry policy (JSON requests and streaming downloads
alike): transport errors back off exponentially; 401 clears the token cache and
refetches; 429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``."""
def __init__(
self, token_provider: MicrosoftGraphTokenProvider, *,
base_url: str = DEFAULT_GRAPH_BASE_URL, timeout: float = 60.0, max_retries: int = 3,
transport: httpx.AsyncBaseTransport | None = None,
sleep: Callable[[float], Awaitable[None]] | None = None,
user_agent: str = "Hermes-Agent/graph-client",
) -> None:
self.token_provider = token_provider
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.max_retries = max(0, int(max_retries))
self._transport = transport
self._sleep = sleep or asyncio.sleep
self.user_agent = user_agent
user_agent: str = "Hermes-Agent/graph-client") -> None:
self.token_provider, self.base_url, self.timeout = token_provider, base_url.rstrip("/"), timeout
self.max_retries, self.user_agent = max(0, int(max_retries)), user_agent
self._transport, self._sleep = transport, sleep or asyncio.sleep
@classmethod
def from_env(cls, **kwargs: Any) -> "MicrosoftGraphClient":
@@ -77,19 +61,20 @@ class MicrosoftGraphClient:
async def patch_json(self, path: str, *, json_body: Any | None = None, headers: Headers = None) -> Any:
response = await self._request("PATCH", path, json_body=json_body, headers=headers)
if response.status_code == 204 or not response.content:
return {}
return self._decode_json(response)
return self._decode_json_or(response, {})
async def delete(self, path: str, *, headers: Headers = None) -> dict[str, Any]:
response = await self._request("DELETE", path, headers=headers)
if response.status_code == 204 or not response.content:
return {"deleted": True, "status_code": response.status_code}
return self._decode_json(response)
return self._decode_json_or(response, {"deleted": True, "status_code": response.status_code})
async def iterate_pages(
self, path: str, *, params: Params = None, headers: Headers = None
) -> AsyncIterator[dict[str, Any]]:
def _decode_json_or(self, response: httpx.Response, empty: Any) -> Any:
"""*empty* for a 204 / bodiless response, else the decoded JSON body."""
return empty if response.status_code == 204 or not response.content else self._decode_json(response)
async def collect_paginated(
self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]:
"""Follow ``@odata.nextLink`` and concatenate every page's ``value`` list."""
items: list[Any] = []
# Query params go on the first request only; @odata.nextLink already embeds them.
next_url: str | None = self._resolve_url(path)
next_params = dict(params or {})
@@ -98,28 +83,17 @@ class MicrosoftGraphClient:
payload = self._decode_json(response)
if not isinstance(payload, dict):
raise MicrosoftGraphClientError(
f"Expected paginated Graph response dict, got {type(payload).__name__}."
)
yield payload
next_url = payload.get("@odata.nextLink")
next_params = {}
async def collect_paginated(
self, path: str, *, params: Params = None, headers: Headers = None
) -> list[Any]:
items: list[Any] = []
async for page in self.iterate_pages(path, params=params, headers=headers):
value = page.get("value")
if isinstance(value, list):
items.extend(value)
f"Expected paginated Graph response dict, got {type(payload).__name__}.")
if isinstance(payload.get("value"), list):
items.extend(payload["value"])
next_url, next_params = payload.get("@odata.nextLink"), {}
return items
async def download_to_file(
self, path: str, destination: str | Path, *, headers: Headers = None, chunk_size: int = 65536
) -> dict[str, Any]:
"""Download a Graph resource to disk, streaming the body chunk-by-chunk
(recordings and other large artifacts never need to fit in memory).
Written to a ``.part`` file and renamed into place only on success."""
"""Stream a Graph resource to disk chunk-by-chunk (large recordings never
fit in memory); written to ``.part`` and renamed into place only on success."""
url = self._resolve_url(path)
target = Path(destination)
target.parent.mkdir(parents=True, exist_ok=True)
@@ -147,8 +121,7 @@ class MicrosoftGraphClient:
async def _request(
self, method: str, path_or_url: str, *,
params: Params = None, json_body: Any | None = None, headers: Headers = None,
) -> httpx.Response:
params: Params = None, json_body: Any | None = None, headers: Headers = None) -> httpx.Response:
url = self._resolve_url(path_or_url)
async def perform(client: httpx.AsyncClient, request_headers: dict[str, str]):
@@ -160,14 +133,10 @@ class MicrosoftGraphClient:
async def _with_retries(
self, method: str, url: str, accept: str, json_body: Any | None, headers: Headers,
perform: Callable[[httpx.AsyncClient, dict[str, str]], Awaitable[tuple[httpx.Response, Any]]],
kind: str,
) -> Any:
"""Run ``perform`` (returning ``(response, result)``) under the retry policy.
``kind`` ("request"/"download") only labels the transport-failure messages.
A ``MicrosoftGraphAPIError`` for the failing status is raised once retries
are exhausted or the status is not retryable; only a 401 forces a token refresh.
"""
kind: str) -> Any:
"""Run ``perform`` (-> ``(response, result)``) under the retry policy. ``kind``
only labels transport-failure messages. Raises ``MicrosoftGraphAPIError`` once
retries are exhausted or the status is not retryable; only 401 forces a token refresh."""
attempt = 0
last_error: Exception | None = None
@@ -175,8 +144,7 @@ class MicrosoftGraphClient:
token = await self.token_provider.get_access_token(
force_refresh=attempt > 0
and isinstance(last_error, MicrosoftGraphAPIError)
and last_error.status_code == 401
)
and last_error.status_code == 401)
request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent}
if json_body is not None:
request_headers["Content-Type"] = "application/json"
@@ -187,27 +155,21 @@ class MicrosoftGraphClient:
async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client:
response, result = await perform(client, request_headers)
except httpx.HTTPError as exc:
last_error = exc
last_error, response = exc, None
if attempt >= self.max_retries:
raise MicrosoftGraphClientError(
f"Microsoft Graph {kind} failed for {method} {url}: {exc}"
) from exc
await self._sleep(self._retry_delay(None, attempt))
attempt += 1
continue
if response.status_code < 400:
return result
api_error = last_error = self._build_api_error(method, url, response)
status = response.status_code
if attempt < self.max_retries and (status in (401, 429) or 500 <= status < 600):
f"Microsoft Graph {kind} failed for {method} {url}: {exc}") from exc
else:
if response.status_code < 400:
return result
last_error = self._build_api_error(method, url, response)
status = response.status_code
if attempt >= self.max_retries or not (status in (401, 429) or 500 <= status < 600):
raise last_error
if status == 401:
self.token_provider.clear_cache()
await self._sleep(self._retry_delay(response, attempt))
attempt += 1
continue
raise api_error
await self._sleep(self._retry_delay(response, attempt))
attempt += 1
raise MicrosoftGraphClientError(f"Microsoft Graph {kind} exhausted retries for {method} {url}.")
@@ -224,29 +186,21 @@ class MicrosoftGraphClient:
except ValueError as exc:
raise MicrosoftGraphClientError(
"Microsoft Graph response was not valid JSON for "
f"{response.request.method} {response.request.url}"
) from exc
f"{response.request.method} {response.request.url}") from exc
@staticmethod
def _retry_delay(response: httpx.Response | None, attempt: int) -> float:
if response is not None:
retry_after = parse_retry_after_seconds(response.headers)
if retry_after is not None:
return retry_after
return min(8.0, 0.5 * (2 ** attempt))
retry_after = parse_retry_after_seconds(response.headers) if response is not None else None
return min(8.0, 0.5 * (2 ** attempt)) if retry_after is None else retry_after
@staticmethod
def _build_api_error(method: str, url: str, response: httpx.Response) -> MicrosoftGraphAPIError:
message = response.text.strip() or "unknown error"
try:
payload: Any = response.json()
except ValueError:
payload = None
if isinstance(payload, dict):
detail = format_graph_error(payload.get("error"))
if detail is not None:
message = detail
detail = format_graph_error(payload.get("error")) if isinstance(payload, dict) else None
return MicrosoftGraphAPIError(
response.status_code, method, url, message,
retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload,
)
response.status_code, method, url,
detail if detail is not None else (response.text.strip() or "unknown error"),
retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload)

View File

@@ -5,35 +5,361 @@ Single-shape tool; the mode is inferred from the args: DISCOVERY (``query``;
FTS5 deduped by lineage, adaptive detail hydrates only the top result),
SCROLL (``session_id`` + ``around_message_id``; ±window around the anchor),
READ (``session_id`` alone; whole session or head/tail), BROWSE (no args).
No LLM calls — every shape returns actual DB messages. Helpers live in
``session_search_tool_common`` / ``_discover`` and are re-exported here.
No LLM calls — every shape returns actual DB messages.
"""
import json
import logging
from typing import Any, List, Optional
from datetime import datetime
from typing import Any, Dict, List, Optional, Union
from tools.session_search_tool_common import ( # noqa: F401 (re-exports)
_COMPACTION_PREFIXES, _DEMOTED_SESSION_SOURCES, _DISCOVER_SCAN_LIMIT,
_DISCOVER_SEARCH_FIELDS, _FRESH_RESET_END_REASONS, _HIDDEN_SESSION_SOURCES,
_annotate_rebuild_status, _format_timestamp, _get_message_storage_state,
_is_compacted_message, _is_compacted_state, _is_compaction_summary,
_ok, _order_for_recall, _quiet, _resolve_lineage, _resolve_to_parent, _session_end_reason,
_session_left_live_context, _session_link, _session_meta_block, _shape_message,
)
from tools.session_search_tool_discover import ( # noqa: F401 (re-exports)
_discover, _normalize_title_query, _title_match_result,
)
from hermes_state_common import _RESET_END_REASONS
# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool),
# delegate subagent runs, kanban workers — not the user's history.
_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool")
# Searchable but DEMOTED below interactive sessions: cron vocabulary dominates bare
# BM25 and starves out the user's own sessions ("recall blindness").
_DEMOTED_SESSION_SOURCES = ("cron",)
# FTS rows scanned before dedup-by-lineage — well above the distinct sessions a query
# returns, so interactive matches buried under cron hits survive the demotion pass.
_DISCOVER_SCAN_LIMIT = 300
# Raw FTS rows are only a plan input; the response hydrates its own window/bookends.
_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started")
# Compaction handoff summaries (agent/context_compressor.py); excluded from bookends.
_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:")
# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the predecessor
# WITHOUT carrying its transcript forward — unlike compression continuations and
# live delegation children. Derived from the gateway set so the two cannot drift.
_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"}
def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
"""``fn()``, or *default* after debug-logging *msg* (exception appended as a
final ``%s`` arg when *with_exc*) on any exception."""
try:
return fn()
except Exception as e:
logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True)
return default
def _loud(fn, log_msg, error_prefix, *log_args):
"""``(value, None)`` from ``fn()``, or ``(None, tool_error_json)`` after an
error-level log — for DB calls whose failure the model must see."""
try:
return fn(), None
except Exception as e:
logging.error(log_msg, *log_args, e, exc_info=True)
return None, tool_error(f"{error_prefix}: {e}", success=False)
def _format_timestamp(ts: Union[int, float, str, None]) -> str:
"""Unix timestamp (number / numeric string) -> readable date; ISO strings pass
through; "unknown" for None; str(ts) if conversion fails."""
if ts is None:
return "unknown"
if isinstance(ts, str) and not ts.replace(".", "").replace("-", "").isdigit():
return ts
try:
return datetime.fromtimestamp(float(ts)).strftime("%B %d, %Y at %I:%M %p")
except Exception as e:
logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True)
return str(ts)
def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]:
"""The ``session_meta`` sub-object shared by read/scroll responses."""
return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"),
"model": meta.get("model"), "title": meta.get("title")}
def _ok(**payload) -> str:
"""Serialize a successful tool result (``success`` first, then *payload* in order)."""
return json.dumps({"success": True, **payload}, ensure_ascii=False)
def _is_compaction_summary(content: str) -> bool:
"""Return True if *content* looks like a generated compaction handoff."""
return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES)
def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]:
"""Walk parent_session_id to the root -> ``(root_id, has_compression_hop)``. The
flag separates a compression-split lineage (parent summarised away) from a
delegation lineage (child still visible to the parent). Errors -> ``(session_id, False)``."""
visited: set[str] = set()
cur, has_compression = session_id, False
while cur and cur not in visited:
visited.add(cur)
s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True)
if not s:
break
has_compression = has_compression or s.get("end_reason") == "compression"
if not s.get("parent_session_id"):
break
cur = s["parent_session_id"]
return cur, has_compression
def _resolve_lineage(db, session_id: str) -> str:
"""Return only the lineage root (ignores the compression hop)."""
return _resolve_to_parent(db, session_id)[0]
def _session_left_live_context(db, session_id: str) -> bool:
"""True when the transcript left everyone's live context: ``compression``
(summarised into the child) or a fresh reset (child starts empty). Live delegation
children (``end_reason is None``) and ``branched`` parents (copied verbatim into
the branch) ARE the current context, so they stay excluded from recall."""
s = session_id and _quiet(lambda: db.get_session(session_id), None, "get_session failed for %s", session_id)
end_reason = (s.get("end_reason") or None) if s else None
return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS
def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]:
"""Return the owning session and visibility flags for *message_id*."""
if not message_id:
return None
def _lookup():
with db._lock:
return db._conn.execute(
"SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,)).fetchone()
row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id)
return dict(row) if row is not None else None
def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool:
"""Compaction archives are ``active=0, compacted=1``; rewind/undo rows are
``active=0, compacted=0`` and must stay hidden."""
return state is not None and state["active"] == 0 and state["compacted"] == 1
def _is_compacted_message(db, message_id) -> bool:
"""True for a compaction-archived row — content no longer in live context, so
it stays discoverable even on the current session. False on any error."""
return _is_compacted_state(_get_message_storage_state(db, message_id))
def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None:
"""Note rebuild progress while the deferred FTS backfill runs so the agent can
explain thin results instead of treating them as ground truth. Never raises."""
try:
status = db.fts_rebuild_status()
except Exception:
return
if status is None:
return
payload["index_rebuild"] = {"percent": status["percent"], "note": (
f"The search index is rebuilding in the background ({status['percent']}% done, "
f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages "
f"may be incomplete until it finishes.")}
def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Stable-sort so interactive sessions rank above automation; BM25 order is
kept within each class, so a cron hit never displaces an interactive one."""
return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0)
def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None,
max_content_len: Optional[int] = None) -> Dict[str, Any]:
"""Slim a message row. Keeps ``content`` even when empty (tool-call-only
assistant turns); with *max_content_len* truncates and flags it."""
content = m.get("content")
if isinstance(content, str) and "\x1b" in content: # archived terminal output carries ANSI
from tools.ansi_strip import strip_ansi
content = strip_ansi(content)
entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")}
entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)})
if anchor_id is not None and m.get("id") == anchor_id:
entry["anchor"] = True
if max_content_len and content and len(content) > max_content_len:
entry.update(content=content[:max_content_len] + "…", content_truncated=True, original_content_chars=len(content))
return {k: v for k, v in entry.items() if v is not None or k == "content"}
def _session_link(session_id: str, profile: str = None) -> str:
"""The reference the agent writes for a session — same value the desktop composer
emits, so it renders as a titled link. The profile segment is omitted when it
can't be named confidently (a bare id still resolves, just not across profiles)."""
name = (profile or "").strip()
if not name:
def _active():
from hermes_cli.profiles import get_active_profile_name
resolved = get_active_profile_name()
return "" if resolved == "custom" else resolved
name = _quiet(_active, "", "get_active_profile_name failed for session link")
return f"@session:{name}/{session_id}" if name else f"@session:{session_id}"
def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]:
"""Return a discovery-shaped result when the query matches a session title."""
title_query = query.strip().strip("`'\"") # models often quote a remembered title
if not title_query:
return None
session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None,
"resolve_session_by_title failed for %r", title_query)
if not session_id:
return None
lineage_root = _resolve_lineage(db, session_id)
# Same-lineage title hits are in-context only while the session is live;
# /new-reset and compression-ended parents are not.
if (current_lineage_root and lineage_root == current_lineage_root
and not _session_left_live_context(db, session_id)):
return None
session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None,
"get_session failed for title match %s", session_id) or {}
if session_meta.get("source") in _HIDDEN_SESSION_SOURCES:
return None
messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id)
anchor_id = messages[0].get("id") if messages else None
view = {} if anchor_id is None else _quiet(
lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {},
"get_anchored_view failed for title match %s/%s", session_id, anchor_id)
title = session_meta.get("title") or title_query
entry = _discovery_entry(
lineage_root, session_id=session_id, when=_format_timestamp(session_meta.get("started_at")),
source=session_meta.get("source", "unknown"), model=session_meta.get("model") or "unknown",
title=title, matched_role="session_title", match_message_id=anchor_id,
snippet=f"Session title matched: {title}",
bookend_start=[_shape_message(m) for m in (view.get("bookend_start") or messages[:3])],
messages=[_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])],
bookend_end=[_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])],
messages_before=view.get("messages_before", 0),
messages_after=view.get("messages_after", max(len(messages) - 5, 0)), detail="full")
entry["_lineage_root"] = lineage_root
return entry
def _discovery_entry(lineage_root: Optional[str], **fields) -> Dict[str, Any]:
"""One discovery result in canonical key order; ``parent_session_id`` is set
when the hit lives in a child of its lineage root."""
entry = {k: fields[k] for k in (
"session_id", "when", "source", "model", "title", "matched_role", "match_message_id", "snippet",
"bookend_start", "messages", "bookend_end", "messages_before", "messages_after", "detail")}
if lineage_root and lineage_root != entry["session_id"]:
entry["parent_session_id"] = lineage_root
return entry
def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str:
payload = {"success": True, "mode": "discover", "query": query, "detail": detail,
"results": results, "count": len(results), **extra}
_annotate_rebuild_status(db, payload)
return json.dumps(payload, ensure_ascii=False)
def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None:
"""Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*.
The raw owning session_id stays on the row — only it pairs validly with the FTS
match id. Current-lineage hits are skipped UNLESS the transcript left live
context (compression-ended, /new-reset predecessor, or an in-place compacted row
on the SAME session); a live delegation child (end_reason=None) stays excluded."""
for r in raw_results:
if len(seen_sessions) >= limit:
break
raw_sid = r["session_id"]
resolved_sid = _resolve_lineage(db, raw_sid)
is_compacted_hit = _is_compacted_message(db, r.get("id"))
in_current_lineage = bool(current_lineage_root) and resolved_sid == current_lineage_root
if in_current_lineage and not (_session_left_live_context(db, raw_sid) or is_compacted_hit):
continue
if current_session_id and raw_sid == current_session_id and not is_compacted_hit:
continue
seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid})
def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]:
return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or [])
if not _is_compaction_summary(m.get("content", ""))]
def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]:
"""One discovery result from a surviving FTS row; None (hit dropped) if the
anchored view can't be loaded."""
hit_sid = match_info.get("session_id") or lineage_root
msg_id = match_info.get("id")
try:
view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3)
except Exception as e:
logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True)
return None
session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {}
full = result_detail == "full"
window_messages = [m for m in (view.get("window") or []) if full or m.get("id") == msg_id]
return _discovery_entry(
lineage_root, session_id=hit_sid,
when=_format_timestamp(session_meta.get("started_at") or match_info.get("session_started")),
source=session_meta.get("source") or match_info.get("source", "unknown"),
model=session_meta.get("model") or match_info.get("model") or "unknown",
title=session_meta.get("title") or None, matched_role=match_info.get("role"),
match_message_id=msg_id, snippet=match_info.get("snippet") or "",
bookend_start=_bookend(view, "bookend_start") if full else [],
messages=[_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages],
bookend_end=_bookend(view, "bookend_end") if full else [],
messages_before=view.get("messages_before", 0), messages_after=view.get("messages_after", 0),
detail=result_detail)
def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str],
detail: str, current_session_id: str = None, link_profile: str = None) -> str:
"""Discovery shape: FTS5 plus adaptive or full result hydration."""
current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None
title_result = _title_match_result(db, query, current_lineage_root)
raw_results, err = _loud(lambda: db.search_messages(
query=query, role_filter=role_filter or ["user", "assistant"],
exclude_sources=list(_HIDDEN_SESSION_SOURCES), limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort,
fields=_DISCOVER_SEARCH_FIELDS), "FTS5 search failed: %s", "Search failed")
if err:
return err
# Demote cron rows below interactive ones BEFORE dedup so a high-volume cron
# corpus can't starve the user's own sessions out of the top `limit`.
raw_results = _order_for_recall(raw_results)
if not raw_results and not title_result:
return _discover_payload(db, query, detail, [], message=(
"No matching sessions found. FTS5 ANDs all terms by default — "
"broaden with OR (`alpha OR beta`), exact-match with quoted "
"phrases, exclude with NOT, or prefix-match with `deploy*`."))
seen_sessions: Dict[str, Dict[str, Any]] = {}
results = []
if title_result:
title_lineage = title_result.pop("_lineage_root", None)
if title_lineage:
seen_sessions[title_lineage] = {"_title_only": True}
results.append(title_result)
_dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root)
for lineage_root, match_info in seen_sessions.items():
if match_info.get("_title_only"):
continue
# Adaptive: only the top-ranked result is fully hydrated.
entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact")
if entry is not None:
results.append(entry)
for entry in results:
entry["link"] = _session_link(entry["session_id"], link_profile)
return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=(
"When referring the user to a session, write its `link` value "
"verbatim inline mid-sentence (it renders as a titled link) — never "
"as markdown, in backticks, on its own line, or next to the "
"title/id/date. To read more around a compact result, scroll: "
"session_search(session_id=..., around_message_id=match_message_id)."))
def _resolve_profile_db(profile: str):
"""Open another profile's ``state.db`` read-only (no write lock — safe on a
live DB), or None for the current profile."""
"""Another profile's ``state.db`` opened read-only (safe on a live DB), or None
for the current profile."""
if profile is None or not str(profile).strip():
return None
from hermes_cli import profiles as profiles_mod
from hermes_state import SessionDB
canon = profiles_mod.normalize_profile_name(profile)
profiles_mod.validate_profile_name(canon)
if not profiles_mod.profile_exists(canon):
@@ -42,10 +368,9 @@ def _resolve_profile_db(profile: str):
def _locate_session_db(session_id: str):
"""Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)``
or ``(None, None)``. Ids are globally unique, so the first hit is authoritative."""
"""Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)`` or
``(None, None)``. Ids are globally unique, so the first hit is authoritative."""
from pathlib import Path
try:
from hermes_cli import profiles as profiles_mod
from hermes_state import SessionDB
@@ -60,14 +385,12 @@ def _locate_session_db(session_id: str):
if str(db_path) in seen or not db_path.exists():
continue
seen.add(str(db_path))
try:
pdb = SessionDB(db_path=db_path, read_only=True)
except Exception:
continue
if _quiet(lambda: pdb.get_session(session_id), None,
"get_session probe failed for %s in %s", session_id, name):
pdb = _quiet(lambda: SessionDB(db_path=db_path, read_only=True), None, "open %s failed", db_path)
if pdb and _quiet(lambda: pdb.get_session(session_id), None,
"get_session probe failed for %s in %s", session_id, name):
return pdb, name
pdb.close()
if pdb:
pdb.close()
return None, None
@@ -78,18 +401,14 @@ def _get_session_meta(db, session_id: str) -> dict:
def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_profile: str = None) -> str:
"""Read shape: whole session, or first ``head`` + last ``tail`` messages
with a pointer to scroll the middle."""
"""Read shape: whole session, or ``head`` + ``tail`` messages with a scroll pointer."""
meta = _get_session_meta(db, session_id)
if not meta:
return tool_error(f"session_id not found: {session_id}", success=False)
try:
rows = db.get_messages(session_id)
except Exception as e:
logging.error("get_messages failed for %s: %s", session_id, e, exc_info=True)
return tool_error(f"failed to load session: {e}", success=False)
rows, err = _loud(lambda: db.get_messages(session_id), "get_messages failed for %s: %s", "failed to load session",
session_id)
if err:
return err
shaped = [_shape_message(m) for m in rows]
total = len(shaped)
truncated = total > head + tail
@@ -102,86 +421,67 @@ def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_prof
def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_profile: str = None) -> str:
"""Browse shape: metadata for the most recent sessions (no LLM, no FTS5)."""
try:
# list_sessions_rich (include_children=False) already applies the
# canonical child classifier: roots, /branch children and /new-reset
# children are admitted, delegation/compression children hidden.
# Re-classifying here re-hid legacy reset children — trust the query.
def _browse():
# list_sessions_rich already applies the canonical child classifier (roots,
# /branch and /new-reset children admitted; delegation/compression children
# hidden). Re-classifying here re-hid legacy reset children — trust the query.
sessions = db.list_sessions_rich(
limit=limit + 15, # extra so we can skip current / compression roots
exclude_sources=list(_HIDDEN_SESSION_SOURCES),
order_by_last_active=True,
)
exclude_sources=list(_HIDDEN_SESSION_SOURCES), order_by_last_active=True)
current_root, has_compression_hop = (
_resolve_to_parent(db, current_session_id)
if current_session_id else (None, False)
)
_resolve_to_parent(db, current_session_id) if current_session_id else (None, False))
results = []
for s in sessions:
sid = s.get("id", "")
if sid == current_session_id:
continue
# Compression continuation: the root's turns were summarised into
# the live child, so hide the root. /new-reset children share a
# root but carry no transcript — keep that root browsable.
if has_compression_hop and current_root and sid == current_root:
# Compression continuation: the root was summarised into the live child, so
# hide it. /new-reset children carry no transcript — keep that root browsable.
if sid == current_session_id or (has_compression_hop and current_root and sid == current_root):
continue
results.append({
"session_id": sid, "link": _session_link(sid, link_profile), "title": s.get("title") or None,
"source": s.get("source", ""), "started_at": s.get("started_at", ""),
"last_active": s.get("last_active", ""), "message_count": s.get("message_count", 0),
"preview": s.get("preview", ""),
})
**{k: s.get(k, "") for k in ("source", "started_at", "last_active")},
"message_count": s.get("message_count", 0), "preview": s.get("preview", "")})
if len(results) >= limit:
break
return _ok(mode="browse", results=results, count=len(results), message=(
f"Showing {len(results)} most recent sessions. Pass a query= to search, "
"or session_id+around_message_id to scroll."))
except Exception as e:
logging.error("Error listing recent sessions: %s", e, exc_info=True)
return tool_error(f"Failed to list recent sessions: {e}", success=False)
out, err = _loud(_browse, "Error listing recent sessions: %s", "Failed to list recent sessions")
return err or out
def _clamp_int(value, default: int, lo: int, hi: int) -> int:
if not isinstance(value, int):
try:
value = int(value)
except (TypeError, ValueError):
value = default
try:
value = int(value)
except (TypeError, ValueError):
value = default
return max(lo, min(value, hi))
def _anchor_in_live_context(db, anchor_state, anchor_session_id: str, current_session_id: str) -> bool:
"""True when the scroll anchor is still in the caller's active context and
must be rejected. Same-lineage history that has LEFT live context (compacted
rows, compression-ended parents, /new-reset predecessors) is allowed, so
scroll never rejects a result discovery just returned."""
a_root = _resolve_lineage(db, anchor_session_id)
c_root = _resolve_lineage(db, current_session_id)
if not (a_root and c_root and a_root == c_root):
return False
if _is_compacted_state(anchor_state):
"""True when the scroll anchor is still in the caller's active context (reject).
Same-lineage history that LEFT live context (compacted rows, compression-ended
parents, /new-reset predecessors) is allowed, so scroll never rejects a result
discovery just returned."""
if not _same_lineage(db, anchor_session_id, current_session_id) or _is_compacted_state(anchor_state):
return False
# Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history.
is_inactive_non_compacted = (
anchor_state is not None
and anchor_state["active"] == 0
and anchor_state["compacted"] != 1
)
return is_inactive_non_compacted or not _session_left_live_context(db, anchor_session_id)
inactive_non_compacted = anchor_state is not None and anchor_state["active"] == 0 and anchor_state["compacted"] != 1
return inactive_non_compacted or not _session_left_live_context(db, anchor_session_id)
def _same_lineage(db, a: str, b: str) -> bool:
a_root, b_root = _resolve_lineage(db, a), _resolve_lineage(db, b)
return bool(a_root and b_root and a_root == b_root)
def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, window: int):
"""Lineage rebind: the caller paired a parent session_id with a message id
that lives in a descendant (compaction / delegation create child sessions).
Returns ``(view, warning)`` from the owning session, or ``(None, None)``."""
a_root = _resolve_lineage(db, session_id)
o_root = _resolve_lineage(db, owning)
if not (a_root and o_root and a_root == o_root):
return None, None
rebind_view = _quiet(lambda: db.get_messages_around(owning, around_message_id, window=window),
None, "rebind get_messages_around failed: %s", with_exc=True)
"""Lineage rebind when the caller paired a parent session_id with a message id
living in a descendant. ``(view, warning)`` from the owner, or ``(None, None)``."""
rebind_view = _same_lineage(db, session_id, owning) and _quiet(
lambda: db.get_messages_around(owning, around_message_id, window=window),
None, "rebind get_messages_around failed: %s", with_exc=True)
if not (rebind_view and rebind_view.get("window")):
return None, None
return rebind_view, f"around_message_id {around_message_id} lives in {owning} (child of {session_id}); rebound transparently"
@@ -189,8 +489,8 @@ def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, w
def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
current_session_id: str = None) -> str:
"""Scroll shape: a window of messages centered on an anchor (no FTS5, no
bookends). Rebinds silently if the anchor lives in a same-lineage child."""
"""Scroll shape: a window centered on an anchor (no FTS5, no bookends);
rebinds silently if the anchor lives in a same-lineage child."""
if not isinstance(session_id, str) or not session_id.strip():
return tool_error("scroll requires session_id", success=False)
session_id = session_id.strip()
@@ -199,41 +499,29 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
except (TypeError, ValueError):
return tool_error("scroll requires integer around_message_id", success=False)
window = _clamp_int(window, 5, 1, 20)
# Locate the anchor BEFORE the current-lineage guard (see _anchor_in_live_context).
anchor_state = _get_message_storage_state(db, around_message_id)
owning_session_id = anchor_state.get("session_id") if anchor_state is not None else None
if current_session_id and _anchor_in_live_context(
db, anchor_state, owning_session_id or session_id, current_session_id
):
db, anchor_state, owning_session_id or session_id, current_session_id):
return tool_error("scroll rejected: anchor lives in the current session lineage (already in your active context)", success=False)
session_meta = _get_session_meta(db, session_id)
if not session_meta:
return tool_error(f"session_id not found: {session_id}", success=False)
try:
view = db.get_messages_around(session_id, around_message_id, window=window)
except Exception as e:
logging.error("get_messages_around failed: %s", e, exc_info=True)
return tool_error(f"failed to load messages: {e}", success=False)
view, err = _loud(lambda: db.get_messages_around(session_id, around_message_id, window=window),
"get_messages_around failed: %s", "failed to load messages")
if err:
return err
messages = view.get("window") or []
rebind_warning = None
if not messages and owning_session_id and owning_session_id != session_id:
rebind_view, rebind_warning = _rebind_to_owner(
db, session_id, owning_session_id, around_message_id, window
)
rebind_view, rebind_warning = _rebind_to_owner(db, session_id, owning_session_id, around_message_id, window)
if rebind_view is not None:
view = rebind_view
messages = view["window"]
view, messages = rebind_view, rebind_view["window"]
session_meta = _get_session_meta(db, owning_session_id) or session_meta
session_id = owning_session_id
if not messages:
return tool_error(f"around_message_id {around_message_id} not in session_id {session_id}", success=False)
return _ok(
mode="scroll", session_id=session_id, around_message_id=around_message_id,
session_meta=_session_meta_block(session_meta), window=window,
@@ -243,101 +531,84 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
"id; backward: the FIRST message's id (the boundary message repeats "
"as an orientation marker). messages_before/messages_after < window "
"means you've hit that end of the session."),
**({"warning": rebind_warning} if rebind_warning else {}),
)
**({"warning": rebind_warning} if rebind_warning else {}))
def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str:
"""Read shape. On a miss in the target profile, scan every profile (the
model may have dropped the owning profile from the link) and tag the result
with the profile it was found in."""
"""Read shape; on a miss scan every profile (the model may have dropped the
owning profile from the link) and tag the result with where it was found."""
result = _read_session(db, sid, link_profile=profile)
if json.loads(result).get("success"):
return result
located, owner = _locate_session_db(sid)
if located is not None:
try:
found = json.loads(_read_session(located, sid, link_profile=owner))
finally:
located.close()
if found.get("success"):
found["profile"] = owner
return json.dumps(found, ensure_ascii=False)
return result
if located is None:
return result
try:
found = json.loads(_read_session(located, sid, link_profile=owner))
finally:
located.close()
if not found.get("success"):
return result
found["profile"] = owner
return json.dumps(found, ensure_ascii=False)
def _dispatch(query, role_filter, limit, db, current_session_id, session_id,
around_message_id, window, sort, profile, detail, owned_dbs) -> str:
"""Mode dispatch (see module docstring). Scroll wins over read/discovery when
an anchor is set — the agent asked for a specific slice. Profile DBs opened
here are appended to *owned_dbs* for the caller to close."""
# A raw `@session:<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:
session_id = emb_id
if emb_profile and (profile is None or not str(profile).strip()):
profile = emb_profile
# Cross-profile read: swap in the named profile's DB (read-only) for every
# shape. Current-lineage guards key off ids that won't collide, so they
# stay inert.
if profile is not None and str(profile).strip():
try:
profile_db = _resolve_profile_db(profile)
except Exception as e:
return tool_error(f"profile '{profile}': {e}", success=False)
if profile_db is not None:
db, current_session_id = profile_db, None
owned_dbs.append(profile_db)
has_session = isinstance(session_id, str) and bool(session_id.strip())
if has_session and around_message_id is not None:
return _scroll(db, session_id, around_message_id, window, current_session_id)
if has_session:
# Cross-profile: swap in the named profile's DB (read-only) for every shape;
# current-lineage guards key off ids that won't collide, so they stay inert.
try:
profile_db = _resolve_profile_db(profile)
except Exception as e:
return tool_error(f"profile '{profile}': {e}", success=False)
if profile_db is not None:
db, current_session_id = profile_db, None
owned_dbs.append(profile_db)
if isinstance(session_id, str) and session_id.strip():
if around_message_id is not None:
return _scroll(db, session_id, around_message_id, window, current_session_id)
return _read_with_profile_fallback(db, session_id.strip(), profile)
limit = _clamp_int(limit, 3, 1, 10)
if not query or not isinstance(query, str) or not query.strip():
return _list_recent_sessions(db, limit, current_session_id, link_profile=profile)
role_list = ([r.strip() for r in role_filter.split(",") if r.strip()] or None) if isinstance(role_filter, str) else None
sort_norm = sort.strip().lower() if isinstance(sort, str) else None
if sort_norm not in ("newest", "oldest"):
sort_norm = None
sort_norm = sort_norm if sort_norm in ("newest", "oldest") else None
detail_norm = "full" if isinstance(detail, str) and detail.strip().lower() == "full" else "adaptive"
return _discover(
db=db, query=query.strip(), role_filter=role_list, limit=limit, sort=sort_norm,
detail=detail_norm, current_session_id=current_session_id, link_profile=profile,
)
detail=detail_norm, current_session_id=current_session_id, link_profile=profile)
def session_search(query: str = "", role_filter: str = None, limit: int = 3, db=None,
current_session_id: str = None, session_id: str = None, around_message_id: int = None,
window: int = 5, sort: str = None, profile: str = None, detail: str = "adaptive") -> str:
"""Run session search and close databases opened by this invocation.
Parameter order is positional-compatible with older callers."""
"""Run session search, closing DBs opened here. Positional order is frozen for old callers."""
owned_dbs: List[Any] = []
if db is None:
try:
from hermes_state import get_shared_session_db
db = get_shared_session_db()
owned_dbs.append(db)
except Exception:
logging.debug("SessionDB unavailable for session_search", exc_info=True)
from hermes_state import format_session_db_unavailable
return tool_error(format_session_db_unavailable(), success=False)
try:
return _dispatch(query, role_filter, limit, db, current_session_id, session_id,
around_message_id, window, sort, profile, detail, owned_dbs)
finally:
from hermes_state import release_or_close
for owned_db in reversed(owned_dbs):
_quiet(lambda: release_or_close(owned_db), None, "Failed to close session_search SessionDB")
@@ -465,18 +736,8 @@ registry.register(
toolset="session_search",
schema=SESSION_SEARCH_SCHEMA,
handler=lambda args, **kw: session_search(
query=args.get("query") or "",
role_filter=args.get("role_filter"),
limit=args.get("limit", 3),
session_id=args.get("session_id"),
around_message_id=args.get("around_message_id"),
window=args.get("window", 5),
sort=args.get("sort"),
detail=args.get("detail", "adaptive"),
profile=args.get("profile"),
db=kw.get("db"),
current_session_id=kw.get("current_session_id"),
),
query=args.get("query") or "", limit=args.get("limit", 3), window=args.get("window", 5),
detail=args.get("detail", "adaptive"), db=kw.get("db"), current_session_id=kw.get("current_session_id"),
**{k: args.get(k) for k in ("role_filter", "session_id", "around_message_id", "sort", "profile")}),
check_fn=check_session_search_requirements,
emoji="🔍",
)
emoji="🔍")

View File

@@ -1,230 +0,0 @@
"""Shared helpers for the session_search tool: source classification, lineage
resolution, message storage state, and response shaping. Imported by
``tools.session_search_tool`` (which re-exports the names) and
``tools.session_search_tool_discover``."""
import json
import logging
from datetime import datetime
from typing import Any, Dict, List, Optional, Union
from hermes_state_common import _RESET_END_REASONS
# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool),
# delegate subagent runs, kanban workers — not the user's history.
_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool")
# Searchable but DEMOTED below interactive sessions: cron sessions' repetitive
# vocabulary dominates bare BM25 and starves out the user's own sessions
# ("recall blindness"). Demoting keeps them reachable when they're the only match.
_DEMOTED_SESSION_SOURCES = ("cron",)
# FTS rows scanned before dedup-by-lineage — well above the handful of distinct
# sessions a query returns, so interactive matches buried under cron hits are
# still in hand for the demotion pass.
_DISCOVER_SCAN_LIMIT = 300
# Raw FTS rows are only a discovery-plan input; the response hydrates its own
# anchored window and bookends after lineage dedup.
_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started")
# Generated context-compaction handoff summaries (agent/context_compressor.py);
# excluded from bookends so huge compaction payloads aren't re-introduced.
_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:")
# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the
# predecessor WITHOUT carrying its transcript forward — unlike compression
# continuations and live delegation children. Derived from the gateway set so
# this tool and the recovery fence cannot drift.
_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"}
def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
"""Call ``fn()``; on any exception debug-log *msg* (appending the exception
as a final ``%s`` arg when *with_exc*) and return *default*."""
try:
return fn()
except Exception as e:
logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True)
return default
def _format_timestamp(ts: Union[int, float, str, None]) -> str:
"""Unix timestamp (number or numeric string) or ISO string -> readable date.
"unknown" for None; str(ts) if conversion fails."""
if ts is None:
return "unknown"
try:
value = ts
if isinstance(ts, str):
if not ts.replace(".", "").replace("-", "").isdigit():
return ts
value = float(ts)
if isinstance(value, (int, float)):
return datetime.fromtimestamp(value).strftime("%B %d, %Y at %I:%M %p")
except (ValueError, OSError, OverflowError) as e:
logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True)
except Exception as e:
logging.debug("Unexpected error formatting timestamp %s: %s", ts, e, exc_info=True)
return str(ts)
def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]:
"""The ``session_meta`` sub-object shared by read/scroll responses."""
return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"),
"model": meta.get("model"), "title": meta.get("title")}
def _ok(**payload) -> str:
"""Serialize a successful tool result (``success`` first, then *payload* in order)."""
return json.dumps({"success": True, **payload}, ensure_ascii=False)
def _is_compaction_summary(content: str) -> bool:
"""Return True if *content* looks like a generated compaction handoff."""
return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES)
def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]:
"""Walk parent_session_id to the lineage root -> ``(root_id, has_compression_hop)``.
The flag distinguishes a compression-split lineage (parent content summarised
away) from a delegation lineage (child content still visible to the parent).
Falls back to ``(session_id, False)`` on errors."""
if not session_id:
return session_id, False
visited: set[str] = set()
cur, has_compression = session_id, False
while cur and cur not in visited:
visited.add(cur)
s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True)
if not s:
break
if s.get("end_reason") == "compression":
has_compression = True
if not s.get("parent_session_id"):
break
cur = s["parent_session_id"]
return cur, has_compression
def _resolve_lineage(db, session_id: str) -> str:
"""Return only the lineage root (ignores the compression hop)."""
return _resolve_to_parent(db, session_id)[0]
def _session_end_reason(db, session_id: str) -> Optional[str]:
"""Return the session's ``end_reason``, or None if missing/unended/error."""
if not session_id:
return None
try:
s = db.get_session(session_id)
return (s.get("end_reason") or None) if s else None
except Exception:
return None
def _session_left_live_context(db, session_id: str) -> bool:
"""True when *session_id*'s transcript is no longer in anyone's live context:
``compression`` (summarised into the continuation child) or a fresh reset
(child starts empty). Everything else stays excluded from same-lineage
recall — live delegation children (``end_reason is None``) are visible to
the parent agent, and ``branched`` parents were copied verbatim into the
branch child, so their content IS the current context."""
end_reason = _session_end_reason(db, session_id)
return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS
def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]:
"""Return the owning session and visibility flags for *message_id*."""
if not message_id:
return None
def _lookup():
with db._lock:
return db._conn.execute(
"SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,)
).fetchone()
row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id)
return dict(row) if row is not None else None
def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool:
"""Compaction archives are ``active=0, compacted=1`` (content summarised
away by archive_and_compact). Rewind/undo rows are ``active=0, compacted=0``
and must stay hidden."""
return state is not None and state["active"] == 0 and state["compacted"] == 1
def _is_compacted_message(db, message_id) -> bool:
"""True if *message_id* is a compaction-archived row — pre-compaction content
no longer in live context, so it should stay discoverable even on the
current session. False on any error (caller falls back to skipping)."""
return _is_compacted_state(_get_message_storage_state(db, message_id))
def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None:
"""Add a rebuild-progress note while the deferred FTS backfill is running,
so the agent can explain thin/slow results instead of treating them as
ground truth. No-op (never raises) when no rebuild is pending."""
try:
status = db.fts_rebuild_status()
except Exception:
return
if status is None:
return
payload["index_rebuild"] = {"percent": status["percent"], "note": (
f"The search index is rebuilding in the background ({status['percent']}% done, "
f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages "
f"may be incomplete until it finishes."
)}
def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Stable-sort FTS rows so interactive sessions rank above automation.
BM25 order is preserved within each class; only cross-class order changes,
so a cron hit never displaces an interactive hit during lineage dedup."""
return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0)
def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None,
max_content_len: Optional[int] = None) -> Dict[str, Any]:
"""Slim a message row for the tool response. Keeps content even if empty
(absent content is meaningful — tool-call-only assistant turns). With
*max_content_len*, content is truncated and ``content_truncated`` /
``original_content_chars`` added."""
content = m.get("content")
if isinstance(content, str) and "\x1b" in content:
# Recalled messages can carry ANSI escapes (archived terminal output).
from tools.ansi_strip import strip_ansi
content = strip_ansi(content)
original_chars = None
if max_content_len and content and len(content) > max_content_len:
original_chars = len(content)
content = content[:max_content_len] + "…"
entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")}
entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)})
if anchor_id is not None and m.get("id") == anchor_id:
entry["anchor"] = True
if original_chars is not None:
entry["content_truncated"] = True
entry["original_content_chars"] = original_chars
return {k: v for k, v in entry.items() if v is not None or k == "content"}
def _session_link(session_id: str, profile: str = None) -> str:
"""The reference the agent writes to point the user at a session — same
value the desktop composer emits, so it renders as a titled link. The
profile segment is omitted when it can't be named confidently (a bare id
still resolves, it just can't disambiguate across profiles)."""
name = (profile or "").strip()
if not name:
def _active():
from hermes_cli.profiles import get_active_profile_name
resolved = get_active_profile_name()
return "" if resolved == "custom" else resolved
name = _quiet(_active, "", "get_active_profile_name failed for session link")
return f"@session:{name}/{session_id}" if name else f"@session:{session_id}"

View File

@@ -1,192 +0,0 @@
"""Discovery shape of session_search: FTS5 query, title match, lineage dedup
and adaptive/full hydration of the surviving results."""
import json
import logging
from typing import Any, Dict, List, Optional
from tools.registry import tool_error
from tools.session_search_tool_common import (
_DISCOVER_SCAN_LIMIT, _DISCOVER_SEARCH_FIELDS, _HIDDEN_SESSION_SOURCES, _annotate_rebuild_status,
_format_timestamp, _is_compacted_message, _is_compaction_summary, _order_for_recall, _quiet,
_resolve_lineage, _resolve_to_parent, _session_left_live_context, _session_link, _shape_message,
)
def _normalize_title_query(query: str) -> str:
"""Strip common quoting the model may include around a remembered title."""
return query.strip().strip("`'\"")
def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]:
"""Return a discovery-shaped result when the query matches a session title."""
title_query = _normalize_title_query(query)
if not title_query:
return None
session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None,
"resolve_session_by_title failed for %r", title_query)
if not session_id:
return None
lineage_root = _resolve_lineage(db, session_id)
# Same-lineage title hits are in-context only while the session is live;
# /new-reset and compression-ended parents are not.
if (
current_lineage_root
and lineage_root == current_lineage_root
and not _session_left_live_context(db, session_id)
):
return None
session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None,
"get_session failed for title match %s", session_id) or {}
if session_meta.get("source") in _HIDDEN_SESSION_SOURCES:
return None
messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id)
anchor_id = messages[0].get("id") if messages else None
view = {}
if anchor_id is not None:
view = _quiet(lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {},
"get_anchored_view failed for title match %s/%s", session_id, anchor_id)
entry = {
"session_id": session_id, "when": _format_timestamp(session_meta.get("started_at")),
"source": session_meta.get("source", "unknown"), "model": session_meta.get("model") or "unknown",
"title": session_meta.get("title") or title_query, "matched_role": "session_title",
"match_message_id": anchor_id,
"snippet": f"Session title matched: {session_meta.get('title') or title_query}",
"bookend_start": [_shape_message(m) for m in (view.get("bookend_start") or messages[:3])],
"messages": [_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])],
"bookend_end": [_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])],
"messages_before": view.get("messages_before", 0),
"messages_after": view.get("messages_after", max(len(messages) - 5, 0)),
"detail": "full", "_lineage_root": lineage_root,
}
if lineage_root and lineage_root != session_id:
entry["parent_session_id"] = lineage_root
return entry
def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str:
payload = {"success": True, "mode": "discover", "query": query, "detail": detail,
"results": results, "count": len(results), **extra}
_annotate_rebuild_status(db, payload)
return json.dumps(payload, ensure_ascii=False)
def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None:
"""Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*.
The raw owning session_id stays on the row — only it pairs validly with the
FTS match id for the anchored window. Current-lineage hits are skipped
UNLESS the transcript left live context: compression-ended session, /new-
reset predecessor (hiding it made gateway recall blind after every /new),
or an in-place compacted row on the SAME session_id. A live delegation
child has end_reason=None, so it stays excluded.
"""
for r in raw_results:
if len(seen_sessions) >= limit:
break
raw_sid = r["session_id"]
resolved_sid, _ = _resolve_to_parent(db, raw_sid)
is_compacted_hit = _is_compacted_message(db, r.get("id"))
is_ended_session = _session_left_live_context(db, raw_sid)
if (
current_lineage_root
and resolved_sid == current_lineage_root
and not (is_ended_session or is_compacted_hit)
):
continue
if current_session_id and raw_sid == current_session_id and not is_compacted_hit:
continue
seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid})
def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]:
return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or [])
if not _is_compaction_summary(m.get("content", ""))]
def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]:
"""Build one discovery result from a surviving FTS row; None if the anchored
view can't be loaded (the hit is dropped)."""
hit_sid = match_info.get("session_id") or lineage_root
msg_id = match_info.get("id")
try:
view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3)
except Exception as e:
logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True)
return None
session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {}
full = result_detail == "full"
window_messages = view.get("window") or []
if not full:
window_messages = [m for m in window_messages if m.get("id") == msg_id]
entry = {
"session_id": hit_sid,
"when": _format_timestamp(session_meta.get("started_at") or match_info.get("session_started")),
"source": session_meta.get("source") or match_info.get("source", "unknown"),
"model": session_meta.get("model") or match_info.get("model") or "unknown",
"title": session_meta.get("title") or None, "matched_role": match_info.get("role"),
"match_message_id": msg_id, "snippet": match_info.get("snippet") or "",
"bookend_start": _bookend(view, "bookend_start") if full else [],
"messages": [_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages],
"bookend_end": _bookend(view, "bookend_end") if full else [],
"messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", 0),
"detail": result_detail,
}
if lineage_root and lineage_root != hit_sid:
entry["parent_session_id"] = lineage_root
return entry
def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str],
detail: str, current_session_id: str = None, link_profile: str = None) -> str:
"""Discovery shape: FTS5 plus adaptive or full result hydration."""
role_list = role_filter if role_filter else ["user", "assistant"]
current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None
title_result = _title_match_result(db, query, current_lineage_root)
try:
raw_results = db.search_messages(
query=query, role_filter=role_list, exclude_sources=list(_HIDDEN_SESSION_SOURCES),
limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS,
)
except Exception as e:
logging.error("FTS5 search failed: %s", e, exc_info=True)
return tool_error(f"Search failed: {e}", success=False)
# Demote cron rows below interactive ones BEFORE dedup so a high-volume
# cron corpus can't starve the user's own sessions out of the top `limit`.
raw_results = _order_for_recall(raw_results)
if not raw_results and not title_result:
return _discover_payload(db, query, detail, [], message=(
"No matching sessions found. FTS5 ANDs all terms by default — "
"broaden with OR (`alpha OR beta`), exact-match with quoted "
"phrases, exclude with NOT, or prefix-match with `deploy*`."
))
seen_sessions: Dict[str, Dict[str, Any]] = {}
results = []
if title_result:
title_lineage = title_result.pop("_lineage_root", None)
if title_lineage:
seen_sessions[title_lineage] = {"_title_only": True}
results.append(title_result)
_dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root)
for lineage_root, match_info in seen_sessions.items():
if match_info.get("_title_only"):
continue
# Adaptive: only the top-ranked result is fully hydrated.
entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact")
if entry is not None:
results.append(entry)
for entry in results:
entry["link"] = _session_link(entry["session_id"], link_profile)
return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=(
"When referring the user to a session, write its `link` value "
"verbatim inline mid-sentence (it renders as a titled link) — never "
"as markdown, in backticks, on its own line, or next to the "
"title/id/date. To read more around a compact result, scroll: "
"session_search(session_id=..., around_message_id=match_message_id)."
))