refactor(memory,langfuse): table-driven hook registration, compact signatures, fold tool-result backfill and finalize paths

This commit is contained in:
Teknium
2026-09-02 22:01:49 -07:00
parent 10c1457847
commit 1be4e79ed9
5 changed files with 46 additions and 105 deletions
+8 -27
View File
@@ -98,9 +98,7 @@ def _iter_entry_points():
eps = importlib.metadata.entry_points()
if hasattr(eps, "select"):
return list(eps.select(group=ENTRY_POINTS_GROUP))
if isinstance(eps, dict):
return list(eps.get(ENTRY_POINTS_GROUP, []))
return [ep for ep in eps if ep.group == ENTRY_POINTS_GROUP]
return list(eps.get(ENTRY_POINTS_GROUP, [])) if isinstance(eps, dict) else [ep for ep in eps if ep.group == ENTRY_POINTS_GROUP]
except Exception as exc:
logger.debug("Memory provider entry-point scan failed: %s", exc)
return []
@@ -174,11 +172,7 @@ def discover_memory_providers() -> List[Tuple[str, str, bool]]:
return results
def load_memory_provider(
name: str,
*,
register_skills: Optional[bool] = None,
) -> Optional["MemoryProvider"]:
def load_memory_provider(name: str, *, register_skills: Optional[bool] = None) -> Optional["MemoryProvider"]:
"""Load a MemoryProvider by name (bundled, user, project, then pip entry point);
None if not found or failing to load. Skills register only for the configured
active provider unless ``register_skills`` is explicit, so inspecting inactive
@@ -214,31 +208,25 @@ def _instantiate_subclass(namespace) -> Optional["MemoryProvider"]:
return None
def _load_provider_from_entry_point(
entry_point,
*,
register_skills: bool = True,
) -> Optional["MemoryProvider"]:
"""Import a provider entry point and extract the MemoryProvider instance."""
def _load_provider_from_entry_point(entry_point, *, register_skills: bool = True) -> Optional["MemoryProvider"]:
"""Import a provider entry point and extract the MemoryProvider instance: an
instance, a subclass, a module with ``register(ctx)``, a factory / ``register``
callable, or a namespace holding a subclass — in that order."""
from agent.memory_provider import MemoryProvider
loaded = entry_point.load()
if isinstance(loaded, MemoryProvider):
return loaded
if isinstance(loaded, type) and issubclass(loaded, MemoryProvider):
try:
return loaded()
except Exception:
pass
if hasattr(loaded, "register"):
collector = _ProviderCollector(entry_point.name, register_skills=register_skills)
loaded.register(collector)
if collector.provider:
return collector.provider
if callable(loaded):
try:
provider = loaded()
@@ -246,7 +234,6 @@ def _load_provider_from_entry_point(
return provider
except TypeError:
pass
collector = _ProviderCollector(entry_point.name, register_skills=register_skills)
loaded(collector)
return collector.provider
@@ -257,11 +244,7 @@ def _load_provider_from_entry_point(
return provider
def _load_provider_from_dir(
provider_dir: Path,
*,
register_skills: bool = True,
) -> Optional["MemoryProvider"]:
def _load_provider_from_dir(provider_dir: Path, *, register_skills: bool = True) -> Optional["MemoryProvider"]:
"""Import a provider module; ``register(ctx)`` first, else a top-level subclass."""
name = provider_dir.name
mod = _loader.load_plugin_module(
@@ -368,9 +351,7 @@ def _get_active_memory_provider() -> Optional[str]:
return None
def _prune_inactive_memory_provider_skills(
active_provider: Optional[str] = None,
) -> None:
def _prune_inactive_memory_provider_skills(active_provider: Optional[str] = None) -> None:
"""Remove tracked skills that no longer belong to the active provider."""
if active_provider is None:
active_provider = _get_active_memory_provider()
+8 -25
View File
@@ -885,9 +885,8 @@ class HindsightMemoryProvider(MemoryProvider):
def _daemon_start_worker(self) -> None:
import traceback
log_dir = get_hermes_home() / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
log_path = log_dir / "hindsight-embed.log"
log_path = get_hermes_home() / "logs" / "hindsight-embed.log"
log_path.parent.mkdir(parents=True, exist_ok=True)
def _log(text: str, exc: bool = False) -> None:
with open(log_path, "a", encoding="utf-8") as f:
@@ -1046,16 +1045,9 @@ class HindsightMemoryProvider(MemoryProvider):
metadata.update({name: value for name in _METADATA_ATTRS if (value := getattr(self, f"_{name}"))})
return metadata
def _build_retain_kwargs(
self,
content: str,
*,
context: str | None = None,
metadata: Dict[str, str] | None = None,
tags: List[str] | None = None,
occurred_at: str | None = None,
update_mode: str | None = None,
) -> Dict[str, Any]:
def _build_retain_kwargs(self, content: str, *, context: str | None = None,
metadata: Dict[str, str] | None = None, tags: List[str] | None = None,
occurred_at: str | None = None, update_mode: str | None = None) -> Dict[str, Any]:
"""Build one aretain_batch item. The server resolves occurred_start/end (incl.
relative phrases in content) from the item timestamp: explicit occurred_at
wins, else the configured event clock."""
@@ -1220,14 +1212,8 @@ class HindsightMemoryProvider(MemoryProvider):
# -- session lifecycle -------------------------------------------------------
def on_session_switch(
self,
new_session_id: str,
*,
parent_session_id: str = "",
reset: bool = False,
**kwargs,
) -> None:
def on_session_switch(self, new_session_id: str, *, parent_session_id: str = "",
reset: bool = False, **kwargs) -> None:
"""Rotate per-session state (/resume, /branch, /reset, /new, compression) so
writes don't land in the previous session's document. Always: flush buffered
turns under the OLD ids first (``retain_every_n_turns > 1`` would silently
@@ -1301,10 +1287,7 @@ class HindsightMemoryProvider(MemoryProvider):
# bounded join keeps shutdown predictable even if the daemon is wedged.
writer = self._writer_thread
if writer is not None and writer.is_alive():
try:
self._retain_queue.put(_WRITER_SENTINEL)
except Exception:
pass
self._retain_queue.put(_WRITER_SENTINEL)
writer.join(timeout=10.0)
if writer.is_alive():
logger.warning("Hindsight writer did not stop within 10s; abandoning %d pending retain(s)",
+1 -2
View File
@@ -137,9 +137,8 @@ def _resolve_bank_id_template(template: str, fallback: str, **placeholders: str)
collapsed (``hermes-{user}`` -> ``hermes``). Empty/invalid template -> *fallback*."""
if not template:
return fallback
sanitized = {k: _sanitize_bank_segment(v) for k, v in placeholders.items()}
try:
rendered = template.format(**sanitized)
rendered = template.format(**{k: _sanitize_bank_segment(v) for k, v in placeholders.items()})
except (KeyError, IndexError) as exc:
logger.warning("Invalid bank_id_template %r: %s — using fallback %r",
template, exc, fallback)
+1 -9
View File
@@ -85,15 +85,7 @@ def probe_existing_customization(api_url: str, bank_id: str, api_key: str | None
return bool(data.get("bank") or data.get("mental_models") or data.get("directives"))
def run_template_step(
*,
api_url: str,
bank_id: str,
api_key: str | None,
select,
cancelled,
log=print,
) -> str | None:
def run_template_step(*, api_url: str, bank_id: str, api_key: str | None, select, cancelled, log=print) -> str | None:
"""Wizard starter-template step. ``select(title, items, default, cancel_returns)``
is the picker (injected: testable without curses). Returns the applied template
id, or None if skipped/blank/failed. Never raises — a template is a nice-to-have."""
+28 -42
View File
@@ -404,10 +404,8 @@ def _safe_value(value: Any, *, max_chars: Optional[int] = None, depth: int = 0,
def _extract_last_user_message(messages: Any) -> Any:
if not isinstance(messages, list):
return None
for message in reversed(messages):
if isinstance(message, dict) and message.get("role") == "user":
return {"role": "user", "content": _capture_content(message.get("content"))}
return None
last = next((m for m in reversed(messages) if isinstance(m, dict) and m.get("role") == "user"), None)
return None if last is None else {"role": "user", "content": _capture_content(last.get("content"))}
def _coerce_request_messages(*, request_messages: Any = None, messages: Any = None,
@@ -415,9 +413,7 @@ def _coerce_request_messages(*, request_messages: Any = None, messages: Any = No
for candidate in (request_messages, messages, conversation_history):
if isinstance(candidate, list):
return candidate
if user_message is None:
return []
return [{"role": "user", "content": user_message}]
return [] if user_message is None else [{"role": "user", "content": user_message}]
def _serialize_system_prompt(system_prompt: Any) -> Optional[dict[str, Any]]:
@@ -435,9 +431,7 @@ def _serialize_system_prompt(system_prompt: Any) -> Optional[dict[str, Any]]:
text = "\n\n".join(parts)
else:
return None
if not text:
return None
return {"role": "system", "content": _capture_content(text)}
return {"role": "system", "content": _capture_content(text)} if text else None
def _messages_for_langfuse_input(*, request_messages: Any = None, messages: Any = None,
@@ -707,10 +701,10 @@ def _finalize_all_traces() -> None:
for key, state in states:
try:
_end_children(state, include_subagents=True)
state.root_span.end()
_exit_root_ctx(state)
except Exception as exc: # pragma: no cover - fail-open
_debug(f"atexit finalize failed for {key}: {exc}")
else:
_end_root(state, f"atexit finalize for {key}")
if states:
_flush(_get_langfuse())
@@ -744,11 +738,7 @@ def _finish_trace(task_key: str, *, output: Any = None) -> None:
getattr(state.root_span, method)(output=final_output)
except Exception as exc:
_debug(f"{label} failed: {exc}")
try:
state.root_span.end()
except Exception as exc:
_debug(f"root end() failed: {exc}")
_exit_root_ctx(state)
_end_root(state, "root end()")
except Exception as exc: # pragma: no cover - fail-open
_debug(f"finish trace failed: {exc}")
# Last-chance end so an unexpected error still exports the root.
@@ -991,12 +981,11 @@ def on_post_tool_call(*, tool_name: str = "", args: Any = None, result: Any = No
if state is None:
return
observation = state.tools.pop(tool_call_id, None) if tool_call_id else None
if observation is None:
queue = state.pending_tools_by_name.get(tool_name)
if queue:
observation = queue.pop(0)
if not queue:
state.pending_tools_by_name.pop(tool_name, None)
queue = state.pending_tools_by_name.get(tool_name) if observation is None else None
if queue:
observation = queue.pop(0)
if not queue:
state.pending_tools_by_name.pop(tool_name, None)
if observation is None:
return
@@ -1006,14 +995,12 @@ def on_post_tool_call(*, tool_name: str = "", args: Any = None, result: Any = No
if tool_call_id:
with _STATE_LOCK:
state = _TRACE_STATE.get(task_key)
if state is not None:
for tool_call in reversed(state.turn_tool_calls):
if tool_call.get("id") == tool_call_id:
tool_call["output"] = safe_result_value
function_payload = tool_call.get("function")
if isinstance(function_payload, dict):
function_payload["output"] = safe_result_value
break
calls = state.turn_tool_calls if state is not None else []
tool_call = next((tc for tc in reversed(calls) if tc.get("id") == tool_call_id), None)
if tool_call is not None:
tool_call["output"] = safe_result_value
if isinstance(tool_call.get("function"), dict):
tool_call["function"]["output"] = safe_result_value
_end_observation(
observation, output=safe_result_value,
@@ -1137,14 +1124,13 @@ def on_subagent_stop(*, parent_turn_id: str = "", child_session_id: Any = None,
def register(ctx) -> None:
# Both hook-name variants so the plugin works across Hermes versions:
# *_api_request fire per API call (preferred); *_llm_call once per turn.
ctx.register_hook("pre_api_request", on_pre_llm_request)
ctx.register_hook("post_api_request", on_post_llm_call)
ctx.register_hook("api_request_error", on_api_request_error)
ctx.register_hook("pre_llm_call", on_pre_llm_call)
ctx.register_hook("post_llm_call", on_post_llm_call)
ctx.register_hook("pre_tool_call", on_pre_tool_call)
ctx.register_hook("post_tool_call", on_post_tool_call)
ctx.register_hook("on_session_finalize", on_session_finalize)
ctx.register_hook("on_session_end", on_session_finalize)
ctx.register_hook("subagent_start", on_subagent_start)
ctx.register_hook("subagent_stop", on_subagent_stop)
hooks = (
("pre_api_request", on_pre_llm_request), ("post_api_request", on_post_llm_call),
("api_request_error", on_api_request_error), ("pre_llm_call", on_pre_llm_call),
("post_llm_call", on_post_llm_call), ("pre_tool_call", on_pre_tool_call),
("post_tool_call", on_post_tool_call), ("on_session_finalize", on_session_finalize),
("on_session_end", on_session_finalize), ("subagent_start", on_subagent_start),
("subagent_stop", on_subagent_stop),
)
for name, fn in hooks:
ctx.register_hook(name, fn)