Просмотр исходного кода

feat: route event names through event agent

zhenyu.hu 3 недель назад
Родитель
Сommit
48556b799d

+ 21 - 3
src/agent_lab/application/event_agent.py

@@ -1,5 +1,6 @@
 import json
-from collections.abc import Iterable
+from collections.abc import Iterable, Sequence
+from dataclasses import dataclass
 from typing import Any
 
 from agent_lab.application.tools import ToolRegistry, build_default_tool_registry
@@ -7,6 +8,12 @@ from agent_lab.domain.events import ToolCallEvent
 from agent_lab.domain.messages import ChatMessage
 
 
+@dataclass(frozen=True)
+class EventAgentRequest:
+    events: list[ToolCallEvent]
+    history: list[ChatMessage]
+
+
 class EventAgent:
     def __init__(
         self,
@@ -16,7 +23,11 @@ class EventAgent:
         self.enabled_tools = set(enabled_tools)
         self.registry = registry or build_default_tool_registry()
 
-    async def handle(self, event: ToolCallEvent) -> ChatMessage:
+    async def handle(
+        self,
+        event: ToolCallEvent,
+        history: Sequence[ChatMessage] = (),
+    ) -> ChatMessage:
         if event.name not in self.enabled_tools:
             return self._tool_reply(
                 event,
@@ -27,7 +38,7 @@ class EventAgent:
             )
 
         try:
-            payload = self.registry.handle(event)
+            payload = self.registry.handle(event, history=history)
         except Exception as exc:
             payload = {
                 "tool": event.name,
@@ -35,6 +46,13 @@ class EventAgent:
             }
         return self._tool_reply(event, payload)
 
+    async def handle_many(
+        self,
+        events: Sequence[ToolCallEvent],
+        history: Sequence[ChatMessage],
+    ) -> list[ChatMessage]:
+        return [await self.handle(event, history=history) for event in events]
+
     def _tool_reply(self, event: ToolCallEvent, payload: dict[str, Any]) -> ChatMessage:
         return ChatMessage(
             role="tool",

+ 4 - 2
src/agent_lab/application/queues.py

@@ -2,7 +2,7 @@ import asyncio
 from dataclasses import dataclass, field
 from typing import Any
 
-from agent_lab.domain.events import ToolCallEvent
+from agent_lab.application.event_agent import EventAgentRequest
 from agent_lab.domain.messages import ChatMessage
 
 
@@ -10,4 +10,6 @@ from agent_lab.domain.messages import ChatMessage
 class RuntimeQueues:
     input: asyncio.Queue[ChatMessage] = field(default_factory=asyncio.Queue)
     output: asyncio.Queue[dict[str, Any]] = field(default_factory=asyncio.Queue)
-    events: asyncio.Queue[ToolCallEvent | None] = field(default_factory=asyncio.Queue)
+    events: asyncio.Queue[EventAgentRequest | None] = field(
+        default_factory=asyncio.Queue
+    )

+ 59 - 31
src/agent_lab/application/runtime.py

@@ -4,9 +4,10 @@ from collections.abc import AsyncIterator, Callable
 from typing import Any, Protocol
 
 from agent_lab.application.contracts import AgentParams, DebugRunRequest
-from agent_lab.application.event_agent import EventAgent
+from agent_lab.application.event_agent import EventAgent, EventAgentRequest
 from agent_lab.application.queues import RuntimeQueues
 from agent_lab.application.tools import ToolRegistry, build_default_tool_registry
+from agent_lab.domain.events import ToolCallEvent
 from agent_lab.domain.messages import ChatMessage, StreamItem
 
 
@@ -89,7 +90,6 @@ class DebugRuntime:
     ) -> None:
         messages = self._build_initial_messages(request)
         await queues.input.put(ChatMessage(role="user", content=request.user_message))
-        tools = self.registry.chat_tools(request.event_agent.enabled_tools)
 
         await queues.output.put({"type": "session_started"})
 
@@ -107,7 +107,13 @@ class DebugRuntime:
             saw_event = False
             assistant_content: list[str] = []
             assistant_tool_calls: list[dict[str, Any]] = []
+            events: list[ToolCallEvent] = []
             tool_replies: list[ChatMessage] = []
+            tools = (
+                self.registry.chat_tools(request.event_agent.enabled_tools)
+                if event_loops < request.event_agent.max_event_loops
+                else []
+            )
 
             async for item in self.chat_client.stream_chat(
                 messages=messages,
@@ -135,34 +141,46 @@ class DebugRuntime:
 
                 if item.kind == "event" and item.event is not None:
                     saw_event = True
+                    event = self._event_name_only(item.event)
+                    events.append(event)
                     assistant_tool_calls.append(
                         {
-                            "id": item.event.id,
+                            "id": event.id,
                             "type": "function",
                             "function": {
-                                "name": item.event.name,
-                                "arguments": item.event.raw_arguments,
+                                "name": event.name,
+                                "arguments": event.raw_arguments,
                             },
                         }
                     )
-                    await queues.events.put(item.event)
+
+            if assistant_content or assistant_tool_calls:
+                assistant_message = ChatMessage(
+                    role="assistant",
+                    content="".join(assistant_content),
+                    tool_calls=assistant_tool_calls or None,
+                )
+                messages.append(assistant_message)
+
+            if events:
+                for event in events:
                     await queues.output.put(
-                        {"type": "event", "event": item.event.model_dump()}
+                        {"type": "event", "event": event.model_dump()}
+                    )
+                await queues.events.put(
+                    EventAgentRequest(
+                        events=events,
+                        history=list(messages),
                     )
-                    reply = await self._wait_for_tool_reply(queues, item.event.id)
-                    tool_replies.append(reply)
+                )
+                tool_replies = await self._wait_for_tool_replies(
+                    queues,
+                    [event.id for event in events],
+                )
+                for reply in tool_replies:
                     await queues.output.put(
                         {"type": "tool_result", "message": reply.model_dump()}
                     )
-
-            if assistant_content or assistant_tool_calls:
-                messages.append(
-                    ChatMessage(
-                        role="assistant",
-                        content="".join(assistant_content),
-                        tool_calls=assistant_tool_calls or None,
-                    )
-                )
             messages.extend(tool_replies)
 
             await queues.output.put(
@@ -183,8 +201,6 @@ class DebugRuntime:
                 break
 
             event_loops += 1
-            if event_loops >= request.event_agent.max_event_loops:
-                break
 
         await queues.output.put({"type": "done"})
 
@@ -194,24 +210,36 @@ class DebugRuntime:
         event_agent: EventAgent,
     ) -> None:
         while True:
-            event = await queues.events.get()
-            if event is None:
+            request = await queues.events.get()
+            if request is None:
                 return
-            await queues.input.put(await event_agent.handle(event))
+            for reply in await event_agent.handle_many(
+                request.events,
+                history=request.history,
+            ):
+                await queues.input.put(reply)
 
-    async def _wait_for_tool_reply(
+    async def _wait_for_tool_replies(
         self,
         queues: RuntimeQueues,
-        tool_call_id: str,
-    ) -> ChatMessage:
+        tool_call_ids: list[str],
+    ) -> list[ChatMessage]:
+        pending = set(tool_call_ids)
+        replies: dict[str, ChatMessage] = {}
         deferred: list[ChatMessage] = []
-        while True:
+        while pending:
             message = await queues.input.get()
-            if message.role == "tool" and message.tool_call_id == tool_call_id:
-                for deferred_message in deferred:
-                    await queues.input.put(deferred_message)
-                return message
+            if message.role == "tool" and message.tool_call_id in pending:
+                replies[message.tool_call_id] = message
+                pending.remove(message.tool_call_id)
+                continue
             deferred.append(message)
+        for deferred_message in deferred:
+            await queues.input.put(deferred_message)
+        return [replies[tool_call_id] for tool_call_id in tool_call_ids]
+
+    def _event_name_only(self, event: ToolCallEvent) -> ToolCallEvent:
+        return event.model_copy(update={"arguments": {}, "raw_arguments": "{}"})
 
     def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]:
         messages = [

+ 59 - 4
src/agent_lab/application/tools.py

@@ -1,12 +1,23 @@
-from collections.abc import Callable, Iterable
+from collections.abc import Callable, Iterable, Sequence
 from copy import deepcopy
 from dataclasses import dataclass
 from typing import Any
 
 from agent_lab.domain.events import ToolCallEvent
+from agent_lab.domain.messages import ChatMessage
 
 
 ToolHandler = Callable[[ToolCallEvent], dict[str, Any]]
+ToolArgumentResolver = Callable[
+    [ToolCallEvent, Sequence[ChatMessage]],
+    dict[str, Any],
+]
+
+EVENT_NAME_ONLY_PARAMETERS: dict[str, Any] = {
+    "type": "object",
+    "properties": {},
+    "additionalProperties": False,
+}
 
 
 @dataclass(frozen=True)
@@ -15,6 +26,7 @@ class ToolDefinition:
     description: str
     parameters: dict[str, Any]
     handler: ToolHandler
+    argument_resolver: ToolArgumentResolver | None = None
 
 
 class ToolRegistry:
@@ -39,21 +51,42 @@ class ToolRegistry:
                 "function": {
                     "name": definition.name,
                     "description": definition.description,
-                    "parameters": deepcopy(definition.parameters),
+                    "parameters": deepcopy(EVENT_NAME_ONLY_PARAMETERS),
                 },
             }
             for definition in self._definitions.values()
             if definition.name in enabled
         ]
 
