|
6 | 6 | from typing import Any, AsyncIterator |
7 | 7 |
|
8 | 8 | import pytest |
| 9 | +from pydantic_ai.run import AgentRunResult, AgentRunResultEvent |
9 | 10 | from pydantic_ai.messages import ( |
10 | 11 | TextPart, |
11 | 12 | PartEndEvent, |
@@ -481,3 +482,67 @@ async def test_author_is_agent(self, events: list[Any]): |
481 | 482 | content = getattr(e, "content", None) |
482 | 483 | if content is not None and hasattr(content, "author"): |
483 | 484 | assert content.author == "agent" |
| 485 | + |
| 486 | + |
| 487 | +class TestOnResultCallback: |
| 488 | + """on_result callback: captures the terminal AgentRunResultEvent without |
| 489 | + altering streaming output.""" |
| 490 | + |
| 491 | + def _make_result_event(self, output: Any = "hello") -> AgentRunResultEvent: |
| 492 | + result = AgentRunResult(output=output, _output_tool_name=None) |
| 493 | + return AgentRunResultEvent(result=result) |
| 494 | + |
| 495 | + async def test_callback_invoked_once_with_result_event(self): |
| 496 | + """on_result is called exactly once, with the AgentRunResultEvent.""" |
| 497 | + captured: list[AgentRunResultEvent] = [] |
| 498 | + |
| 499 | + def on_result(event: AgentRunResultEvent) -> None: |
| 500 | + captured.append(event) |
| 501 | + |
| 502 | + result_event = self._make_result_event("the answer") |
| 503 | + events = [result_event] |
| 504 | + await _collect(convert_pydantic_ai_to_agentex_events(_aiter(events), on_result=on_result)) |
| 505 | + |
| 506 | + assert len(captured) == 1 |
| 507 | + assert captured[0] is result_event |
| 508 | + assert captured[0].result.output == "the answer" |
| 509 | + |
| 510 | + async def test_streaming_output_unchanged_with_callback(self): |
| 511 | + """Yielded StreamTaskMessage* sequence is identical whether on_result is set or not.""" |
| 512 | + result_event = self._make_result_event() |
| 513 | + events = [ |
| 514 | + PartStartEvent(index=0, part=TextPart(content="")), |
| 515 | + PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="hi")), |
| 516 | + PartEndEvent(index=0, part=TextPart(content="hi")), |
| 517 | + result_event, |
| 518 | + ] |
| 519 | + |
| 520 | + captured: list[AgentRunResultEvent] = [] |
| 521 | + out_with = await _collect(convert_pydantic_ai_to_agentex_events(_aiter(events), on_result=captured.append)) |
| 522 | + out_without = await _collect(convert_pydantic_ai_to_agentex_events(_aiter(events))) |
| 523 | + |
| 524 | + assert len(out_with) == len(out_without) |
| 525 | + for a, b in zip(out_with, out_without): |
| 526 | + assert type(a) is type(b) |
| 527 | + assert len(captured) == 1 |
| 528 | + |
| 529 | + async def test_no_callback_no_error(self): |
| 530 | + """AgentRunResultEvent is silently ignored when on_result is None.""" |
| 531 | + result_event = self._make_result_event() |
| 532 | + events = [result_event] |
| 533 | + out = await _collect(convert_pydantic_ai_to_agentex_events(_aiter(events))) |
| 534 | + assert out == [] |
| 535 | + |
| 536 | + async def test_async_callback_is_awaited(self): |
| 537 | + """An async on_result callable is properly awaited.""" |
| 538 | + captured: list[AgentRunResultEvent] = [] |
| 539 | + |
| 540 | + async def on_result_async(event: AgentRunResultEvent) -> None: |
| 541 | + captured.append(event) |
| 542 | + |
| 543 | + result_event = self._make_result_event("async_output") |
| 544 | + events = [result_event] |
| 545 | + await _collect(convert_pydantic_ai_to_agentex_events(_aiter(events), on_result=on_result_async)) |
| 546 | + |
| 547 | + assert len(captured) == 1 |
| 548 | + assert captured[0].result.output == "async_output" |
0 commit comments