|
|
@@ -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(
|