Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions docs/ai-python/content/docs/basics/streaming.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,29 @@ if __name__ == "__main__":
`stream.output` returns text by default. When you pass `output_type`, it returns
an instance of that Pydantic model after the stream finishes.

### Stream partial objects

While streaming with `output_type`, the stream also emits
`ai.events.PartialOutput` events: a best-effort parse of the JSON generated so
far, updated as deltas arrive. The value is a plain dict and does not validate
against the model -- use it for progressive rendering:

```python
async with ai.stream(model, messages, output_type=UprisingForecast) as stream:
async for event in stream:
match event:
case ai.events.PartialOutput(value=partial):
print(partial.get("eta", "..."))
case ai.events.TextDelta():
pass

forecast = stream.output # validated instance, available after the stream
```

The latest snapshot is also available as `stream.partial_output`. Partials are
not validated; `stream.output` remains the validated result once the stream has
ended.

## Pass tools

`ai.stream` accepts a list of `ai.Tool`, however, it does not execute function
Expand Down
7 changes: 7 additions & 0 deletions docs/ai-python/content/docs/reference/events.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,13 @@ Text events:
- `TextEnd`
- `StreamEnd`

Structured-output events:

- `PartialOutput`: emitted while streaming with `output_type` set. `value` is
the best-effort parse of the JSON generated so far as a plain dict; it does
not validate against `output_type` (the validated instance is
`Stream.output` after the stream ends).

`StreamEnd` carries the final response metadata. New in 0.4.0:

