refactor(agent): add turn hook factories

This commit is contained in:
chengyongru
2026-07-08 21:02:12 +08:00
committed by Xubin Ren
parent 04dbf17426
commit 8559458258
7 changed files with 248 additions and 32 deletions
+73 -3
View File
@@ -6,7 +6,13 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext, CompositeHook
from nanobot.agent.hook import (
AgentHook,
AgentHookContext,
AgentRunHookContext,
AgentTurnHookContext,
CompositeHook,
)
def _ctx() -> AgentHookContext:
@@ -348,7 +354,7 @@ async def test_composite_can_wrap_another_composite():
# ---------------------------------------------------------------------------
def _make_loop(tmp_path, hooks=None):
def _make_loop(tmp_path, hooks=None, hook_factories=None):
from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus
@@ -363,7 +369,11 @@ def _make_loop(tmp_path, hooks=None):
patch("nanobot.agent.loop.Consolidator"):
mock_sub_mgr.return_value.cancel_by_session = AsyncMock(return_value=0)
loop = AgentLoop(
bus=bus, provider=provider, workspace=tmp_path, hooks=hooks,
bus=bus,
provider=provider,
workspace=tmp_path,
hooks=hooks,
hook_factories=hook_factories,
)
return loop
@@ -405,6 +415,66 @@ async def test_agent_loop_extra_hook_receives_calls(tmp_path):
assert "after_run:completed" in events
@pytest.mark.asyncio
async def test_agent_loop_turn_hook_factories_receive_context(tmp_path):
"""Turn-scoped hooks can be supplied externally and see turn-local context."""
from nanobot.providers.base import LLMResponse
captured: list[tuple[str, AgentTurnHookContext]] = []
events: list[str] = []
class TrackingHook(AgentHook):
def __init__(self, label: str) -> None:
super().__init__()
self._label = label
async def before_iteration(self, context):
events.append(f"{self._label}:{context.iteration}")
def factory(label: str):
def _create(context: AgentTurnHookContext) -> AgentHook:
captured.append((label, context))
return TrackingHook(label)
return _create
loop = _make_loop(tmp_path, hook_factories=[factory("registered")])
loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="done", tool_calls=[], usage={})
)
loop.tools.get_definitions = MagicMock(return_value=[])
async def on_progress(*args, **kwargs):
pass
await loop._run_agent_loop(
[{"role": "user", "content": "hi"}],
on_progress=on_progress,
channel="websocket",
chat_id="chat-1",
message_id="msg-1",
metadata={"source": "test"},
session_key="websocket:chat-1",
hook_factories=[factory("turn")],
)
assert events == ["registered:0", "turn:0"]
assert [label for label, _ in captured] == ["registered", "turn"]
assert [context.on_progress for _, context in captured] == [on_progress, on_progress]
assert [context.workspace for _, context in captured] == [tmp_path, tmp_path]
assert [context.channel for _, context in captured] == ["websocket", "websocket"]
assert [context.chat_id for _, context in captured] == ["chat-1", "chat-1"]
assert [context.message_id for _, context in captured] == ["msg-1", "msg-1"]
assert [context.session_key for _, context in captured] == [
"websocket:chat-1",
"websocket:chat-1",
]
assert [context.metadata for _, context in captured] == [
{"source": "test"},
{"source": "test"},
]
@pytest.mark.asyncio
async def test_agent_loop_extra_hook_error_isolation(tmp_path):
"""A faulty extra hook does not crash the agent loop."""
+60 -2
View File
@@ -1,6 +1,6 @@
import pytest
from nanobot.agent.hook import AgentHook, AgentHookContext
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentTurnHookContext
from nanobot.agent.turn_hooks import AgentTurnHookSpec, build_agent_turn_hook
@@ -44,10 +44,67 @@ async def test_turn_hook_builder_runs_registered_hooks_before_turn_hooks() -> No
@pytest.mark.asyncio
async def test_turn_hook_builder_skips_extra_hooks_for_ephemeral_turns_by_default() -> None:
async def test_turn_hook_builder_runs_factories_with_matching_registration_order(
tmp_path,
) -> None:
events: list[str] = []
captured: list[AgentTurnHookContext] = []
def factory(label: str):
def _create(context: AgentTurnHookContext) -> AgentHook:
captured.append(context)
return RecordingHook(events, label)
return _create
hook = build_agent_turn_hook(AgentTurnHookSpec(
on_iteration=lambda iteration: events.append(f"progress:{iteration}"),
channel="websocket",
chat_id="chat-1",
message_id="msg-1",
session_key="websocket:chat-1",
workspace=tmp_path,
metadata={"source": "test"},
registered_hook_factories=[factory("registered_factory")],
registered_hooks=[RecordingHook(events, "registered")],
turn_hook_factories=[factory("turn_factory")],
turn_hooks=[RecordingHook(events, "turn")],
))
await hook.before_iteration(AgentHookContext(iteration=2, messages=[]))
assert events == [
"progress:2",
"registered_factory:2",
"registered:2",
"turn_factory:2",
"turn:2",
]
assert [context.workspace for context in captured] == [tmp_path, tmp_path]
assert [context.channel for context in captured] == ["websocket", "websocket"]
assert [context.chat_id for context in captured] == ["chat-1", "chat-1"]
assert [context.message_id for context in captured] == ["msg-1", "msg-1"]
assert [context.session_key for context in captured] == [
"websocket:chat-1",
"websocket:chat-1",
]
assert [context.metadata for context in captured] == [
{"source": "test"},
{"source": "test"},
]
@pytest.mark.asyncio
async def test_turn_hook_builder_skips_extra_hooks_for_ephemeral_turns_by_default() -> None:
events: list[str] = []
factory_calls: list[str] = []
def factory(context: AgentTurnHookContext) -> AgentHook:
factory_calls.append(context.channel)
return RecordingHook(events, "factory")
hook = build_agent_turn_hook(AgentTurnHookSpec(
registered_hook_factories=[factory],
registered_hooks=[RecordingHook(events)],
ephemeral=True,
))
@@ -55,6 +112,7 @@ async def test_turn_hook_builder_skips_extra_hooks_for_ephemeral_turns_by_defaul
await hook.before_iteration(AgentHookContext(iteration=1, messages=[]))
assert events == []
assert factory_calls == []
@pytest.mark.asyncio