pi-python-core 0.8.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pi_python/__init__.py +160 -0
- pi_python/_version.py +1 -0
- pi_python/agent.py +396 -0
- pi_python/cancellation.py +24 -0
- pi_python/data/models.json +3315 -0
- pi_python/errors.py +49 -0
- pi_python/estimate.py +144 -0
- pi_python/events.py +138 -0
- pi_python/function_tools.py +438 -0
- pi_python/hooks.py +44 -0
- pi_python/limits.py +28 -0
- pi_python/loop.py +431 -0
- pi_python/lowlevel.py +179 -0
- pi_python/mcp.py +187 -0
- pi_python/messages.py +405 -0
- pi_python/models.py +155 -0
- pi_python/provider.py +123 -0
- pi_python/providers/__init__.py +21 -0
- pi_python/providers/anthropic.py +673 -0
- pi_python/providers/common.py +201 -0
- pi_python/providers/completions.py +1149 -0
- pi_python/providers/oauth.py +542 -0
- pi_python/providers/openai.py +681 -0
- pi_python/providers/transport.py +574 -0
- pi_python/proxy.py +304 -0
- pi_python/py.typed +0 -0
- pi_python/queues.py +76 -0
- pi_python/recovery.py +209 -0
- pi_python/run.py +419 -0
- pi_python/stream.py +251 -0
- pi_python/sync.py +78 -0
- pi_python/testing.py +25 -0
- pi_python/tools.py +546 -0
- pi_python/transcript.py +167 -0
- pi_python_core-0.8.1.dist-info/METADATA +119 -0
- pi_python_core-0.8.1.dist-info/RECORD +39 -0
- pi_python_core-0.8.1.dist-info/WHEEL +4 -0
- pi_python_core-0.8.1.dist-info/licenses/LICENSE +21 -0
- pi_python_core-0.8.1.dist-info/licenses/NOTICE +8 -0
pi_python/run.py
ADDED
|
@@ -0,0 +1,419 @@
|
|
|
1
|
+
"""One prompt or continuation: its tasks, cancellation, tool outcomes and events."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from .limits import RunLimits
|
|
5
|
+
from .provider import Provider
|
|
6
|
+
import asyncio
|
|
7
|
+
from collections.abc import Awaitable
|
|
8
|
+
from copy import deepcopy
|
|
9
|
+
from dataclasses import asdict, dataclass, field
|
|
10
|
+
from typing import TYPE_CHECKING, Any, Callable
|
|
11
|
+
from uuid import uuid4
|
|
12
|
+
|
|
13
|
+
from .cancellation import CancelToken
|
|
14
|
+
from .errors import ConfigurationError
|
|
15
|
+
from .hooks import AgentConfigUpdate, Hooks, RunContext, TurnUpdate
|
|
16
|
+
from .messages import (
|
|
17
|
+
AssistantMessage,
|
|
18
|
+
Message,
|
|
19
|
+
ToolResultMessage,
|
|
20
|
+
message_to_dict,
|
|
21
|
+
validate_history,
|
|
22
|
+
)
|
|
23
|
+
from .errors import SubscriptionError
|
|
24
|
+
from .tools import ToolOutcome, error_result, invoke
|
|
25
|
+
from .transcript import declare_tool_changes
|
|
26
|
+
from .loop import abort_error, failure_message, run_loop
|
|
27
|
+
|
|
28
|
+
if TYPE_CHECKING:
|
|
29
|
+
from .agent import Agent
|
|
30
|
+
|
|
31
|
+
# Token reason for cancellation of the caller's Python task. Unlike an explicit abort,
|
|
32
|
+
# it unwinds immediately and re-raises CancelledError to the caller.
|
|
33
|
+
CALLER_CANCELLED = "caller_cancelled"
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
@dataclass
|
|
37
|
+
class RunResult:
|
|
38
|
+
status: str
|
|
39
|
+
messages: list[Message]
|
|
40
|
+
usage: dict[str, Any] = field(default_factory=dict)
|
|
41
|
+
stop_reason: str | None = None
|
|
42
|
+
errors: list[str] = field(default_factory=list)
|
|
43
|
+
reconciliation_required: bool = False
|
|
44
|
+
tool_outcomes: list[ToolOutcome] = field(default_factory=list)
|
|
45
|
+
cleanup_complete: bool = True
|
|
46
|
+
queued_steering: int = 0
|
|
47
|
+
queued_follow_up: int = 0
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def merge(target: Any, update: AgentConfigUpdate) -> None:
|
|
51
|
+
"""Field-wise replacement; mutable values are copied, never deep-merged."""
|
|
52
|
+
for key in ("model", "options", "tools"):
|
|
53
|
+
value = getattr(update, key)
|
|
54
|
+
if value is not None:
|
|
55
|
+
setattr(target, key, deepcopy(value))
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
class Run:
|
|
59
|
+
"""State that lives for one run. The Agent keeps everything that outlives it."""
|
|
60
|
+
|
|
61
|
+
def __init__(self, agent: Agent, skip_initial_steering: bool = False):
|
|
62
|
+
self.agent = agent
|
|
63
|
+
self.token = CancelToken()
|
|
64
|
+
self.id = str(uuid4())
|
|
65
|
+
self.turn = 0
|
|
66
|
+
defaults = agent._defaults
|
|
67
|
+
self.context = RunContext(
|
|
68
|
+
deepcopy(agent._messages),
|
|
69
|
+
defaults.model or "mock",
|
|
70
|
+
deepcopy(defaults.options or {}),
|
|
71
|
+
deepcopy(defaults.tools or []),
|
|
72
|
+
)
|
|
73
|
+
self.start = len(agent._messages)
|
|
74
|
+
self.skip_initial_steering = skip_initial_steering
|
|
75
|
+
self.partial: Any = None
|
|
76
|
+
self.started = False # message_start was published for the current response
|
|
77
|
+
self.outcomes: list[ToolOutcome] = []
|
|
78
|
+
self.active_batch: list[ToolOutcome] = []
|
|
79
|
+
self._committed: set[int] = set()
|
|
80
|
+
self._ended: set[int] = set()
|
|
81
|
+
self._owned: set[asyncio.Future] = set()
|
|
82
|
+
self.driver: asyncio.Task | None = None
|
|
83
|
+
self.loop: asyncio.AbstractEventLoop | None = None
|
|
84
|
+
self.driver_started = False
|
|
85
|
+
self.finalizing = False
|
|
86
|
+
|
|
87
|
+
# Agent configuration seen by the loop.
|
|
88
|
+
@property
|
|
89
|
+
def provider(self) -> Provider:
|
|
90
|
+
return self.agent.provider
|
|
91
|
+
|
|
92
|
+
@property
|
|
93
|
+
def hooks(self) -> Hooks:
|
|
94
|
+
return self.agent.hooks
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def limits(self) -> RunLimits:
|
|
98
|
+
return self.agent.limits
|
|
99
|
+
|
|
100
|
+
@property
|
|
101
|
+
def execution_mode(self) -> str:
|
|
102
|
+
return self.agent.execution_mode
|
|
103
|
+
|
|
104
|
+
@property
|
|
105
|
+
def pending_calls(self) -> tuple[str, ...]:
|
|
106
|
+
return tuple(o.call.id for o in self.active_batch if id(o) not in self._committed)
|
|
107
|
+
|
|
108
|
+
def fail(self, error: str) -> None:
|
|
109
|
+
self.agent._last_error = error
|
|
110
|
+
|
|
111
|
+
def aborted(self) -> bool:
|
|
112
|
+
"""An explicit abort, seen from the driver or a tool worker.
|
|
113
|
+
|
|
114
|
+
Pi's abort lets the run reach its normal end: the current response or tool
|
|
115
|
+
batch settles, then the run stops with an aborted assistant message. Caller
|
|
116
|
+
task cancellation instead unwinds at once.
|
|
117
|
+
"""
|
|
118
|
+
task = asyncio.current_task()
|
|
119
|
+
return (
|
|
120
|
+
self.token.cancelled
|
|
121
|
+
and self.token.reason != CALLER_CANCELLED
|
|
122
|
+
and not (task is not None and task.cancelling())
|
|
123
|
+
)
|
|
124
|
+
|
|
125
|
+
def checkpoint(self) -> None:
|
|
126
|
+
if self.token.cancelled and not self.aborted():
|
|
127
|
+
raise asyncio.CancelledError(self.token.reason)
|
|
128
|
+
|
|
129
|
+
def cancelled_outcome(self) -> tuple[str, str, AssistantMessage | None]:
|
|
130
|
+
"""Status, reason and, for an explicit abort, Pi's closing aborted message."""
|
|
131
|
+
status = "limit_reached" if self.token.reason == "run_timeout" else "cancelled"
|
|
132
|
+
failure = failure_message(self, "aborted", abort_error(self)) if self.aborted() else None
|
|
133
|
+
return status, self.token.reason or "cancelled", failure
|
|
134
|
+
|
|
135
|
+
def stuck(self) -> bool:
|
|
136
|
+
"""A managed operation ignored cancellation past the cleanup deadline."""
|
|
137
|
+
return any(not task.done() for task in self._owned)
|
|
138
|
+
|
|
139
|
+
async def await_owned(
|
|
140
|
+
self,
|
|
141
|
+
awaitable: Awaitable,
|
|
142
|
+
timeout: float | None = None,
|
|
143
|
+
on_timeout: Callable[[], None] | None = None,
|
|
144
|
+
) -> Any:
|
|
145
|
+
if timeout is None and self.token.cancelled:
|
|
146
|
+
timeout = self.limits.cleanup_timeout
|
|
147
|
+
task = asyncio.ensure_future(awaitable)
|
|
148
|
+
self._owned.add(task)
|
|
149
|
+
|
|
150
|
+
def done(future: asyncio.Future) -> None:
|
|
151
|
+
self._owned.discard(future)
|
|
152
|
+
if not future.cancelled():
|
|
153
|
+
future.exception() # retrieved even when the owner has already been cancelled
|
|
154
|
+
|
|
155
|
+
task.add_done_callback(done)
|
|
156
|
+
# A cancellation request also wakes an awaited end subscriber. The driver
|
|
157
|
+
# remains the cleanup owner even when the caller cancels repeatedly.
|
|
158
|
+
cancel_wait = None if self.token.cancelled else asyncio.create_task(self.token.wait())
|
|
159
|
+
watched = {task} if cancel_wait is None else {task, cancel_wait}
|
|
160
|
+
try:
|
|
161
|
+
done_tasks, _ = await asyncio.wait(
|
|
162
|
+
watched, timeout=timeout, return_when=asyncio.FIRST_COMPLETED
|
|
163
|
+
)
|
|
164
|
+
if task in done_tasks:
|
|
165
|
+
return task.result()
|
|
166
|
+
cancelled = cancel_wait is not None and cancel_wait in done_tasks
|
|
167
|
+
if not cancelled and on_timeout is not None:
|
|
168
|
+
on_timeout()
|
|
169
|
+
# Task cancellation is Python's abort signal. As in Pi, the operation may
|
|
170
|
+
# finish its own way; it gets the cleanup deadline to do so.
|
|
171
|
+
task.cancel()
|
|
172
|
+
await asyncio.wait({task}, timeout=self.limits.cleanup_timeout)
|
|
173
|
+
if (
|
|
174
|
+
cancelled
|
|
175
|
+
and task.done()
|
|
176
|
+
and not task.cancelled()
|
|
177
|
+
and self.token.reason != CALLER_CANCELLED
|
|
178
|
+
):
|
|
179
|
+
return task.result()
|
|
180
|
+
if cancelled:
|
|
181
|
+
raise asyncio.CancelledError(self.token.reason)
|
|
182
|
+
raise TimeoutError("Operation deadline exceeded")
|
|
183
|
+
finally:
|
|
184
|
+
if cancel_wait is not None:
|
|
185
|
+
cancel_wait.cancel()
|
|
186
|
+
await asyncio.gather(cancel_wait, return_exceptions=True)
|
|
187
|
+
|
|
188
|
+
async def emit(self, kind: str, call_id: str | None = None, **data: Any) -> None:
|
|
189
|
+
await self.agent._events.emit(kind, self.id, self.turn, data, call_id, self.await_owned)
|
|
190
|
+
|
|
191
|
+
async def hook(self, name: str, *args: Any) -> Any:
|
|
192
|
+
callback = getattr(self.hooks, name)
|
|
193
|
+
if callback is None:
|
|
194
|
+
return None
|
|
195
|
+
return await self.await_owned(invoke(callback, *args))
|
|
196
|
+
|
|
197
|
+
async def poll_queue(self, steering: bool) -> list[Message]:
|
|
198
|
+
agent = self.agent
|
|
199
|
+
callback = agent._get_steering_messages if steering else agent._get_follow_up_messages
|
|
200
|
+
if callback is not None:
|
|
201
|
+
values = await self.await_owned(invoke(callback)) or []
|
|
202
|
+
validate_history(values)
|
|
203
|
+
return deepcopy(values)
|
|
204
|
+
return agent._queues.take(steering)
|
|
205
|
+
|
|
206
|
+
async def append(self, message: Message, *, started: bool = False) -> None:
|
|
207
|
+
data = message_to_dict(message)
|
|
208
|
+
# Commit before publishing; a recorder fault cannot erase an observed execution.
|
|
209
|
+
self.agent._messages.append(deepcopy(message))
|
|
210
|
+
self.context.messages.append(deepcopy(message))
|
|
211
|
+
self.context.new_messages.append(deepcopy(message))
|
|
212
|
+
if not started:
|
|
213
|
+
await self.emit("message_start", message=data)
|
|
214
|
+
await self.emit("message_end", message=data)
|
|
215
|
+
|
|
216
|
+
async def pending(self, pending: list[Message]) -> None:
|
|
217
|
+
"""Commit new input, declaring tool changes in a system message first."""
|
|
218
|
+
messages = declare_tool_changes(
|
|
219
|
+
self.context.messages, pending, [t.declaration() for t in self.context.tools]
|
|
220
|
+
)
|
|
221
|
+
for message in messages:
|
|
222
|
+
before = len(self.agent._messages)
|
|
223
|
+
try:
|
|
224
|
+
await self.append(message)
|
|
225
|
+
finally:
|
|
226
|
+
if len(self.agent._messages) > before:
|
|
227
|
+
self.agent._queues.committed(message, pending)
|
|
228
|
+
|
|
229
|
+
async def prepare(self, name: str) -> list[Message]:
|
|
230
|
+
"""Run a prepare hook; its update persists for the rest of this run."""
|
|
231
|
+
update = await self.hook(name, deepcopy(self.context), self.token)
|
|
232
|
+
if update is None:
|
|
233
|
+
return []
|
|
234
|
+
if not isinstance(update, TurnUpdate):
|
|
235
|
+
raise ConfigurationError(f"{name} must return TurnUpdate or None")
|
|
236
|
+
self.agent._validate_update(update)
|
|
237
|
+
merge(self.context, update)
|
|
238
|
+
if update.context is not None:
|
|
239
|
+
validate_history(update.context)
|
|
240
|
+
self.context.messages = deepcopy(update.context)
|
|
241
|
+
validate_history(update.messages)
|
|
242
|
+
return deepcopy(update.messages)
|
|
243
|
+
|
|
244
|
+
async def apply_updates(self) -> None:
|
|
245
|
+
"""Explicit update_config calls made while running take effect at a turn boundary."""
|
|
246
|
+
updates = self.agent._updates
|
|
247
|
+
while updates:
|
|
248
|
+
update = updates.pop(0)
|
|
249
|
+
merge(self.agent._defaults, update)
|
|
250
|
+
merge(self.context, update)
|
|
251
|
+
model = update.model
|
|
252
|
+
await self.emit(
|
|
253
|
+
"config_update",
|
|
254
|
+
model=getattr(model, "id", model),
|
|
255
|
+
options=update.options,
|
|
256
|
+
tools=None if update.tools is None else [t.name for t in update.tools],
|
|
257
|
+
)
|
|
258
|
+
|
|
259
|
+
async def cleanup(self) -> None:
|
|
260
|
+
tasks = set(self._owned)
|
|
261
|
+
for task in tasks:
|
|
262
|
+
task.cancel()
|
|
263
|
+
if tasks:
|
|
264
|
+
_, pending = await asyncio.wait(tasks, timeout=self.limits.cleanup_timeout)
|
|
265
|
+
if pending:
|
|
266
|
+
agent = self.agent
|
|
267
|
+
agent._cleanup_complete = False
|
|
268
|
+
self.fail("CleanupTimeoutError: managed tasks did not finish")
|
|
269
|
+
|
|
270
|
+
# A blocking tool cannot be interrupted, but it does finish eventually;
|
|
271
|
+
# the instance becomes usable again once nothing it owns is running.
|
|
272
|
+
def finished(task: asyncio.Future) -> None:
|
|
273
|
+
pending.discard(task)
|
|
274
|
+
if not pending:
|
|
275
|
+
agent._cleanup_finished()
|
|
276
|
+
|
|
277
|
+
for task in list(pending):
|
|
278
|
+
task.add_done_callback(finished)
|
|
279
|
+
|
|
280
|
+
async def commit_outcome(self, outcome: ToolOutcome) -> None:
|
|
281
|
+
if id(outcome) in self._committed:
|
|
282
|
+
return
|
|
283
|
+
if outcome.result is None:
|
|
284
|
+
if outcome.execution_status == "running":
|
|
285
|
+
outcome.result = error_result("not_stopped", "Tool did not stop after cancellation")
|
|
286
|
+
elif outcome.raw_result is not None:
|
|
287
|
+
outcome.result = error_result(
|
|
288
|
+
"finalization_cancelled", "Execution finished; finalization interrupted"
|
|
289
|
+
)
|
|
290
|
+
else:
|
|
291
|
+
outcome.result = error_result("cancelled", "Tool call did not start")
|
|
292
|
+
if outcome.execution_status == "unknown":
|
|
293
|
+
self.agent._unknown = True
|
|
294
|
+
result = outcome.result
|
|
295
|
+
outcome._settled = True
|
|
296
|
+
self._committed.add(id(outcome))
|
|
297
|
+
message = ToolResultMessage(
|
|
298
|
+
outcome.call.id,
|
|
299
|
+
outcome.call.name,
|
|
300
|
+
deepcopy(result.content),
|
|
301
|
+
result.is_error,
|
|
302
|
+
details=deepcopy(result.details),
|
|
303
|
+
usage=deepcopy(result.usage),
|
|
304
|
+
nested_calls=deepcopy(result.nested_calls),
|
|
305
|
+
)
|
|
306
|
+
# Persist even when delivery fails; then stop further external execution.
|
|
307
|
+
await self.append(message)
|
|
308
|
+
|
|
309
|
+
async def end_tool(self, outcome: ToolOutcome) -> None:
|
|
310
|
+
if id(outcome) in self._ended:
|
|
311
|
+
return
|
|
312
|
+
self._ended.add(id(outcome))
|
|
313
|
+
await self.emit(
|
|
314
|
+
"tool_execution_end",
|
|
315
|
+
outcome.call.id,
|
|
316
|
+
result=asdict(outcome.result) if outcome.result else None,
|
|
317
|
+
execution_status=outcome.execution_status,
|
|
318
|
+
)
|
|
319
|
+
|
|
320
|
+
async def drive(self, pending: list[Message]) -> RunResult:
|
|
321
|
+
agent = self.agent
|
|
322
|
+
self.driver_started = True
|
|
323
|
+
status, reason = "completed", "stop"
|
|
324
|
+
errors: list[str] = []
|
|
325
|
+
# Pi records a run that stopped outside a model response as one more failed
|
|
326
|
+
# or aborted assistant message, followed by turn_end.
|
|
327
|
+
failure: AssistantMessage | None = None
|
|
328
|
+
try:
|
|
329
|
+
status, reason = await run_loop(self, pending)
|
|
330
|
+
except asyncio.CancelledError:
|
|
331
|
+
status, reason, failure = self.cancelled_outcome()
|
|
332
|
+
except Exception as exc:
|
|
333
|
+
if isinstance(exc, TimeoutError) and self.aborted():
|
|
334
|
+
# A bounded step that outlived the cleanup deadline after abort is part of the abort.
|
|
335
|
+
status, reason, failure = self.cancelled_outcome()
|
|
336
|
+
else:
|
|
337
|
+
status, reason = "failed", type(exc).__name__
|
|
338
|
+
self.fail(f"{type(exc).__name__}: {exc}")
|
|
339
|
+
errors.append(agent._last_error or reason)
|
|
340
|
+
if not isinstance(exc, SubscriptionError): # the recorder is broken; stop
|
|
341
|
+
failure = failure_message(self, "error", agent._last_error)
|
|
342
|
+
self.finalizing = True
|
|
343
|
+
# Finalization has an independent owner: repeated abort cannot interrupt it.
|
|
344
|
+
self.token.cancel(reason) if status != "completed" else None
|
|
345
|
+
await self.cleanup()
|
|
346
|
+
for outcome in self.active_batch:
|
|
347
|
+
try:
|
|
348
|
+
await self.commit_outcome(outcome)
|
|
349
|
+
await self.end_tool(outcome)
|
|
350
|
+
except asyncio.CancelledError:
|
|
351
|
+
status = "cancelled"
|
|
352
|
+
errors.append("Result event delivery cancelled")
|
|
353
|
+
except Exception as exc:
|
|
354
|
+
status = "failed"
|
|
355
|
+
errors.append(f"{type(exc).__name__}: {exc}")
|
|
356
|
+
tail = agent._messages[-1] if len(agent._messages) > self.start else None
|
|
357
|
+
if failure is not None and not (
|
|
358
|
+
isinstance(tail, AssistantMessage) and tail.stop_reason in {"error", "aborted"}
|
|
359
|
+
):
|
|
360
|
+
try:
|
|
361
|
+
if failure.stop_reason == "aborted":
|
|
362
|
+
self.fail(failure.error or "aborted")
|
|
363
|
+
await self.append(failure)
|
|
364
|
+
await self.emit("turn_end", message=message_to_dict(failure), tool_results=[])
|
|
365
|
+
except asyncio.CancelledError:
|
|
366
|
+
status = "cancelled" if status != "limit_reached" else status
|
|
367
|
+
errors.append("Failure message delivery cancelled")
|
|
368
|
+
except Exception as exc:
|
|
369
|
+
status = "failed"
|
|
370
|
+
errors.append(f"{type(exc).__name__}: {exc}")
|
|
371
|
+
if agent._unknown:
|
|
372
|
+
status, reason = ("cancelled" if status == "cancelled" else "failed"), "outcome_unknown"
|
|
373
|
+
if agent._last_error and agent._last_error not in errors:
|
|
374
|
+
errors.append(agent._last_error)
|
|
375
|
+
if not agent._cleanup_complete and agent._last_error not in errors:
|
|
376
|
+
errors.append(agent._last_error or "CleanupTimeoutError")
|
|
377
|
+
self.partial = None
|
|
378
|
+
try:
|
|
379
|
+
# End delivery is itself bounded on cancellation/error paths.
|
|
380
|
+
if status == "completed" or status == "limit_reached" and reason != "run_timeout":
|
|
381
|
+
await self.emit("agent_end", status=status, stop_reason=reason)
|
|
382
|
+
elif agent._cleanup_complete:
|
|
383
|
+
await self.await_owned(
|
|
384
|
+
self.emit("agent_end", status=status, stop_reason=reason),
|
|
385
|
+
self.limits.cleanup_timeout,
|
|
386
|
+
)
|
|
387
|
+
except asyncio.CancelledError:
|
|
388
|
+
status = "limit_reached" if self.token.reason == "run_timeout" else "cancelled"
|
|
389
|
+
reason = self.token.reason or "cancelled"
|
|
390
|
+
await self.cleanup()
|
|
391
|
+
except Exception as exc:
|
|
392
|
+
status = "failed" if status != "cancelled" else status
|
|
393
|
+
errors.append(f"{type(exc).__name__}: {exc}")
|
|
394
|
+
self.fail(errors[-1])
|
|
395
|
+
await self.cleanup()
|
|
396
|
+
# Explicit config updates persist even if no further model request occurred.
|
|
397
|
+
for update in agent._updates:
|
|
398
|
+
merge(agent._defaults, update)
|
|
399
|
+
agent._updates.clear()
|
|
400
|
+
usage: dict[str, Any] = {}
|
|
401
|
+
for m in agent._messages[self.start :]:
|
|
402
|
+
if isinstance(m, AssistantMessage):
|
|
403
|
+
for key, value in m.usage.items():
|
|
404
|
+
if type(value) in (int, float):
|
|
405
|
+
usage[key] = usage.get(key, 0) + value
|
|
406
|
+
# Inputs selected for a later boundary but never committed remain queued.
|
|
407
|
+
agent._queues.restore()
|
|
408
|
+
return RunResult(
|
|
409
|
+
status,
|
|
410
|
+
deepcopy(agent._messages[self.start :]),
|
|
411
|
+
usage,
|
|
412
|
+
reason,
|
|
413
|
+
errors,
|
|
414
|
+
agent._unknown,
|
|
415
|
+
deepcopy(self.outcomes),
|
|
416
|
+
agent._cleanup_complete,
|
|
417
|
+
len(agent._queues.steering),
|
|
418
|
+
len(agent._queues.follow_up),
|
|
419
|
+
)
|
pi_python/stream.py
ADDED
|
@@ -0,0 +1,251 @@
|
|
|
1
|
+
"""Full Pi-style model events with defensive snapshots and strict terminal validation."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from collections.abc import AsyncGenerator, AsyncIterator, Callable
|
|
5
|
+
from typing import Any
|
|
6
|
+
from .cancellation import CancelToken
|
|
7
|
+
from .provider import ModelRequest
|
|
8
|
+
import asyncio
|
|
9
|
+
import json
|
|
10
|
+
from copy import deepcopy
|
|
11
|
+
from functools import wraps
|
|
12
|
+
from .errors import ProviderProtocolError
|
|
13
|
+
from .messages import AssistantMessage, TextContent, ThinkingContent, ToolCall, message_to_dict
|
|
14
|
+
from .provider import ModelEvent
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def partial_json(value: str) -> dict:
|
|
18
|
+
"""Best-effort UI preview only. Final arguments always use strict JSON parsing."""
|
|
19
|
+
try:
|
|
20
|
+
result = json.loads(value)
|
|
21
|
+
return result if isinstance(result, dict) else {}
|
|
22
|
+
except ValueError:
|
|
23
|
+
pass
|
|
24
|
+
# Bound preview work. This limit never relaxes strict final JSON validation.
|
|
25
|
+
if len(value) > 262144:
|
|
26
|
+
return {}
|
|
27
|
+
candidates = [len(value)] + [i for i in range(len(value) - 1, -1, -1) if value[i] in ",:{["][
|
|
28
|
+
:31
|
|
29
|
+
]
|
|
30
|
+
for end in candidates:
|
|
31
|
+
prefix = value[:end].rstrip().rstrip(",")
|
|
32
|
+
stack = []
|
|
33
|
+
quoted = False
|
|
34
|
+
escaped = False
|
|
35
|
+
for char in prefix:
|
|
36
|
+
if quoted:
|
|
37
|
+
if escaped:
|
|
38
|
+
escaped = False
|
|
39
|
+
elif char == "\\":
|
|
40
|
+
escaped = True
|
|
41
|
+
elif char == '"':
|
|
42
|
+
quoted = False
|
|
43
|
+
elif char == '"':
|
|
44
|
+
quoted = True
|
|
45
|
+
elif char in "{[":
|
|
46
|
+
stack.append("}" if char == "{" else "]")
|
|
47
|
+
elif char in "}]" and stack:
|
|
48
|
+
stack.pop()
|
|
49
|
+
suffix = ('"' if quoted else "") + "".join(reversed(stack))
|
|
50
|
+
try:
|
|
51
|
+
parsed = json.loads(prefix + suffix)
|
|
52
|
+
if isinstance(parsed, dict):
|
|
53
|
+
return parsed
|
|
54
|
+
except ValueError:
|
|
55
|
+
pass
|
|
56
|
+
return {}
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def terminal_content_matches(
|
|
60
|
+
ended: list[TextContent | ThinkingContent | ToolCall],
|
|
61
|
+
final: list[TextContent | ThinkingContent | ToolCall],
|
|
62
|
+
) -> bool:
|
|
63
|
+
"""Allow only Pi's late encrypted-reasoning backfill; never alter executable data."""
|
|
64
|
+
if len(ended) != len(final):
|
|
65
|
+
return False
|
|
66
|
+
for old, new in zip(ended, final):
|
|
67
|
+
if old == new:
|
|
68
|
+
continue
|
|
69
|
+
if not isinstance(old, ThinkingContent) or not isinstance(new, ThinkingContent):
|
|
70
|
+
return False
|
|
71
|
+
if old.thinking != new.thinking or old.redacted != new.redacted:
|
|
72
|
+
return False
|
|
73
|
+
try:
|
|
74
|
+
before, after = (
|
|
75
|
+
json.loads(old.thinking_signature or ""),
|
|
76
|
+
json.loads(new.thinking_signature or ""),
|
|
77
|
+
)
|
|
78
|
+
except ValueError:
|
|
79
|
+
return False
|
|
80
|
+
if not isinstance(before, dict) or not isinstance(after, dict):
|
|
81
|
+
return False
|
|
82
|
+
if before.get("encrypted_content") or not after.get("encrypted_content"):
|
|
83
|
+
return False
|
|
84
|
+
before.pop("encrypted_content", None)
|
|
85
|
+
after.pop("encrypted_content", None)
|
|
86
|
+
if before != after:
|
|
87
|
+
return False
|
|
88
|
+
return True
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
_BLOCK_STARTS = {"text_start", "thinking_start", "toolcall_start"}
|
|
92
|
+
_BLOCK_ENDS = {"text_end", "thinking_end", "toolcall_end"}
|
|
93
|
+
_DELTAS = {"text_delta": TextContent, "thinking_delta": ThinkingContent, "toolcall_delta": ToolCall}
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
async def checked_events(
|
|
97
|
+
source: AsyncIterator[ModelEvent], partial: AssistantMessage
|
|
98
|
+
) -> AsyncGenerator[ModelEvent, None]:
|
|
99
|
+
"""Enforce the ModelEvent contract and attach an independent `partial` snapshot.
|
|
100
|
+
|
|
101
|
+
Block boundaries must pair up, deltas must fall inside an open block of their own
|
|
102
|
+
kind, and the final message must equal the ended blocks (only Pi's late encrypted
|
|
103
|
+
reasoning backfill may differ). `done` is released only after the source ends
|
|
104
|
+
cleanly. Exceptions propagate; `error` events from a checked source pass through.
|
|
105
|
+
"""
|
|
106
|
+
opened: set[int] = set()
|
|
107
|
+
closed: set[int] = set()
|
|
108
|
+
arguments: dict[int, str] = {}
|
|
109
|
+
announced = False
|
|
110
|
+
terminal = None
|
|
111
|
+
async for source_event in source:
|
|
112
|
+
if terminal is not None:
|
|
113
|
+
raise ProviderProtocolError("Event after final message")
|
|
114
|
+
if not isinstance(source_event, ModelEvent):
|
|
115
|
+
raise ProviderProtocolError("Expected ModelEvent")
|
|
116
|
+
event = deepcopy(source_event)
|
|
117
|
+
kind = event.type
|
|
118
|
+
if kind == "start":
|
|
119
|
+
if announced:
|
|
120
|
+
raise ProviderProtocolError("Duplicate start event")
|
|
121
|
+
announced = True
|
|
122
|
+
if event.partial is not None:
|
|
123
|
+
partial = deepcopy(event.partial)
|
|
124
|
+
event.partial = deepcopy(partial)
|
|
125
|
+
yield event
|
|
126
|
+
continue
|
|
127
|
+
if kind == "error":
|
|
128
|
+
yield event
|
|
129
|
+
return
|
|
130
|
+
if kind == "done":
|
|
131
|
+
message = event.message
|
|
132
|
+
if not isinstance(message, AssistantMessage) or message.stop_reason == "pending":
|
|
133
|
+
raise ProviderProtocolError("Invalid terminal message")
|
|
134
|
+
message_to_dict(message)
|
|
135
|
+
if opened != closed or (opened and len(message.content) != len(partial.content)):
|
|
136
|
+
raise ProviderProtocolError("Terminal with unfinished blocks")
|
|
137
|
+
if opened and not terminal_content_matches(partial.content, message.content):
|
|
138
|
+
raise ProviderProtocolError("End/final content mismatch")
|
|
139
|
+
terminal = ModelEvent("done", message=deepcopy(message), reason=message.stop_reason)
|
|
140
|
+
continue # Do not publish success before the source closes cleanly.
|
|
141
|
+
if not announced:
|
|
142
|
+
announced = True
|
|
143
|
+
yield ModelEvent("start", partial=deepcopy(partial))
|
|
144
|
+
if kind in _BLOCK_STARTS:
|
|
145
|
+
if event.index != len(partial.content) or event.block is None:
|
|
146
|
+
raise ProviderProtocolError("Invalid content block start")
|
|
147
|
+
opened.add(event.index)
|
|
148
|
+
partial.content.append(deepcopy(event.block))
|
|
149
|
+
elif kind in _DELTAS:
|
|
150
|
+
if not isinstance(event.delta, str):
|
|
151
|
+
raise ProviderProtocolError("Delta must be text")
|
|
152
|
+
if event.index not in opened or event.index in closed:
|
|
153
|
+
raise ProviderProtocolError("Delta outside open block")
|
|
154
|
+
block = partial.content[event.index]
|
|
155
|
+
if not isinstance(block, _DELTAS[kind]):
|
|
156
|
+
raise ProviderProtocolError("Delta type mismatch")
|
|
157
|
+
if isinstance(block, TextContent):
|
|
158
|
+
block.text += event.delta
|
|
159
|
+
elif isinstance(block, ThinkingContent):
|
|
160
|
+
block.thinking += event.delta
|
|
161
|
+
else:
|
|
162
|
+
arguments[event.index] = arguments.get(event.index, "") + event.delta
|
|
163
|
+
block.arguments = partial_json(arguments[event.index])
|
|
164
|
+
event.call_id, event.name = block.id, block.name
|
|
165
|
+
elif kind in _BLOCK_ENDS:
|
|
166
|
+
if event.index not in opened or event.index in closed or event.block is None:
|
|
167
|
+
raise ProviderProtocolError("Invalid block end")
|
|
168
|
+
old = partial.content[event.index]
|
|
169
|
+
new = event.block
|
|
170
|
+
if type(old) is not type(new):
|
|
171
|
+
raise ProviderProtocolError("Block type changed")
|
|
172
|
+
if isinstance(old, TextContent) and isinstance(new, TextContent):
|
|
173
|
+
if old.text != new.text:
|
|
174
|
+
raise ProviderProtocolError("Text delta/end mismatch")
|
|
175
|
+
elif isinstance(old, ThinkingContent) and isinstance(new, ThinkingContent):
|
|
176
|
+
if old.thinking != new.thinking:
|
|
177
|
+
raise ProviderProtocolError("Thinking delta/end mismatch")
|
|
178
|
+
elif isinstance(old, ToolCall) and isinstance(new, ToolCall):
|
|
179
|
+
if old.id != new.id or old.name != new.name:
|
|
180
|
+
raise ProviderProtocolError("Tool identity changed")
|
|
181
|
+
if event.index in arguments:
|
|
182
|
+
try:
|
|
183
|
+
assembled = json.loads(arguments[event.index])
|
|
184
|
+
except ValueError as exc:
|
|
185
|
+
raise ProviderProtocolError("Incomplete tool arguments") from exc
|
|
186
|
+
if assembled != new.arguments:
|
|
187
|
+
raise ProviderProtocolError("Tool delta/end mismatch")
|
|
188
|
+
partial.content[event.index] = deepcopy(new)
|
|
189
|
+
closed.add(event.index)
|
|
190
|
+
event.content = (
|
|
191
|
+
new.text
|
|
192
|
+
if isinstance(new, TextContent)
|
|
193
|
+
else new.thinking
|
|
194
|
+
if isinstance(new, ThinkingContent)
|
|
195
|
+
else None
|
|
196
|
+
)
|
|
197
|
+
else:
|
|
198
|
+
raise ProviderProtocolError(f"Unknown model event: {kind}")
|
|
199
|
+
event.partial = deepcopy(partial)
|
|
200
|
+
yield event
|
|
201
|
+
if terminal is None:
|
|
202
|
+
raise ProviderProtocolError("Missing terminal message")
|
|
203
|
+
yield terminal
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def event_contract(function: Callable[..., AsyncGenerator[ModelEvent, None]]) -> Any:
|
|
207
|
+
"""Decorate a remote provider: checked events, and failures as an `error` event."""
|
|
208
|
+
|
|
209
|
+
@wraps(function)
|
|
210
|
+
async def wrapped(
|
|
211
|
+
self: Any, request: ModelRequest, cancel: CancelToken
|
|
212
|
+
) -> AsyncGenerator[ModelEvent, None]:
|
|
213
|
+
partial = AssistantMessage(
|
|
214
|
+
[],
|
|
215
|
+
stop_reason="pending",
|
|
216
|
+
provider=self.name,
|
|
217
|
+
model=request.model,
|
|
218
|
+
api=getattr(self, "api", None),
|
|
219
|
+
)
|
|
220
|
+
snapshot = partial
|
|
221
|
+
iterator = function(self, request, cancel)
|
|
222
|
+
events = checked_events(iterator, partial)
|
|
223
|
+
try:
|
|
224
|
+
async for event in events:
|
|
225
|
+
if event.partial is not None:
|
|
226
|
+
snapshot = event.partial
|
|
227
|
+
yield event
|
|
228
|
+
except asyncio.CancelledError:
|
|
229
|
+
failed = deepcopy(snapshot)
|
|
230
|
+
failed.stop_reason, failed.error = "aborted", "Request cancelled"
|
|
231
|
+
yield ModelEvent("error", message=failed, reason="aborted")
|
|
232
|
+
except Exception as exc:
|
|
233
|
+
failed = deepcopy(snapshot)
|
|
234
|
+
failed.stop_reason, failed.error = "error", f"{type(exc).__name__}: {exc}"
|
|
235
|
+
if hasattr(exc, "status") and hasattr(exc, "category"):
|
|
236
|
+
failed.diagnostics = [
|
|
237
|
+
{
|
|
238
|
+
"type": "provider_http_error",
|
|
239
|
+
"status": exc.status,
|
|
240
|
+
"category": exc.category,
|
|
241
|
+
"retry_after": getattr(exc, "retry_after", None),
|
|
242
|
+
"request_id": getattr(exc, "request_id", None),
|
|
243
|
+
}
|
|
244
|
+
]
|
|
245
|
+
yield ModelEvent("error", message=failed, reason="error")
|
|
246
|
+
finally:
|
|
247
|
+
await events.aclose()
|
|
248
|
+
await iterator.aclose()
|
|
249
|
+
|
|
250
|
+
wrapped.checked = True # type: ignore[attr-defined]
|
|
251
|
+
return wrapped
|