ソースを参照

Add explicit runtime queues

zhenyu.hu 3 週間 前
コミット
b2bba51093

+ 13 - 0
src/agent_lab/application/queues.py

@@ -0,0 +1,13 @@
+import asyncio
+from dataclasses import dataclass, field
+from typing import Any
+
+from agent_lab.domain.events import ToolCallEvent
+from agent_lab.domain.messages import ChatMessage
+
+
+@dataclass
+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] = field(default_factory=asyncio.Queue)

+ 48 - 12
src/agent_lab/application/runtime.py

@@ -3,6 +3,7 @@ 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.queues import RuntimeQueues
 from agent_lab.domain.messages import ChatMessage, StreamItem
 
 
@@ -17,22 +18,29 @@ class ChatClient(Protocol):
 
 
 class DebugRuntime:
-    def __init__(self, chat_client: ChatClient) -> None:
+    def __init__(
+        self,
+        chat_client: ChatClient,
+        queues: RuntimeQueues | None = None,
+    ) -> None:
         self.chat_client = chat_client
+        self.queues = queues
 
     async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
+        queues = self.queues or RuntimeQueues()
         messages = self._build_initial_messages(request)
+        await queues.input.put(ChatMessage(role="user", content=request.user_message))
         tools = self._build_tools(request.event_agent.enabled_tools)
         event_agent = EventAgent(request.event_agent.enabled_tools)
 
-        yield {"type": "session_started"}
+        yield await self._emit(queues, {"type": "session_started"})
 
         event_loops = 0
         while True:
+            await self._drain_input(queues, messages)
             saw_event = False
             assistant_content: list[str] = []
             assistant_tool_calls: list[dict[str, Any]] = []
-            pending_tool_replies: list[ChatMessage] = []
 
             async for item in self.chat_client.stream_chat(
                 messages=messages,
@@ -41,11 +49,17 @@ class DebugRuntime:
             ):
                 if item.kind == "message_delta":
                     assistant_content.append(item.content or "")
-                    yield {"type": "message_delta", "content": item.content}
+                    yield await self._emit(
+                        queues,
+                        {"type": "message_delta", "content": item.content},
+                    )
                     continue
 
                 if item.kind == "usage" and item.usage is not None:
-                    yield {"type": "usage", "usage": item.usage.model_dump()}
+                    yield await self._emit(
+                        queues,
+                        {"type": "usage", "usage": item.usage.model_dump()},
+                    )
                     continue
 
                 if item.kind == "event" and item.event is not None:
@@ -60,10 +74,18 @@ class DebugRuntime:
                             },
                         }
                     )
-                    yield {"type": "event", "event": item.event.model_dump()}
-                    reply = await event_agent.handle(item.event)
-                    pending_tool_replies.append(reply)
-                    yield {"type": "tool_result", "message": reply.model_dump()}
+                    await queues.events.put(item.event)
+                    yield await self._emit(
+                        queues,
+                        {"type": "event", "event": item.event.model_dump()},
+                    )
+                    event = await queues.events.get()
+                    reply = await event_agent.handle(event)
+                    await queues.input.put(reply)
+                    yield await self._emit(
+                        queues,
+                        {"type": "tool_result", "message": reply.model_dump()},
+                    )
 
             if assistant_content or assistant_tool_calls:
                 messages.append(
@@ -73,7 +95,6 @@ class DebugRuntime:
                         tool_calls=assistant_tool_calls or None,
                     )
                 )
-            messages.extend(pending_tool_replies)
 
             if not saw_event:
                 break
@@ -82,7 +103,7 @@ class DebugRuntime:
             if event_loops >= request.event_agent.max_event_loops:
                 break
 
-        yield {"type": "done"}
+        yield await self._emit(queues, {"type": "done"})
 
     def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]:
         messages = [
@@ -90,9 +111,24 @@ class DebugRuntime:
             for prompt in request.system_prompts
         ]
         messages.extend(request.pre_messages)
-        messages.append(ChatMessage(role="user", content=request.user_message))
         return messages
 
+    async def _drain_input(
+        self,
+        queues: RuntimeQueues,
+        messages: list[ChatMessage],
+    ) -> None:
+        while not queues.input.empty():
+            messages.append(await queues.input.get())
+
+    async def _emit(
+        self,
+        queues: RuntimeQueues,
+        payload: dict[str, Any],
+    ) -> dict[str, Any]:
+        await queues.output.put(payload)
+        return await queues.output.get()
+
     def _build_tools(self, enabled_tools: list[str]) -> list[dict[str, Any]]:
         tools: list[dict[str, Any]] = []
         if "handoff_note" in enabled_tools:

+ 134 - 0
tests/test_debug_runtime.py

