-
Notifications
You must be signed in to change notification settings - Fork 9
Expand file tree
/
Copy path_pydantic_ai_sync.py
More file actions
329 lines (296 loc) · 14.3 KB
/
Copy path_pydantic_ai_sync.py
File metadata and controls
329 lines (296 loc) · 14.3 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
"""Pydantic AI streaming integration for Agentex.
Converts a Pydantic AI ``AgentStreamEvent`` stream (as yielded by
``agent.run_stream_events(...)`` or via an ``event_stream_handler``) into the
Agentex ``StreamTaskMessage*`` events that the Agentex server understands.
Typical sync usage:
from pydantic_ai import Agent
from agentex.lib.adk import convert_pydantic_ai_to_agentex_events
agent = Agent("openai:gpt-4o", system_prompt="...")
@acp.on_message_send
async def handle_message_send(params):
async with agent.run_stream_events(params.content.content) as stream:
async for event in convert_pydantic_ai_to_agentex_events(stream):
yield event
Recommended: unified surface
-----------------------------
For new handlers, prefer ``UnifiedEmitter`` + ``PydanticAITurn`` over the
bare converter. The unified surface wires tracing automatically when a
``trace_id`` is provided, so tool and reasoning spans are derived from the
same event stream with no extra setup:
from agentex.lib.core.harness import UnifiedEmitter
from agentex.lib.adk._modules._pydantic_ai_turn import PydanticAITurn
emitter = UnifiedEmitter(task_id=task_id, trace_id=trace_id, parent_span_id=parent_span_id)
turn = PydanticAITurn(agent.run_stream_events(prompt), model="openai:gpt-4o")
async for event in emitter.yield_turn(turn):
yield event # forwarded over the ACP streaming response; spans derived automatically
``convert_pydantic_ai_to_agentex_events`` remains the low-level tap for
callers that manage their own tracing or need direct access to the raw
converted stream.
"""
from __future__ import annotations
import json
import inspect
from typing import Any, Callable, AsyncIterator
from pydantic_ai.run import AgentRunResultEvent
from pydantic_ai.messages import (
TextPart,
PartEndEvent,
ThinkingPart,
ToolCallPart,
TextPartDelta,
PartDeltaEvent,
PartStartEvent,
ToolReturnPart,
FinalResultEvent,
ThinkingPartDelta,
ToolCallPartDelta,
FunctionToolCallEvent,
FunctionToolResultEvent,
)
from agentex.lib.utils.logging import make_logger
from agentex.types.reasoning_content import ReasoningContent
from agentex.types.task_message_delta import TextDelta
from agentex.types.tool_request_delta import ToolRequestDelta
from agentex.types.task_message_update import (
StreamTaskMessageDone,
StreamTaskMessageFull,
StreamTaskMessageDelta,
StreamTaskMessageStart,
)
from agentex.types.task_message_content import TextContent
from agentex.types.tool_request_content import ToolRequestContent
from agentex.types.tool_response_content import ToolResponseContent
from agentex.types.reasoning_content_delta import ReasoningContentDelta
logger = make_logger(__name__)
def _args_delta_to_str(args_delta: str | dict[str, Any] | None) -> str:
"""Normalize a Pydantic AI ``ToolCallPartDelta.args_delta`` to a string fragment.
Pydantic AI emits string fragments for providers that stream JSON tokens
(OpenAI, Anthropic) and dicts for providers that emit one-shot tool calls.
Agentex's ``ToolRequestDelta.arguments_delta`` is concatenated server-side
and parsed as a single JSON object on completion, so we always produce a
string. For dict deltas this is a one-shot dump; subsequent dict deltas
will not compose correctly, but in practice dict deltas arrive as a single
final fragment.
"""
if args_delta is None:
return ""
if isinstance(args_delta, str):
return args_delta
return json.dumps(args_delta)
def _tool_return_content(result: ToolReturnPart | Any) -> Any:
"""Best-effort extraction of the user-visible content from a tool result.
``FunctionToolResultEvent.part`` is ``ToolReturnPart | RetryPromptPart``.
For ``ToolReturnPart`` we surface ``.content`` directly; for ``RetryPromptPart``
(a retry signal back to the model) we surface a string description so the
UI sees the failure reason.
"""
content = getattr(result, "content", None)
if content is None:
return str(result)
if isinstance(content, (str, int, float, bool, list, dict)):
return content
if hasattr(content, "model_dump"):
try:
return content.model_dump()
except Exception:
return str(content)
return str(content)
async def convert_pydantic_ai_to_agentex_events(
stream_response: AsyncIterator[Any],
on_result: Callable[[AgentRunResultEvent], Any] | None = None,
) -> AsyncIterator[StreamTaskMessageStart | StreamTaskMessageDelta | StreamTaskMessageFull | StreamTaskMessageDone]:
"""Convert a Pydantic AI agent event stream into Agentex stream events.
Mapping:
PartStartEvent(TextPart) -> StreamTaskMessageStart(TextContent)
PartStartEvent(ThinkingPart) -> StreamTaskMessageStart(ReasoningContent)
PartStartEvent(ToolCallPart) -> StreamTaskMessageStart(ToolRequestContent)
PartDeltaEvent(TextPartDelta) -> StreamTaskMessageDelta(TextDelta)
PartDeltaEvent(ThinkingPart..) -> StreamTaskMessageDelta(ReasoningContentDelta)
PartDeltaEvent(ToolCallPart..) -> StreamTaskMessageDelta(ToolRequestDelta)
PartEndEvent -> StreamTaskMessageDone
FunctionToolResultEvent -> StreamTaskMessageFull(ToolResponseContent)
FunctionToolCallEvent -> (ignored — already covered by Start/Delta/End)
FinalResultEvent -> (ignored — informational; the run-level
AgentRunResultEvent terminates the stream)
AgentRunResultEvent -> (ignored — Agentex closes the per-message
stream via PartEndEvent already)
Args:
stream_response: The async iterator yielded by Pydantic AI's
``agent.run_stream_events(...)`` context manager (or a stream of
``AgentStreamEvent`` items received in an ``event_stream_handler``).
on_result: Optional callback invoked with the terminal
``AgentRunResultEvent`` when the run completes. Both sync and
async callables are accepted. No ``StreamTaskMessage*`` events are
yielded for this terminal event; the callback is the only side
effect. Useful for capturing run-level usage without altering the
streaming output.
Yields:
Agentex ``StreamTaskMessage*`` events suitable for forwarding back over
the ACP streaming response.
"""
next_message_index = 0
# Maps Pydantic AI's per-response part index to our absolute message index.
# Part indices restart at 0 on each new model response in a multi-step run,
# so we always overwrite the entry on PartStartEvent.
part_to_message_index: dict[int, int] = {}
# Tool-call metadata indexed by Pydantic AI part index (so deltas can
# surface the tool_call_id even when ToolCallPartDelta.tool_call_id is None).
tool_call_meta: dict[int, tuple[str, str]] = {}
async for event in stream_response:
if isinstance(event, PartStartEvent):
message_index = next_message_index
next_message_index += 1
part_to_message_index[event.index] = message_index
if isinstance(event.part, TextPart):
yield StreamTaskMessageStart(
type="start",
index=message_index,
content=TextContent(
type="text",
author="agent",
content="",
),
)
if event.part.content:
yield StreamTaskMessageDelta(
type="delta",
index=message_index,
delta=TextDelta(type="text", text_delta=event.part.content),
)
elif isinstance(event.part, ThinkingPart):
yield StreamTaskMessageStart(
type="start",
index=message_index,
content=ReasoningContent(
type="reasoning",
author="agent",
summary=[],
content=[],
style="active",
),
)
if event.part.content:
yield StreamTaskMessageDelta(
type="delta",
index=message_index,
delta=ReasoningContentDelta(
type="reasoning_content",
content_index=0,
content_delta=event.part.content,
),
)
elif isinstance(event.part, ToolCallPart):
tool_call_meta[event.index] = (event.part.tool_call_id, event.part.tool_name)
# Pydantic AI may already have a fully-formed args dict at start
# when the provider returns the tool call in one shot; surface it
# directly so clients see the complete arguments without waiting
# for deltas.
initial_args: dict[str, Any] = {}
if isinstance(event.part.args, dict):
# dict(...) materializes a fresh dict[str, Any]; pydantic-ai's
# ToolCallPart.args includes TypedDict-style variants that
# pyright doesn't narrow to plain dict[str, Any] via isinstance.
initial_args = dict(event.part.args)
yield StreamTaskMessageStart(
type="start",
index=message_index,
content=ToolRequestContent(
type="tool_request",
author="agent",
tool_call_id=event.part.tool_call_id,
name=event.part.tool_name,
arguments=initial_args,
),
)
if isinstance(event.part.args, str) and event.part.args:
yield StreamTaskMessageDelta(
type="delta",
index=message_index,
delta=ToolRequestDelta(
type="tool_request",
tool_call_id=event.part.tool_call_id,
name=event.part.tool_name,
arguments_delta=event.part.args,
),
)
else:
logger.debug("Unhandled PartStartEvent part type: %r", type(event.part).__name__)
elif isinstance(event, PartDeltaEvent):
message_index = part_to_message_index.get(event.index)
if message_index is None:
logger.debug("PartDeltaEvent for unknown part index %s; skipping", event.index)
continue
if isinstance(event.delta, TextPartDelta):
yield StreamTaskMessageDelta(
type="delta",
index=message_index,
delta=TextDelta(type="text", text_delta=event.delta.content_delta),
)
elif isinstance(event.delta, ThinkingPartDelta):
if event.delta.content_delta:
yield StreamTaskMessageDelta(
type="delta",
index=message_index,
delta=ReasoningContentDelta(
type="reasoning_content",
content_index=0,
content_delta=event.delta.content_delta,
),
)
elif isinstance(event.delta, ToolCallPartDelta):
meta = tool_call_meta.get(event.index)
if meta is None:
# First time we've seen this part; the provider didn't emit
# a PartStartEvent first. Synthesize one from the delta if
# we have enough information.
tool_call_id = event.delta.tool_call_id or ""
tool_name = event.delta.tool_name_delta or ""
tool_call_meta[event.index] = (tool_call_id, tool_name)
else:
tool_call_id, tool_name = meta
yield StreamTaskMessageDelta(
type="delta",
index=message_index,
delta=ToolRequestDelta(
type="tool_request",
tool_call_id=tool_call_id,
name=tool_name,
arguments_delta=_args_delta_to_str(event.delta.args_delta),
),
)
else:
logger.debug("Unhandled PartDeltaEvent delta type: %r", type(event.delta).__name__)
elif isinstance(event, PartEndEvent):
message_index = part_to_message_index.get(event.index)
if message_index is None:
continue
yield StreamTaskMessageDone(type="done", index=message_index)
elif isinstance(event, FunctionToolResultEvent):
result = event.part
tool_call_id = result.tool_call_id
tool_name = getattr(result, "tool_name", "") or ""
message_index = next_message_index
next_message_index += 1
content_payload = _tool_return_content(result)
yield StreamTaskMessageFull(
type="full",
index=message_index,
content=ToolResponseContent(
type="tool_response",
author="agent",
tool_call_id=tool_call_id,
name=tool_name,
content=content_payload,
),
)
elif isinstance(event, (FunctionToolCallEvent, FinalResultEvent, AgentRunResultEvent)):
# Already covered by PartStart/PartDelta/PartEnd events above, or
# informational only (FinalResultEvent / AgentRunResultEvent signal
# run-level state, not new content to surface).
if isinstance(event, AgentRunResultEvent) and on_result is not None:
ret = on_result(event)
if inspect.iscoroutine(ret):
await ret
continue
else:
logger.debug("Unhandled Pydantic AI event type: %r", type(event).__name__)