-    def handle(self, event: ToolCallEvent) -> dict[str, Any]:
+    def handle(
+        self,
+        event: ToolCallEvent,
+        history: Sequence[ChatMessage] = (),
+    ) -> dict[str, Any]:
         definition = self._definitions.get(event.name)
         if definition is None:
             return {
                 "tool": event.name,
                 "error": "unknown tool",
             }
-        return definition.handler(event)
+        resolved_event = event.model_copy(
+            update={
+                "arguments": self._resolve_arguments(definition, event, history),
+            }
+        )
+        return definition.handler(resolved_event)
+
+    def _resolve_arguments(
+        self,
+        definition: ToolDefinition,
+        event: ToolCallEvent,
+        history: Sequence[ChatMessage],
+    ) -> dict[str, Any]:
+        if definition.argument_resolver is not None:
+            return definition.argument_resolver(event, history)
+        if history:
+            return {}
+        return deepcopy(event.arguments)
 
 
 def build_default_tool_registry() -> ToolRegistry:
@@ -70,11 +103,33 @@ def build_default_tool_registry() -> ToolRegistry:
                     "required": ["message"],
                 },
                 handler=_handle_handoff_note,
+                argument_resolver=_resolve_handoff_note_arguments,
             )
         ]
     )
 
 
+def _resolve_handoff_note_arguments(
+    event: ToolCallEvent,
+    history: Sequence[ChatMessage],
+) -> dict[str, Any]:
+    content = _latest_content(history, preferred_roles=("assistant", "user"))
+    if not content and not history:
+        content = str(event.arguments.get("message", ""))
+    return {"message": content}
+
+
+def _latest_content(
+    history: Sequence[ChatMessage],
+    preferred_roles: tuple[str, ...],
+) -> str:
+    for role in preferred_roles:
+        for message in reversed(history):
+            if message.role == role and message.content.strip():
+                return message.content.strip()
+    return ""
+
+
 def _handle_handoff_note(event: ToolCallEvent) -> dict[str, Any]:
     return {
         "tool": "handoff_note",

+ 1 - 1
src/agent_lab/presentation/static/app.js

@@ -328,7 +328,7 @@ function handleServerMessage(message) {
     return;
   }
   if (message.type === "event") {
-    appendLog("event", `${message.event.name}: ${JSON.stringify(message.event.arguments)}`);
+    appendLog("event", message.event.name);
     return;
   }
   if (message.type === "tool_result") {

+ 210 - 19
tests/test_debug_runtime.py

@@ -7,6 +7,7 @@ from typing import Any
 import pytest
 
 from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
+from agent_lab.application.event_agent import EventAgentRequest
 from agent_lab.application.runtime import DebugRuntime
 from agent_lab.application.tools import ToolDefinition, ToolRegistry
 from agent_lab.domain.events import ToolCallEvent
@@ -44,6 +45,9 @@ class RecordingQueue(asyncio.Queue):
             return item.role
         if isinstance(item, ToolCallEvent):
             return f"event:{item.name}:{item.id}"
+        if isinstance(item, EventAgentRequest):
+            events = ",".join(f"{event.name}:{event.id}" for event in item.events)
+            return f"event_request:{events}"
         if isinstance(item, dict):
             return f"output:{item.get('type')}"
         return type(item).__name__
@@ -65,8 +69,8 @@ class FakeChatClient:
                 ToolCallEvent(
                     id="call_1",
                     name="handoff_note",
-                    arguments={"message": "need event agent"},
-                    raw_arguments='{"message":"need event agent"}',
+                    arguments={},
+                    raw_arguments="{}",
                 )
             )
             return
@@ -92,8 +96,8 @@ class StrictHistoryChatClient:
                 ToolCallEvent(
                     id="call_1",
                     name="handoff_note",
-                    arguments={"message": "need event agent"},
-                    raw_arguments='{"message":"need event agent"}',
+                    arguments={},
+                    raw_arguments="{}",
                 )
             )
             return
