1313from .. import middleware as middleware_
1414from .. import models , types
1515from ..types import builders
16+ from . import events as events_
1617from . import runtime
1718
1819
@@ -199,21 +200,23 @@ def resolve(self, tool_parts: list[types.ToolCallPart]) -> list[ToolCall]:
199200 ]
200201
201202
202- StreamItem = types . Event | types .Message
203+ StreamItem = events_ . AgentEvent | types .Message
203204
204205
205206class LoopFn (Protocol ):
206207 def __call__ (self , context : Context ) -> AsyncGenerator [StreamItem ]: ...
207208
208209
209- async def _message_events (message : types .Message ) -> AsyncGenerator [types .Event ]:
210- yield types .MessageStart (message = message )
211- yield types .MessageEnd (message = message )
210+ async def _message_events (
211+ message : types .Message ,
212+ ) -> AsyncGenerator [events_ .AgentEvent ]:
213+ yield events_ .MessageStart (message = message )
214+ yield events_ .MessageEnd (message = message )
212215
213216
214217async def _coerce_events (
215218 source : AsyncIterable [StreamItem ],
216- ) -> AsyncGenerator [types . Event ]:
219+ ) -> AsyncGenerator [events_ . AgentEvent ]:
217220 async for item in source :
218221 if isinstance (item , types .Message ):
219222 async for event in _message_events (item ):
@@ -222,15 +225,23 @@ async def _coerce_events(
222225 yield item
223226
224227
225- async def _default_loop (context : Context ) -> AsyncGenerator [types . Event ]:
228+ async def _default_loop (context : Context ) -> AsyncGenerator [events_ . AgentEvent ]:
226229 while True :
227230 stream = models .stream (
228231 context .model ,
229232 context .messages ,
230233 tools = context .tools ,
231234 )
232- async for event in stream :
233- yield event
235+ async for stream_event in stream :
236+ yield stream_event
237+
238+ # Bridge: emit MessageStart/MessageEnd around the assistant message
239+ # the model stream just produced, so _collect_messages and downstream
240+ # consumers (AI-SDK outbound, label stamping) see the same boundary
241+ # events they did under the previous adapter contract.
242+ if stream .message is not None and stream .message .parts :
243+ async for boundary in _message_events (stream .message ):
244+ yield boundary
234245
235246 tool_calls = context .resolve (stream .tool_calls )
236247 if not tool_calls :
@@ -244,14 +255,14 @@ async def _default_loop(context: Context) -> AsyncGenerator[types.Event]:
244255 # Left un-stamped: the tool result is the input of the *next* turn,
245256 # so the next stream() call will stamp it with that turn's id.
246257 tool_msg = builders .tool_message (* (t .result () for t in tasks ))
247- async for event in _message_events (tool_msg ):
248- yield event
258+ async for boundary in _message_events (tool_msg ):
259+ yield boundary
249260
250261
251262async def _collect_messages (
252263 source : AsyncIterable [StreamItem ],
253264 messages : list [types .Message ],
254- ) -> AsyncGenerator [types . Event ]:
265+ ) -> AsyncGenerator [events_ . AgentEvent ]:
255266 """Intercept yielded events and collect MessageEnd messages into *messages*.
256267
257268 This runs on the **producer** side (same coroutine as the loop function),
@@ -260,7 +271,7 @@ async def _collect_messages(
260271 happened on the consumer side of the runtime queue.
261272 """
262273 async for event in _coerce_events (source ):
263- if isinstance (event , types .MessageEnd ):
274+ if isinstance (event , events_ .MessageEnd ):
264275 message = event .message
265276 for i , existing in enumerate (messages ):
266277 if existing .id == message .id :
@@ -292,7 +303,7 @@ async def yield_from(source: AsyncIterable[StreamItem]) -> str:
292303 last : types .Message | None = None
293304 async for item in _coerce_events (source ):
294305 await rt .put_event (item )
295- if isinstance (item , types .MessageEnd ):
306+ if isinstance (item , events_ .MessageEnd ):
296307 last = item .message
297308 return last .text if last else ""
298309
@@ -325,7 +336,7 @@ async def run(
325336 * ,
326337 label : str | None = None ,
327338 middleware : list [middleware_ .Middleware ] | None = None ,
328- ) -> AsyncGenerator [types . Event ]:
339+ ) -> AsyncGenerator [events_ . AgentEvent ]:
329340 """Run the agent loop, yielding events to the consumer.
330341
331342 Args:
@@ -349,7 +360,7 @@ async def run(
349360
350361 async def _real (
351362 call : middleware_ .AgentRunContext ,
352- ) -> AsyncGenerator [types . Event ]:
363+ ) -> AsyncGenerator [events_ . AgentEvent ]:
353364 context = Context (
354365 model = call .model ,
355366 messages = list (call .messages ),
@@ -359,8 +370,8 @@ async def _real(
359370 async for event in runtime .run (source ):
360371 if call .label is not None :
361372 event_message : types .Message | None = None
362- if isinstance (event , types .MessageEnd ) or (
363- isinstance (event , types .MessageStart )
373+ if isinstance (event , events_ .MessageEnd ) or (
374+ isinstance (event , events_ .MessageStart )
364375 and event .message is not None
365376 ):
366377 event_message = event .message
0 commit comments