test_event_agent.py 11 KB

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