@@ -102,6 +106,42 @@ class StrictHistoryChatClient:
         yield StreamItem.message_delta("final answer")
 
 
+class MultiEventChatClient:
+    def __init__(self) -> None:
+        self.calls = 0
+        self.second_call_messages: list[ChatMessage] = []
+
+    async def stream_chat(
+        self,
+        messages: list[ChatMessage],
+        tools: list[dict],
+        params: AgentParams,
+    ) -> AsyncIterator[StreamItem]:
+        self.calls += 1
+        if self.calls == 1:
+            yield StreamItem.message_delta("Checking events.")
+            yield StreamItem.event(
+                ToolCallEvent(
+                    id="call_1",
+                    name="handoff_note",
+                    arguments={"message": "ignored chat argument"},
+                    raw_arguments='{"message":"ignored chat argument"}',
+                )
+            )
+            yield StreamItem.event(
+                ToolCallEvent(
+                    id="call_2",
+                    name="audit_note",
+                    arguments={"message": "ignored chat argument"},
+                    raw_arguments='{"message":"ignored chat argument"}',
+                )
+            )
+            return
+
+        self.second_call_messages = list(messages)
+        yield StreamItem.message_delta("Final answer.")
+
+
 class ToolCapturingChatClient:
     def __init__(self) -> None:
         self.tools: list[dict[str, Any]] = []
