| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304 |
- 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,
- ) -> AsyncIterator[StreamItem]:
- self.calls.append(
- {
- "messages": list(messages),
- "tools": list(tools),
- "params": params,
- }
- )
- 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.event(
- 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,
- ) -> AsyncIterator[StreamItem]:
- self.calls.append(
- {
- "messages": list(messages),
- "tools": list(tools),
- "params": params,
- }
- )
- yield StreamItem.message_delta("I should have called the tool.")
- @pytest.mark.asyncio
- async def test_event_agent_resolves_tool_arguments_with_llm_tool_call():
- 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": "LLM generated handoff",
- }
- assert chat_client.calls[0]["tools"] == [
- {
- "type": "function",
- "function": {
- "name": "handoff_note",
- "description": "Send a note to the event agent.",
- "parameters": {
- "type": "object",
- "properties": {
- "message": {"type": "string"},
- },
- "required": ["message"],
- },
- },
- }
- ]
- assert chat_client.calls[0]["params"].model == "event-model"
- assert agent.raw_model_chunks([event]) == [
- {
- "event_id": "call_1",
- "event_name": "handoff_note",
- "chunks": [
- {
- "choices": [
- {
- "delta": {
- "tool_calls": [
- {
- "index": 0,
- "function": {"name": "handoff_note"},
- }
- ]
- },
- "finish_reason": None,
- }
- ]
- }
- ],
- }
- ]
- @pytest.mark.asyncio
- async def test_event_agent_falls_back_to_context_arguments_when_llm_returns_no_tool_call():
- 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[0]["tools"][0]["function"]["name"] == "mock_search"
- @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"},
- 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",
- }
|