Skip to content

Commit 6281761

Browse files
committed
Make middleware composable in the nested case, add stream result proto
1 parent 553e455 commit 6281761

8 files changed

Lines changed: 225 additions & 83 deletions

File tree

examples/temporal-middleware/main.py

Lines changed: 21 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -130,66 +130,39 @@ async def llm_call_activity(params: LLMParams) -> LLMResult:
130130
# wrap_tool, same as without middleware.
131131

132132

133-
class _BufferedStreamResult:
134-
"""Wraps a single buffered Message to look like a StreamResult."""
135-
136-
def __init__(self, message: ai.Message) -> None:
137-
self._message = message
138-
139-
def __aiter__(self) -> AsyncGenerator[ai.Message]:
140-
return self._generate()
141-
142-
async def _generate(self) -> AsyncGenerator[ai.Message]:
143-
yield self._message
144-
145-
@property
146-
def tool_calls(self) -> list[ai.ToolCallPart]:
147-
return self._message.tool_calls
148-
149-
@property
150-
def text(self) -> str:
151-
return self._message.text
152-
153-
@property
154-
def usage(self) -> ai.Usage | None:
155-
return self._message.usage
156-
157-
@property
158-
def output(self) -> Any:
159-
return self._message.output
160-
161-
162133
class TemporalMiddleware(ai.Middleware):
163-
"""Routes LLM calls and tool executions through Temporal activities.
164-
165-
The middleware tracks messages itself so it can serialize the full
166-
conversation to the LLM activity at each step.
167-
"""
134+
"""Routes LLM calls and tool executions through Temporal activities."""
168135

169-
def __init__(
170-
self,
171-
initial_messages: list[ai.Message],
172-
tool_schemas: list[dict[str, Any]],
173-
) -> None:
174-
self._messages = list(initial_messages)
136+
def __init__(self, tool_schemas: list[dict[str, Any]]) -> None:
175137
self._tool_schemas = tool_schemas
176138

177-
async def wrap_model(self, call: ai.middleware.ModelContext, next: Any) -> Any:
139+
async def wrap_model(
140+
self,
141+
call: ai.middleware.ModelContext,
142+
next: Any,
143+
) -> ai.StreamResultLike:
178144
"""LLM call → Temporal activity."""
179145
result = await temporalio.workflow.execute_activity(
180146
llm_call_activity,
181147
LLMParams(
182-
messages=[m.model_dump() for m in self._messages],
148+
messages=[m.model_dump() for m in call.messages],
183149
tool_schemas=self._tool_schemas,
184150
),
185151
start_to_close_timeout=datetime.timedelta(minutes=5),
186152
retry_policy=temporalio.common.RetryPolicy(maximum_attempts=3),
187153
)
188154
msg = ai.Message.model_validate(result.message)
189-
self._messages.append(msg)
190-
return _BufferedStreamResult(msg)
191155

192-
async def wrap_tool(self, call: ai.middleware.ToolContext, next: Any) -> Any:
156+
async def _single() -> AsyncGenerator[ai.Message]:
157+
yield msg
158+
159+
return ai.StreamResult.from_generator(_single())
160+
161+
async def wrap_tool(
162+
self,
163+
call: ai.middleware.ToolContext,
164+
next: Any,
165+
) -> ai.Message:
193166
"""Tool execution → Temporal activity."""
194167
result = await temporalio.workflow.execute_activity(
195168
tool_dispatch_activity,
@@ -199,16 +172,14 @@ async def wrap_tool(self, call: ai.middleware.ToolContext, next: Any) -> Any:
199172
),
200173
start_to_close_timeout=datetime.timedelta(minutes=2),
201174
)
202-
tool_msg = ai.tool_message(
175+
return ai.tool_message(
203176
ai.ToolResultPart(
204177
tool_call_id=call.tool_call_id,
205178
tool_name=call.tool_name,
206179
result=result.result,
207180
is_error=result.is_error,
208181
)
209182
)
210-
self._messages.append(tool_msg)
211-
return tool_msg
212183

213184

214185
# ── Agent (default loop — no customization) ──────────────────────
@@ -237,9 +208,9 @@ async def run(self, user_query: str) -> str:
237208
"description": t.description,
238209
"param_schema": t.param_schema,
239210
}
240-
for t in weather_agent._tools
211+
for t in weather_agent.tools
241212
]
242-
mw = TemporalMiddleware(messages, tool_schemas)
213+
mw = TemporalMiddleware(tool_schemas)
243214