@@ -116,6 +156,33 @@ class ToolCapturingChatClient:
         yield StreamItem.message_delta("final answer")
 
 
+class EventLoopLimitChatClient:
+    def __init__(self) -> None:
+        self.calls = 0
+        self.tools_by_call: list[list[dict[str, Any]]] = []
+
+    async def stream_chat(
+        self,
+        messages: list[ChatMessage],
+        tools: list[dict],
+        params: AgentParams,
+    ) -> AsyncIterator[StreamItem]:
+        self.calls += 1
+        self.tools_by_call.append(list(tools))
+        if self.calls == 1:
+            yield StreamItem.event(
+                ToolCallEvent(
+                    id="call_1",
+                    name="handoff_note",
+                    arguments={},
+                    raw_arguments="{}",
+                )
+            )
+            return
+
+        yield StreamItem.message_delta("final after event limit")
+
+
 class RoundStatsChatClient:
     async def stream_chat(
         self,
@@ -150,8 +217,8 @@ class EventRoundStatsChatClient:
                 ToolCallEvent(
                     id="call_1",
                     name="handoff_note",
-                    arguments={"message": "need event agent"},
-                    raw_arguments='{"message":"need event agent"}',
+                    arguments={},
+                    raw_arguments="{}",
                 )
             )
             yield StreamItem.usage_item(
@@ -206,6 +273,104 @@ async def test_runtime_routes_chat_events_through_event_agent_then_continues_cha
     assert outputs[4]["content"] == "final answer"
 
 
+@pytest.mark.asyncio
+async def test_runtime_batches_round_events_before_continuing_chat_agent():
+    def resolve_from_history(
+        event: ToolCallEvent,
+        history: list[ChatMessage],
+    ) -> dict[str, Any]:
+        return {"message": history[-1].content, "event": event.name}
+
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="handoff_note",
+                description="Send a handoff note.",
+                parameters={
+                    "type": "object",
+                    "properties": {"message": {"type": "string"}},
+                    "required": ["message"],
+                },
+                handler=lambda event: {
+                    "tool": event.name,
+                    "message": event.arguments["message"],
+                },
+                argument_resolver=resolve_from_history,
+            ),
+            ToolDefinition(
+                name="audit_note",
+                description="Send an audit note.",
+                parameters={
+                    "type": "object",
+                    "properties": {"message": {"type": "string"}},
+                    "required": ["message"],
+                },
+                handler=lambda event: {
+                    "tool": event.name,
+                    "message": event.arguments["message"],
+                },
+                argument_resolver=resolve_from_history,
+            ),
+        ]
+    )
+    request = DebugRunRequest(
+        user_message="debug this",
+        system_prompts=[],
+        pre_messages=[],
+        chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
+        event_agent=EventAgentParams(
+            enabled_tools=["handoff_note", "audit_note"],
+            max_event_loops=2,
+        ),
+    )
+    client = MultiEventChatClient()
+    runtime = DebugRuntime(client, registry=registry)
+
+    outputs = [message async for message in runtime.run(request)]
+
+    assert [message["type"] for message in outputs] == [
+        "session_started",
+        "message_delta",
+        "event",
+        "event",
+        "tool_result",
+        "tool_result",
+        "round_stats",
+        "message_delta",
+        "round_stats",
+        "done",
+    ]
+    assert client.calls == 2
+    assert [message.role for message in client.second_call_messages] == [
+        "user",
+        "assistant",
+        "tool",
+        "tool",
+    ]
+    assistant_message = client.second_call_messages[1]
+    assert assistant_message.content == "Checking events."
+    assert assistant_message.tool_calls == [
+        {
+            "id": "call_1",
+            "type": "function",
+            "function": {"name": "handoff_note", "arguments": "{}"},
+        },
+        {
+            "id": "call_2",
+            "type": "function",
+            "function": {"name": "audit_note", "arguments": "{}"},
+        },
+    ]
+    assert [
+        json.loads(message.content)
+        for message in client.second_call_messages
+        if message.role == "tool"
+    ] == [
+        {"tool": "handoff_note", "message": "Checking events."},
+        {"tool": "audit_note", "message": "Checking events."},
+    ]
+
+
 @pytest.mark.asyncio
 async def test_runtime_start_returns_queues_for_downstream_output_consumer():
     request = DebugRunRequest(
@@ -234,6 +399,35 @@ async def test_runtime_start_returns_queues_for_downstream_output_consumer():
     ]
 
 
