Skip to content
Merged
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
8 changes: 8 additions & 0 deletions docs/ai-python/content/docs/reference/ai/stream.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,14 @@ If the last input message has `replay=True`, `ai.stream` does not call the
provider. It emits replay tool-end events from the existing assistant message
so resumable agent flows can dispatch the same tool calls again.

`Stream.replay_message` synthesizes stream events from a complete `Message`.

```python
async with ai.Stream.replay_message(message) as stream:
async for event in stream:
...
```

## Errors

If the provider stream ends before a finish event, iteration raises
Expand Down
26 changes: 5 additions & 21 deletions docs/ai-python/content/docs/reference/events.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ event-model namespace, and events are Pydantic models.

## Stream Events

Model stream events inherit from `BaseEvent`.
Model stream events inherit from `ModelEvent`.

```python
event.message
Expand Down Expand Up @@ -42,16 +42,13 @@ Tool events:
- `BuiltinToolEnd`
- `BuiltinToolResult`

File and hook events:
File events:

- `FileEvent`
- `HookSuspension`
- `HookResolution`

Stream event unions:
Stream event union:

- `Event`
- `DiscriminatedEvent`

## Agent Events

Expand All @@ -61,12 +58,11 @@ Agents can emit event values that are not raw provider stream events.
- `PartialToolCallResult`: Carries values yielded by streaming tools and
`yield_from`.
- `HookEvent`: Carries an internal hook message and `HookPart`.
- `RunBlocked`: Emitted when the run is blocked on deferred hooks.

Agent event unions:
Agent event union:

- `AgentMessageEvent`
- `AgentEvent`
- `TerminalEvent`

## Aggregator

Expand All @@ -84,15 +80,3 @@ class MyAggregator(ai.events.Aggregator[Item, Result, ModelInput]):

`snapshot()` is the rich value stored in tool results. `get_model_input()` is
the value sent back to the model.

## replay_message_events

`replay_message_events` synthesizes stream events from a complete `Message`.

```python
async for event in ai.events.replay_message_events(message):
...
```

Use it when cached, durable, or test messages need to flow through code that
expects stream events.
4 changes: 1 addition & 3 deletions src/ai/agents/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -1394,9 +1394,7 @@ async def _real(call: Context) -> AsyncGenerator[events_.AgentEvent]:
# signal for the loop's tool dispatcher (which already
# ran by the time we see the event here), not user-
# facing output.
if not (
isinstance(event, events_.BaseEvent) and event.replay
):
if not event.replay:
yield event
if transition is not None:
yield transition
Expand Down
83 changes: 41 additions & 42 deletions src/ai/types/events.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,19 +11,14 @@
# serialization border in the case of durable execution


# Placeholder so BaseEvent.message is typed as Message (not Message | None).
# Placeholder so ModelEvent.message is typed as Message (not Message | None).
# Stream.__anext__ stamps the real in-progress message before yielding,
# so consumers never see this value.
_DUMMY_MESSAGE = messages.Message(id="<unset>", role="assistant", parts=[])


class BaseEvent(pydantic.BaseModel):
"""Common fields stamped onto every event by the streaming wrapper.

``message`` carries the in-progress (or final) assistant message; the
streaming layer aggregates parts into it as deltas arrive and stamps
a reference onto each yielded event. ``usage`` carries the latest
usage value reported by the provider (latest-wins across the stream).
"""Anything ``ai.stream`` or ``Agent.run`` yields.

``replay`` is set on synthetic events emitted when ``models.stream``
short-circuits an existing assistant turn (resume-after-approval
Expand All @@ -32,103 +27,115 @@ class BaseEvent(pydantic.BaseModel):
internally. Excluded from JSON: it's a control flag, not data.
"""

message: messages.Message = _DUMMY_MESSAGE
usage: usage_.Usage | None = None
provider_metadata: dict[str, Any] | None = None
replay: bool = pydantic.Field(default=False, exclude=True, repr=False)

model_config = pydantic.ConfigDict(frozen=True)


class StreamStart(BaseEvent):
class ModelEvent(BaseEvent):
"""Streamed out of a model request (``ai.stream``).

``message`` carries the in-progress (or final) assistant message; the
streaming layer aggregates parts into it as deltas arrive and stamps
a reference onto each yielded event (``Stream.__anext__``). ``usage``
carries the latest usage value reported by the provider (latest-wins
across the stream).
"""

message: messages.Message = _DUMMY_MESSAGE
usage: usage_.Usage | None = None
provider_metadata: dict[str, Any] | None = None


class StreamStart(ModelEvent):
kind: Literal["stream_start"] = "stream_start"


class StreamEnd(BaseEvent):
class StreamEnd(ModelEvent):
kind: Literal["stream_end"] = "stream_end"


class TextStart(BaseEvent):
class TextStart(ModelEvent):
block_id: str = ""

kind: Literal["text_start"] = "text_start"


class TextDelta(BaseEvent):
class TextDelta(ModelEvent):
chunk: str
block_id: str = ""

kind: Literal["text_delta"] = "text_delta"


class TextEnd(BaseEvent):
class TextEnd(ModelEvent):
block_id: str = ""

kind: Literal["text_end"] = "text_end"


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

kind: Literal["reasoning_start"] = "reasoning_start"


class ReasoningDelta(BaseEvent):
class ReasoningDelta(ModelEvent):
chunk: str
block_id: str = ""

kind: Literal["reasoning_delta"] = "reasoning_delta"


class ReasoningEnd(BaseEvent):
class ReasoningEnd(ModelEvent):
block_id: str = ""

kind: Literal["reasoning_end"] = "reasoning_end"


class ToolStart(BaseEvent):
class ToolStart(ModelEvent):
tool_call_id: str = ""
tool_name: str = ""

kind: Literal["tool_start"] = "tool_start"


class ToolDelta(BaseEvent):
class ToolDelta(ModelEvent):
chunk: str
tool_call_id: str = ""

kind: Literal["tool_delta"] = "tool_delta"


class ToolEnd(BaseEvent):
class ToolEnd(ModelEvent):
tool_call: messages.ToolCallPart
tool_call_id: str = ""

kind: Literal["tool_end"] = "tool_end"


class BuiltinToolStart(BaseEvent):
class BuiltinToolStart(ModelEvent):
tool_call_id: str = ""
tool_name: str = ""

kind: Literal["builtin_tool_start"] = "builtin_tool_start"


class BuiltinToolDelta(BaseEvent):
class BuiltinToolDelta(ModelEvent):
chunk: str
tool_call_id: str = ""

kind: Literal["builtin_tool_delta"] = "builtin_tool_delta"


class BuiltinToolEnd(BaseEvent):
class BuiltinToolEnd(ModelEvent):
tool_call: messages.BuiltinToolCallPart
tool_call_id: str = ""

kind: Literal["builtin_tool_end"] = "builtin_tool_end"


class BuiltinToolResult(BaseEvent):
class BuiltinToolResult(ModelEvent):
"""Provider returned a result for a built-in tool call."""

result: messages.BuiltinToolReturnPart
Expand All @@ -137,7 +144,7 @@ class BuiltinToolResult(BaseEvent):
kind: Literal["builtin_tool_result"] = "builtin_tool_result"


class FileEvent(BaseEvent):
class FileEvent(ModelEvent):
"""A complete generated file from the LLM."""

block_id: str = ""
Expand Down Expand Up @@ -167,11 +174,6 @@ class FileEvent(BaseEvent):
| FileEvent
)

DiscriminatedEvent = Annotated[
Event,
pydantic.Field(discriminator="kind"),
]


async def _replay_message_events(
msg: messages.Message,
Expand Down Expand Up @@ -274,7 +276,7 @@ def to_model_input(cls, snapshot: Result) -> ModelInput:
...


class PartialToolCallResult(pydantic.BaseModel):
class PartialToolCallResult(BaseEvent):
"""Emitted when tool calls or other yield_from callers yield values."""

tool_call_id: str | None = None
Expand All @@ -292,7 +294,7 @@ def key(self) -> object:
kind: Literal["partial_tool_call_result"] = "partial_tool_call_result"


class ToolCallResult(pydantic.BaseModel):
class ToolCallResult(BaseEvent):
"""Emitted after tool calls execute — carries the result message.

When the framework auto-catches an exception raised by the tool,
Expand All @@ -313,7 +315,7 @@ class ToolCallResult(pydantic.BaseModel):
kind: Literal["tool_call_result"] = "tool_call_result"


class HookEvent(pydantic.BaseModel):
class HookEvent(BaseEvent):
"""Emitted when a hook suspends, resolves, or is cancelled."""

message: messages.Message
Expand All @@ -322,7 +324,7 @@ class HookEvent(pydantic.BaseModel):
kind: Literal["hook"] = "hook"


class RunBlocked(pydantic.BaseModel):
class RunBlocked(BaseEvent):
"""The run is blocked on hooks.

Emitted when the run stops being able to make progress without
Expand All @@ -348,13 +350,10 @@ class RunBlocked(pydantic.BaseModel):
kind: Literal["run_blocked"] = "run_blocked"


AgentMessageEvent = Event | ToolCallResult | HookEvent

AgentEvent = (
Event | ToolCallResult | HookEvent | PartialToolCallResult | RunBlocked
)

TerminalEvent = StreamEnd | ToolCallResult | HookEvent
AgentEvent = Annotated[
Event | ToolCallResult | HookEvent | PartialToolCallResult | RunBlocked,
pydantic.Field(discriminator="kind"),
]


class RunStateTracker:
Expand Down
7 changes: 6 additions & 1 deletion tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,7 +201,12 @@ async def collect_messages(
"""Collect terminal messages from an event stream."""
result: list[messages_.Message] = []
async for event in source:
if isinstance(event, agent_events_.TerminalEvent):
if isinstance(
event,
agent_events_.StreamEnd
| agent_events_.ToolCallResult
| agent_events_.HookEvent,
):
result.append(event.message)
return result

Expand Down
Loading