244215
final_text = ""
245216
async for msg in weather_agent.run(model, messages, middleware=[mw]):

src/ai/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
ImageParams,
2222
Model,
2323
ModelCost,
24+
StreamResult,
2425
VideoParams,
2526
check_connection,
2627
generate,
@@ -36,6 +37,7 @@
3637
Part,
3738
PartState,
3839
ReasoningPart,
40+
StreamResultLike,
3941
StructuredOutputPart,
4042
TextPart,
4143
ToolCallPart,
@@ -85,6 +87,8 @@
8587
"ImageParams",
8688
"VideoParams",
8789
"Client",
90+
"StreamResult",
91+
"StreamResultLike",
8892
"check_connection",
8993
"model",
9094
"models",

src/ai/agents/agent.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -281,6 +281,11 @@ def __init__(
281281
self._tools: list[Tool[..., Any]] = tools or []
282282
self._loop_fn: LoopFn = _default_loop
283283

284+
@property
285+
def tools(self) -> list[Tool[..., Any]]:
286+
"""The agent's registered tools (read-only copy)."""
287+
return list(self._tools)
288+
284289
def loop(self, fn: LoopFn) -> LoopFn:
285290
"""Decorator: override the default loop function."""
286291
self._loop_fn = fn
@@ -332,9 +337,13 @@ async def _real(
332337
# Activate middleware for this run (and everything it calls).
333338
# When middleware is None (default), inherit the parent's middleware
334339
# from the context var — this lets nested agents share middleware.
340+
# When middleware is explicitly provided, *extend* the parent stack
341+
# so that outer cross-cutting concerns (tracing, durability) are
342+
# preserved. Pass ``middleware=[]`` to clear the stack entirely.
335343
mw_token: middleware_.Token | None = None
336344
if middleware is not None:
337-
mw_token = middleware_.activate(middleware)
345+
parent = middleware_.get()
346+
mw_token = middleware_.activate(parent + middleware)
338347
try:
339348
chain = middleware_._build_agent_run_chain(_real)
340349
async for message in chain(call):

src/ai/middleware.py

Lines changed: 25 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424

2525
from .types import messages as messages_
2626
from .types import tools as tools_
27+
from .types.stream import StreamResultLike
2728

2829
# ---------------------------------------------------------------------------
2930
# Call context objects — frozen dataclasses with isolated mutable fields.
@@ -115,15 +116,7 @@ def __post_init__(self) -> None:
115116
# Middleware base class — override the methods you care about.
116117
# ---------------------------------------------------------------------------
117118

118-
# Forward-reference type aliases to avoid importing heavy modules at the
119-
# top level. The actual types are checked at call sites only.
120-
#
121-
# These are intentionally broad (Any) so that the middleware module does
122-
# not import from models/ or agents/ — which would create circular deps.
123-
124-
# StreamResult from models/__init__.py — async-iterable of Message snapshots.
125-
_StreamResult = Any
126-
# Message
119+
# Message alias for brevity in signatures.
127120
_Message = messages_.Message
128121

129122
# Agent run next-function type: call -> async generator of messages.
@@ -159,13 +152,24 @@ async def wrap_agent_run(self, call, next):
159152
async def wrap_model(
160153
self,
161154
call: ModelContext,
162-
next: Callable[[ModelContext], Awaitable[_StreamResult]],
163-
) -> _StreamResult:
155+
next: Callable[[ModelContext], Awaitable[StreamResultLike]],
156+
) -> StreamResultLike:
164157
"""Wrap a model streaming call.
165158
166-
``next(call)`` returns a :class:`~ai.models.StreamResult` that is
167-
async-iterable over ``Message`` snapshots. You can do work before,
168-
iterate / transform the stream, or do cleanup after.
159+
``next(call)`` returns a :class:`~ai.types.StreamResultLike` that
160+
is async-iterable over ``Message`` snapshots. You can do work
161+
before, iterate / transform the stream, or do cleanup after.
162+
163+
To transform the stream, use
164+
:meth:`~ai.models.StreamResult.from_generator`::
165+
166+
async def wrap_model(self, call, next):
167+
stream = await next(call)
168+
async def _add_suffix():
169+
async for msg in stream:
170+
yield msg
171+
from ai.models import StreamResult
172+
return StreamResult.from_generator(_add_suffix())
169173
"""
170174
return await next(call)
171175

@@ -242,8 +246,8 @@ def deactivate(token: Token) -> None:
242246

243247

244248
def _build_model_chain(
245-
real: Callable[[ModelContext], Awaitable[_StreamResult]],
246-
) -> Callable[[ModelContext], Awaitable[_StreamResult]]:
249+
real: Callable[[ModelContext], Awaitable[StreamResultLike]],
250+
) -> Callable[[ModelContext], Awaitable[StreamResultLike]]:
247251
mw = get()
248252
if not mw:
249253
return real
@@ -252,9 +256,10 @@ def _build_model_chain(
252256
for m in reversed(mw):
253257

254258
def _make(
255-
m: Middleware, nxt: Callable[[ModelContext], Awaitable[_StreamResult]]
256-
) -> Callable[[ModelContext], Awaitable[_StreamResult]]:
257-
async def _wrapped(call: ModelContext) -> _StreamResult:
259+
m: Middleware,
260+
nxt: Callable[[ModelContext], Awaitable[StreamResultLike]],
261+
) -> Callable[[ModelContext], Awaitable[StreamResultLike]]:
262+
async def _wrapped(call: ModelContext) -> StreamResultLike:
258263
return await m.wrap_model(call, nxt)
259264