- `finish_reason`: Why the model stopped: `stop`, `length`, `content_filter`,
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ classifiers = [
requires-python = ">=3.12"
dependencies = [
"httpx>=0.28.1",
"json-repair==0.*",
"modelsdotdev==0.*",
"pydantic>=2.13",
"typing-extensions>=4.15.0",
Expand Down
66 changes: 53 additions & 13 deletions src/ai/models/core/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import contextlib
import dataclasses
from collections import deque
from contextlib import AbstractAsyncContextManager
from typing import (
TYPE_CHECKING,
Expand All @@ -14,6 +15,7 @@
runtime_checkable,
)

import json_repair
import pydantic

# ``typing.TypeVar`` lacks the ``default=`` kwarg on Python <3.13.
Expand Down Expand Up @@ -125,6 +127,11 @@ def __init__(
# (``Stream(gen)``, ``Stream.replay_message``).
self._span: telemetry.Span[telemetry.AiStreamSpanData] | None = None
self._first_output_seen = False
# Synthetic PartialOutput events waiting to be yielded ahead of the
# next provider event, plus the latest parsed snapshot behind
# ``Stream.partial_output``. Only active when ``output_type`` is set.
self._pending_events: deque[types.events.Event] = deque()
self._last_partial: dict[str, Any] | None = None

@classmethod
def replay_message(
Expand Down Expand Up @@ -192,19 +199,41 @@ def __aiter__(self) -> Self:
return self

async def __anext__(self: Self) -> types.events.Event:
try:
event = await self._gen.__anext__()
except StopAsyncIteration:
if not self._hydrator.ended:
raise errors.ProviderIncompleteResponseError(
"provider stream ended without a finish event; "
"the response is incomplete",
# Premature termination is a transient transport or
# provider failure: worth retrying.
is_retryable=True,
) from None
raise
event = self._hydrator.feed(event)
if self._pending_events:
# Synthetic events (e.g. PartialOutput) bypass the hydrator,
# so stamp the live in-progress message ourselves — every
# yielded ModelEvent must carry it, never the dummy default.
event = self._pending_events.popleft().model_copy(
update={"message": self.message}
)
else:
try:
event = await self._gen.__anext__()
except StopAsyncIteration:
if not self._hydrator.ended:
raise errors.ProviderIncompleteResponseError(
"provider stream ended without a finish event; "
"the response is incomplete",
# Premature termination is a transient transport or
# provider failure: worth retrying.
is_retryable=True,
) from None
raise
event = self._hydrator.feed(event)
# Structured-output streaming: after each text delta, re-parse the
# accumulated JSON and queue a best-effort snapshot ahead of the
# next provider event. ``stream_stable`` keeps truncated values as
# verbatim strings instead of creatively repairing them, so
# successive snapshots only grow toward the final object.
if self._output_type is not None and isinstance(
event, types.events.TextDelta
):
parsed = json_repair.loads(self.message.text, stream_stable=True)
if isinstance(parsed, dict) and parsed != self._last_partial:
self._last_partial = parsed
self._pending_events.append(
types.events.PartialOutput(value=parsed)
)
# Milestones on the live span: replayed work gets no synthetic
# timings (span.replay covers the replay branch of ``stream()``,
# event.replay covers individual synthetic events).
Expand Down Expand Up @@ -273,6 +302,17 @@ def output(self) -> StreamOutputT:
"""
return cast("StreamOutputT", self.message.get_output(self._output_type))

@property
def partial_output(self) -> dict[str, Any] | None:
"""Latest best-effort parse of the structured output so far.

Only populated when ``output_type`` was set and at least one
text delta has arrived. The value is a plain dict and does not
validate against ``output_type``; use :attr:`output` for the
validated instance once the stream has ended.
"""
return self._last_partial


async def _replay_tool_calls(
msg: types.messages.Message,
Expand Down
16 changes: 16 additions & 0 deletions src/ai/types/events.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,21 @@ class TextEnd(ModelEvent):
kind: Literal["text_end"] = "text_end"


class PartialOutput(ModelEvent):
"""Best-effort snapshot of structured output generated so far.

Emitted while streaming with ``output_type`` set, after each text
delta that changes the parse. ``value`` is the partial parse of the
JSON generated so far as a plain dict -- it does not validate
against ``output_type`` and grows toward the final object, which is
available validated via ``Stream.output`` when the stream ends.
"""

value: dict[str, Any]

kind: Literal["partial_output"] = "partial_output"


class ReasoningStart(ModelEvent):
block_id: str = ""

Expand Down Expand Up @@ -179,6 +194,7 @@ class FileEvent(ModelEvent):
| TextStart
| TextDelta
| TextEnd
| PartialOutput
| ReasoningStart
| ReasoningDelta
| ReasoningEnd
Expand Down
157 changes: 156 additions & 1 deletion tests/models/core/test_api.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
from __future__ import annotations

import asyncio
from collections.abc import AsyncGenerator, Sequence
from collections.abc import AsyncGenerator, Callable, Sequence
from typing import Any, Literal, cast

import pydantic
Expand Down Expand Up @@ -822,3 +822,158 @@ async def test_replayed_turn_gets_replay_span(recorder: Recorder) -> None:
assert isinstance(call.data, ai.experimental_telemetry.AiStreamSpanData)
assert call.data.message is not None
assert call.data.message.text == "prior turn"


# -- Streaming partial output ----------------------------------------------


class _Answer(pydantic.BaseModel):
a: str
b: int = 0


class _PinsOutput(pydantic.BaseModel):
a: int
b: int = 0
c: str = ""


async def _scripted_deltas(chunks: list[str]) -> AsyncGenerator[events_.Event]:
yield events_.StreamStart()
yield events_.TextStart(block_id="t1")
for chunk in chunks:
yield events_.TextDelta(block_id="t1", chunk=chunk)
yield events_.TextEnd(block_id="t1")
yield events_.StreamEnd()


def _scripted_impl(
chunks: list[str],
) -> Callable[
[models.Model, list[messages_.Message]], AsyncGenerator[events_.Event]
]:
"""Provider stream seam replaying hand-built text deltas."""

async def impl(
model: models.Model,
messages: list[messages_.Message],
**kwargs: Any,
) -> AsyncGenerator[events_.Event]:
async for event in _scripted_deltas(chunks):
yield event

return impl


async def test_stream_emits_partial_output_snapshots() -> None:
MOCK_PROVIDER._stream_impl = _scripted_impl(['{"a": "hel', 'lo"}'])

seen: list[dict[str, Any]] = []
order: list[str] = []
async with models.stream(
MOCK_MODEL, [ai.user_message("go")], output_type=_Answer
) as stream:
async for event in stream:
if isinstance(event, events_.PartialOutput):
seen.append(event.value)
order.append("partial")
elif isinstance(event, events_.TextDelta):
order.append("delta")

# Every yielded ModelEvent carries the live in-progress
# message, never the dummy default — synthetic PartialOutput
# events bypass the hydrator so Stream stamps them itself.
assert event.message.id != "<unset>"

assert stream.partial_output is None or isinstance(
stream.partial_output, dict
)

# One snapshot per parse-changing delta: truncated string stays
# stable under stream_stable, then the completed object.
assert seen == [{"a": "hel"}, {"a": "hello"}]
assert order == ["delta", "partial", "delta", "partial"]
assert stream.output == _Answer(a="hello")


async def test_stream_partial_output_property_tracks_latest() -> None:
MOCK_PROVIDER._stream_impl = _scripted_impl(['{"a": "x', '"}'])

latest: dict[str, Any] | None = None
async with models.stream(
MOCK_MODEL, [ai.user_message("go")], output_type=_Answer
) as stream:
assert stream.partial_output is None
async for _ in stream:
if stream.partial_output is not None:
latest = stream.partial_output

assert latest == {"a": "x"}


async def test_stream_without_output_type_never_emits_partials() -> None:
MOCK_PROVIDER._stream_impl = _scripted_impl(["plain text"])

partials: list[events_.PartialOutput] = []
async with models.stream(MOCK_MODEL, [ai.user_message("go")]) as stream:
async for event in stream:
if isinstance(event, events_.PartialOutput):
partials.append(event)
elif isinstance(event, events_.TextDelta):
assert stream.partial_output is None

assert partials == []


async def test_stream_partial_output_pins_json_repair_behaviors() -> None:
# Pins observed json_repair stream_stable behaviors so a minor bump
# cannot silently change what consumers see mid-stream.
MOCK_PROVIDER._stream_impl = _scripted_impl(
[
'{"a": 1, "b"',
': 2, "c": "trunc',
'ated"}',
]
)

snapshots: list[dict[str, Any]] = []
async with models.stream(
MOCK_MODEL,
[ai.user_message("go")],
output_type=_PinsOutput,
) as stream:
async for event in stream:
if isinstance(event, events_.PartialOutput):
snapshots.append(event.value)

assert snapshots == [
{"a": 1}, # dangling key without a value is dropped
{"a": 1, "b": 2, "c": "trunc"}, # truncated string kept verbatim
{"a": 1, "b": 2, "c": "truncated"},
]
assert stream.output == _PinsOutput(a=1, b=2, c="truncated")


async def test_agent_stream_forwards_partial_output() -> None:
# Two text parts -> two deltas (emit_events_for_messages emits one
# delta per part).
answer = ai.messages.Message(
id="msg-1",
role="assistant",
parts=[
ai.messages.TextPart(text='{"a": "hel'),
ai.messages.TextPart(text='lo"}'),
],
)
mock_llm([[answer]])

my_agent = ai.Agent()
partials: list[dict[str, Any]] = []
async with my_agent.run(
MOCK_MODEL, [ai.user_message("go")], output_type=_Answer
) as stream:
async for event in stream:
if isinstance(event, events_.PartialOutput):
partials.append(event.value)

assert partials == [{"a": "hel"}, {"a": "hello"}]
13 changes: 12 additions & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading