66import logging
77import uuid
88from abc import ABC , abstractmethod
9- from collections .abc import AsyncGenerator
9+ from collections .abc import AsyncGenerator , Sequence
1010from typing import TYPE_CHECKING , Any
1111
1212from ag_ui .core import (
5353 merge_tools ,
5454 register_additional_client_tools ,
5555)
56- from ._utils import convert_agui_tools_to_agent_framework , generate_event_id , get_role_value
56+ from ._utils import (
57+ convert_agui_tools_to_agent_framework ,
58+ generate_event_id ,
59+ get_conversation_id_from_update ,
60+ get_role_value ,
61+ )
5762
5863if TYPE_CHECKING :
5964 from ._agent import AgentConfig
6065 from ._confirmation_strategies import ConfirmationStrategy
66+ from ._events import AgentFrameworkEventBridge
67+ from ._orchestration ._state_manager import StateManager
6168
6269
6370logger = logging .getLogger (__name__ )
@@ -92,6 +99,8 @@ def __init__(
9299 self ._last_message = None
93100 self ._run_id : str | None = None
94101 self ._thread_id : str | None = None
102+ self ._supplied_run_id : str | None = None
103+ self ._supplied_thread_id : str | None = None
95104
96105 @property
97106 def messages (self ):
@@ -125,26 +134,66 @@ def last_message(self):
125134 self ._last_message = self .messages [- 1 ]
126135 return self ._last_message
127136
137+ @property
138+ def supplied_run_id (self ) -> str | None :
139+ """Get the supplied run ID, if any."""
140+ if self ._supplied_run_id is None :
141+ self ._supplied_run_id = self .input_data .get ("run_id" ) or self .input_data .get ("runId" )
142+ return self ._supplied_run_id
143+
128144 @property
129145 def run_id (self ) -> str :
130- """Get or generate run ID."""
146+ """Get supplied run ID or generate a new run ID."""
147+ if self ._run_id :
148+ return self ._run_id
149+
150+ if self .supplied_run_id :
151+ self ._run_id = self .supplied_run_id
152+
131153 if self ._run_id is None :
132- self ._run_id = self .input_data .get ("run_id" ) or self .input_data .get ("runId" ) or str (uuid .uuid4 ())
133- # This should never be None after the if block above, but satisfy type checkers
134- if self ._run_id is None : # pragma: no cover
135- raise RuntimeError ("Failed to initialize run_id" )
154+ self ._run_id = str (uuid .uuid4 ())
155+
136156 return self ._run_id
137157
158+ @property
159+ def supplied_thread_id (self ) -> str | None :
160+ """Get the supplied thread ID, if any."""
161+ if self ._supplied_thread_id is None :
162+ self ._supplied_thread_id = self .input_data .get ("thread_id" ) or self .input_data .get ("threadId" )
163+ return self ._supplied_thread_id
164+
138165 @property
139166 def thread_id (self ) -> str :
140- """Get or generate thread ID."""
167+ """Get supplied thread ID or generate a new thread ID."""
168+ if self ._thread_id :
169+ return self ._thread_id
170+
171+ if self .supplied_thread_id :
172+ self ._thread_id = self .supplied_thread_id
173+
141174 if self ._thread_id is None :
142- self ._thread_id = self .input_data .get ("thread_id" ) or self .input_data .get ("threadId" ) or str (uuid .uuid4 ())
143- # This should never be None after the if block above, but satisfy type checkers
144- if self ._thread_id is None : # pragma: no cover
145- raise RuntimeError ("Failed to initialize thread_id" )
175+ self ._thread_id = str (uuid .uuid4 ())
176+
146177 return self ._thread_id
147178
179+ def update_run_id (self , new_run_id : str ) -> None :
180+ """Update the run ID in the context.
181+
182+ Args:
183+ new_run_id: The new run ID to set
184+ """
185+ self ._supplied_run_id = new_run_id
186+ self ._run_id = new_run_id
187+
188+ def update_thread_id (self , new_thread_id : str ) -> None :
189+ """Update the thread ID in the context.
190+
191+ Args:
192+ new_thread_id: The new thread ID to set
193+ """
194+ self ._supplied_thread_id = new_thread_id
195+ self ._thread_id = new_thread_id
196+
148197
149198class Orchestrator (ABC ):
150199 """Base orchestrator for agent execution flows."""
@@ -297,6 +346,28 @@ def can_handle(self, context: ExecutionContext) -> bool:
297346 """
298347 return True
299348
349+ def _create_initial_events (
350+ self , event_bridge : "AgentFrameworkEventBridge" , state_manager : "StateManager"
351+ ) -> Sequence [BaseEvent ]:
352+ """Generate initial events for the run.
353+
354+ Args:
355+ event_bridge: Event bridge for creating events
356+ Returns:
357+ Initial AG-UI events
358+ """
359+ events : list [BaseEvent ] = [event_bridge .create_run_started_event ()]
360+
361+ predict_event = state_manager .predict_state_event ()
362+ if predict_event :
363+ events .append (predict_event )
364+
365+ snapshot_event = state_manager .initial_snapshot_event (event_bridge )
366+ if snapshot_event :
367+ events .append (snapshot_event )
368+
369+ return events
370+
300371 async def run (
301372 self ,
302373 context : ExecutionContext ,
@@ -342,17 +413,11 @@ async def run(
342413 approval_tool_name = approval_tool_name ,
343414 )
344415
345- yield event_bridge .create_run_started_event ()
346-
347- predict_event = state_manager .predict_state_event ()
348- if predict_event :
349- yield predict_event
350-
351- snapshot_event = state_manager .initial_snapshot_event (event_bridge )
352- if snapshot_event :
353- yield snapshot_event
416+ if context .config .use_service_thread :
417+ thread = AgentThread (service_thread_id = context .supplied_thread_id )
418+ else :
419+ thread = AgentThread ()
354420
355- thread = AgentThread ()
356421 thread .metadata = { # type: ignore[attr-defined]
357422 "ag_ui_thread_id" : context .thread_id ,
358423 "ag_ui_run_id" : context .run_id ,
@@ -363,6 +428,8 @@ async def run(
363428 provider_messages = context .messages or []
364429 snapshot_messages = context .snapshot_messages
365430 if not provider_messages :
431+ for event in self ._create_initial_events (event_bridge , state_manager ):
432+ yield event
366433 logger .warning ("No messages provided in AG-UI input" )
367434 yield event_bridge .create_run_finished_event ()
368435 return
@@ -554,13 +621,41 @@ def _build_messages_snapshot(tool_message_id: str | None = None) -> MessagesSnap
554621 confirmation_message = strategy .on_state_rejected ()
555622
556623 message_id = generate_event_id ()
624+ for event in self ._create_initial_events (event_bridge , state_manager ):
625+ yield event
557626 yield TextMessageStartEvent (message_id = message_id , role = "assistant" )
558627 yield TextMessageContentEvent (message_id = message_id , delta = confirmation_message )
559628 yield TextMessageEndEvent (message_id = message_id )
560629 yield event_bridge .create_run_finished_event ()
561630 return
562631
632+ should_recreate_event_bridge = False
563633 async for update in context .agent .run_stream (messages_to_run , ** run_kwargs ):
634+ conv_id = get_conversation_id_from_update (update )
635+ if conv_id and conv_id != context .thread_id :
636+ context .update_thread_id (conv_id )
637+ should_recreate_event_bridge = True
638+
639+ if update .response_id and update .response_id != context .run_id :
640+ context .update_run_id (update .response_id )
641+ should_recreate_event_bridge = True
642+
643+ if should_recreate_event_bridge :
644+ event_bridge = AgentFrameworkEventBridge (
645+ run_id = context .run_id ,
646+ thread_id = context .thread_id ,
647+ predict_state_config = context .config .predict_state_config ,
648+ current_state = current_state ,
649+ skip_text_content = skip_text_content ,
650+ require_confirmation = context .config .require_confirmation ,
651+ approval_tool_name = approval_tool_name ,
652+ )
653+ should_recreate_event_bridge = False
654+
655+ if update_count == 0 :
656+ for event in self ._create_initial_events (event_bridge , state_manager ):
657+ yield event
658+
564659 update_count += 1
565660 logger .info (f"[STREAM] Received update #{ update_count } from agent" )
566661 if all_updates is not None :
@@ -672,6 +767,11 @@ def _build_messages_snapshot(tool_message_id: str | None = None) -> MessagesSnap
672767 yield TextMessageEndEvent (message_id = message_id )
673768 logger .info (f"Emitted conversational message with length={ len (response_dict ['message' ])} " )
674769
770+ if all_updates is not None and len (all_updates ) == 0 :
771+ logger .info ("No updates received from agent - emitting initial events" )
772+ for event in self ._create_initial_events (event_bridge , state_manager ):
773+ yield event
774+
675775 logger .info (f"[FINALIZE] Checking for unclosed message. current_message_id={ event_bridge .current_message_id } " )
676776 if event_bridge .current_message_id :
677777 logger .info (f"[FINALIZE] Emitting TextMessageEndEvent for message_id={ event_bridge .current_message_id } " )
0 commit comments