Tool handlers and lifecycle hooks now use the native async client API (aretain, arecall, areflect, acreate_bank) instead of sync wrappers that call loop.run_until_complete(), which deadlocks in async contexts like Discord/Telegram gateways.
This commit is contained in:
parent
0bcbf8491b
commit
35dfd3aa0c
6 changed files with 1079 additions and 827 deletions
|
|
@ -4,7 +4,12 @@ Hindsight memory integration for [Hermes Agent](https://github.com/NousResearch/
|
||||||
|
|
||||||
## What it does
|
## What it does
|
||||||
|
|
||||||
This package registers three tools into Hermes via its plugin system:
|
**Automatic memory on every turn** — no tool calls required:
|
||||||
|
|
||||||
|
- **`pre_llm_call` hook** — Before each LLM call, recalls relevant memories and injects them into the system prompt. The model sees cross-session context automatically.
|
||||||
|
- **`post_llm_call` hook** — After each turn, retains the user/assistant exchange so it can be recalled in future sessions.
|
||||||
|
|
||||||
|
**Three explicit tools** for when the model wants direct control:
|
||||||
|
|
||||||
- **`hindsight_retain`** — Stores information to long-term memory. Hermes calls this when the user shares facts, preferences, or anything worth remembering.
|
- **`hindsight_retain`** — Stores information to long-term memory. Hermes calls this when the user shares facts, preferences, or anything worth remembering.
|
||||||
- **`hindsight_recall`** — Searches long-term memory for relevant information. Returns a numbered list of matching memories.
|
- **`hindsight_recall`** — Searches long-term memory for relevant information. Returns a numbered list of matching memories.
|
||||||
|
|
@ -12,6 +17,8 @@ This package registers three tools into Hermes via its plugin system:
|
||||||
|
|
||||||
These tools appear under the `[hindsight]` toolset in Hermes's `/tools` list.
|
These tools appear under the `[hindsight]` toolset in Hermes's `/tools` list.
|
||||||
|
|
||||||
|
> **Note:** The lifecycle hooks require hermes-agent with [PR #2823](https://github.com/NousResearch/hermes-agent/pull/2823) or later. On older versions, only the tools are registered — hooks are silently skipped.
|
||||||
|
|
||||||
## Setup
|
## Setup
|
||||||
|
|
||||||
### 1. Install hindsight-hermes into the Hermes venv
|
### 1. Install hindsight-hermes into the Hermes venv
|
||||||
|
|
@ -230,7 +237,10 @@ This exposes the same retain/recall/reflect operations through Hermes's MCP inte
|
||||||
| `api_key` | `HINDSIGHT_API_KEY` | — | API key for authentication |
|
| `api_key` | `HINDSIGHT_API_KEY` | — | API key for authentication |
|
||||||
| `bank_id` | `HINDSIGHT_BANK_ID` | — | Memory bank ID |
|
| `bank_id` | `HINDSIGHT_BANK_ID` | — | Memory bank ID |
|
||||||
| `budget` | `HINDSIGHT_BUDGET` | `mid` | Recall budget (low/mid/high) |
|
| `budget` | `HINDSIGHT_BUDGET` | `mid` | Recall budget (low/mid/high) |
|
||||||
| `max_tokens` | — | `4096` | Max tokens for recall results |
|
| — | `HINDSIGHT_AUTO_RETAIN` | `true` | Auto-retain conversation turns via `post_llm_call` hook |
|
||||||
|
| — | `HINDSIGHT_RECALL_BUDGET` | same as `budget` | Budget for the `pre_llm_call` recall hook |
|
||||||
|
| — | `HINDSIGHT_RECALL_MAX_TOKENS` | `4096` | Max tokens for the `pre_llm_call` recall hook |
|
||||||
|
| `max_tokens` | — | `4096` | Max tokens for recall results (tools) |
|
||||||
| `tags` | — | — | Tags applied when storing memories |
|
| `tags` | — | — | Tags applied when storing memories |
|
||||||
| `recall_tags` | — | — | Tags to filter recall results |
|
| `recall_tags` | — | — | Tags to filter recall results |
|
||||||
| `recall_tags_match` | — | `any` | Tag matching mode (any/all/any_strict/all_strict) |
|
| `recall_tags_match` | — | `any` | Tag matching mode (any/all/any_strict/all_strict) |
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,13 @@
|
||||||
"""Hindsight-Hermes: Persistent memory tools for Hermes agents.
|
"""Hindsight-Hermes: Persistent memory for Hermes agents.
|
||||||
|
|
||||||
Provides Hindsight retain/recall/reflect as native Hermes tools via the
|
Provides Hindsight retain/recall/reflect as native Hermes tools via the
|
||||||
plugin system or manual ``register_tools()`` call.
|
plugin system or manual ``register_tools()`` call. When running on a
|
||||||
|
Hermes build that supports lifecycle hooks, the plugin also:
|
||||||
|
|
||||||
|
- **pre_llm_call** — recalls relevant memories and injects them into the
|
||||||
|
system prompt so the model has cross-session context on every turn.
|
||||||
|
- **post_llm_call** — retains the user/assistant exchange so it can be
|
||||||
|
recalled in future sessions.
|
||||||
|
|
||||||
Plugin usage (auto-discovery)::
|
Plugin usage (auto-discovery)::
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -180,29 +180,29 @@ def register_tools(
|
||||||
resolved_client = _resolve_client(client, hindsight_api_url, api_key)
|
resolved_client = _resolve_client(client, hindsight_api_url, api_key)
|
||||||
created_banks: set[str] = set()
|
created_banks: set[str] = set()
|
||||||
|
|
||||||
def _ensure_bank(bid: str) -> None:
|
async def _ensure_bank(bid: str) -> None:
|
||||||
if bid in created_banks:
|
if bid in created_banks:
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
resolved_client.create_bank(bank_id=bid, name=bid)
|
await resolved_client.acreate_bank(bank_id=bid, name=bid)
|
||||||
created_banks.add(bid)
|
created_banks.add(bid)
|
||||||
except Exception:
|
except Exception:
|
||||||
created_banks.add(bid)
|
created_banks.add(bid)
|
||||||
|
|
||||||
def handle_retain(args: dict[str, Any], **kwargs: Any) -> str:
|
async def handle_retain(args: dict[str, Any], **kwargs: Any) -> str:
|
||||||
try:
|
try:
|
||||||
bid = _resolve_bank_id(args, bank_id, bank_resolver)
|
bid = _resolve_bank_id(args, bank_id, bank_resolver)
|
||||||
_ensure_bank(bid)
|
await _ensure_bank(bid)
|
||||||
retain_kwargs: dict[str, Any] = {"bank_id": bid, "content": args["content"]}
|
retain_kwargs: dict[str, Any] = {"bank_id": bid, "content": args["content"]}
|
||||||
if tags:
|
if tags:
|
||||||
retain_kwargs["tags"] = tags
|
retain_kwargs["tags"] = tags
|
||||||
resolved_client.retain(**retain_kwargs)
|
await resolved_client.aretain(**retain_kwargs)
|
||||||
return json.dumps({"result": "Memory stored successfully."})
|
return json.dumps({"result": "Memory stored successfully."})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Retain failed: {e}")
|
logger.error(f"Retain failed: {e}")
|
||||||
return json.dumps({"error": str(e)})
|
return json.dumps({"error": str(e)})
|
||||||
|
|
||||||
def handle_recall(args: dict[str, Any], **kwargs: Any) -> str:
|
async def handle_recall(args: dict[str, Any], **kwargs: Any) -> str:
|
||||||
try:
|
try:
|
||||||
bid = _resolve_bank_id(args, bank_id, bank_resolver)
|
bid = _resolve_bank_id(args, bank_id, bank_resolver)
|
||||||
recall_kwargs: dict[str, Any] = {
|
recall_kwargs: dict[str, Any] = {
|
||||||
|
|
@ -214,7 +214,7 @@ def register_tools(
|
||||||
if recall_tags:
|
if recall_tags:
|
||||||
recall_kwargs["tags"] = recall_tags
|
recall_kwargs["tags"] = recall_tags
|
||||||
recall_kwargs["tags_match"] = recall_tags_match
|
recall_kwargs["tags_match"] = recall_tags_match
|
||||||
response = resolved_client.recall(**recall_kwargs)
|
response = await resolved_client.arecall(**recall_kwargs)
|
||||||
if not response.results:
|
if not response.results:
|
||||||
return json.dumps({"result": "No relevant memories found."})
|
return json.dumps({"result": "No relevant memories found."})
|
||||||
lines = []
|
lines = []
|
||||||
|
|
@ -225,7 +225,7 @@ def register_tools(
|
||||||
logger.error(f"Recall failed: {e}")
|
logger.error(f"Recall failed: {e}")
|
||||||
return json.dumps({"error": str(e)})
|
return json.dumps({"error": str(e)})
|
||||||
|
|
||||||
def handle_reflect(args: dict[str, Any], **kwargs: Any) -> str:
|
async def handle_reflect(args: dict[str, Any], **kwargs: Any) -> str:
|
||||||
try:
|
try:
|
||||||
bid = _resolve_bank_id(args, bank_id, bank_resolver)
|
bid = _resolve_bank_id(args, bank_id, bank_resolver)
|
||||||
reflect_kwargs: dict[str, Any] = {
|
reflect_kwargs: dict[str, Any] = {
|
||||||
|
|
@ -233,7 +233,7 @@ def register_tools(
|
||||||
"query": args["query"],
|
"query": args["query"],
|
||||||
"budget": budget,
|
"budget": budget,
|
||||||
}
|
}
|
||||||
response = resolved_client.reflect(**reflect_kwargs)
|
response = await resolved_client.areflect(**reflect_kwargs)
|
||||||
return json.dumps(
|
return json.dumps(
|
||||||
{"result": response.text or "No relevant memories found."}
|
{"result": response.text or "No relevant memories found."}
|
||||||
)
|
)
|
||||||
|
|
@ -284,29 +284,29 @@ def register(ctx: Any) -> None:
|
||||||
resolved_client = _resolve_client(None, hindsight_api_url, api_key)
|
resolved_client = _resolve_client(None, hindsight_api_url, api_key)
|
||||||
created_banks: set[str] = set()
|
created_banks: set[str] = set()
|
||||||
|
|
||||||
def _ensure_bank(bid: str) -> None:
|
async def _ensure_bank(bid: str) -> None:
|
||||||
if bid in created_banks:
|
if bid in created_banks:
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
resolved_client.create_bank(bank_id=bid, name=bid)
|
await resolved_client.acreate_bank(bank_id=bid, name=bid)
|
||||||
created_banks.add(bid)
|
created_banks.add(bid)
|
||||||
except Exception:
|
except Exception:
|
||||||
created_banks.add(bid)
|
created_banks.add(bid)
|
||||||
|
|
||||||
def handle_retain(args: dict[str, Any], **kwargs: Any) -> str:
|
async def handle_retain(args: dict[str, Any], **kwargs: Any) -> str:
|
||||||
try:
|
try:
|
||||||
bid = _resolve_bank_id(args, bank_id, None)
|
bid = _resolve_bank_id(args, bank_id, None)
|
||||||
_ensure_bank(bid)
|
await _ensure_bank(bid)
|
||||||
resolved_client.retain(bank_id=bid, content=args["content"])
|
await resolved_client.aretain(bank_id=bid, content=args["content"])
|
||||||
return json.dumps({"result": "Memory stored successfully."})
|
return json.dumps({"result": "Memory stored successfully."})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Retain failed: {e}")
|
logger.error(f"Retain failed: {e}")
|
||||||
return json.dumps({"error": str(e)})
|
return json.dumps({"error": str(e)})
|
||||||
|
|
||||||
def handle_recall(args: dict[str, Any], **kwargs: Any) -> str:
|
async def handle_recall(args: dict[str, Any], **kwargs: Any) -> str:
|
||||||
try:
|
try:
|
||||||
bid = _resolve_bank_id(args, bank_id, None)
|
bid = _resolve_bank_id(args, bank_id, None)
|
||||||
response = resolved_client.recall(
|
response = await resolved_client.arecall(
|
||||||
bank_id=bid, query=args["query"], budget=budget
|
bank_id=bid, query=args["query"], budget=budget
|
||||||
)
|
)
|
||||||
if not response.results:
|
if not response.results:
|
||||||
|
|
@ -317,10 +317,10 @@ def register(ctx: Any) -> None:
|
||||||
logger.error(f"Recall failed: {e}")
|
logger.error(f"Recall failed: {e}")
|
||||||
return json.dumps({"error": str(e)})
|
return json.dumps({"error": str(e)})
|
||||||
|
|
||||||
def handle_reflect(args: dict[str, Any], **kwargs: Any) -> str:
|
async def handle_reflect(args: dict[str, Any], **kwargs: Any) -> str:
|
||||||
try:
|
try:
|
||||||
bid = _resolve_bank_id(args, bank_id, None)
|
bid = _resolve_bank_id(args, bank_id, None)
|
||||||
response = resolved_client.reflect(
|
response = await resolved_client.areflect(
|
||||||
bank_id=bid, query=args["query"], budget=budget
|
bank_id=bid, query=args["query"], budget=budget
|
||||||
)
|
)
|
||||||
return json.dumps(
|
return json.dumps(
|
||||||
|
|
@ -349,6 +349,72 @@ def register(ctx: Any) -> None:
|
||||||
handler=handle_reflect,
|
handler=handle_reflect,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ── Lifecycle hooks ──────────────────────────────────────────────────
|
||||||
|
# These require hermes-agent ≥ the version that invokes pre/post_llm_call.
|
||||||
|
# When running on an older hermes-agent the hooks are simply never called,
|
||||||
|
# so registering them is always safe.
|
||||||
|
|
||||||
|
recall_budget = os.environ.get("HINDSIGHT_RECALL_BUDGET", budget)
|
||||||
|
recall_max_tokens = int(os.environ.get("HINDSIGHT_RECALL_MAX_TOKENS", "4096"))
|
||||||
|
retain_enabled = os.environ.get("HINDSIGHT_AUTO_RETAIN", "true").lower() in {"1", "true", "yes", "on"}
|
||||||
|
|
||||||
|
async def _on_pre_llm_call(
|
||||||
|
*,
|
||||||
|
session_id: str = "",
|
||||||
|
user_message: str = "",
|
||||||
|
conversation_history: list | None = None,
|
||||||
|
is_first_turn: bool = False,
|
||||||
|
model: str = "",
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> dict[str, str] | None:
|
||||||
|
"""Recall relevant memories and inject them as system prompt context."""
|
||||||
|
if not user_message or not bank_id:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
await _ensure_bank(bank_id)
|
||||||
|
response = await resolved_client.arecall(
|
||||||
|
bank_id=bank_id,
|
||||||
|
query=user_message,
|
||||||
|
budget=recall_budget,
|
||||||
|
max_tokens=recall_max_tokens,
|
||||||
|
)
|
||||||
|
if not response.results:
|
||||||
|
return None
|
||||||
|
lines = [f"- {r.text}" for r in response.results]
|
||||||
|
context = (
|
||||||
|
"# Hindsight Memory (persistent cross-session context)\n"
|
||||||
|
"Use this to answer questions about the user and prior sessions. "
|
||||||
|
"Do not call tools to look up information that is already present here.\n\n"
|
||||||
|
+ "\n".join(lines)
|
||||||
|
)
|
||||||
|
return {"context": context}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Hindsight pre_llm_call recall failed: %s", exc)
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def _on_post_llm_call(
|
||||||
|
*,
|
||||||
|
session_id: str = "",
|
||||||
|
user_message: str = "",
|
||||||
|
assistant_response: str = "",
|
||||||
|
model: str = "",
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> None:
|
||||||
|
"""Retain the conversation turn so it can be recalled in future sessions."""
|
||||||
|
if not retain_enabled or not bank_id:
|
||||||
|
return
|
||||||
|
if not user_message or not assistant_response:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
await _ensure_bank(bank_id)
|
||||||
|
content = f"User: {user_message}\nAssistant: {assistant_response}"
|
||||||
|
await resolved_client.aretain(bank_id=bank_id, content=content)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Hindsight post_llm_call retain failed: %s", exc)
|
||||||
|
|
||||||
|
ctx.register_hook("pre_llm_call", _on_pre_llm_call)
|
||||||
|
ctx.register_hook("post_llm_call", _on_post_llm_call)
|
||||||
|
|
||||||
|
|
||||||
def memory_instructions(
|
def memory_instructions(
|
||||||
*,
|
*,
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,7 @@ dependencies = [
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
dev = [
|
dev = [
|
||||||
"pytest>=7.0.0",
|
"pytest>=7.0.0",
|
||||||
|
"pytest-asyncio>=0.23.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.entry-points."hermes_agent.plugins"]
|
[project.entry-points."hermes_agent.plugins"]
|
||||||
|
|
@ -52,8 +53,10 @@ packages = ["hindsight_hermes"]
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
|
asyncio_mode = "auto"
|
||||||
|
|
||||||
[dependency-groups]
|
[dependency-groups]
|
||||||
dev = [
|
dev = [
|
||||||
"pytest>=9.0.2",
|
"pytest>=9.0.2",
|
||||||
|
"pytest-asyncio>=0.23.0",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@
|
||||||
import json
|
import json
|
||||||
import sys
|
import sys
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
|
@ -35,6 +35,7 @@ def _clean_config():
|
||||||
@pytest.fixture()
|
@pytest.fixture()
|
||||||
def mock_client():
|
def mock_client():
|
||||||
client = MagicMock()
|
client = MagicMock()
|
||||||
|
# Sync methods (used by memory_instructions and register_tools sync path)
|
||||||
client.create_bank = MagicMock()
|
client.create_bank = MagicMock()
|
||||||
client.retain = MagicMock()
|
client.retain = MagicMock()
|
||||||
client.recall = MagicMock(
|
client.recall = MagicMock(
|
||||||
|
|
@ -46,6 +47,18 @@ def mock_client():
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
client.reflect = MagicMock(return_value=SimpleNamespace(text="Synthesized answer"))
|
client.reflect = MagicMock(return_value=SimpleNamespace(text="Synthesized answer"))
|
||||||
|
# Async methods (used by tool handlers and hooks)
|
||||||
|
client.acreate_bank = AsyncMock()
|
||||||
|
client.aretain = AsyncMock()
|
||||||
|
client.arecall = AsyncMock(
|
||||||
|
return_value=SimpleNamespace(
|
||||||
|
results=[
|
||||||
|
SimpleNamespace(text="Memory 1"),
|
||||||
|
SimpleNamespace(text="Memory 2"),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
client.areflect = AsyncMock(return_value=SimpleNamespace(text="Synthesized answer"))
|
||||||
return client
|
return client
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -133,60 +146,68 @@ class TestRegisterTools:
|
||||||
names = {call.kwargs["name"] for call in mock_registry.register.call_args_list}
|
names = {call.kwargs["name"] for call in mock_registry.register.call_args_list}
|
||||||
assert names == {"hindsight_retain", "hindsight_recall", "hindsight_reflect"}
|
assert names == {"hindsight_retain", "hindsight_recall", "hindsight_reflect"}
|
||||||
|
|
||||||
def test_retain_handler_success(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
|
async def test_retain_handler_success(self, mock_client, mock_registry):
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
||||||
result = json.loads(handler({"content": "hello"}))
|
result = json.loads(await handler({"content": "hello"}))
|
||||||
assert result["result"] == "Memory stored successfully."
|
assert result["result"] == "Memory stored successfully."
|
||||||
mock_client.retain.assert_called_once_with(bank_id="b", content="hello")
|
mock_client.aretain.assert_called_once_with(bank_id="b", content="hello")
|
||||||
|
|
||||||
def test_retain_handler_with_tags(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
|
async def test_retain_handler_with_tags(self, mock_client, mock_registry):
|
||||||
register_tools(bank_id="b", client=mock_client, tags=["tag1"])
|
register_tools(bank_id="b", client=mock_client, tags=["tag1"])
|
||||||
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
||||||
handler({"content": "hello"})
|
await handler({"content": "hello"})
|
||||||
mock_client.retain.assert_called_once_with(bank_id="b", content="hello", tags=["tag1"])
|
mock_client.aretain.assert_called_once_with(bank_id="b", content="hello", tags=["tag1"])
|
||||||
|
|
||||||
def test_recall_handler_success(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
|
async def test_recall_handler_success(self, mock_client, mock_registry):
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[1].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[1].kwargs["handler"]
|
||||||
result = json.loads(handler({"query": "test"}))
|
result = json.loads(await handler({"query": "test"}))
|
||||||
assert "Memory 1" in result["result"]
|
assert "Memory 1" in result["result"]
|
||||||
assert "Memory 2" in result["result"]
|
assert "Memory 2" in result["result"]
|
||||||
|
|
||||||
def test_recall_handler_no_results(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
mock_client.recall.return_value = SimpleNamespace(results=[])
|
async def test_recall_handler_no_results(self, mock_client, mock_registry):
|
||||||
|
mock_client.arecall.return_value = SimpleNamespace(results=[])
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[1].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[1].kwargs["handler"]
|
||||||
result = json.loads(handler({"query": "test"}))
|
result = json.loads(await handler({"query": "test"}))
|
||||||
assert result["result"] == "No relevant memories found."
|
assert result["result"] == "No relevant memories found."
|
||||||
|
|
||||||
def test_reflect_handler_success(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
|
async def test_reflect_handler_success(self, mock_client, mock_registry):
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[2].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[2].kwargs["handler"]
|
||||||
result = json.loads(handler({"query": "test"}))
|
result = json.loads(await handler({"query": "test"}))
|
||||||
assert result["result"] == "Synthesized answer"
|
assert result["result"] == "Synthesized answer"
|
||||||
|
|
||||||
def test_handler_returns_error_on_exception(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
mock_client.retain.side_effect = RuntimeError("boom")
|
async def test_handler_returns_error_on_exception(self, mock_client, mock_registry):
|
||||||
|
mock_client.aretain.side_effect = RuntimeError("boom")
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
||||||
result = json.loads(handler({"content": "hello"}))
|
result = json.loads(await handler({"content": "hello"}))
|
||||||
assert "error" in result
|
assert "error" in result
|
||||||
assert "boom" in result["error"]
|
assert "boom" in result["error"]
|
||||||
|
|
||||||
def test_ensure_bank_called(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
|
async def test_ensure_bank_called(self, mock_client, mock_registry):
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
||||||
handler({"content": "hello"})
|
await handler({"content": "hello"})
|
||||||
mock_client.create_bank.assert_called_once_with(bank_id="b", name="b")
|
mock_client.acreate_bank.assert_called_once_with(bank_id="b", name="b")
|
||||||
|
|
||||||
def test_ensure_bank_idempotent(self, mock_client, mock_registry):
|
@pytest.mark.asyncio
|
||||||
|
async def test_ensure_bank_idempotent(self, mock_client, mock_registry):
|
||||||
register_tools(bank_id="b", client=mock_client)
|
register_tools(bank_id="b", client=mock_client)
|
||||||
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
handler = mock_registry.register.call_args_list[0].kwargs["handler"]
|
||||||
handler({"content": "first"})
|
await handler({"content": "first"})
|
||||||
handler({"content": "second"})
|
await handler({"content": "second"})
|
||||||
# create_bank should only be called once
|
# acreate_bank should only be called once
|
||||||
mock_client.create_bank.assert_called_once()
|
mock_client.acreate_bank.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
# --- register (plugin entry point) tests ---
|
# --- register (plugin entry point) tests ---
|
||||||
|
|
@ -208,6 +229,123 @@ class TestRegisterPlugin:
|
||||||
register(ctx)
|
register(ctx)
|
||||||
ctx.register_tool.assert_not_called()
|
ctx.register_tool.assert_not_called()
|
||||||
|
|
||||||
|
def test_register_hooks(self, monkeypatch, mock_client):
|
||||||
|
monkeypatch.setenv("HINDSIGHT_API_URL", "http://localhost:8888")
|
||||||
|
monkeypatch.setenv("HINDSIGHT_BANK_ID", "test-bank")
|
||||||
|
ctx = MagicMock()
|
||||||
|
with patch("hindsight_hermes.tools._resolve_client", return_value=mock_client):
|
||||||
|
register(ctx)
|
||||||
|
hook_names = {call.args[0] for call in ctx.register_hook.call_args_list}
|
||||||
|
assert hook_names == {"pre_llm_call", "post_llm_call"}
|
||||||
|
|
||||||
|
|
||||||
|
# --- lifecycle hook tests ---
|
||||||
|
|
||||||
|
|
||||||
|
class TestLifecycleHooks:
|
||||||
|
"""Tests for pre_llm_call and post_llm_call hook callbacks."""
|
||||||
|
|
||||||
|
def _get_hook(self, ctx_mock, hook_name: str):
|
||||||
|
"""Extract the registered hook callback by name from the mock ctx."""
|
||||||
|
for call in ctx_mock.register_hook.call_args_list:
|
||||||
|
if call.args[0] == hook_name:
|
||||||
|
return call.args[1]
|
||||||
|
raise AssertionError(f"Hook {hook_name!r} not registered")
|
||||||
|
|
||||||
|
def _register_with_hooks(self, monkeypatch, mock_client, **env_overrides):
|
||||||
|
monkeypatch.setenv("HINDSIGHT_API_URL", "http://localhost:8888")
|
||||||
|
monkeypatch.setenv("HINDSIGHT_BANK_ID", "test-bank")
|
||||||
|
for k, v in env_overrides.items():
|
||||||
|
monkeypatch.setenv(k, v)
|
||||||
|
ctx = MagicMock()
|
||||||
|
with patch("hindsight_hermes.tools._resolve_client", return_value=mock_client):
|
||||||
|
register(ctx)
|
||||||
|
return ctx
|
||||||
|
|
||||||
|
# -- pre_llm_call --
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_pre_llm_call_returns_context(self, monkeypatch, mock_client):
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
|
result = await hook(
|
||||||
|
session_id="s1",
|
||||||
|
user_message="what color do I like?",
|
||||||
|
conversation_history=[],
|
||||||
|
is_first_turn=True,
|
||||||
|
model="test",
|
||||||
|
)
|
||||||
|
assert result is not None
|
||||||
|
assert "context" in result
|
||||||
|
assert "Memory 1" in result["context"]
|
||||||
|
assert "Memory 2" in result["context"]
|
||||||
|
mock_client.arecall.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_pre_llm_call_returns_none_on_no_results(self, monkeypatch, mock_client):
|
||||||
|
mock_client.arecall.return_value = SimpleNamespace(results=[])
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
|
result = await hook(user_message="hello")
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_pre_llm_call_returns_none_on_empty_message(self, monkeypatch, mock_client):
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
|
result = await hook(user_message="")
|
||||||
|
assert result is None
|
||||||
|
mock_client.arecall.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_pre_llm_call_returns_none_on_error(self, monkeypatch, mock_client):
|
||||||
|
mock_client.arecall.side_effect = RuntimeError("connection failed")
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "pre_llm_call")
|
||||||
|
result = await hook(user_message="hello")
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
# -- post_llm_call --
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_post_llm_call_retains_turn(self, monkeypatch, mock_client):
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
|
await hook(
|
||||||
|
session_id="s1",
|
||||||
|
user_message="remember I like green",
|
||||||
|
assistant_response="Got it, you like green!",
|
||||||
|
model="test",
|
||||||
|
)
|
||||||
|
mock_client.aretain.assert_called_once()
|
||||||
|
content = mock_client.aretain.call_args.kwargs["content"]
|
||||||
|
assert "remember I like green" in content
|
||||||
|
assert "Got it, you like green!" in content
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_post_llm_call_skips_empty_messages(self, monkeypatch, mock_client):
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
|
await hook(user_message="", assistant_response="hello")
|
||||||
|
mock_client.aretain.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_post_llm_call_skips_when_disabled(self, monkeypatch, mock_client):
|
||||||
|
ctx = self._register_with_hooks(
|
||||||
|
monkeypatch, mock_client, HINDSIGHT_AUTO_RETAIN="false"
|
||||||
|
)
|
||||||
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
|
await hook(user_message="hi", assistant_response="hello")
|
||||||
|
mock_client.aretain.assert_not_called()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_post_llm_call_does_not_raise_on_error(self, monkeypatch, mock_client):
|
||||||
|
mock_client.aretain.side_effect = RuntimeError("boom")
|
||||||
|
ctx = self._register_with_hooks(monkeypatch, mock_client)
|
||||||
|
hook = self._get_hook(ctx, "post_llm_call")
|
||||||
|
# Should not raise
|
||||||
|
await hook(user_message="hi", assistant_response="hello")
|
||||||
|
|
||||||
|
|
||||||
# --- memory_instructions tests ---
|
# --- memory_instructions tests ---
|
||||||
|
|
||||||
|
|
|
||||||
File diff suppressed because it is too large
Load diff
Loading…
Reference in a new issue