+@pytest.mark.asyncio
+async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
+    request = DebugRunRequest(
+        user_message="debug this",
+        system_prompts=[],
+        pre_messages=[],
+        chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
+        event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=1),
+    )
+    client = EventLoopLimitChatClient()
+    runtime = DebugRuntime(client)
+
+    outputs = [message async for message in runtime.run(request)]
+
+    assert client.calls == 2
+    assert client.tools_by_call[0][0]["function"]["name"] == "handoff_note"
+    assert client.tools_by_call[1] == []
+    assert [message["type"] for message in outputs] == [
+        "session_started",
+        "event",
+        "tool_result",
+        "round_stats",
+        "message_delta",
+        "round_stats",
+        "done",
+    ]
+    assert outputs[4]["content"] == "final after event limit"
+
+
 @pytest.mark.asyncio
 async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
     RuntimeQueues = _runtime_queues_class()
@@ -342,12 +536,12 @@ async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
     assert queue_log.index(("input", "put", "user")) < queue_log.index(
         ("input", "get", "user")
     )
-    assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
-        ("events", "get", "event:handoff_note:call_1")
-    )
-    assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
-        ("input", "put", "tool:call_1")
-    )
+    assert queue_log.index(
+        ("events", "put", "event_request:handoff_note:call_1")
+    ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
+    assert queue_log.index(
+        ("events", "get", "event_request:handoff_note:call_1")
+    ) < queue_log.index(("input", "put", "tool:call_1"))
     assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
         ("input", "get", "tool:call_1")
     )
