import asyncio import json from collections.abc import AsyncIterator import pytest from agent_lab.application.contracts import AgentParams from agent_lab.application.event_agent import EventAgent from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry from agent_lab.domain.events import ToolCallEvent from agent_lab.domain.messages import ChatMessage, StreamItem class ToolCallingChatClient: def __init__(self, arguments: dict) -> None: self.arguments = arguments self.calls: list[dict] = [] async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, tool_choice: dict | None = None, ) -> AsyncIterator[StreamItem]: self.calls.append( { "messages": list(messages), "tools": list(tools), "params": params, "tool_choice": tool_choice, } ) tool_name = tools[0]["function"]["name"] yield StreamItem.raw_response_chunk( { "choices": [ { "delta": { "tool_calls": [ { "index": 0, "function": {"name": tool_name}, } ] }, "finish_reason": None, } ] } ) yield StreamItem.provider_tool_call( ToolCallEvent( id="llm_call_1", name=tool_name, arguments=self.arguments, raw_arguments=json.dumps(self.arguments), ) ) class NoToolCallChatClient: def __init__(self) -> None: self.calls: list[dict] = [] async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, tool_choice: dict | None = None, ) -> AsyncIterator[StreamItem]: self.calls.append( { "messages": list(messages), "tools": list(tools), "params": params, "tool_choice": tool_choice, } ) yield StreamItem.message_delta("I should have called the tool.") class WrongSourceToolCallChatClient: def __init__(self, item: StreamItem) -> None: self.item = item async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, tool_choice: dict | None = None, ) -> AsyncIterator[StreamItem]: yield self.item class WrongNameToolCallingChatClient(ToolCallingChatClient): async def stream_chat( self, messages: list[ChatMessage], tools: list[dict], params: AgentParams, tool_choice: dict | None = None, ) -> AsyncIterator[StreamItem]: self.calls.append({"messages": list(messages), "tools": list(tools)}) yield StreamItem.provider_tool_call( ToolCallEvent( id="wrong-call", name="wrong.tool", arguments=self.arguments, raw_arguments=json.dumps(self.arguments), ) ) @pytest.mark.asyncio async def test_event_agent_uses_deterministic_resolver_before_llm_fallback(): chat_client = ToolCallingChatClient({"message": "LLM generated handoff"}) agent = EventAgent( enabled_tools=["handoff_note"], chat_client=chat_client, params=AgentParams(model="event-model", temperature=0, max_tokens=80), ) event = ToolCallEvent( id="call_1", name="handoff_note", arguments={"message": "chat agent argument should be ignored"}, raw_arguments='{"message":"chat agent argument should be ignored"}', ) history = [ ChatMessage(role="user", content="debug this event flow"), ChatMessage(role="assistant", content="I need the event agent."), ] reply = await agent.handle(event, history=history) assert reply.role == "tool" assert reply.tool_call_id == "call_1" assert reply.name == "handoff_note" assert json.loads(reply.content) == { "tool": "handoff_note", "message": "I need the event agent.", } assert chat_client.calls == [] assert agent.raw_model_chunks([event]) == [ { "event_id": "call_1", "event_name": "handoff_note", "chunks": [], } ] @pytest.mark.asyncio async def test_event_agent_deterministic_mock_resolver_skips_llm(): chat_client = NoToolCallChatClient() agent = EventAgent( enabled_tools=["mock_search"], chat_client=chat_client, params=AgentParams(model="event-model", temperature=0, max_tokens=80), ) event = ToolCallEvent( id="call_1", name="mock_search", arguments={}, raw_arguments="{}", ) history = [ ChatMessage(role="user", content="Find latency docs"), ChatMessage(role="assistant", content="Need a search for latency docs"), ] reply = await agent.handle(event, history=history) payload = json.loads(reply.content) assert payload["tool"] == "mock_search" assert payload["query"] == "Need a search for latency docs" assert "event agent did not return arguments" not in reply.content assert chat_client.calls == [] @pytest.mark.asyncio async def test_event_agent_calls_llm_once_for_incomplete_deterministic_arguments(): chat_client = ToolCallingChatClient({"message": "LLM generated handoff"}) registry = ToolRegistry( [ ToolDefinition( name="ambiguous_handoff", description="Resolve an ambiguous handoff.", parameters={ "type": "object", "properties": {"message": {"type": "string"}}, "required": ["message"], }, handler=lambda event: { "tool": event.name, "message": event.arguments["message"], }, argument_resolver=lambda event, context: {}, ) ] ) event = ToolCallEvent( id="call_1", name="ambiguous_handoff", arguments={}, raw_arguments="{}", ) reply = await EventAgent( enabled_tools=["ambiguous_handoff"], registry=registry, chat_client=chat_client, ).handle(event, history=[ChatMessage(role="user", content="ambiguous")]) assert json.loads(reply.content) == { "tool": "ambiguous_handoff", "message": "LLM generated handoff", } assert len(chat_client.calls) == 1 assert chat_client.calls[0]["tool_choice"] == { "type": "function", "function": {"name": "ambiguous_handoff"}, } @pytest.mark.asyncio async def test_event_agent_rejects_fallback_provider_call_for_wrong_tool_name(): chat_client = WrongNameToolCallingChatClient({"message": "wrong"}) registry = ToolRegistry( [ ToolDefinition( name="expected.tool", description="Expected tool.", parameters={ "type": "object", "properties": {"message": {"type": "string"}}, "required": ["message"], }, handler=lambda event: {"tool": event.name}, argument_resolver=lambda event, context: {}, ) ] ) reply = await EventAgent( enabled_tools=["expected.tool"], registry=registry, chat_client=chat_client, ).handle( ToolCallEvent( id="call-1", name="expected.tool", arguments={}, raw_arguments="{}", ) ) assert json.loads(reply.content) == { "tool": "expected.tool", "error": "fallback returned tool wrong.tool for expected.tool", } @pytest.mark.asyncio async def test_event_agent_handle_many_starts_independent_handlers_concurrently(): started: list[str] = [] both_started = asyncio.Event() async def handler(event: ToolCallEvent) -> dict: started.append(event.name) if len(started) == 2: both_started.set() await asyncio.wait_for(both_started.wait(), timeout=0.2) return {"tool": event.name} registry = ToolRegistry( [ ToolDefinition( name=name, description=f"Handle {name}.", parameters={"type": "object"}, handler=handler, ) for name in ("event.first", "event.second") ] ) events = [ ToolCallEvent(id=f"call-{index}", name=name, arguments={}, raw_arguments="{}") for index, name in enumerate(("event.first", "event.second"), start=1) ] replies = await asyncio.wait_for( EventAgent( enabled_tools=[event.name for event in events], registry=registry, ).handle_many(events, history=[]), timeout=0.5, ) assert started == ["event.first", "event.second"] assert [reply.name for reply in replies] == ["event.first", "event.second"] assert [json.loads(reply.content) for reply in replies] == [ {"tool": "event.first"}, {"tool": "event.second"}, ] @pytest.mark.asyncio @pytest.mark.parametrize( "item", [ StreamItem.text_event( ToolCallEvent( id="text_event_1", name="mock_search", arguments={"query": "wrong source"}, raw_arguments='{"query":"wrong source"}', ) ), StreamItem.event( ToolCallEvent( id="legacy_event_1", name="mock_search", arguments={"query": "legacy arguments"}, raw_arguments='{"query":"legacy arguments"}', ) ), ], ids=["text_event", "legacy_event"], ) async def test_event_agent_ignores_non_provider_tool_call_sources(item: StreamItem): event = ToolCallEvent( id="call_1", name="mock_search", arguments={}, raw_arguments="{}", ) registry = ToolRegistry( [ ToolDefinition( name="mock_search", description="Search mock external knowledge for the current turn.", parameters={ "type": "object", "properties": {"query": {"type": "string"}}, "required": ["query"], }, handler=lambda resolved: { "tool": resolved.name, "query": resolved.arguments["query"], }, argument_resolver=lambda event, context: {}, ) ] ) agent = EventAgent( enabled_tools=["mock_search"], registry=registry, chat_client=WrongSourceToolCallChatClient(item), ) reply = await agent.handle( event, history=[ChatMessage(role="assistant", content="fallback query")], ) payload = json.loads(reply.content) assert payload["tool"] == "mock_search" assert payload == { "tool": "mock_search", "error": "missing required arguments: query", } @pytest.mark.asyncio async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools(): registry = ToolRegistry( [ ToolDefinition( name="handoff_note", description="Send a note to the event agent.", parameters={"type": "object"}, handler=lambda event: {"tool": event.name, "message": "handled"}, ) ] ) disabled_reply = await EventAgent( enabled_tools=[], registry=registry, ).handle( ToolCallEvent( id="call_1", name="handoff_note", arguments={"message": "inspect this event"}, raw_arguments='{"message":"inspect this event"}', ) ) unknown_reply = await EventAgent( enabled_tools=["missing_tool"], registry=registry, ).handle( ToolCallEvent( id="call_2", name="missing_tool", arguments={}, raw_arguments="{}", ) ) assert json.loads(disabled_reply.content) == { "tool": "handoff_note", "error": "tool disabled", } assert json.loads(unknown_reply.content) == { "tool": "missing_tool", "error": "unknown tool", } @pytest.mark.asyncio async def test_event_agent_returns_structured_error_when_tool_handler_raises(): def fail_tool(event: ToolCallEvent) -> dict: raise RuntimeError("boom") registry = ToolRegistry( [ ToolDefinition( name="handoff_note", description="Send a note to the event agent.", parameters={"type": "object"}, handler=fail_tool, ) ] ) event = ToolCallEvent( id="call_1", name="handoff_note", arguments={"message": "inspect this event"}, raw_arguments='{"message":"inspect this event"}', ) reply = await EventAgent( enabled_tools=["handoff_note"], registry=registry, ).handle(event) assert reply.role == "tool" assert reply.tool_call_id == "call_1" assert json.loads(reply.content) == { "tool": "handoff_note", "error": "tool handler failed: boom", } @pytest.mark.asyncio async def test_event_agent_llm_receives_history_and_agent_config_context(): chat_client = ToolCallingChatClient( {"message": "Use strict tool parameters.", "thinking": "disabled"} ) registry = ToolRegistry( [ ToolDefinition( name="handoff_note", description="Send a note to the event agent.", parameters={ "type": "object", "properties": { "message": {"type": "string"}, "thinking": {"type": "string"}, }, "required": ["message", "thinking"], }, handler=lambda event: { "tool": event.name, "message": event.arguments["message"], "thinking": event.arguments["thinking"], }, ) ] ) event = ToolCallEvent( id="call_1", name="handoff_note", arguments={}, raw_arguments="{}", ) reply = await EventAgent( enabled_tools=["handoff_note"], registry=registry, chat_client=chat_client, params=AgentParams(model="event-model", temperature=0.4, max_tokens=120), ).handle( event, history=[ChatMessage(role="user", content="debug this")], system_prompt="Use strict tool parameters.", extra_body={"thinking": {"type": "disabled"}}, ) messages = chat_client.calls[0]["messages"] assert any(message.content == "debug this" for message in messages) assert any(message.content == "Use strict tool parameters." for message in messages) assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}} assert json.loads(reply.content) == { "tool": "handoff_note", "message": "Use strict tool parameters.", "thinking": "disabled", } @pytest.mark.asyncio async def test_event_agent_projects_complete_tool_round_to_visible_history(): chat_client = ToolCallingChatClient({"message": "projected history"}) registry = ToolRegistry( [ ToolDefinition( name="handoff_note", description="Send a note to the event agent.", parameters={ "type": "object", "properties": {"message": {"type": "string"}}, "required": ["message"], }, handler=lambda resolved: { "tool": resolved.name, "message": resolved.arguments["message"], }, argument_resolver=lambda event, context: {}, ) ] ) event = ToolCallEvent( id="call_current", name="handoff_note", arguments={}, raw_arguments="{}", ) prior_call = ToolCallEvent( id="call_prior", name="mock_search", arguments={"query": "latency"}, raw_arguments='{"query":"latency"}', ) await EventAgent( enabled_tools=["handoff_note"], registry=registry, chat_client=chat_client, ).handle( event, history=[ ChatMessage(role="user", content="Find latency docs"), ChatMessage( role="assistant", content="I checked the latency sources.", tool_calls=[prior_call], ), ChatMessage( role="tool", content='{"results":["doc"]}', tool_call_id="call_prior", ), ChatMessage(role="user", content="Prepare a handoff"), ], ) projected_messages = chat_client.calls[0]["messages"] visible_assistant = next( message for message in projected_messages if message.content == "I checked the latency sources." ) assert visible_assistant.tool_calls == [] assert not any(message.role == "tool" for message in projected_messages)