Merge branch 'simp/r3-33-I' into simp/r3-33
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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="🔍")
|
||||
|
||||
@@ -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}"
|
||||
@@ -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)."
|
||||
))
|
||||
Reference in New Issue
Block a user