refactor(hindsight): inline one-site helpers, tuple recall result, merged tool table
This commit is contained in:
@@ -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 -------------------------------------------------------
|
||||
|
||||
|
||||
Reference in New Issue
Block a user