| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577 |
- 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)
|