@@ -430,7 +624,7 @@ async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
             "type": "function",
             "function": {
                 "name": "handoff_note",
-                "arguments": '{"message":"need event agent"}',
+                "arguments": "{}",
             },
         }
     ]
@@ -439,7 +633,7 @@ async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
 
 
 @pytest.mark.asyncio
-async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent():
+async def test_runtime_passes_event_names_without_tool_parameters_to_chat_agent():
     registry = ToolRegistry(
         [
             ToolDefinition(
@@ -478,11 +672,8 @@ async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent()
                 "description": "Registry-owned handoff tool.",
                 "parameters": {
                     "type": "object",
-                    "properties": {
-                        "message": {"type": "string"},
-                        "priority": {"type": "number"},
-                    },
-                    "required": ["message"],
+                    "properties": {},
+                    "additionalProperties": False,
                 },
             },
         }

+ 10 - 5
tests/test_event_agent.py

@@ -5,26 +5,31 @@ import pytest
 from agent_lab.application.event_agent import EventAgent
 from agent_lab.application.tools import ToolDefinition, ToolRegistry
 from agent_lab.domain.events import ToolCallEvent
+from agent_lab.domain.messages import ChatMessage
 
 
 @pytest.mark.asyncio
-async def test_event_agent_executes_enabled_tool_as_tool_reply_message():
+async def test_event_agent_resolves_tool_arguments_from_history():
     agent = EventAgent(enabled_tools=["handoff_note"])
     event = ToolCallEvent(
         id="call_1",
         name="handoff_note",
-        arguments={"message": "inspect this event"},
-        raw_arguments='{"message":"inspect this event"}',
+        arguments={"message": "chat agent argument should be ignored"},
+        raw_arguments='{"message":"chat agent argument should be ignored"}',
     )
+    history = [
+        ChatMessage(role="user", content="debug this event flow"),
+        ChatMessage(role="assistant", content="I need the event agent."),
+    ]
 
-    reply = await agent.handle(event)
+    reply = await agent.handle(event, history=history)
 
     assert reply.role == "tool"
     assert reply.tool_call_id == "call_1"
     assert reply.name == "handoff_note"
     assert json.loads(reply.content) == {
         "tool": "handoff_note",
-        "message": "inspect this event",
+        "message": "I need the event agent.",
     }
 
 

+ 2 - 2
tests/test_websocket_api.py

@@ -208,8 +208,8 @@ class HistoryCapturingChatClient:
                 ToolCallEvent(
                     id="call_1",
                     name="handoff_note",
-                    arguments={"message": "inspect this"},
-                    raw_arguments='{"message":"inspect this"}',
+                    arguments={},
+                    raw_arguments="{}",
                 )
             )
             return