import asyncio import importlib import json from collections.abc import AsyncIterator 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 from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage def _runtime_queues_class(): module = importlib.import_module("agent_lab.application.queues") return module.RuntimeQueues async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]: return [message async for message in stream] 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, 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__ 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={}, raw_arguments="{}", ) ) 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={}, raw_arguments="{}", ) ) return self.second_call_messages = list(messages) 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]] = [] 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 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, 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={}, raw_arguments="{}", ) ) 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_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( 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), ) runtime = DebugRuntime(RoundStatsChatClient()) queues = runtime.start(request) outputs: list[dict[str, Any]] = [] while True: message = await asyncio.wait_for(queues.output.get(), timeout=1) outputs.append(message) if message["type"] == "done": break assert [message["type"] for message in outputs] == [ "session_started", "message_delta", "usage", "round_stats", "done", ] @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() queues = RuntimeQueues() 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 = StrictHistoryChatClient() runtime = DebugRuntime(client, queues=queues) stream = runtime.run(request) assert await anext(stream) == {"type": "session_started"} event_message = await anext(stream) assert event_message["type"] == "event" await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs")) remaining = [message async for message in stream] assert remaining[-1] == {"type": "done"} assert [message.role for message in client.second_call_messages] == [ "user", "assistant", "tool", "user", ] assert client.second_call_messages[2].tool_call_id == "call_1" assert client.second_call_messages[3].content == "follow-up while tool runs" @pytest.mark.asyncio async def test_runtime_continues_when_event_agent_tool_handler_raises(): def fail_tool(event: ToolCallEvent) -> dict[str, Any]: raise RuntimeError("boom") registry = ToolRegistry( [ ToolDefinition( name="handoff_note", description="Broken handoff tool.", parameters={"type": "object"}, handler=fail_tool, ) ] ) 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), ) runtime = DebugRuntime(FakeChatClient(), registry=registry) outputs = await asyncio.wait_for( _collect_outputs(runtime.run(request)), timeout=1, ) assert [message["type"] for message in outputs] == [ "session_started", "event", "tool_result", "round_stats", "message_delta", "round_stats", "done", ] assert json.loads(outputs[2]["message"]["content"]) == { "tool": "handoff_note", "error": "tool handler failed: boom", } @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_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") ) @pytest.mark.asyncio async def test_runtime_run_consumes_output_queue_in_stream_order(): 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", ] output_puts = [ entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put" ] output_gets = [ entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get" ] assert output_puts == [ "output:session_started", "output:event", "output:tool_result", "output:round_stats", "output:message_delta", "output:round_stats", "output:done", ] assert output_gets == output_puts @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": "{}", }, } ] assert tool_message.tool_call_id == "call_1" assert outputs[-1] == {"type": "done"} @pytest.mark.asyncio async def test_runtime_passes_event_names_without_tool_parameters_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": {}, "additionalProperties": False, }, }, } ] @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, }, ]