from collections.abc import AsyncIterator 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 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") @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_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"}