|
| 1 | +"""Tests for the root turn span on the sync OpenAI-Agents path. |
| 2 | +
|
| 3 | +The defect these guard: with no root span and no parenting, a trace is a flat |
| 4 | +list of fragments whose durations do not sum to the turn, so tool time simply |
| 5 | +goes missing. These assert a root span exists, that tool spans hang off it, and |
| 6 | +that tracing failures never stop the turn from streaming. |
| 7 | +""" |
| 8 | + |
| 9 | +from __future__ import annotations |
| 10 | + |
| 11 | +from types import SimpleNamespace |
| 12 | +from unittest.mock import AsyncMock, MagicMock |
| 13 | + |
| 14 | +import pytest |
| 15 | + |
| 16 | +from agentex.lib.adk._modules import _sync_turn as turn_mod |
| 17 | + |
| 18 | + |
| 19 | +def _adk(root_span=None): |
| 20 | + return SimpleNamespace( |
| 21 | + tracing=SimpleNamespace( |
| 22 | + start_span=AsyncMock(return_value=root_span), |
| 23 | + end_span=AsyncMock(), |
| 24 | + ) |
| 25 | + ) |
| 26 | + |
| 27 | + |
| 28 | +def _runner_yielding(*events): |
| 29 | + async def _stream(): |
| 30 | + for e in events: |
| 31 | + yield e |
| 32 | + |
| 33 | + result = MagicMock() |
| 34 | + result.stream_events = _stream |
| 35 | + return MagicMock(return_value=result) |
| 36 | + |
| 37 | + |
| 38 | +def _agent(name="analyst"): |
| 39 | + agent = MagicMock() |
| 40 | + agent.name = name |
| 41 | + return agent |
| 42 | + |
| 43 | + |
| 44 | +async def _drain(gen): |
| 45 | + return [e async for e in gen] |
| 46 | + |
| 47 | + |
| 48 | +@pytest.mark.asyncio |
| 49 | +async def test_opens_a_root_span_and_closes_it(monkeypatch): |
| 50 | + root = SimpleNamespace(id="root-1") |
| 51 | + adk = _adk(root) |
| 52 | + monkeypatch.setattr(turn_mod, "_get_adk", lambda: adk) |
| 53 | + monkeypatch.setattr(turn_mod.Runner, "run_streamed", _runner_yielding("a", "b")) |
| 54 | + |
| 55 | + events = await _drain(turn_mod.run_turn_streamed(_agent(), [], trace_id="tr1", task_id="task1")) |
| 56 | + |
| 57 | + assert events == ["a", "b"] |
| 58 | + assert adk.tracing.start_span.await_args.kwargs["name"] == "turn" |
| 59 | + adk.tracing.end_span.assert_awaited_once() |
| 60 | + assert adk.tracing.end_span.await_args.kwargs["span"] is root |
| 61 | + |
| 62 | + |
| 63 | +@pytest.mark.asyncio |
| 64 | +async def test_tool_spans_are_parented_to_the_root(monkeypatch): |
| 65 | + root = SimpleNamespace(id="root-1") |
| 66 | + monkeypatch.setattr(turn_mod, "_get_adk", lambda: _adk(root)) |
| 67 | + |
| 68 | + captured = {} |
| 69 | + |
| 70 | + class _Hooks: |
| 71 | + def __init__(self, **kwargs): |
| 72 | + captured.update(kwargs) |
| 73 | + |
| 74 | + async def close_open_tool_spans(self): |
| 75 | + return None |
| 76 | + |
| 77 | + monkeypatch.setattr(turn_mod, "SyncTracingHooks", _Hooks) |
| 78 | + monkeypatch.setattr(turn_mod.Runner, "run_streamed", _runner_yielding("a")) |
| 79 | + |
| 80 | + await _drain(turn_mod.run_turn_streamed(_agent(), [], trace_id="tr1")) |
| 81 | + |
| 82 | + assert captured["parent_span_id"] == "root-1" |
| 83 | + |
| 84 | + |
| 85 | +@pytest.mark.asyncio |
| 86 | +async def test_no_trace_id_still_streams(monkeypatch): |
| 87 | + adk = _adk() |
| 88 | + monkeypatch.setattr(turn_mod, "_get_adk", lambda: adk) |
| 89 | + monkeypatch.setattr(turn_mod.Runner, "run_streamed", _runner_yielding("a", "b")) |
| 90 | + |
| 91 | + events = await _drain(turn_mod.run_turn_streamed(_agent(), [])) |
| 92 | + |
| 93 | + assert events == ["a", "b"] |
| 94 | + adk.tracing.start_span.assert_not_awaited() |
| 95 | + |
| 96 | + |
| 97 | +@pytest.mark.asyncio |
| 98 | +async def test_start_span_failure_still_streams(monkeypatch): |
| 99 | + adk = _adk() |
| 100 | + adk.tracing.start_span.side_effect = RuntimeError("tracing backend down") |
| 101 | + monkeypatch.setattr(turn_mod, "_get_adk", lambda: adk) |
| 102 | + monkeypatch.setattr(turn_mod.Runner, "run_streamed", _runner_yielding("a")) |
| 103 | + |
| 104 | + assert await _drain(turn_mod.run_turn_streamed(_agent(), [], trace_id="tr1")) == ["a"] |
| 105 | + |
| 106 | + |
| 107 | +@pytest.mark.asyncio |
| 108 | +async def test_root_span_closes_when_the_run_raises(monkeypatch): |
| 109 | + root = SimpleNamespace(id="root-1") |
| 110 | + adk = _adk(root) |
| 111 | + monkeypatch.setattr(turn_mod, "_get_adk", lambda: adk) |
| 112 | + |
| 113 | + async def _boom(): |
| 114 | + yield "a" |
| 115 | + raise RuntimeError("max turns exceeded") |
| 116 | + |
| 117 | + result = MagicMock() |
| 118 | + result.stream_events = _boom |
| 119 | + monkeypatch.setattr(turn_mod.Runner, "run_streamed", MagicMock(return_value=result)) |
| 120 | + |
| 121 | + with pytest.raises(RuntimeError): |
| 122 | + await _drain(turn_mod.run_turn_streamed(_agent(), [], trace_id="tr1")) |
| 123 | + |
| 124 | + adk.tracing.end_span.assert_awaited_once() |
0 commit comments