test_event_agent.py 8.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264
  1. import json
  2. from collections.abc import AsyncIterator
  3. import pytest
  4. from agent_lab.application.contracts import AgentParams
  5. from agent_lab.application.event_agent import EventAgent
  6. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  7. from agent_lab.domain.events import ToolCallEvent
  8. from agent_lab.domain.messages import ChatMessage, StreamItem
  9. class ToolCallingChatClient:
  10. def __init__(self, arguments: dict) -> None:
  11. self.arguments = arguments
  12. self.calls: list[dict] = []
  13. async def stream_chat(
  14. self,
  15. messages: list[ChatMessage],
  16. tools: list[dict],
  17. params: AgentParams,
  18. ) -> AsyncIterator[StreamItem]:
  19. self.calls.append(
  20. {
  21. "messages": list(messages),
  22. "tools": list(tools),
  23. "params": params,
  24. }
  25. )
  26. tool_name = tools[0]["function"]["name"]
  27. yield StreamItem.event(
  28. ToolCallEvent(
  29. id="llm_call_1",
  30. name=tool_name,
  31. arguments=self.arguments,
  32. raw_arguments=json.dumps(self.arguments),
  33. )
  34. )
  35. class NoToolCallChatClient:
  36. def __init__(self) -> None:
  37. self.calls: list[dict] = []
  38. async def stream_chat(
  39. self,
  40. messages: list[ChatMessage],
  41. tools: list[dict],
  42. params: AgentParams,
  43. ) -> AsyncIterator[StreamItem]:
  44. self.calls.append(
  45. {
  46. "messages": list(messages),
  47. "tools": list(tools),
  48. "params": params,
  49. }
  50. )
  51. yield StreamItem.message_delta("I should have called the tool.")
  52. @pytest.mark.asyncio
  53. async def test_event_agent_resolves_tool_arguments_with_llm_tool_call():
  54. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  55. agent = EventAgent(
  56. enabled_tools=["handoff_note"],
  57. chat_client=chat_client,
  58. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  59. )
  60. event = ToolCallEvent(
  61. id="call_1",
  62. name="handoff_note",
  63. arguments={"message": "chat agent argument should be ignored"},
  64. raw_arguments='{"message":"chat agent argument should be ignored"}',
  65. )
  66. history = [
  67. ChatMessage(role="user", content="debug this event flow"),
  68. ChatMessage(role="assistant", content="I need the event agent."),
  69. ]
  70. reply = await agent.handle(event, history=history)
  71. assert reply.role == "tool"
  72. assert reply.tool_call_id == "call_1"
  73. assert reply.name == "handoff_note"
  74. assert json.loads(reply.content) == {
  75. "tool": "handoff_note",
  76. "message": "LLM generated handoff",
  77. }
  78. assert chat_client.calls[0]["tools"] == [
  79. {
  80. "type": "function",
  81. "function": {
  82. "name": "handoff_note",
  83. "description": "Send a note to the event agent.",
  84. "parameters": {
  85. "type": "object",
  86. "properties": {
  87. "message": {"type": "string"},
  88. },
  89. "required": ["message"],
  90. },
  91. },
  92. }
  93. ]
  94. assert chat_client.calls[0]["params"].model == "event-model"
  95. @pytest.mark.asyncio
  96. async def test_event_agent_falls_back_to_context_arguments_when_llm_returns_no_tool_call():
  97. chat_client = NoToolCallChatClient()
  98. agent = EventAgent(
  99. enabled_tools=["mock_search"],
  100. chat_client=chat_client,
  101. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  102. )
  103. event = ToolCallEvent(
  104. id="call_1",
  105. name="mock_search",
  106. arguments={},
  107. raw_arguments="{}",
  108. )
  109. history = [
  110. ChatMessage(role="user", content="Find latency docs"),
  111. ChatMessage(role="assistant", content="Need a search for latency docs"),
  112. ]
  113. reply = await agent.handle(event, history=history)
  114. payload = json.loads(reply.content)
  115. assert payload["tool"] == "mock_search"
  116. assert payload["query"] == "Need a search for latency docs"
  117. assert "event agent did not return arguments" not in reply.content
  118. assert chat_client.calls[0]["tools"][0]["function"]["name"] == "mock_search"
  119. @pytest.mark.asyncio
  120. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  121. registry = ToolRegistry(
  122. [
  123. ToolDefinition(
  124. name="handoff_note",
  125. description="Send a note to the event agent.",
  126. parameters={"type": "object"},
  127. handler=lambda event: {"tool": event.name, "message": "handled"},
  128. )
  129. ]
  130. )
  131. disabled_reply = await EventAgent(
  132. enabled_tools=[],
  133. registry=registry,
  134. ).handle(
  135. ToolCallEvent(
  136. id="call_1",
  137. name="handoff_note",
  138. arguments={"message": "inspect this event"},
  139. raw_arguments='{"message":"inspect this event"}',
  140. )
  141. )
  142. unknown_reply = await EventAgent(
  143. enabled_tools=["missing_tool"],
  144. registry=registry,
  145. ).handle(
  146. ToolCallEvent(
  147. id="call_2",
  148. name="missing_tool",
  149. arguments={},
  150. raw_arguments="{}",
  151. )
  152. )
  153. assert json.loads(disabled_reply.content) == {
  154. "tool": "handoff_note",
  155. "error": "tool disabled",
  156. }
  157. assert json.loads(unknown_reply.content) == {
  158. "tool": "missing_tool",
  159. "error": "unknown tool",
  160. }
  161. @pytest.mark.asyncio
  162. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  163. def fail_tool(event: ToolCallEvent) -> dict:
  164. raise RuntimeError("boom")
  165. registry = ToolRegistry(
  166. [
  167. ToolDefinition(
  168. name="handoff_note",
  169. description="Send a note to the event agent.",
  170. parameters={"type": "object"},
  171. handler=fail_tool,
  172. )
  173. ]
  174. )
  175. event = ToolCallEvent(
  176. id="call_1",
  177. name="handoff_note",
  178. arguments={"message": "inspect this event"},
  179. raw_arguments='{"message":"inspect this event"}',
  180. )
  181. reply = await EventAgent(
  182. enabled_tools=["handoff_note"],
  183. registry=registry,
  184. ).handle(event)
  185. assert reply.role == "tool"
  186. assert reply.tool_call_id == "call_1"
  187. assert json.loads(reply.content) == {
  188. "tool": "handoff_note",
  189. "error": "tool handler failed: boom",
  190. }
  191. @pytest.mark.asyncio
  192. async def test_event_agent_llm_receives_history_and_agent_config_context():
  193. chat_client = ToolCallingChatClient(
  194. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  195. )
  196. registry = ToolRegistry(
  197. [
  198. ToolDefinition(
  199. name="handoff_note",
  200. description="Send a note to the event agent.",
  201. parameters={"type": "object"},
  202. handler=lambda event: {
  203. "tool": event.name,
  204. "message": event.arguments["message"],
  205. "thinking": event.arguments["thinking"],
  206. },
  207. )
  208. ]
  209. )
  210. event = ToolCallEvent(
  211. id="call_1",
  212. name="handoff_note",
  213. arguments={},
  214. raw_arguments="{}",
  215. )
  216. reply = await EventAgent(
  217. enabled_tools=["handoff_note"],
  218. registry=registry,
  219. chat_client=chat_client,
  220. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  221. ).handle(
  222. event,
  223. history=[ChatMessage(role="user", content="debug this")],
  224. system_prompt="Use strict tool parameters.",
  225. extra_body={"thinking": {"type": "disabled"}},
  226. )
  227. messages = chat_client.calls[0]["messages"]
  228. assert any(message.content == "debug this" for message in messages)
  229. assert any(message.content == "Use strict tool parameters." for message in messages)
  230. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  231. assert json.loads(reply.content) == {
  232. "tool": "handoff_note",
  233. "message": "Use strict tool parameters.",
  234. "thinking": "disabled",
  235. }