260265
return _wrapped
@@ -357,6 +362,7 @@ async def _wrapped(call: AgentRunContext) -> AsyncGenerator[_Message]:
357362
"HookContext",
358363
"Middleware",
359364
"ModelContext",
365+
"StreamResultLike",
360366
"ToolContext",
361367
"activate",
362368
"deactivate",

src/ai/models/__init__.py

Lines changed: 30 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@
3939
from .. import middleware as middleware_
4040
from ..types import messages as messages_
4141
from ..types import tools as tools_
42+
from ..types.stream import StreamResultLike
4243
from .ai_gateway.types import GenerateParams, ImageParams, VideoParams
4344
from .core.catalog import get_models, get_providers, register_catalog
4445
from .core.catalog import model as model
@@ -157,12 +158,32 @@ class StreamResult:
157158
158159
Properties like ``.text`` and ``.tool_calls`` delegate to the final
159160
``Message`` snapshot and are available after iteration completes.
161+
162+
Satisfies :class:`~ai.types.StreamResultLike`.
160163
"""
161164

162165
def __init__(self, gen: AsyncGenerator[messages_.Message]) -> None:
163166
self._gen = gen
164167
self._final: messages_.Message | None = None
165168

169+
@classmethod
170+
def from_generator(cls, gen: AsyncGenerator[messages_.Message]) -> StreamResult:
171+
"""Create a :class:`StreamResult` from an async generator.
172+
173+
This is the public API for middleware that needs to transform or
174+
replace the stream returned by ``wrap_model``::
175+
176+
async def wrap_model(self, call, next):
177+
original = await next(call)
178+
179+
async def _transformed():
180+
async for msg in original:
181+
yield modify(msg)
182+
183+
return StreamResult.from_generator(_transformed())
184+
"""
185+
return cls(gen)
186+
166187
def __aiter__(self) -> AsyncGenerator[messages_.Message]:
167188
return self._iterate()
168189

@@ -197,12 +218,15 @@ async def stream(
197218
output_type: type[pydantic.BaseModel] | None = None,
198219
client: Client | None = None,
199220
**kwargs: Any,
200-
) -> StreamResult:
221+
) -> StreamResultLike:
201222
"""Stream an LLM response.
202223
203-
Returns a :class:`StreamResult` that is async-iterable and collects
204-
the final ``Message``. After iteration, access ``.text``,
224+
Returns a :class:`StreamResultLike` that is async-iterable and
225+
collects the final ``Message``. After iteration, access ``.text``,
205226
``.tool_calls``, ``.usage``, etc.
227+
228+
Without middleware the concrete type is :class:`StreamResult`; with
229+
middleware it may be any :class:`~ai.StreamResultLike`.
206230
"""
207231
call = middleware_.ModelContext(
208232
model=model,
@@ -213,7 +237,7 @@ async def stream(
213237
kwargs=kwargs,
214238
)
215239

216-
async def _real(call: middleware_.ModelContext) -> StreamResult:
240+
async def _real(call: middleware_.ModelContext) -> StreamResultLike:
217241
_ensure_adapters()
218242
c = call.client or _auto_client(call.model)
219243
adapter_fn = _stream_adapters.get(call.model.adapter)
@@ -235,8 +259,7 @@ async def _real(call: middleware_.ModelContext) -> StreamResult:
235259
)
236260

237261
chain = middleware_._build_model_chain(_real)
238-
result: StreamResult = await chain(call)
239-
return result
262+
return await chain(call)
240263

241264

242265
async def generate(
@@ -337,6 +360,7 @@ async def buffer(gen: AsyncGenerator[messages_.Message]) -> messages_.Message:
337360
"ModelCost",
338361
"StreamFn",
339362
"StreamResult",
363+
"StreamResultLike",
340364
"VideoParams",
341365
# Catalog
342366
"get_models",

src/ai/types/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
Usage,
2020
generate_id,
2121
)
22+
from .stream import StreamResultLike
2223
from .tools import ToolLike, ToolSchema
2324

2425
__all__ = [
@@ -28,6 +29,7 @@
2829
"Part",
2930
"PartState",
3031
"ReasoningPart",
32+
"StreamResultLike",
3133
"StructuredOutputPart",
3234
"TextPart",
3335
"ToolCallPart",

src/ai/types/stream.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
"""StreamResultLike — structural protocol for stream results.
2+
3+
Middleware authors can type-check against this protocol without depending
4+
on the concrete ``StreamResult`` class in ``ai.models``.
5+
"""
6+
7+
from __future__ import annotations
8+
9+
from collections.abc import AsyncGenerator
10+
from typing import Any, Protocol, runtime_checkable
11+
12+
from . import messages as messages_
13+
14+
15+
@runtime_checkable
16+
class StreamResultLike(Protocol):
17+
"""Structural protocol satisfied by :class:`ai.models.StreamResult`.
18+
19+
Middleware that transforms or replaces the stream returned by
20+
``wrap_model`` should return an object satisfying this protocol.
21+
The easiest way is ``StreamResult.from_generator(gen)``.
22+
"""
23+
24+
def __aiter__(self) -> AsyncGenerator[messages_.Message]: ...
25+
26+
@property
27+
def text(self) -> str: ...
28+
29+
@property
30+
def tool_calls(self) -> list[messages_.ToolCallPart]: ...
31+
32+
@property
33+
def usage(self) -> messages_.Usage | None: ...
34+
35+
@property
36+
def output(self) -> Any: ...

0 commit comments

Comments
 (0)