refactor(agent): add turn hook factories
This commit is contained in:
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user