@@ -1,4 +1,7 @@
+import asyncio
+import importlib
 from collections.abc import AsyncIterator
+from typing import Any
 
 import pytest
 
@@ -8,6 +11,38 @@ from agent_lab.domain.events import ToolCallEvent
 from agent_lab.domain.messages import ChatMessage, StreamItem
 
 
+def _runtime_queues_class():
+    module = importlib.import_module("agent_lab.application.queues")
+    return module.RuntimeQueues
+
+
+class RecordingQueue(asyncio.Queue):
+    def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
+        super().__init__()
+        self.name = name
+        self.log = log
+
+    async def put(self, item: Any) -> None:
+        self.log.append((self.name, "put", self._describe(item)))
+        await super().put(item)
+
+    async def get(self) -> Any:
+        item = await super().get()
+        self.log.append((self.name, "get", self._describe(item)))
+        return item
+
+    def _describe(self, item: Any) -> str:
+        if isinstance(item, ChatMessage):
+            if item.role == "tool":
+                return f"tool:{item.tool_call_id}"
+            return item.role
+        if isinstance(item, ToolCallEvent):
+            return f"event:{item.name}:{item.id}"
+        if isinstance(item, dict):
+            return f"output:{item.get('type')}"
+        return type(item).__name__
+
+
 class FakeChatClient:
     def __init__(self) -> None:
         self.calls = 0
@@ -61,6 +96,19 @@ class StrictHistoryChatClient:
         yield StreamItem.message_delta("final answer")
 
 
+def test_runtime_queues_exposes_input_output_and_events_queues():
+    RuntimeQueues = _runtime_queues_class()
+
+    queues = RuntimeQueues()
+
+    assert isinstance(queues.input, asyncio.Queue)
+    assert isinstance(queues.output, asyncio.Queue)
+    assert isinstance(queues.events, asyncio.Queue)
+    assert queues.input is not queues.output
+    assert queues.input is not queues.events
+    assert queues.output is not queues.events
+
+
 @pytest.mark.asyncio
 async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
     request = DebugRunRequest(
@@ -87,6 +135,92 @@ async def test_runtime_routes_chat_events_through_event_agent_then_continues_cha
     assert outputs[3]["content"] == "final answer"
 
 
+@pytest.mark.asyncio
+async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
+    RuntimeQueues = _runtime_queues_class()
+    queue_log: list[tuple[str, str, str]] = []
+    queues = RuntimeQueues(
+        input=RecordingQueue("input", queue_log),
+        output=RecordingQueue("output", queue_log),
+        events=RecordingQueue("events", queue_log),
+    )
+    request = DebugRunRequest(
+        user_message="debug this",
+        system_prompts=["You are a debugger."],
+        pre_messages=[],
+        chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
+        event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
+    )
+    runtime = DebugRuntime(FakeChatClient(), queues=queues)
+
+    outputs = [message async for message in runtime.run(request)]
+
+    assert [message["type"] for message in outputs] == [
+        "session_started",
+        "event",
+        "tool_result",
+        "message_delta",
+        "done",
+    ]
+    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(("input", "put", "tool:call_1")) < queue_log.index(
+        ("input", "get", "tool:call_1")
+    )
+
+
+@pytest.mark.asyncio
+async def test_runtime_yields_existing_output_order_from_output_queue():
+    RuntimeQueues = _runtime_queues_class()
+    queue_log: list[tuple[str, str, str]] = []
+    queues = RuntimeQueues(
+        input=RecordingQueue("input", queue_log),
+        output=RecordingQueue("output", queue_log),
+        events=RecordingQueue("events", queue_log),
+    )
+    request = DebugRunRequest(
+        user_message="debug this",
+        system_prompts=["You are a debugger."],
+        pre_messages=[],
+        chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
+        event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
+    )
+    runtime = DebugRuntime(FakeChatClient(), queues=queues)
+
+    outputs = [message async for message in runtime.run(request)]
+
+    assert [message["type"] for message in outputs] == [
+        "session_started",
+        "event",
+        "tool_result",
+        "message_delta",
+        "done",
+    ]
+    assert [
+        entry
+        for entry in queue_log
+        if entry[0] == "output" and entry[1] in {"put", "get"}
+    ] == [
+        ("output", "put", "output:session_started"),
+        ("output", "get", "output:session_started"),
+        ("output", "put", "output:event"),
+        ("output", "get", "output:event"),
+        ("output", "put", "output:tool_result"),
+        ("output", "get", "output:tool_result"),
+        ("output", "put", "output:message_delta"),
+        ("output", "get", "output:message_delta"),
+        ("output", "put", "output:done"),
+        ("output", "get", "output:done"),
+    ]
+
+
 @pytest.mark.asyncio
 async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
     request = DebugRunRequest(