import asyncio import importlib from collections.abc import AsyncIterator from typing import Any import pytest from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams from agent_lab.application.runtime import DebugRuntime 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 async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, ) -> AsyncIterator[StreamItem]: self.calls += 1 if self.calls == 1: yield StreamItem.event( ToolCallEvent( id="call_1", name="handoff_note", arguments={"message": "need event agent"}, raw_arguments='{"message":"need event agent"}', ) ) return assert any(message.role == "tool" for message in messages) yield StreamItem.message_delta("final answer") class StrictHistoryChatClient: 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.event( ToolCallEvent( id="call_1", name="handoff_note", arguments={"message": "need event agent"}, raw_arguments='{"message":"need event agent"}', ) ) return self.second_call_messages = list(messages) 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( 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), ) client = FakeChatClient() runtime = DebugRuntime(client) outputs = [message async for message in runtime.run(request)] assert client.calls == 2 assert [message["type"] for message in outputs] == [ "session_started", "event", "tool_result", "message_delta", "done", ] assert outputs[1]["event"]["name"] == "handoff_note" 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( 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=2), ) client = StrictHistoryChatClient() runtime = DebugRuntime(client) outputs = [message async for message in runtime.run(request)] assert client.calls == 2 assert [message.role for message in client.second_call_messages] == [ "system", "user", "assistant", "tool", ] assistant_message = client.second_call_messages[2] tool_message = client.second_call_messages[3] assert assistant_message.content == "" assert assistant_message.tool_calls == [ { "id": "call_1", "type": "function", "function": { "name": "handoff_note", "arguments": '{"message":"need event agent"}', }, } ] assert tool_message.tool_call_id == "call_1" assert outputs[-1] == {"type": "done"}