From a30092e517be6eb4552be7faa8238cbdcc1cd5d2 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:41:53 -0700 Subject: [PATCH] refactor(hindsight): inline one-site helpers, tuple recall result, merged tool table --- plugins/memory/hindsight/__init__.py | 145 +++++++++------------------ 1 file changed, 48 insertions(+), 97 deletions(-) diff --git a/plugins/memory/hindsight/__init__.py b/plugins/memory/hindsight/__init__.py index a750191a35..a838d87d77 100644 --- a/plugins/memory/hindsight/__init__.py +++ b/plugins/memory/hindsight/__init__.py @@ -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 ``/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 -------------------------------------------------------