refactor(memory,langfuse): table-driven hook registration, compact signatures, fold tool-result backfill and finalize paths
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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)",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user