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.application.tools import ToolDefinition, ToolRegistry from agent_lab.domain.events import ToolCallEvent from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage 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") class ToolCapturingChatClient: def __init__(self) -> None: self.tools: list[dict[str, Any]] = [] async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, ) -> AsyncIterator[StreamItem]: self.tools = list(tools) yield StreamItem.message_delta("final answer") class RoundStatsChatClient: async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, ) -> AsyncIterator[StreamItem]: yield StreamItem.message_delta("hello") yield StreamItem.usage_item( TokenUsage( prompt_tokens=10, completion_tokens=20, total_tokens=30, cached_tokens=5, ) ) class EventRoundStatsChatClient: 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"}', ) ) yield StreamItem.usage_item( TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3) ) return yield StreamItem.message_delta("final answer") yield StreamItem.usage_item( TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10) ) 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", "round_stats", "message_delta", "round_stats", "done", ] assert outputs[1]["event"]["name"] == "handoff_note" assert outputs[4]["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", "round_stats", "message_delta", "round_stats", "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", "round_stats", "message_delta", "round_stats", "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:round_stats"), ("output", "get", "output:round_stats"), ("output", "put", "output:message_delta"), ("output", "get", "output:message_delta"), ("output", "put", "output:round_stats"), ("output", "get", "output:round_stats"), ("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"} @pytest.mark.asyncio async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent(): registry = ToolRegistry( [ ToolDefinition( name="handoff_note", description="Registry-owned handoff tool.", parameters={ "type": "object", "properties": { "message": {"type": "string"}, "priority": {"type": "number"}, }, "required": ["message"], }, handler=lambda event: {"tool": event.name, "message": "handled"}, ) ] ) 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=2), ) client = ToolCapturingChatClient() runtime = DebugRuntime(client, registry=registry) outputs = [message async for message in runtime.run(request)] assert outputs[-1] == {"type": "done"} assert client.tools == [ { "type": "function", "function": { "name": "handoff_note", "description": "Registry-owned handoff tool.", "parameters": { "type": "object", "properties": { "message": {"type": "string"}, "priority": {"type": "number"}, }, "required": ["message"], }, }, } ] @pytest.mark.asyncio async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn(): 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=[], max_event_loops=1), ) ticks = iter([1.0, 1.123, 1.456]) runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks)) outputs = [message async for message in runtime.run(request)] assert [message["type"] for message in outputs] == [ "session_started", "message_delta", "usage", "round_stats", "done", ] assert outputs[3] == { "type": "round_stats", "round_index": 1, "ttft_ms": 123, "elapsed_ms": 456, "prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30, "cached_tokens": 5, "had_event": False, } @pytest.mark.asyncio async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff(): 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=2), ) ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2]) client = EventRoundStatsChatClient() runtime = DebugRuntime(client, clock=lambda: next(ticks)) outputs = [message async for message in runtime.run(request)] stats = [message for message in outputs if message["type"] == "round_stats"] assert client.calls == 2 assert stats == [ { "type": "round_stats", "round_index": 1, "ttft_ms": None, "elapsed_ms": 250, "prompt_tokens": 3, "completion_tokens": 0, "total_tokens": 3, "cached_tokens": 0, "had_event": True, }, { "type": "round_stats", "round_index": 2, "ttft_ms": 50, "elapsed_ms": 200, "prompt_tokens": 4, "completion_tokens": 6, "total_tokens": 10, "cached_tokens": 0, "had_event": False, }, ]