test_event_agent.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418
  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. tool_choice: dict | None = None,
  19. ) -> AsyncIterator[StreamItem]:
  20. self.calls.append(
  21. {
  22. "messages": list(messages),
  23. "tools": list(tools),
  24. "params": params,
  25. "tool_choice": tool_choice,
  26. }
  27. )
  28. tool_name = tools[0]["function"]["name"]
  29. yield StreamItem.raw_response_chunk(
  30. {
  31. "choices": [
  32. {
  33. "delta": {
  34. "tool_calls": [
  35. {
  36. "index": 0,
  37. "function": {"name": tool_name},
  38. }
  39. ]
  40. },
  41. "finish_reason": None,
  42. }
  43. ]
  44. }
  45. )
  46. yield StreamItem.provider_tool_call(
  47. ToolCallEvent(
  48. id="llm_call_1",
  49. name=tool_name,
  50. arguments=self.arguments,
  51. raw_arguments=json.dumps(self.arguments),
  52. )
  53. )
  54. class NoToolCallChatClient:
  55. def __init__(self) -> None:
  56. self.calls: list[dict] = []
  57. async def stream_chat(
  58. self,
  59. messages: list[ChatMessage],
  60. tools: list[dict],
  61. params: AgentParams,
  62. tool_choice: dict | None = None,
  63. ) -> AsyncIterator[StreamItem]:
  64. self.calls.append(
  65. {
  66. "messages": list(messages),
  67. "tools": list(tools),
  68. "params": params,
  69. "tool_choice": tool_choice,
  70. }
  71. )
  72. yield StreamItem.message_delta("I should have called the tool.")
  73. class WrongSourceToolCallChatClient:
  74. def __init__(self, item: StreamItem) -> None:
  75. self.item = item
  76. async def stream_chat(
  77. self,
  78. messages: list[ChatMessage],
  79. tools: list[dict],
  80. params: AgentParams,
  81. tool_choice: dict | None = None,
  82. ) -> AsyncIterator[StreamItem]:
  83. yield self.item
  84. @pytest.mark.asyncio
  85. async def test_event_agent_resolves_tool_arguments_with_llm_tool_call():
  86. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  87. agent = EventAgent(
  88. enabled_tools=["handoff_note"],
  89. chat_client=chat_client,
  90. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  91. )
  92. event = ToolCallEvent(
  93. id="call_1",
  94. name="handoff_note",
  95. arguments={"message": "chat agent argument should be ignored"},
  96. raw_arguments='{"message":"chat agent argument should be ignored"}',
  97. )
  98. history = [
  99. ChatMessage(role="user", content="debug this event flow"),
  100. ChatMessage(role="assistant", content="I need the event agent."),
  101. ]
  102. reply = await agent.handle(event, history=history)
  103. assert reply.role == "tool"
  104. assert reply.tool_call_id == "call_1"
  105. assert reply.name == "handoff_note"
  106. assert json.loads(reply.content) == {
  107. "tool": "handoff_note",
  108. "message": "LLM generated handoff",
  109. }
  110. assert chat_client.calls[0]["tools"] == [
  111. {
  112. "type": "function",
  113. "function": {
  114. "name": "handoff_note",
  115. "description": "Send a note to the event agent.",
  116. "parameters": {
  117. "type": "object",
  118. "properties": {
  119. "message": {"type": "string"},
  120. },
  121. "required": ["message"],
  122. },
  123. },
  124. }
  125. ]
  126. assert chat_client.calls[0]["tool_choice"] == {
  127. "type": "function",
  128. "function": {"name": "handoff_note"},
  129. }
  130. assert chat_client.calls[0]["params"].model == "event-model"
  131. assert agent.raw_model_chunks([event]) == [
  132. {
  133. "event_id": "call_1",
  134. "event_name": "handoff_note",
  135. "chunks": [
  136. {
  137. "choices": [
  138. {
  139. "delta": {
  140. "tool_calls": [
  141. {
  142. "index": 0,
  143. "function": {"name": "handoff_note"},
  144. }
  145. ]
  146. },
  147. "finish_reason": None,
  148. }
  149. ]
  150. }
  151. ],
  152. }
  153. ]
  154. @pytest.mark.asyncio
  155. async def test_event_agent_falls_back_to_context_arguments_when_llm_returns_no_tool_call():
  156. chat_client = NoToolCallChatClient()
  157. agent = EventAgent(
  158. enabled_tools=["mock_search"],
  159. chat_client=chat_client,
  160. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  161. )
  162. event = ToolCallEvent(
  163. id="call_1",
  164. name="mock_search",
  165. arguments={},
  166. raw_arguments="{}",
  167. )
  168. history = [
  169. ChatMessage(role="user", content="Find latency docs"),
  170. ChatMessage(role="assistant", content="Need a search for latency docs"),
  171. ]
  172. reply = await agent.handle(event, history=history)
  173. payload = json.loads(reply.content)
  174. assert payload["tool"] == "mock_search"
  175. assert payload["query"] == "Need a search for latency docs"
  176. assert "event agent did not return arguments" not in reply.content
  177. assert chat_client.calls[0]["tools"][0]["function"]["name"] == "mock_search"
  178. @pytest.mark.asyncio
  179. @pytest.mark.parametrize(
  180. "item",
  181. [
  182. StreamItem.text_event(
  183. ToolCallEvent(
  184. id="text_event_1",
  185. name="mock_search",
  186. arguments={"query": "wrong source"},
  187. raw_arguments='{"query":"wrong source"}',
  188. )
  189. ),
  190. StreamItem.event(
  191. ToolCallEvent(
  192. id="legacy_event_1",
  193. name="mock_search",
  194. arguments={"query": "legacy arguments"},
  195. raw_arguments='{"query":"legacy arguments"}',
  196. )
  197. ),
  198. ],
  199. ids=["text_event", "legacy_event"],
  200. )
  201. async def test_event_agent_ignores_non_provider_tool_call_sources(item: StreamItem):
  202. event = ToolCallEvent(
  203. id="call_1",
  204. name="mock_search",
  205. arguments={},
  206. raw_arguments="{}",
  207. )
  208. agent = EventAgent(
  209. enabled_tools=["mock_search"],
  210. chat_client=WrongSourceToolCallChatClient(item),
  211. )
  212. reply = await agent.handle(
  213. event,
  214. history=[ChatMessage(role="assistant", content="fallback query")],
  215. )
  216. payload = json.loads(reply.content)
  217. assert payload["tool"] == "mock_search"
  218. assert payload["query"] == "fallback query"
  219. @pytest.mark.asyncio
  220. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  221. registry = ToolRegistry(
  222. [
  223. ToolDefinition(
  224. name="handoff_note",
  225. description="Send a note to the event agent.",
  226. parameters={"type": "object"},
  227. handler=lambda event: {"tool": event.name, "message": "handled"},
  228. )
  229. ]
  230. )
  231. disabled_reply = await EventAgent(
  232. enabled_tools=[],
  233. registry=registry,
  234. ).handle(
  235. ToolCallEvent(
  236. id="call_1",
  237. name="handoff_note",
  238. arguments={"message": "inspect this event"},
  239. raw_arguments='{"message":"inspect this event"}',
  240. )
  241. )
  242. unknown_reply = await EventAgent(
  243. enabled_tools=["missing_tool"],
  244. registry=registry,
  245. ).handle(
  246. ToolCallEvent(
  247. id="call_2",
  248. name="missing_tool",
  249. arguments={},
  250. raw_arguments="{}",
  251. )
  252. )
  253. assert json.loads(disabled_reply.content) == {
  254. "tool": "handoff_note",
  255. "error": "tool disabled",
  256. }
  257. assert json.loads(unknown_reply.content) == {
  258. "tool": "missing_tool",
  259. "error": "unknown tool",
  260. }
  261. @pytest.mark.asyncio
  262. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  263. def fail_tool(event: ToolCallEvent) -> dict:
  264. raise RuntimeError("boom")
  265. registry = ToolRegistry(
  266. [
  267. ToolDefinition(
  268. name="handoff_note",
  269. description="Send a note to the event agent.",
  270. parameters={"type": "object"},
  271. handler=fail_tool,
  272. )
  273. ]
  274. )
  275. event = ToolCallEvent(
  276. id="call_1",
  277. name="handoff_note",
  278. arguments={"message": "inspect this event"},
  279. raw_arguments='{"message":"inspect this event"}',
  280. )
  281. reply = await EventAgent(
  282. enabled_tools=["handoff_note"],
  283. registry=registry,
  284. ).handle(event)
  285. assert reply.role == "tool"
  286. assert reply.tool_call_id == "call_1"
  287. assert json.loads(reply.content) == {
  288. "tool": "handoff_note",
  289. "error": "tool handler failed: boom",
  290. }
  291. @pytest.mark.asyncio
  292. async def test_event_agent_llm_receives_history_and_agent_config_context():
  293. chat_client = ToolCallingChatClient(
  294. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  295. )
  296. registry = ToolRegistry(
  297. [
  298. ToolDefinition(
  299. name="handoff_note",
  300. description="Send a note to the event agent.",
  301. parameters={"type": "object"},
  302. handler=lambda event: {
  303. "tool": event.name,
  304. "message": event.arguments["message"],
  305. "thinking": event.arguments["thinking"],
  306. },
  307. )
  308. ]
  309. )
  310. event = ToolCallEvent(
  311. id="call_1",
  312. name="handoff_note",
  313. arguments={},
  314. raw_arguments="{}",
  315. )
  316. reply = await EventAgent(
  317. enabled_tools=["handoff_note"],
  318. registry=registry,
  319. chat_client=chat_client,
  320. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  321. ).handle(
  322. event,
  323. history=[ChatMessage(role="user", content="debug this")],
  324. system_prompt="Use strict tool parameters.",
  325. extra_body={"thinking": {"type": "disabled"}},
  326. )
  327. messages = chat_client.calls[0]["messages"]
  328. assert any(message.content == "debug this" for message in messages)
  329. assert any(message.content == "Use strict tool parameters." for message in messages)
  330. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  331. assert json.loads(reply.content) == {
  332. "tool": "handoff_note",
  333. "message": "Use strict tool parameters.",
  334. "thinking": "disabled",
  335. }
  336. @pytest.mark.asyncio
  337. async def test_event_agent_projects_complete_tool_round_to_visible_history():
  338. chat_client = ToolCallingChatClient({"message": "projected history"})
  339. event = ToolCallEvent(
  340. id="call_current",
  341. name="handoff_note",
  342. arguments={},
  343. raw_arguments="{}",
  344. )
  345. prior_call = ToolCallEvent(
  346. id="call_prior",
  347. name="mock_search",
  348. arguments={"query": "latency"},
  349. raw_arguments='{"query":"latency"}',
  350. )
  351. await EventAgent(
  352. enabled_tools=["handoff_note"],
  353. chat_client=chat_client,
  354. ).handle(
  355. event,
  356. history=[
  357. ChatMessage(role="user", content="Find latency docs"),
  358. ChatMessage(
  359. role="assistant",
  360. content="I checked the latency sources.",
  361. tool_calls=[prior_call],
  362. ),
  363. ChatMessage(
  364. role="tool",
  365. content='{"results":["doc"]}',
  366. tool_call_id="call_prior",
  367. ),
  368. ChatMessage(role="user", content="Prepare a handoff"),
  369. ],
  370. )
  371. projected_messages = chat_client.calls[0]["messages"]
  372. visible_assistant = next(
  373. message
  374. for message in projected_messages
  375. if message.content == "I checked the latency sources."
  376. )
  377. assert visible_assistant.tool_calls == []
  378. assert not any(message.role == "tool" for message in projected_messages)