Skip to content

Commit 3e2df6e

Browse files
declan-scaleclaude
andcommitted
feat(pydantic-ai): PydanticAITurn HarnessTurn + usage normalization
Adds PydanticAITurn, a HarnessTurn wrapping a pydantic-ai event stream, with pydantic_ai_usage_to_turn_usage mapping verified RunUsage fields (requests, input_tokens, output_tokens, cache_read_tokens, total_tokens) onto TurnUsage via defensive getattr; usage() populates after events exhaustion. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
1 parent b68b8d4 commit 3e2df6e

2 files changed

Lines changed: 274 additions & 0 deletions

File tree

Lines changed: 110 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,110 @@
1+
"""PydanticAITurn: a HarnessTurn wrapping a pydantic-ai event stream.
2+
3+
Adapts a pydantic-ai ``AgentStreamEvent`` stream into the canonical
4+
``StreamTaskMessage*`` stream while capturing run-level usage from the
5+
terminal ``AgentRunResultEvent``.
6+
7+
Typical usage::
8+
9+
async with agent.run_stream_events(user_msg) as stream:
10+
turn = PydanticAITurn(stream, model="openai:gpt-4o")
11+
async for event in turn.events:
12+
yield event
13+
span.set_attributes(turn.usage().model_dump())
14+
"""
15+
16+
from __future__ import annotations
17+
18+
from typing import Any, AsyncIterator
19+
20+
from pydantic_ai.run import AgentRunResultEvent
21+
22+
from agentex.lib.core.harness.types import TurnUsage
23+
from agentex.types.task_message_update import (
24+
StreamTaskMessageDone,
25+
StreamTaskMessageFull,
26+
StreamTaskMessageDelta,
27+
StreamTaskMessageStart,
28+
)
29+
from agentex.lib.adk._modules._pydantic_ai_sync import convert_pydantic_ai_to_agentex_events
30+
31+
StreamTaskMessage = StreamTaskMessageStart | StreamTaskMessageDelta | StreamTaskMessageFull | StreamTaskMessageDone
32+
33+
34+
def pydantic_ai_usage_to_turn_usage(usage: Any, model: str | None) -> TurnUsage:
35+
"""Map a pydantic-ai ``RunUsage`` onto ``TurnUsage``.
36+
37+
Uses defensive ``getattr(..., None)`` so a future field rename in
38+
pydantic-ai degrades to ``None`` rather than raising ``AttributeError``.
39+
40+
RunUsage fields (verified against pydantic-ai in this repo):
41+
input_tokens, cache_write_tokens, cache_read_tokens, output_tokens,
42+
input_audio_tokens, cache_audio_read_tokens, output_audio_tokens,
43+
details, requests, tool_calls.
44+
``total_tokens`` is a computed property.
45+
46+
Mapping:
47+
requests -> num_llm_calls
48+
input_tokens -> input_tokens
49+
output_tokens -> output_tokens
50+
cache_read_tokens -> cached_input_tokens
51+
total_tokens -> total_tokens
52+
"""
53+
raw_input = getattr(usage, "input_tokens", None)
54+
raw_output = getattr(usage, "output_tokens", None)
55+
raw_cache_read = getattr(usage, "cache_read_tokens", None)
56+
raw_total = getattr(usage, "total_tokens", None)
57+
raw_requests = getattr(usage, "requests", None)
58+
59+
return TurnUsage(
60+
model=model,
61+
input_tokens=raw_input if raw_input else None,
62+
output_tokens=raw_output if raw_output else None,
63+
cached_input_tokens=raw_cache_read if raw_cache_read else None,
64+
total_tokens=raw_total if raw_total else None,
65+
num_llm_calls=raw_requests if raw_requests is not None else 0,
66+
)
67+
68+
69+
class PydanticAITurn:
70+
"""A single harness turn backed by a pydantic-ai event stream.
71+
72+
Satisfies the ``HarnessTurn`` protocol: ``events`` async-generates the
73+
canonical ``StreamTaskMessage*`` stream; ``usage()`` returns a normalized
74+
``TurnUsage`` (valid only after ``events`` is exhausted).
75+
"""
76+
77+
def __init__(self, stream: AsyncIterator[Any], model: str | None = None) -> None:
78+
self._stream = stream
79+
self._model = model
80+
self._usage = TurnUsage(model=model)
81+
82+
@property
83+
def events(self) -> AsyncIterator[StreamTaskMessage]:
84+
return self._generate_events()
85+
86+
async def _generate_events(self) -> AsyncIterator[StreamTaskMessage]:
87+
def _capture(result_event: AgentRunResultEvent) -> None:
88+
run_result = getattr(result_event, "result", None)
89+
if run_result is None:
90+
return
91+
usage_attr = getattr(run_result, "usage", None)
92+
if usage_attr is None:
93+
return
94+
# In newer pydantic-ai, .usage is a DeprecatedCallableRunUsage —
95+
# it's both a property value and callable (emitting a deprecation
96+
# warning when called). Access it as a plain attribute to avoid the
97+
# warning; it already IS the RunUsage instance.
98+
usage_obj = usage_attr
99+
self._usage = pydantic_ai_usage_to_turn_usage(usage_obj, self._model)
100+
101+
async for ev in convert_pydantic_ai_to_agentex_events(self._stream, on_result=_capture):
102+
yield ev
103+
104+
def usage(self) -> TurnUsage:
105+
"""Return the normalized usage for this turn.
106+
107+
Valid only after ``events`` is exhausted (single-pass contract).
108+
Before exhaustion the model field is set but token fields are None.
109+
"""
110+
return self._usage
Lines changed: 164 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,164 @@
1+
"""Tests for PydanticAITurn and pydantic_ai_usage_to_turn_usage."""
2+
3+
from __future__ import annotations
4+
5+
from typing import Any, AsyncIterator
6+
7+
from pydantic_ai.run import AgentRunResult, AgentRunResultEvent
8+
from pydantic_ai.usage import RunUsage
9+
from pydantic_ai.messages import (
10+
TextPart,
11+
PartEndEvent,
12+
TextPartDelta,
13+
PartDeltaEvent,
14+
PartStartEvent,
15+
)
16+
17+
from agentex.lib.core.harness import HarnessTurn
18+
from agentex.lib.adk._modules._pydantic_ai_turn import (
19+
PydanticAITurn,
20+
pydantic_ai_usage_to_turn_usage,
21+
)
22+
23+
24+
async def _aiter(events: list[Any]) -> AsyncIterator[Any]:
25+
for e in events:
26+
yield e
27+
28+
29+
async def _collect(stream: AsyncIterator[Any]) -> list[Any]:
30+
return [e async for e in stream]
31+
32+
33+
def _make_result_event(output: Any = "done", usage: RunUsage | None = None) -> AgentRunResultEvent:
34+
result = AgentRunResult(output=output, _output_tool_name=None)
35+
if usage is not None:
36+
result._state.usage = usage
37+
return AgentRunResultEvent(result=result)
38+
39+
40+
class TestUsageNormalization:
41+
def test_usage_normalization_maps_fields(self):
42+
"""Real RunUsage fields map correctly onto TurnUsage."""
43+
usage = RunUsage(
44+
requests=3,
45+
input_tokens=200,
46+
output_tokens=80,
47+
cache_read_tokens=25,
48+
)
49+
turn_usage = pydantic_ai_usage_to_turn_usage(usage, model="openai:gpt-4o")
50+
51+
assert turn_usage.model == "openai:gpt-4o"
52+
assert turn_usage.input_tokens == 200
53+
assert turn_usage.output_tokens == 80
54+
assert turn_usage.num_llm_calls == 3
55+
56+
def test_total_tokens_is_computed(self):
57+
"""RunUsage.total_tokens is a computed property; we surface it correctly."""
58+
usage = RunUsage(input_tokens=100, output_tokens=50)
59+
turn_usage = pydantic_ai_usage_to_turn_usage(usage, model="openai:gpt-4o")
60+
assert turn_usage.total_tokens == 150
61+
62+
def test_cache_read_tokens_mapped_to_cached_input_tokens(self):
63+
usage = RunUsage(input_tokens=100, output_tokens=50, cache_read_tokens=20)
64+
turn_usage = pydantic_ai_usage_to_turn_usage(usage, model="openai:gpt-4o")
65+
assert turn_usage.cached_input_tokens == 20
66+
67+
def test_none_model(self):
68+
"""model=None is preserved."""
69+
usage = RunUsage()
70+
turn_usage = pydantic_ai_usage_to_turn_usage(usage, model=None)
71+
assert turn_usage.model is None
72+
73+
def test_empty_usage_produces_zero_counts(self):
74+
"""An empty RunUsage maps to 0 counts and None tokens."""
75+
usage = RunUsage()
76+
turn_usage = pydantic_ai_usage_to_turn_usage(usage, model="openai:gpt-4o")
77+
assert turn_usage.num_llm_calls == 0
78+
assert turn_usage.input_tokens is None
79+
assert turn_usage.output_tokens is None
80+
81+
82+
class TestPydanticAITurn:
83+
async def test_turn_satisfies_harness_turn_protocol(self):
84+
"""PydanticAITurn is structurally compatible with HarnessTurn."""
85+
turn = PydanticAITurn(_aiter([]), model="openai:gpt-4o")
86+
assert isinstance(turn, HarnessTurn)
87+
88+
async def test_usage_before_exhaustion_returns_default(self):
89+
"""usage() before iterating events returns default TurnUsage (model set, tokens None)."""
90+
result_event = _make_result_event(usage=RunUsage(requests=1, input_tokens=100, output_tokens=40))
91+
events = [result_event]
92+
turn = PydanticAITurn(_aiter(events), model="openai:gpt-4o")
93+
94+
# Do NOT exhaust events — check usage pre-run
95+
pre_usage = turn.usage()
96+
assert pre_usage.model == "openai:gpt-4o"
97+
assert pre_usage.input_tokens is None
98+
assert pre_usage.output_tokens is None
99+
assert pre_usage.num_llm_calls == 0
100+
101+
async def test_turn_events_and_usage(self):
102+
"""Driving events to exhaustion populates usage from the terminal event."""
103+
known_usage = RunUsage(
104+
requests=2,
105+
input_tokens=300,
106+
output_tokens=120,
107+
cache_read_tokens=30,
108+
)
109+
result_event = _make_result_event(usage=known_usage)
110+
events = [
111+
PartStartEvent(index=0, part=TextPart(content="")),
112+
PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="hi")),
113+
PartEndEvent(index=0, part=TextPart(content="hi")),
114+
result_event,
115+
]
116+
turn = PydanticAITurn(_aiter(events), model="openai:gpt-4o")
117+
118+
collected = await _collect(turn.events)
119+
120+
# Events match bare converter output (Start + Delta + Done = 3 events)
121+
assert len(collected) == 3
122+
123+
# Usage is populated after exhaustion
124+
usage = turn.usage()
125+
assert usage.model == "openai:gpt-4o"
126+
assert usage.input_tokens == 300
127+
assert usage.output_tokens == 120
128+
assert usage.cached_input_tokens == 30
129+
assert usage.num_llm_calls == 2
130+
assert usage.total_tokens == 420
131+
132+
async def test_events_match_bare_converter(self):
133+
"""Yielded events are identical to bare convert_pydantic_ai_to_agentex_events output."""
134+
from agentex.lib.adk._modules._pydantic_ai_sync import convert_pydantic_ai_to_agentex_events
135+
136+
text_events = [
137+
PartStartEvent(index=0, part=TextPart(content="")),
138+
PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="Hello")),
139+
PartEndEvent(index=0, part=TextPart(content="Hello")),
140+
]
141+
142+
turn = PydanticAITurn(_aiter(text_events), model="openai:gpt-4o")
143+
turn_out = await _collect(turn.events)
144+
145+
bare_out = await _collect(convert_pydantic_ai_to_agentex_events(_aiter(text_events)))
146+
147+
assert len(turn_out) == len(bare_out)
148+
for a, b in zip(turn_out, bare_out):
149+
assert type(a) is type(b)
150+
assert a.model_dump() == b.model_dump()
151+
152+
async def test_no_usage_event_leaves_default_usage(self):
153+
"""If the stream has no AgentRunResultEvent, usage() returns the default (tokens None)."""
154+
events = [
155+
PartStartEvent(index=0, part=TextPart(content="")),
156+
PartEndEvent(index=0, part=TextPart(content="")),
157+
]
158+
turn = PydanticAITurn(_aiter(events), model="openai:gpt-4o")
159+
await _collect(turn.events)
160+
161+
usage = turn.usage()
162+
assert usage.model == "openai:gpt-4o"
163+
assert usage.input_tokens is None
164+
assert usage.num_llm_calls == 0

0 commit comments

Comments
 (0)