refactor(hindsight): inline one-site helpers, tuple recall result, merged tool table

This commit is contained in:
Teknium
2026-09-02 23:41:53 -07:00
parent b2d09087c8
commit a30092e517

View File

@@ -20,7 +20,6 @@ import queue
import sys
import threading
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Callable, Dict, List, Optional
@@ -53,14 +52,6 @@ _LOCAL_MODES = {"local", "local_embedded"}
_RETAIN_CONTEXT_DEFAULT = "conversation between Hermes Agent and the User"
@dataclass(frozen=True)
class _RecallResult:
"""Recall text + memory count for the recall indicator (0 for reflect/error)."""
text: str
count: int
def _ensure_client_dependency() -> None:
"""Lazily install the Hindsight client (``tools.lazy_deps``) before importing it."""
try:
@@ -105,17 +96,6 @@ _append_capability_cache: Dict[str, bool] = {}
_append_capability_lock = threading.Lock()
def _meets_minimum_version(actual: str | None, required: str) -> bool:
"""True if *actual* >= *required* (semver). False on missing/invalid."""
if not actual:
return False
try:
from packaging.version import Version
return Version(actual) >= Version(required)
except Exception:
return False
def _fetch_hindsight_api_version(api_url: str, api_key: str | None = None,
timeout: float = 5.0) -> str | None:
"""GET ``<api_url>/version`` -> version string, or None on any failure (= legacy API)."""
@@ -145,7 +125,11 @@ def _check_api_supports_update_mode_append(api_url: str, api_key: str | None = N
if api_url in _append_capability_cache:
return _append_capability_cache[api_url]
version = _fetch_hindsight_api_version(api_url, api_key)
supported = _meets_minimum_version(version, _MIN_VERSION_FOR_UPDATE_MODE_APPEND)
try: # missing/invalid version -> unsupported
from packaging.version import Version
supported = bool(version) and Version(version) >= Version(_MIN_VERSION_FOR_UPDATE_MODE_APPEND)
except Exception:
supported = False
with _append_capability_lock:
# A concurrent probe may have filled the cache meanwhile; its answer wins.
supported = _append_capability_cache.setdefault(api_url, supported)
@@ -179,13 +163,10 @@ def _get_loop() -> asyncio.AbstractEventLoop:
with _loop_lock:
if _loop is not None and _loop.is_running():
return _loop
_loop = asyncio.new_event_loop()
def _run():
asyncio.set_event_loop(_loop)
_loop.run_forever()
_loop_thread = threading.Thread(target=_run, daemon=True, name="hindsight-loop")
loop = _loop = asyncio.new_event_loop()
_loop_thread = threading.Thread(
target=lambda: (asyncio.set_event_loop(loop), loop.run_forever()), daemon=True, name="hindsight-loop",
)
_loop_thread.start()
return _loop
@@ -295,11 +276,6 @@ def _load_config() -> dict:
}
def _utc_timestamp() -> str:
"""UTC write/audit time for retain metadata."""
return datetime.now(timezone.utc).isoformat(timespec="milliseconds").replace("+00:00", "Z")
def _event_timestamp() -> str:
"""Configured Hermes event time with an explicit UTC offset."""
event_time = _hermes_now()
@@ -333,11 +309,6 @@ _SYSTEM_PROMPT_TAILS = {
"Use hindsight_recall to search, hindsight_reflect for synthesis, "
"hindsight_retain to store facts."),
}
_TOOL_ERRORS = {
"hindsight_retain": "Failed to store memory",
"hindsight_recall": "Failed to search memory",
"hindsight_reflect": "Failed to reflect",
}
class HindsightMemoryProvider(MemoryProvider):
@@ -400,16 +371,7 @@ class HindsightMemoryProvider(MemoryProvider):
self._prefetch_lock = threading.Lock()
self._prefetch_thread = None
self._last_recall_returned, self._last_recall_count = False, 0
self._auto_recall = self._recall_indicator = True
self._recall_sync = False
self._recall_tags: list[str] | None = None
self._recall_tags_match = "any"
self._recall_max_tokens, self._recall_max_input_chars = 4096, 800
# Observation-only by default: observations are Hindsight's consolidated,
# deduplicated layer; raw world/experience facts re-ship the evidence
# they summarize and burn the recall_max_tokens budget.
self._recall_types: list[str] = ["observation"]
self._recall_prompt_preamble = ""
self._apply_recall_settings({}) # pure-config defaults (no env/secret reads)
@property
def name(self) -> str:
@@ -712,17 +674,15 @@ class HindsightMemoryProvider(MemoryProvider):
# -- retain target -----------------------------------------------------------
def _probe_url(self) -> str:
"""/version probe URL: the embedded client's dynamic per-profile port when running, else api_url."""
url = getattr(self._client, "url", None) if self._mode == "local_embedded" else None
return str(url) if url else (self._api_url or "")
def _resolve_retain_target(self, fallback_document_id: str) -> tuple[str, str | None]:
"""(document_id, update_mode) from live API capability: >= 0.5.0 reuses the
stable session-scoped id with ``update_mode='append'``; older APIs get
*fallback_document_id* (per-process unique) and no update_mode — the only
way the resume-overwrite fix works there."""
if self._session_id and _check_api_supports_update_mode_append(self._probe_url(), self._api_key):
way the resume-overwrite fix works there. The /version probe targets the
embedded client's dynamic per-profile port when running, else api_url."""
url = getattr(self._client, "url", None) if self._mode == "local_embedded" else None
probe_url = str(url) if url else (self._api_url or "")
if self._session_id and _check_api_supports_update_mode_append(probe_url, self._api_key):
return self._session_id, "append"
return fallback_document_id, None
@@ -840,8 +800,10 @@ class HindsightMemoryProvider(MemoryProvider):
self._recall_sync = bool(cfg.get("recall_sync", False))
self._recall_max_tokens = int(cfg.get("recall_max_tokens", 4096))
self._recall_max_input_chars = int(cfg.get("recall_max_input_chars", 800))
# None -> observation-only; a comma-separated string is accepted for
# parity with recall_tags; an explicit list broadens or disables the filter.
# None -> observation-only (Hindsight's consolidated, deduplicated layer; raw
# world/experience facts re-ship the evidence they summarize and burn the
# recall_max_tokens budget); a comma-separated string is accepted for parity
# with recall_tags; an explicit list broadens or disables the filter.
configured_types = cfg.get("recall_types")
if configured_types is None:
self._recall_types = ["observation"]
@@ -920,17 +882,13 @@ class HindsightMemoryProvider(MemoryProvider):
return True
return False
def _recall_kwargs(self, query: str) -> dict:
def _recall(self, query: str) -> list:
kwargs: dict = {"bank_id": self._bank_id, "query": query, "budget": self._budget, "max_tokens": self._recall_max_tokens}
if self._recall_tags:
kwargs.update(tags=self._recall_tags, tags_match=self._recall_tags_match)
if self._recall_types:
kwargs["types"] = self._recall_types
return kwargs
def _recall(self, query: str) -> list:
recall_kwargs = self._recall_kwargs(query)
resp = self._run_hindsight_operation(lambda client: client.arecall(**recall_kwargs))
resp = self._run_hindsight_operation(lambda client: client.arecall(**kwargs))
return resp.results or []
def _reflect(self, query: str) -> str | None:
@@ -939,22 +897,23 @@ class HindsightMemoryProvider(MemoryProvider):
)
return resp.text
def _do_recall(self, query: str) -> _RecallResult:
"""One recall/reflect for *query* (background prefetch and ``recall_sync`` paths)."""
def _do_recall(self, query: str) -> tuple[str, int]:
"""One recall/reflect for *query* (background prefetch and ``recall_sync`` paths)
-> (text, memory count); the count is 0 for reflect (synthesis) and on error."""
if self._recall_max_input_chars and len(query) > self._recall_max_input_chars:
query = query[:self._recall_max_input_chars]
try:
if self._prefetch_method == "reflect":
logger.debug("Recall: calling reflect (bank=%s, query_len=%d)", self._bank_id, len(query))
return _RecallResult(self._reflect(query) or "", 0) # synthesis -> no discrete count
return self._reflect(query) or "", 0
logger.debug("Recall: calling recall (bank=%s, query_len=%d, budget=%s)",
self._bank_id, len(query), self._budget)
results = self._recall(query)
logger.debug("Recall: returned %d results", len(results))
return _RecallResult("\n".join(f"- {r.text}" for r in results if r.text), len(results))
return "\n".join(f"- {r.text}" for r in results if r.text), len(results)
except Exception as e:
logger.debug("Hindsight recall failed: %s", e, exc_info=True)
return _RecallResult("", 0)
return "", 0
def _finish_prefetch(self, result: str, count: int) -> str:
"""Record indicator state (cleared on empty turns, never a stale count); format the block."""
@@ -980,8 +939,7 @@ class HindsightMemoryProvider(MemoryProvider):
# Opt-in: recall synchronously against the *current* message so the
# injected memories match this turn's query, not the previous turn's.
if self._recall_sync:
recalled = _RecallResult("", 0) if self._recall_disabled() else self._do_recall(query)
return self._finish_prefetch(recalled.text, recalled.count)
return self._finish_prefetch(*(("", 0) if self._recall_disabled() else self._do_recall(query)))
# Default: the background worker's result for the previous turn (capped join).
self._join_prefetch(3.0, log=True)
with self._prefetch_lock:
@@ -1005,10 +963,10 @@ class HindsightMemoryProvider(MemoryProvider):
# retain to be recall-visible so the warmed context includes it.
if self._prefetch_waits_for_retain:
self._wait_for_retains_drained(self._prefetch_retain_drain_timeout)
recalled = self._do_recall(query)
if recalled.text:
text, count = self._do_recall(query)
if text:
with self._prefetch_lock:
self._prefetch_result, self._prefetch_count = recalled.text, recalled.count
self._prefetch_result, self._prefetch_count = text, count
self._prefetch_thread = _context_thread(_run, "hindsight-prefetch")
self._prefetch_thread.start()
@@ -1023,7 +981,8 @@ class HindsightMemoryProvider(MemoryProvider):
def _build_metadata(self, *, message_count: int, turn_index: int) -> Dict[str, str]:
metadata: Dict[str, str] = {
"retained_at": _utc_timestamp(),
# UTC write/audit time (event time lives on the item timestamp).
"retained_at": datetime.now(timezone.utc).isoformat(timespec="milliseconds").replace("+00:00", "Z"),
"message_count": str(message_count),
"turn_index": str(turn_index),
}
@@ -1052,10 +1011,6 @@ class HindsightMemoryProvider(MemoryProvider):
item[key] = value
return item
def _lineage_tags(self) -> list[str]:
pairs = (("session", self._session_id), ("parent", self._parent_session_id))
return [f"{kind}:{sid}" for kind, sid in pairs if sid]
def _retain_batch(self, item: dict, *, bank_id: str, document_id: str | None = None,
retain_async: bool | None = None):
"""Dispatch one item via aretain_batch (bank_id/document_id/retain_async are
@@ -1070,7 +1025,8 @@ class HindsightMemoryProvider(MemoryProvider):
writer runs after later sync_turn() calls mutate _session_turns/_turn_index/_session_id."""
content = "[" + ",".join(turns) + "]"
metadata = self._build_metadata(message_count=len(turns) * 2, turn_index=self._turn_index)
tags = self._lineage_tags() or None
lineage = (("session", self._session_id), ("parent", self._parent_session_id))
tags = [f"{kind}:{sid}" for kind, sid in lineage if sid] or None
bank_id, retain_async, retain_context = self._bank_id, self._retain_async, self._retain_context
def _job() -> None:
@@ -1119,7 +1075,12 @@ class HindsightMemoryProvider(MemoryProvider):
job = self._make_turn_retain_job(turns_to_retain, document_id=document_id,
update_mode=update_mode, label="retain")
# Indicator fires only past every skip/buffer gate: solely on turns that persist.
self._emit_saving_indicator()
# Model-independent status line; no-op without retain_indicator/status channel.
if self._retain_indicator and self._status_callback is not None:
try:
self._status_callback(f"{_HINDSIGHT_GLYPH} Hindsight — saving to memory…")
except Exception:
logger.debug("Retain indicator emit failed (non-fatal)", exc_info=True)
self._enqueue_retain(job)
# Advance the watermark only after the delta is queued so a later retain
# doesn't re-ship turns already handed to the writer.
@@ -1132,16 +1093,6 @@ class HindsightMemoryProvider(MemoryProvider):
self._register_atexit()
self._retain_queue.put(job)
def _emit_saving_indicator(self) -> None:
"""Model-independent "saving to memory" status line; no-op without
``retain_indicator``/status channel; never raises."""
if not self._retain_indicator or self._status_callback is None:
return
try:
self._status_callback(f"{_HINDSIGHT_GLYPH} Hindsight — saving to memory…")
except Exception:
logger.debug("Retain indicator emit failed (non-fatal)", exc_info=True)
# -- tools -------------------------------------------------------------------
def get_tool_schemas(self) -> List[Dict[str, Any]]:
@@ -1175,24 +1126,24 @@ class HindsightMemoryProvider(MemoryProvider):
logger.debug("Tool hindsight_reflect: response_len=%d", len(text))
return text or "No relevant memories found."
# tool name -> (required arg, handler)
# tool name -> (required arg, handler, user-facing failure prefix)
_TOOL_HANDLERS = {
"hindsight_retain": ("content", _tool_retain),
"hindsight_recall": ("query", _tool_recall),
"hindsight_reflect": ("query", _tool_reflect),
"hindsight_retain": ("content", _tool_retain, "Failed to store memory"),
"hindsight_recall": ("query", _tool_recall, "Failed to search memory"),
"hindsight_reflect": ("query", _tool_reflect, "Failed to reflect"),
}
def handle_tool_call(self, tool_name: str, args: dict, **kwargs) -> str:
required, handler = self._TOOL_HANDLERS.get(tool_name, ("", None))
if handler is None:
if tool_name not in self._TOOL_HANDLERS:
return tool_error(f"Unknown tool: {tool_name}")
required, handler, failure = self._TOOL_HANDLERS[tool_name]
if not args.get(required, ""):
return tool_error(f"Missing required parameter: {required}")
try:
return json.dumps({"result": handler(self, args)})
except Exception as e:
logger.warning("%s failed: %s", tool_name, e, exc_info=True)
return tool_error(f"{_TOOL_ERRORS[tool_name]}: {e}")
return tool_error(f"{failure}: {e}")
# -- session lifecycle -------------------------------------------------------