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/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