@@ -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-
162133class 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 ]):
0 commit comments