fix(hermes): sync lifecycle hooks for hermes-agent 0.5.0 (#741)
* fix(hermes): convert lifecycle hooks to sync for hermes-agent 0.5.0 compatibility hermes-agent 0.5.0 calls plugin hooks synchronously via invoke_hook(), but our pre_llm_call/post_llm_call were async — coroutines were never awaited, so recall context injection and auto-retain silently did nothing. Switch hooks to sync client methods and add integration tests using the real hermes-agent PluginManager. * fix(hermes): use proper hermes-agent dep with uv source override Replace inline git URL with standard `hermes-agent>=0.5.0` version constraint plus `[tool.uv.sources]` to resolve from the git tag until 0.5.0 lands on PyPI.
This commit is contained in:
parent
2d787c4ffd
commit
e7c9a6832d
4 changed files with 1589 additions and 44 deletions
|
|
@ -349,16 +349,31 @@ def register(ctx: Any) -> None:
|
||||||
)
|
)
|
||||||
|
|
||||||
# ── Lifecycle hooks ──────────────────────────────────────────────────
|
# ── Lifecycle hooks ──────────────────────────────────────────────────
|
||||||
# These require hermes-agent ≥ the version that invokes pre/post_llm_call.
|
# These require hermes-agent ≥ 0.5.0 which invokes pre/post_llm_call.
|
||||||
# When running on an older hermes-agent the hooks are simply never called,
|
# On older hermes-agent the hooks are simply never called, so
|
||||||
# so registering them is always safe.
|
# registering them is always safe.
|
||||||
|
#
|
||||||
|
# IMPORTANT: hermes-agent calls hooks synchronously via invoke_hook(),
|
||||||
|
# so these must be sync functions. We use the sync client methods
|
||||||
|
# (recall / retain / create_bank) rather than the async variants.
|
||||||
|
|
||||||
recall_budget = cfg.get("recallBudget", budget)
|
recall_budget = cfg.get("recallBudget", budget)
|
||||||
recall_max_tokens = cfg.get("recallMaxTokens", 4096)
|
recall_max_tokens = cfg.get("recallMaxTokens", 4096)
|
||||||
retain_enabled = cfg.get("autoRetain", True)
|
retain_enabled = cfg.get("autoRetain", True)
|
||||||
recall_preamble = cfg.get("recallPromptPreamble", "")
|
recall_preamble = cfg.get("recallPromptPreamble", "")
|
||||||
|
|
||||||
async def _on_pre_llm_call(
|
created_banks_sync: set[str] = set()
|
||||||
|
|
||||||
|
def _ensure_bank_sync(bid: str) -> None:
|
||||||
|
if bid in created_banks_sync:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
resolved_client.create_bank(bank_id=bid, name=bid)
|
||||||
|
created_banks_sync.add(bid)
|
||||||
|
except Exception:
|
||||||
|
created_banks_sync.add(bid)
|
||||||
|
|
||||||
|
def _on_pre_llm_call(
|
||||||
*,
|
*,
|
||||||
session_id: str = "",
|
session_id: str = "",
|
||||||
user_message: str = "",
|
user_message: str = "",
|
||||||
|
|
@ -371,8 +386,8 @@ def register(ctx: Any) -> None:
|
||||||
if not user_message or not bank_id:
|
if not user_message or not bank_id:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
await _ensure_bank(bank_id)
|
_ensure_bank_sync(bank_id)
|
||||||
response = await resolved_client.arecall(
|
response = resolved_client.recall(
|
||||||
bank_id=bank_id,
|
bank_id=bank_id,
|
||||||
query=user_message,
|
query=user_message,
|
||||||
budget=recall_budget,
|
budget=recall_budget,
|
||||||
|
|
@ -392,7 +407,7 @@ def register(ctx: Any) -> None:
|
||||||
logger.warning("Hindsight pre_llm_call recall failed: %s", exc)
|
logger.warning("Hindsight pre_llm_call recall failed: %s", exc)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _on_post_llm_call(
|
def _on_post_llm_call(
|
||||||
*,
|
*,
|
||||||
session_id: str = "",
|
session_id: str = "",
|
||||||
user_message: str = "",
|
user_message: str = "",
|
||||||
|
|
@ -406,9 +421,9 @@ def register(ctx: Any) -> None:
|
||||||
if not user_message or not assistant_response:
|
if not user_message or not assistant_response:
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
await _ensure_bank(bank_id)
|
_ensure_bank_sync(bank_id)
|
||||||
content = f"User: {user_message}\nAssistant: {assistant_response}"
|
content = f"User: {user_message}\nAssistant: {assistant_response}"
|
||||||
await resolved_client.aretain(bank_id=bank_id, content=content)
|
resolved_client.retain(bank_id=bank_id, content=content)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.warning("Hindsight post_llm_call retain failed: %s", exc)
|
logger.warning("Hindsight post_llm_call retain failed: %s", exc)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -51,6 +51,10 @@ build-backend = "hatchling.build"
|
||||||
[tool.hatch.build.targets.wheel]
|
[tool.hatch.build.targets.wheel]
|
||||||
packages = ["hindsight_hermes"]
|
packages = ["hindsight_hermes"]
|
||||||
|
|
||||||
|
[tool.uv.sources]
|
||||||
|
# hermes-agent 0.5.0 is not on PyPI yet; pin to the release tag
|
||||||
|
hermes-agent = { git = "https://github.com/NousResearch/hermes-agent.git", tag = "v2026.3.28" }
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
asyncio_mode = "auto"
|
asyncio_mode = "auto"
|
||||||
|
|
@ -59,4 +63,5 @@ asyncio_mode = "auto"
|
||||||
dev = [
|
dev = [
|
||||||
"pytest>=9.0.2",
|
"pytest>=9.0.2",
|
||||||
"pytest-asyncio>=0.23.0",
|
"pytest-asyncio>=0.23.0",
|
||||||
|
"hermes-agent>=0.5.0 ; python_version >= '3.11'",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
"""Tests for hindsight_hermes.tools module."""
|
"""Tests for hindsight_hermes.tools module."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import sys
|
import sys
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
@ -243,7 +244,11 @@ class TestRegisterPlugin:
|
||||||
|
|
||||||
|
|
||||||
class TestLifecycleHooks:
|
class TestLifecycleHooks:
|
||||||
"""Tests for pre_llm_call and post_llm_call hook callbacks."""
|
"""Tests for pre_llm_call and post_llm_call hook callbacks.
|
||||||
|
|
||||||
|
Hooks are synchronous functions (hermes-agent 0.5.0 calls them via
|
||||||
|
invoke_hook which is sync), so these tests call them directly without await.
|
||||||
|
"""
|
||||||
|
|
||||||
def _get_hook(self, ctx_mock, hook_name: str):
|
def _get_hook(self, ctx_mock, hook_name: str):
|
||||||
"""Extract the registered hook callback by name from the mock ctx."""
|
"""Extract the registered hook callback by name from the mock ctx."""
|
||||||
|
|
@ -264,11 +269,10 @@ class TestLifecycleHooks:
|
||||||
|
|
||||||
# -- pre_llm_call --
|
# -- pre_llm_call --
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_pre_llm_call_returns_context(self, monkeypatch, mock_client):
|
||||||
async def test_pre_llm_call_returns_context(self, monkeypatch, mock_client):
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "pre_llm_call")
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
result = await hook(
|
result = hook(
|
||||||
session_id="s1",
|
session_id="s1",
|
||||||
user_message="what color do I like?",
|
user_message="what color do I like?",
|
||||||
conversation_history=[],
|
conversation_history=[],
|
||||||
|
|
@ -279,72 +283,65 @@ class TestLifecycleHooks:
|
||||||
assert "context" in result
|
assert "context" in result
|
||||||
assert "Memory 1" in result["context"]
|
assert "Memory 1" in result["context"]
|
||||||
assert "Memory 2" in result["context"]
|
assert "Memory 2" in result["context"]
|
||||||
mock_client.arecall.assert_called_once()
|
mock_client.recall.assert_called_once()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_pre_llm_call_returns_none_on_no_results(self, monkeypatch, mock_client):
|
||||||
async def test_pre_llm_call_returns_none_on_no_results(self, monkeypatch, mock_client):
|
mock_client.recall.return_value = SimpleNamespace(results=[])
|
||||||
mock_client.arecall.return_value = SimpleNamespace(results=[])
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "pre_llm_call")
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
result = await hook(user_message="hello")
|
result = hook(user_message="hello")
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_pre_llm_call_returns_none_on_empty_message(self, monkeypatch, mock_client):
|
||||||
async def test_pre_llm_call_returns_none_on_empty_message(self, monkeypatch, mock_client):
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "pre_llm_call")
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
result = await hook(user_message="")
|
result = hook(user_message="")
|
||||||
assert result is None
|
assert result is None
|
||||||
mock_client.arecall.assert_not_called()
|
mock_client.recall.assert_not_called()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_pre_llm_call_returns_none_on_error(self, monkeypatch, mock_client):
|
||||||
async def test_pre_llm_call_returns_none_on_error(self, monkeypatch, mock_client):
|
mock_client.recall.side_effect = RuntimeError("connection failed")
|
||||||
mock_client.arecall.side_effect = RuntimeError("connection failed")
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "pre_llm_call")
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
result = await hook(user_message="hello")
|
result = hook(user_message="hello")
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
# -- post_llm_call --
|
# -- post_llm_call --
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_post_llm_call_retains_turn(self, monkeypatch, mock_client):
|
||||||
async def test_post_llm_call_retains_turn(self, monkeypatch, mock_client):
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "post_llm_call")
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
await hook(
|
hook(
|
||||||
session_id="s1",
|
session_id="s1",
|
||||||
user_message="remember I like green",
|
user_message="remember I like green",
|
||||||
assistant_response="Got it, you like green!",
|
assistant_response="Got it, you like green!",
|
||||||
model="test",
|
model="test",
|
||||||
)
|
)
|
||||||
mock_client.aretain.assert_called_once()
|
mock_client.retain.assert_called_once()
|
||||||
content = mock_client.aretain.call_args.kwargs["content"]
|
content = mock_client.retain.call_args.kwargs["content"]
|
||||||
assert "remember I like green" in content
|
assert "remember I like green" in content
|
||||||
assert "Got it, you like green!" in content
|
assert "Got it, you like green!" in content
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_post_llm_call_skips_empty_messages(self, monkeypatch, mock_client):
|
||||||
async def test_post_llm_call_skips_empty_messages(self, monkeypatch, mock_client):
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "post_llm_call")
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
await hook(user_message="", assistant_response="hello")
|
hook(user_message="", assistant_response="hello")
|
||||||
mock_client.aretain.assert_not_called()
|
mock_client.retain.assert_not_called()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_post_llm_call_skips_when_disabled(self, monkeypatch, mock_client):
|
||||||
async def test_post_llm_call_skips_when_disabled(self, monkeypatch, mock_client):
|
|
||||||
ctx = self._register_with_hooks(
|
ctx = self._register_with_hooks(
|
||||||
monkeypatch, mock_client, HINDSIGHT_AUTO_RETAIN="false"
|
monkeypatch, mock_client, HINDSIGHT_AUTO_RETAIN="false"
|
||||||
)
|
)
|
||||||
hook = self._get_hook(ctx, "post_llm_call")
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
await hook(user_message="hi", assistant_response="hello")
|
hook(user_message="hi", assistant_response="hello")
|
||||||
mock_client.aretain.assert_not_called()
|
mock_client.retain.assert_not_called()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def test_post_llm_call_does_not_raise_on_error(self, monkeypatch, mock_client):
|
||||||
async def test_post_llm_call_does_not_raise_on_error(self, monkeypatch, mock_client):
|
mock_client.retain.side_effect = RuntimeError("boom")
|
||||||
mock_client.aretain.side_effect = RuntimeError("boom")
|
|
||||||
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
hook = self._get_hook(ctx, "post_llm_call")
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
# Should not raise
|
# Should not raise
|
||||||
await hook(user_message="hi", assistant_response="hello")
|
hook(user_message="hi", assistant_response="hello")
|
||||||
|
|
||||||
|
|
||||||
# --- memory_instructions tests ---
|
# --- memory_instructions tests ---
|
||||||
|
|
@ -379,3 +376,103 @@ class TestMemoryInstructions:
|
||||||
def test_custom_prefix(self, mock_client):
|
def test_custom_prefix(self, mock_client):
|
||||||
result = memory_instructions(bank_id="b", client=mock_client, prefix="Context:\n")
|
result = memory_instructions(bank_id="b", client=mock_client, prefix="Context:\n")
|
||||||
assert result.startswith("Context:")
|
assert result.startswith("Context:")
|
||||||
|
|
||||||
|
|
||||||
|
# --- Integration tests with real hermes-agent plugin system ---
|
||||||
|
|
||||||
|
|
||||||
|
class TestHermesPluginIntegration:
|
||||||
|
"""Tests that verify hooks work correctly when invoked through
|
||||||
|
hermes-agent's real PluginManager.invoke_hook (sync dispatch)."""
|
||||||
|
|
||||||
|
def _setup_plugin_manager(self, monkeypatch, mock_client):
|
||||||
|
"""Register our plugin via the real PluginContext and return the PluginManager."""
|
||||||
|
from hermes_cli.plugins import PluginManager, PluginManifest, PluginContext
|
||||||
|
|
||||||
|
monkeypatch.setenv("HINDSIGHT_API_URL", "http://localhost:8888")
|
||||||
|
monkeypatch.setenv("HINDSIGHT_BANK_ID", "test-bank")
|
||||||
|
|
||||||
|
manager = PluginManager()
|
||||||
|
manifest = PluginManifest(name="hindsight", source="test")
|
||||||
|
ctx = PluginContext(manifest, manager)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("hindsight_hermes.tools._resolve_client", return_value=mock_client),
|
||||||
|
patch.dict(sys.modules, {"tools": MagicMock(), "tools.registry": MagicMock()}),
|
||||||
|
):
|
||||||
|
register(ctx)
|
||||||
|
|
||||||
|
return manager
|
||||||
|
|
||||||
|
def test_invoke_pre_llm_call_returns_context_dict(self, monkeypatch, mock_client):
|
||||||
|
"""invoke_hook('pre_llm_call') should return a list with a context dict."""
|
||||||
|
manager = self._setup_plugin_manager(monkeypatch, mock_client)
|
||||||
|
|
||||||
|
results = manager.invoke_hook(
|
||||||
|
"pre_llm_call",
|
||||||
|
session_id="s1",
|
||||||
|
user_message="what is my favorite color?",
|
||||||
|
conversation_history=[],
|
||||||
|
is_first_turn=True,
|
||||||
|
model="test",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(results) == 1
|
||||||
|
assert isinstance(results[0], dict)
|
||||||
|
assert "context" in results[0]
|
||||||
|
assert "Memory 1" in results[0]["context"]
|
||||||
|
mock_client.recall.assert_called_once()
|
||||||
|
|
||||||
|
def test_invoke_post_llm_call_retains(self, monkeypatch, mock_client):
|
||||||
|
"""invoke_hook('post_llm_call') should call retain synchronously."""
|
||||||
|
manager = self._setup_plugin_manager(monkeypatch, mock_client)
|
||||||
|
|
||||||
|
results = manager.invoke_hook(
|
||||||
|
"post_llm_call",
|
||||||
|
session_id="s1",
|
||||||
|
user_message="remember I like blue",
|
||||||
|
assistant_response="Noted, you like blue!",
|
||||||
|
model="test",
|
||||||
|
)
|
||||||
|
|
||||||
|
# post_llm_call returns None, so results should be empty
|
||||||
|
assert results == []
|
||||||
|
mock_client.retain.assert_called_once()
|
||||||
|
content = mock_client.retain.call_args.kwargs["content"]
|
||||||
|
assert "remember I like blue" in content
|
||||||
|
|
||||||
|
def test_invoke_pre_llm_call_no_results(self, monkeypatch, mock_client):
|
||||||
|
"""invoke_hook returns empty list when recall finds nothing."""
|
||||||
|
mock_client.recall.return_value = SimpleNamespace(results=[])
|
||||||
|
manager = self._setup_plugin_manager(monkeypatch, mock_client)
|
||||||
|
|
||||||
|
results = manager.invoke_hook(
|
||||||
|
"pre_llm_call",
|
||||||
|
session_id="s1",
|
||||||
|
user_message="hello",
|
||||||
|
conversation_history=[],
|
||||||
|
is_first_turn=True,
|
||||||
|
model="test",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert results == []
|
||||||
|
|
||||||
|
def test_hooks_are_not_coroutines(self, monkeypatch, mock_client):
|
||||||
|
"""Hooks must be plain functions, not async — verify invoke_hook
|
||||||
|
doesn't return coroutine objects."""
|
||||||
|
manager = self._setup_plugin_manager(monkeypatch, mock_client)
|
||||||
|
|
||||||
|
results = manager.invoke_hook(
|
||||||
|
"pre_llm_call",
|
||||||
|
session_id="s1",
|
||||||
|
user_message="test",
|
||||||
|
conversation_history=[],
|
||||||
|
is_first_turn=True,
|
||||||
|
model="test",
|
||||||
|
)
|
||||||
|
|
||||||
|
for r in results:
|
||||||
|
assert not asyncio.iscoroutine(r), (
|
||||||
|
"Hook returned a coroutine — hermes invoke_hook is sync and "
|
||||||
|
"will not await it"
|
||||||
|
)
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
Loading…
Reference in a new issue