test_event_agent.py 28 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928
  1. import asyncio
  2. import json
  3. from collections.abc import AsyncIterator
  4. import pytest
  5. from agent_lab.application.contracts import AgentParams, EventAgentParams
  6. from agent_lab.application.event_agent import EventAgent
  7. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  8. from agent_lab.domain.events import ToolCallEvent
  9. from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
  10. class ToolCallingChatClient:
  11. def __init__(self, arguments: dict) -> None:
  12. self.arguments = arguments
  13. self.calls: list[dict] = []
  14. async def stream_chat(
  15. self,
  16. messages: list[ChatMessage],
  17. tools: list[dict],
  18. params: AgentParams,
  19. tool_choice: dict | None = None,
  20. ) -> AsyncIterator[StreamItem]:
  21. self.calls.append(
  22. {
  23. "messages": list(messages),
  24. "tools": list(tools),
  25. "params": params,
  26. "tool_choice": tool_choice,
  27. }
  28. )
  29. tool_name = tools[0]["function"]["name"]
  30. yield StreamItem.raw_response_chunk(
  31. {
  32. "choices": [
  33. {
  34. "delta": {
  35. "tool_calls": [
  36. {
  37. "index": 0,
  38. "function": {"name": tool_name},
  39. }
  40. ]
  41. },
  42. "finish_reason": None,
  43. }
  44. ]
  45. }
  46. )
  47. yield StreamItem.provider_tool_call(
  48. ToolCallEvent(
  49. id="llm_call_1",
  50. name=tool_name,
  51. arguments=self.arguments,
  52. raw_arguments=json.dumps(self.arguments),
  53. )
  54. )
  55. class NoToolCallChatClient:
  56. def __init__(self) -> None:
  57. self.calls: list[dict] = []
  58. async def stream_chat(
  59. self,
  60. messages: list[ChatMessage],
  61. tools: list[dict],
  62. params: AgentParams,
  63. tool_choice: dict | None = None,
  64. ) -> AsyncIterator[StreamItem]:
  65. self.calls.append(
  66. {
  67. "messages": list(messages),
  68. "tools": list(tools),
  69. "params": params,
  70. "tool_choice": tool_choice,
  71. }
  72. )
  73. yield StreamItem.message_delta("I should have called the tool.")
  74. class WrongSourceToolCallChatClient:
  75. def __init__(self, item: StreamItem) -> None:
  76. self.item = item
  77. async def stream_chat(
  78. self,
  79. messages: list[ChatMessage],
  80. tools: list[dict],
  81. params: AgentParams,
  82. tool_choice: dict | None = None,
  83. ) -> AsyncIterator[StreamItem]:
  84. yield self.item
  85. class WrongNameToolCallingChatClient(ToolCallingChatClient):
  86. async def stream_chat(
  87. self,
  88. messages: list[ChatMessage],
  89. tools: list[dict],
  90. params: AgentParams,
  91. tool_choice: dict | None = None,
  92. ) -> AsyncIterator[StreamItem]:
  93. self.calls.append({"messages": list(messages), "tools": list(tools)})
  94. yield StreamItem.provider_tool_call(
  95. ToolCallEvent(
  96. id="wrong-call",
  97. name="wrong.tool",
  98. arguments=self.arguments,
  99. raw_arguments=json.dumps(self.arguments),
  100. )
  101. )
  102. class TrailingUsageChatClient:
  103. async def stream_chat(
  104. self,
  105. messages: list[ChatMessage],
  106. tools: list[dict],
  107. params: AgentParams,
  108. tool_choice: dict | None = None,
  109. ) -> AsyncIterator[StreamItem]:
  110. tool_name = tools[0]["function"]["name"]
  111. yield StreamItem.raw_response_chunk({"first": tool_name})
  112. yield StreamItem.provider_tool_call(
  113. ToolCallEvent(
  114. id=f"provider-{tool_name}",
  115. name=tool_name,
  116. arguments={"query": tool_name},
  117. raw_arguments=json.dumps({"query": tool_name}),
  118. )
  119. )
  120. yield StreamItem.usage_item(
  121. TokenUsage(
  122. prompt_tokens=3,
  123. completion_tokens=4,
  124. total_tokens=7,
  125. cached_tokens=1,
  126. )
  127. )
  128. class FailedResolutionUsageChatClient:
  129. async def stream_chat(
  130. self,
  131. messages: list[ChatMessage],
  132. tools: list[dict],
  133. params: AgentParams,
  134. tool_choice: dict | None = None,
  135. ) -> AsyncIterator[StreamItem]:
  136. yield StreamItem.message_delta("invalid fallback response")
  137. yield StreamItem.usage_item(
  138. TokenUsage(prompt_tokens=5, completion_tokens=2, total_tokens=7)
  139. )
  140. @pytest.mark.asyncio
  141. async def test_event_agent_uses_deterministic_resolver_before_llm_fallback():
  142. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  143. agent = EventAgent(
  144. enabled_tools=["handoff_note"],
  145. chat_client=chat_client,
  146. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  147. )
  148. event = ToolCallEvent(
  149. id="call_1",
  150. name="handoff_note",
  151. arguments={"message": "chat agent argument should be ignored"},
  152. raw_arguments='{"message":"chat agent argument should be ignored"}',
  153. )
  154. history = [
  155. ChatMessage(role="user", content="debug this event flow"),
  156. ChatMessage(role="assistant", content="I need the event agent."),
  157. ]
  158. reply = await agent.handle(event, history=history)
  159. assert reply.role == "tool"
  160. assert reply.tool_call_id == "call_1"
  161. assert reply.name == "handoff_note"
  162. assert json.loads(reply.content) == {
  163. "tool": "handoff_note",
  164. "message": "I need the event agent.",
  165. }
  166. assert chat_client.calls == []
  167. assert agent.raw_model_chunks([event]) == [
  168. {
  169. "event_id": "call_1",
  170. "event_name": "handoff_note",
  171. "chunks": [],
  172. }
  173. ]
  174. @pytest.mark.asyncio
  175. async def test_event_agent_deterministic_mock_resolver_skips_llm():
  176. chat_client = NoToolCallChatClient()
  177. agent = EventAgent(
  178. enabled_tools=["mock_search"],
  179. chat_client=chat_client,
  180. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  181. )
  182. event = ToolCallEvent(
  183. id="call_1",
  184. name="mock_search",
  185. arguments={},
  186. raw_arguments="{}",
  187. )
  188. history = [
  189. ChatMessage(role="user", content="Find latency docs"),
  190. ChatMessage(role="assistant", content="Need a search for latency docs"),
  191. ]
  192. reply = await agent.handle(event, history=history)
  193. payload = json.loads(reply.content)
  194. assert payload["tool"] == "mock_search"
  195. assert payload["query"] == "Need a search for latency docs"
  196. assert "event agent did not return arguments" not in reply.content
  197. assert chat_client.calls == []
  198. @pytest.mark.asyncio
  199. async def test_event_agent_injects_clock_into_event_kernel_handler_metrics():
  200. clock_values = iter([5.0, 5.030])
  201. registry = ToolRegistry(
  202. [
  203. ToolDefinition(
  204. name="timed.tool",
  205. description="Timed tool.",
  206. parameters={"type": "object"},
  207. handler=lambda event: {"tool": event.name},
  208. )
  209. ]
  210. )
  211. event = ToolCallEvent(
  212. id="timed-event",
  213. name="timed.tool",
  214. arguments={},
  215. raw_arguments="{}",
  216. )
  217. batch = await EventAgent(
  218. enabled_tools=[event.name],
  219. registry=registry,
  220. monotonic_clock=lambda: next(clock_values),
  221. ).handle_many([event], history=[])
  222. assert batch.results[0].tool_latency_ms == 30
  223. @pytest.mark.asyncio
  224. async def test_event_agent_calls_llm_once_for_incomplete_deterministic_arguments():
  225. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  226. registry = ToolRegistry(
  227. [
  228. ToolDefinition(
  229. name="ambiguous_handoff",
  230. description="Resolve an ambiguous handoff.",
  231. parameters={
  232. "type": "object",
  233. "properties": {"message": {"type": "string"}},
  234. "required": ["message"],
  235. },
  236. handler=lambda event: {
  237. "tool": event.name,
  238. "message": event.arguments["message"],
  239. },
  240. argument_resolver=lambda event, context: {},
  241. )
  242. ]
  243. )
  244. event = ToolCallEvent(
  245. id="call_1",
  246. name="ambiguous_handoff",
  247. arguments={},
  248. raw_arguments="{}",
  249. )
  250. reply = await EventAgent(
  251. enabled_tools=["ambiguous_handoff"],
  252. registry=registry,
  253. chat_client=chat_client,
  254. ).handle(event, history=[ChatMessage(role="user", content="ambiguous")])
  255. assert json.loads(reply.content) == {
  256. "tool": "ambiguous_handoff",
  257. "message": "LLM generated handoff",
  258. }
  259. assert len(chat_client.calls) == 1
  260. assert chat_client.calls[0]["tool_choice"] == {
  261. "type": "function",
  262. "function": {"name": "ambiguous_handoff"},
  263. }
  264. @pytest.mark.asyncio
  265. async def test_event_agent_consumes_trailing_usage_and_records_provider_ttft():
  266. clock_values = iter([10.0, 10.011, 10.050, 10.060, 10.080])
  267. registry = ToolRegistry(
  268. [
  269. ToolDefinition(
  270. name="fallback.search",
  271. description="Resolve with fallback.",
  272. parameters={
  273. "type": "object",
  274. "properties": {"query": {"type": "string"}},
  275. "required": ["query"],
  276. },
  277. handler=lambda event: {"query": event.arguments["query"]},
  278. argument_resolver=lambda event, context: {},
  279. )
  280. ]
  281. )
  282. event = ToolCallEvent(
  283. id="event-1",
  284. name="fallback.search",
  285. arguments={},
  286. raw_arguments="{}",
  287. )
  288. agent = EventAgent(
  289. enabled_tools=[event.name],
  290. registry=registry,
  291. chat_client=TrailingUsageChatClient(),
  292. monotonic_clock=lambda: next(clock_values),
  293. )
  294. batch = await agent.handle_many([event], history=[])
  295. assert batch.results[0].status.value == "success"
  296. assert agent.fallback_model_calls([event]) == [
  297. {
  298. "event_id": "event-1",
  299. "event_name": "fallback.search",
  300. "usage": TokenUsage(
  301. prompt_tokens=3,
  302. completion_tokens=4,
  303. total_tokens=7,
  304. cached_tokens=1,
  305. ),
  306. "ttft_ms": 11,
  307. "elapsed_ms": 50,
  308. }
  309. ]
  310. @pytest.mark.asyncio
  311. async def test_event_agent_records_failed_fallback_resolution_attempt_usage():
  312. clock_values = iter([2.0, 2.005, 2.020])
  313. registry = ToolRegistry(
  314. [
  315. ToolDefinition(
  316. name="fallback.search",
  317. description="Resolve with fallback.",
  318. parameters={
  319. "type": "object",
  320. "properties": {"query": {"type": "string"}},
  321. "required": ["query"],
  322. },
  323. handler=lambda event: {"query": event.arguments["query"]},
  324. argument_resolver=lambda event, context: {},
  325. )
  326. ]
  327. )
  328. event = ToolCallEvent(
  329. id="event-failed",
  330. name="fallback.search",
  331. arguments={},
  332. raw_arguments="{}",
  333. )
  334. agent = EventAgent(
  335. enabled_tools=[event.name],
  336. registry=registry,
  337. chat_client=FailedResolutionUsageChatClient(),
  338. monotonic_clock=lambda: next(clock_values),
  339. )
  340. batch = await agent.handle_many([event], history=[])
  341. assert batch.results[0].status.value == "invalid_arguments"
  342. assert agent.fallback_model_calls([event])[0] == {
  343. "event_id": "event-failed",
  344. "event_name": "fallback.search",
  345. "usage": TokenUsage(prompt_tokens=5, completion_tokens=2, total_tokens=7),
  346. "ttft_ms": 5,
  347. "elapsed_ms": 20,
  348. }
  349. @pytest.mark.asyncio
  350. async def test_event_agent_keeps_parallel_fallback_metrics_isolated_by_event():
  351. registry = ToolRegistry(
  352. [
  353. ToolDefinition(
  354. name=name,
  355. description="Resolve with fallback.",
  356. parameters={
  357. "type": "object",
  358. "properties": {"query": {"type": "string"}},
  359. "required": ["query"],
  360. },
  361. handler=lambda event: {"query": event.arguments["query"]},
  362. argument_resolver=lambda event, context: {},
  363. )
  364. for name in ("fallback.first", "fallback.second")
  365. ]
  366. )
  367. events = [
  368. ToolCallEvent(id=f"event-{index}", name=name, arguments={}, raw_arguments="{}")
  369. for index, name in enumerate(
  370. ("fallback.first", "fallback.second"),
  371. start=1,
  372. )
  373. ]
  374. agent = EventAgent(
  375. enabled_tools=[event.name for event in events],
  376. registry=registry,
  377. chat_client=TrailingUsageChatClient(),
  378. params=EventAgentParams(max_parallel_events=2),
  379. )
  380. await agent.handle_many(events, history=[])
  381. metrics = agent.fallback_model_calls(events)
  382. assert [(metric["event_id"], metric["event_name"]) for metric in metrics] == [
  383. ("event-1", "fallback.first"),
  384. ("event-2", "fallback.second"),
  385. ]
  386. assert [metric["usage"].total_tokens for metric in metrics] == [7, 7]
  387. @pytest.mark.asyncio
  388. async def test_event_agent_rejects_fallback_provider_call_for_wrong_tool_name():
  389. chat_client = WrongNameToolCallingChatClient({"message": "wrong"})
  390. registry = ToolRegistry(
  391. [
  392. ToolDefinition(
  393. name="expected.tool",
  394. description="Expected tool.",
  395. parameters={
  396. "type": "object",
  397. "properties": {"message": {"type": "string"}},
  398. "required": ["message"],
  399. },
  400. handler=lambda event: {"tool": event.name},
  401. argument_resolver=lambda event, context: {},
  402. )
  403. ]
  404. )
  405. reply = await EventAgent(
  406. enabled_tools=["expected.tool"],
  407. registry=registry,
  408. chat_client=chat_client,
  409. ).handle(
  410. ToolCallEvent(
  411. id="call-1",
  412. name="expected.tool",
  413. arguments={},
  414. raw_arguments="{}",
  415. )
  416. )
  417. assert json.loads(reply.content) == {
  418. "tool": "expected.tool",
  419. "error": "fallback returned tool wrong.tool for expected.tool",
  420. }
  421. @pytest.mark.asyncio
  422. async def test_event_agent_handle_many_starts_independent_handlers_concurrently():
  423. started: list[str] = []
  424. both_started = asyncio.Event()
  425. async def handler(event: ToolCallEvent) -> dict:
  426. started.append(event.name)
  427. if len(started) == 2:
  428. both_started.set()
  429. await asyncio.wait_for(both_started.wait(), timeout=0.2)
  430. return {"tool": event.name}
  431. registry = ToolRegistry(
  432. [
  433. ToolDefinition(
  434. name=name,
  435. description=f"Handle {name}.",
  436. parameters={"type": "object"},
  437. handler=handler,
  438. )
  439. for name in ("event.first", "event.second")
  440. ]
  441. )
  442. events = [
  443. ToolCallEvent(id=f"call-{index}", name=name, arguments={}, raw_arguments="{}")
  444. for index, name in enumerate(("event.first", "event.second"), start=1)
  445. ]
  446. batch = await asyncio.wait_for(
  447. EventAgent(
  448. enabled_tools=[event.name for event in events],
  449. registry=registry,
  450. params=EventAgentParams(max_parallel_events=2),
  451. ).handle_many(events, history=[]),
  452. timeout=0.5,
  453. )
  454. assert started == ["event.first", "event.second"]
  455. assert [result.event_name for result in batch.results] == [
  456. "event.first",
  457. "event.second",
  458. ]
  459. assert [result.payload for result in batch.results] == [
  460. {"tool": "event.first"},
  461. {"tool": "event.second"},
  462. ]
  463. @pytest.mark.asyncio
  464. @pytest.mark.parametrize(
  465. "item",
  466. [
  467. StreamItem.text_event(
  468. ToolCallEvent(
  469. id="text_event_1",
  470. name="mock_search",
  471. arguments={"query": "wrong source"},
  472. raw_arguments='{"query":"wrong source"}',
  473. )
  474. ),
  475. StreamItem.event(
  476. ToolCallEvent(
  477. id="legacy_event_1",
  478. name="mock_search",
  479. arguments={"query": "legacy arguments"},
  480. raw_arguments='{"query":"legacy arguments"}',
  481. )
  482. ),
  483. ],
  484. ids=["text_event", "legacy_event"],
  485. )
  486. async def test_event_agent_ignores_non_provider_tool_call_sources(item: StreamItem):
  487. event = ToolCallEvent(
  488. id="call_1",
  489. name="mock_search",
  490. arguments={},
  491. raw_arguments="{}",
  492. )
  493. registry = ToolRegistry(
  494. [
  495. ToolDefinition(
  496. name="mock_search",
  497. description="Search mock external knowledge for the current turn.",
  498. parameters={
  499. "type": "object",
  500. "properties": {"query": {"type": "string"}},
  501. "required": ["query"],
  502. },
  503. handler=lambda resolved: {
  504. "tool": resolved.name,
  505. "query": resolved.arguments["query"],
  506. },
  507. argument_resolver=lambda event, context: {},
  508. )
  509. ]
  510. )
  511. agent = EventAgent(
  512. enabled_tools=["mock_search"],
  513. registry=registry,
  514. chat_client=WrongSourceToolCallChatClient(item),
  515. )
  516. reply = await agent.handle(
  517. event,
  518. history=[ChatMessage(role="assistant", content="fallback query")],
  519. )
  520. payload = json.loads(reply.content)
  521. assert payload["tool"] == "mock_search"
  522. assert payload == {
  523. "tool": "mock_search",
  524. "error": "missing required arguments: query",
  525. }
  526. @pytest.mark.asyncio
  527. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  528. registry = ToolRegistry(
  529. [
  530. ToolDefinition(
  531. name="handoff_note",
  532. description="Send a note to the event agent.",
  533. parameters={"type": "object"},
  534. handler=lambda event: {"tool": event.name, "message": "handled"},
  535. )
  536. ]
  537. )
  538. disabled_reply = await EventAgent(
  539. enabled_tools=[],
  540. registry=registry,
  541. ).handle(
  542. ToolCallEvent(
  543. id="call_1",
  544. name="handoff_note",
  545. arguments={"message": "inspect this event"},
  546. raw_arguments='{"message":"inspect this event"}',
  547. )
  548. )
  549. unknown_reply = await EventAgent(
  550. enabled_tools=["missing_tool"],
  551. registry=registry,
  552. ).handle(
  553. ToolCallEvent(
  554. id="call_2",
  555. name="missing_tool",
  556. arguments={},
  557. raw_arguments="{}",
  558. )
  559. )
  560. assert json.loads(disabled_reply.content) == {
  561. "tool": "handoff_note",
  562. "error": "tool disabled",
  563. }
  564. assert json.loads(unknown_reply.content) == {
  565. "tool": "missing_tool",
  566. "error": "unknown tool",
  567. }
  568. @pytest.mark.asyncio
  569. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  570. def fail_tool(event: ToolCallEvent) -> dict:
  571. raise RuntimeError("boom")
  572. registry = ToolRegistry(
  573. [
  574. ToolDefinition(
  575. name="handoff_note",
  576. description="Send a note to the event agent.",
  577. parameters={"type": "object"},
  578. handler=fail_tool,
  579. )
  580. ]
  581. )
  582. event = ToolCallEvent(
  583. id="call_1",
  584. name="handoff_note",
  585. arguments={"message": "inspect this event"},
  586. raw_arguments='{"message":"inspect this event"}',
  587. )
  588. reply = await EventAgent(
  589. enabled_tools=["handoff_note"],
  590. registry=registry,
  591. ).handle(event)
  592. assert reply.role == "tool"
  593. assert reply.tool_call_id == "call_1"
  594. assert json.loads(reply.content) == {
  595. "tool": "handoff_note",
  596. "error": "tool handler failed: boom",
  597. }
  598. @pytest.mark.asyncio
  599. async def test_event_agent_serializes_non_json_handler_payload_as_tool_error():
  600. registry = ToolRegistry(
  601. [
  602. ToolDefinition(
  603. name="bad_payload",
  604. description="Return an invalid payload.",
  605. parameters={"type": "object"},
  606. handler=lambda event: {"invalid": object()},
  607. )
  608. ]
  609. )
  610. reply = await EventAgent(
  611. enabled_tools=["bad_payload"],
  612. registry=registry,
  613. ).handle(
  614. ToolCallEvent(
  615. id="call-1",
  616. name="bad_payload",
  617. arguments={},
  618. raw_arguments="{}",
  619. )
  620. )
  621. assert json.loads(reply.content) == {
  622. "tool": "bad_payload",
  623. "error": "event handler returned non-JSON payload",
  624. }
  625. @pytest.mark.asyncio
  626. async def test_event_agent_normalizes_handler_payload_snapshot_exceptions():
  627. class ExplodingItemsDict(dict):
  628. def items(self):
  629. raise RuntimeError("payload items failed")
  630. registry = ToolRegistry(
  631. [
  632. ToolDefinition(
  633. name="bad_payload",
  634. description="Return a payload that fails during snapshot.",
  635. parameters={"type": "object"},
  636. handler=lambda event: ExplodingItemsDict(ok=True),
  637. )
  638. ]
  639. )
  640. reply = await EventAgent(
  641. enabled_tools=["bad_payload"],
  642. registry=registry,
  643. ).handle(
  644. ToolCallEvent(
  645. id="call-1",
  646. name="bad_payload",
  647. arguments={},
  648. raw_arguments="{}",
  649. )
  650. )
  651. assert reply.role == "tool"
  652. assert reply.tool_call_id == "call-1"
  653. assert json.loads(reply.content) == {
  654. "tool": "bad_payload",
  655. "error": "event handler returned non-JSON payload",
  656. }
  657. @pytest.mark.asyncio
  658. async def test_event_agent_llm_receives_history_and_agent_config_context():
  659. chat_client = ToolCallingChatClient(
  660. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  661. )
  662. registry = ToolRegistry(
  663. [
  664. ToolDefinition(
  665. name="handoff_note",
  666. description="Send a note to the event agent.",
  667. parameters={
  668. "type": "object",
  669. "properties": {
  670. "message": {"type": "string"},
  671. "thinking": {"type": "string"},
  672. },
  673. "required": ["message", "thinking"],
  674. },
  675. handler=lambda event: {
  676. "tool": event.name,
  677. "message": event.arguments["message"],
  678. "thinking": event.arguments["thinking"],
  679. },
  680. )
  681. ]
  682. )
  683. event = ToolCallEvent(
  684. id="call_1",
  685. name="handoff_note",
  686. arguments={},
  687. raw_arguments="{}",
  688. )
  689. reply = await EventAgent(
  690. enabled_tools=["handoff_note"],
  691. registry=registry,
  692. chat_client=chat_client,
  693. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  694. ).handle(
  695. event,
  696. history=[ChatMessage(role="user", content="debug this")],
  697. system_prompt="Use strict tool parameters.",
  698. extra_body={"thinking": {"type": "disabled"}},
  699. )
  700. messages = chat_client.calls[0]["messages"]
  701. assert any(message.content == "debug this" for message in messages)
  702. assert any(message.content == "Use strict tool parameters." for message in messages)
  703. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  704. assert json.loads(reply.content) == {
  705. "tool": "handoff_note",
  706. "message": "Use strict tool parameters.",
  707. "thinking": "disabled",
  708. }
  709. @pytest.mark.asyncio
  710. async def test_event_agent_uses_configured_extra_body_when_override_is_omitted():
  711. chat_client = ToolCallingChatClient({"message": "resolved"})
  712. registry = ToolRegistry(
  713. [
  714. ToolDefinition(
  715. name="handoff_note",
  716. description="Send a note to the event agent.",
  717. parameters={
  718. "type": "object",
  719. "properties": {"message": {"type": "string"}},
  720. "required": ["message"],
  721. },
  722. handler=lambda event: {"tool": event.name},
  723. argument_resolver=lambda event, context: {},
  724. )
  725. ]
  726. )
  727. await EventAgent(
  728. enabled_tools=["handoff_note"],
  729. registry=registry,
  730. chat_client=chat_client,
  731. params=AgentParams(extra_body={"configured": True}),
  732. ).handle(
  733. ToolCallEvent(
  734. id="call-1",
  735. name="handoff_note",
  736. arguments={},
  737. raw_arguments="{}",
  738. )
  739. )
  740. assert chat_client.calls[0]["params"].extra_body == {"configured": True}
  741. @pytest.mark.asyncio
  742. async def test_event_agent_preserves_explicit_empty_extra_body_override():
  743. chat_client = ToolCallingChatClient({"message": "resolved"})
  744. registry = ToolRegistry(
  745. [
  746. ToolDefinition(
  747. name="handoff_note",
  748. description="Send a note to the event agent.",
  749. parameters={
  750. "type": "object",
  751. "properties": {"message": {"type": "string"}},
  752. "required": ["message"],
  753. },
  754. handler=lambda event: {"tool": event.name},
  755. argument_resolver=lambda event, context: {},
  756. )
  757. ]
  758. )
  759. await EventAgent(
  760. enabled_tools=["handoff_note"],
  761. registry=registry,
  762. chat_client=chat_client,
  763. params=AgentParams(extra_body={"configured": True}),
  764. ).handle(
  765. ToolCallEvent(
  766. id="call-1",
  767. name="handoff_note",
  768. arguments={},
  769. raw_arguments="{}",
  770. ),
  771. extra_body={},
  772. )
  773. assert chat_client.calls[0]["params"].extra_body == {}
  774. @pytest.mark.asyncio
  775. async def test_event_agent_projects_complete_tool_round_to_visible_history():
  776. chat_client = ToolCallingChatClient({"message": "projected history"})
  777. registry = ToolRegistry(
  778. [
  779. ToolDefinition(
  780. name="handoff_note",
  781. description="Send a note to the event agent.",
  782. parameters={
  783. "type": "object",
  784. "properties": {"message": {"type": "string"}},
  785. "required": ["message"],
  786. },
  787. handler=lambda resolved: {
  788. "tool": resolved.name,
  789. "message": resolved.arguments["message"],
  790. },
  791. argument_resolver=lambda event, context: {},
  792. )
  793. ]
  794. )
  795. event = ToolCallEvent(
  796. id="call_current",
  797. name="handoff_note",
  798. arguments={},
  799. raw_arguments="{}",
  800. )
  801. prior_call = ToolCallEvent(
  802. id="call_prior",
  803. name="mock_search",
  804. arguments={"query": "latency"},
  805. raw_arguments='{"query":"latency"}',
  806. )
  807. await EventAgent(
  808. enabled_tools=["handoff_note"],
  809. registry=registry,
  810. chat_client=chat_client,
  811. ).handle(
  812. event,
  813. history=[
  814. ChatMessage(role="user", content="Find latency docs"),
  815. ChatMessage(
  816. role="assistant",
  817. content="I checked the latency sources.",
  818. tool_calls=[prior_call],
  819. ),
  820. ChatMessage(
  821. role="tool",
  822. content='{"results":["doc"]}',
  823. tool_call_id="call_prior",
  824. ),
  825. ChatMessage(role="user", content="Prepare a handoff"),
  826. ],
  827. )
  828. projected_messages = chat_client.calls[0]["messages"]
  829. visible_assistant = next(
  830. message
  831. for message in projected_messages
  832. if message.content == "I checked the latency sources."
  833. )
  834. assert visible_assistant.tool_calls == []
  835. assert not any(message.role == "tool" for message in projected_messages)