test_event_agent.py 6.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216
  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. @pytest.mark.asyncio
  36. async def test_event_agent_resolves_tool_arguments_with_llm_tool_call():
  37. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  38. agent = EventAgent(
  39. enabled_tools=["handoff_note"],
  40. chat_client=chat_client,
  41. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  42. )
  43. event = ToolCallEvent(
  44. id="call_1",
  45. name="handoff_note",
  46. arguments={"message": "chat agent argument should be ignored"},
  47. raw_arguments='{"message":"chat agent argument should be ignored"}',
  48. )
  49. history = [
  50. ChatMessage(role="user", content="debug this event flow"),
  51. ChatMessage(role="assistant", content="I need the event agent."),
  52. ]
  53. reply = await agent.handle(event, history=history)
  54. assert reply.role == "tool"
  55. assert reply.tool_call_id == "call_1"
  56. assert reply.name == "handoff_note"
  57. assert json.loads(reply.content) == {
  58. "tool": "handoff_note",
  59. "message": "LLM generated handoff",
  60. }
  61. assert chat_client.calls[0]["tools"] == [
  62. {
  63. "type": "function",
  64. "function": {
  65. "name": "handoff_note",
  66. "description": "Send a note to the event agent.",
  67. "parameters": {
  68. "type": "object",
  69. "properties": {
  70. "message": {"type": "string"},
  71. },
  72. "required": ["message"],
  73. },
  74. },
  75. }
  76. ]
  77. assert chat_client.calls[0]["params"].model == "event-model"
  78. @pytest.mark.asyncio
  79. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  80. registry = ToolRegistry(
  81. [
  82. ToolDefinition(
  83. name="handoff_note",
  84. description="Send a note to the event agent.",
  85. parameters={"type": "object"},
  86. handler=lambda event: {"tool": event.name, "message": "handled"},
  87. )
  88. ]
  89. )
  90. disabled_reply = await EventAgent(
  91. enabled_tools=[],
  92. registry=registry,
  93. ).handle(
  94. ToolCallEvent(
  95. id="call_1",
  96. name="handoff_note",
  97. arguments={"message": "inspect this event"},
  98. raw_arguments='{"message":"inspect this event"}',
  99. )
  100. )
  101. unknown_reply = await EventAgent(
  102. enabled_tools=["missing_tool"],
  103. registry=registry,
  104. ).handle(
  105. ToolCallEvent(
  106. id="call_2",
  107. name="missing_tool",
  108. arguments={},
  109. raw_arguments="{}",
  110. )
  111. )
  112. assert json.loads(disabled_reply.content) == {
  113. "tool": "handoff_note",
  114. "error": "tool disabled",
  115. }
  116. assert json.loads(unknown_reply.content) == {
  117. "tool": "missing_tool",
  118. "error": "unknown tool",
  119. }
  120. @pytest.mark.asyncio
  121. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  122. def fail_tool(event: ToolCallEvent) -> dict:
  123. raise RuntimeError("boom")
  124. registry = ToolRegistry(
  125. [
  126. ToolDefinition(
  127. name="handoff_note",
  128. description="Send a note to the event agent.",
  129. parameters={"type": "object"},
  130. handler=fail_tool,
  131. )
  132. ]
  133. )
  134. event = ToolCallEvent(
  135. id="call_1",
  136. name="handoff_note",
  137. arguments={"message": "inspect this event"},
  138. raw_arguments='{"message":"inspect this event"}',
  139. )
  140. reply = await EventAgent(
  141. enabled_tools=["handoff_note"],
  142. registry=registry,
  143. ).handle(event)
  144. assert reply.role == "tool"
  145. assert reply.tool_call_id == "call_1"
  146. assert json.loads(reply.content) == {
  147. "tool": "handoff_note",
  148. "error": "tool handler failed: boom",
  149. }
  150. @pytest.mark.asyncio
  151. async def test_event_agent_llm_receives_history_and_agent_config_context():
  152. chat_client = ToolCallingChatClient(
  153. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  154. )
  155. registry = ToolRegistry(
  156. [
  157. ToolDefinition(
  158. name="handoff_note",
  159. description="Send a note to the event agent.",
  160. parameters={"type": "object"},
  161. handler=lambda event: {
  162. "tool": event.name,
  163. "message": event.arguments["message"],
  164. "thinking": event.arguments["thinking"],
  165. },
  166. )
  167. ]
  168. )
  169. event = ToolCallEvent(
  170. id="call_1",
  171. name="handoff_note",
  172. arguments={},
  173. raw_arguments="{}",
  174. )
  175. reply = await EventAgent(
  176. enabled_tools=["handoff_note"],
  177. registry=registry,
  178. chat_client=chat_client,
  179. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  180. ).handle(
  181. event,
  182. history=[ChatMessage(role="user", content="debug this")],
  183. system_prompt="Use strict tool parameters.",
  184. extra_body={"thinking": {"type": "disabled"}},
  185. )
  186. messages = chat_client.calls[0]["messages"]
  187. assert any(message.content == "debug this" for message in messages)
  188. assert any(message.content == "Use strict tool parameters." for message in messages)
  189. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  190. assert json.loads(reply.content) == {
  191. "tool": "handoff_note",
  192. "message": "Use strict tool parameters.",
  193. "thinking": "disabled",
  194. }