From 90d9688fb0c7bc720a0cbec369d5927f455bae53 Mon Sep 17 00:00:00 2001 From: Christian Glessner Date: Sat, 19 Sep 2026 08:55:06 +0200 Subject: [PATCH] Python: emit AG-UI RUN_STARTED before the agent runs when thread and run IDs are supplied RUN_STARTED waited for the first agent update so that service-generated IDs could be used. When the request already carries both threadId and runId there is nothing to wait for, so emit it (with PredictState, the initial state snapshot and approval results) before calling the agent. Requests missing either ID keep the previous behaviour. Fixes #8540 --- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 72 +++++++++++-------- .../tests/ag_ui/test_service_thread_id.py | 53 +++++++++++++- 2 files changed, 96 insertions(+), 29 deletions(-) diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py index 1baf001c640..a7bc4bb7ac4 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py @@ -3341,13 +3341,40 @@ async def _run_agent_stream( if state_schema and flow.current_state: messages = _inject_state_context(messages, flow.current_state, state_schema) - # Stream from agent - emit RunStarted after first update to get service IDs + # RunStarted waits for the first update only when the request omitted an ID the + # service could supply. With both IDs supplied it is emitted before the agent runs, + # so clients see the run start while context providers and the first model call + # are still working. run_started_emitted = False + first_update_seen = False provider_thread_id: str | None = None all_updates: list[Any] = [] # Collect for structured output processing latest_state_snapshot: dict[str, Any] | None = ( cast(dict[str, Any], make_json_safe(flow.current_state)) if flow.current_state else None ) + + def _run_start_events() -> list[BaseEvent]: + """RunStarted and the events that follow it once the run IDs are final.""" + nonlocal latest_state_snapshot + events: list[BaseEvent] = [RunStartedEvent(run_id=run_id, thread_id=thread_id)] + # Emit PredictState custom event if configured + if predict_state_config: + predict_state_value = [ + { + "state_key": state_key, + "tool": cfg["tool"], + "tool_argument": cfg["tool_argument"], + } + for state_key, cfg in predict_state_config.items() + ] + events.append(CustomEvent(name="PredictState", value=predict_state_value)) + # Emit initial state snapshot only if we have both state_schema and state + if state_schema and flow.current_state: + latest_state_snapshot = cast(dict[str, Any], make_json_safe(flow.current_state)) + events.append(StateSnapshotEvent(snapshot=flow.current_state)) + events.extend(_make_approval_tool_result_events(resolved_approval_results)) + return events + # Agent middleware can defer the inner run until streaming begins, so the # telemetry override must cover construction, stream resolution, and every pull. # Drive the A2UI runner when one is active (see the gate above); the original agent @@ -3357,6 +3384,11 @@ async def _run_agent_stream( stream_completed = False native_approval_flow_result_ids: set[int] = set() try: + if supplied_thread_id is not None and supplied_run_id is not None: + for event in _run_start_events(): + yield event + run_started_emitted = True + with telemetry_context(): for queued_executions in forwarded_executions.values(): for owner, intent, _ in queued_executions: @@ -3373,36 +3405,20 @@ async def _run_agent_stream( # Use service-generated IDs only when the AG-UI request omitted them. Client-supplied # IDs remain authoritative for lifecycle correlation and thread-scoped persistence. - if not run_started_emitted: + if not first_update_seen: + first_update_seen = True conv_id = get_conversation_id_from_update(update) if conv_id: provider_thread_id = conv_id - if supplied_thread_id is None and conv_id: - thread_id = conv_id - snapshot_session.rebind_thread_id(thread_id) - if supplied_run_id is None and update.response_id: - run_id = update.response_id - # NOW emit RunStarted with proper IDs - yield RunStartedEvent(run_id=run_id, thread_id=thread_id) - # Emit PredictState custom event if configured - if predict_state_config: - predict_state_value = [ - { - "state_key": state_key, - "tool": cfg["tool"], - "tool_argument": cfg["tool_argument"], - } - for state_key, cfg in predict_state_config.items() - ] - yield CustomEvent(name="PredictState", value=predict_state_value) - # Emit initial state snapshot only if we have both state_schema and state - if state_schema and flow.current_state: - latest_state_snapshot = cast(dict[str, Any], make_json_safe(flow.current_state)) - yield StateSnapshotEvent(snapshot=flow.current_state) - run_started_emitted = True - - for event in _make_approval_tool_result_events(resolved_approval_results): - yield event + if not run_started_emitted: + if supplied_thread_id is None and conv_id: + thread_id = conv_id + snapshot_session.rebind_thread_id(thread_id) + if supplied_run_id is None and update.response_id: + run_id = update.response_id + for event in _run_start_events(): + yield event + run_started_emitted = True # Feature #4: Detect tool-only messages (no text content) # Emit TextMessageStartEvent to create message context for tool calls diff --git a/python/packages/ag-ui/tests/ag_ui/test_service_thread_id.py b/python/packages/ag-ui/tests/ag_ui/test_service_thread_id.py index 1c17911a1bb..e3dc079d1f3 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_service_thread_id.py +++ b/python/packages/ag-ui/tests/ag_ui/test_service_thread_id.py @@ -6,7 +6,7 @@ from ag_ui.core import RunFinishedEvent, RunStartedEvent from agent_framework import Content -from agent_framework._types import AgentResponseUpdate, ChatResponseUpdate +from agent_framework._types import AgentResponse, AgentResponseUpdate, ChatResponseUpdate, ResponseStream async def test_service_thread_id_when_there_are_updates(stub_agent): @@ -80,3 +80,54 @@ async def test_service_thread_id_when_user_supplied_thread_id(stub_agent): assert isinstance(events[0], RunStartedEvent) assert events[0].thread_id == "conv_12345" assert isinstance(events[-1], RunFinishedEvent) + + +async def test_run_started_is_emitted_before_the_first_update_when_ids_are_supplied(stub_agent): + """With both IDs supplied, RunStarted does not wait for the agent's first update.""" + import asyncio + + from agent_framework.ag_ui import AgentFrameworkAgent + + gate = asyncio.Event() + + class SlowStartAgent(stub_agent): # type: ignore[misc, valid-type] + """Stands in for context providers and a first model call that take a while.""" + + def run(self, messages: Any = None, *, stream: bool = False, **kwargs: Any) -> Any: + if not stream: + return super().run(messages, stream=stream, **kwargs) + + async def _stream() -> Any: + await gate.wait() + for update in self.updates: + yield update + + return ResponseStream(_stream(), finalizer=AgentResponse.from_updates) + + wrapper = AgentFrameworkAgent(agent=SlowStartAgent()) + input_data = {"messages": [{"role": "user", "content": "Hi"}], "threadId": "thread_1", "runId": "run_1"} + + events = wrapper.run(input_data) + first = await asyncio.wait_for(anext(events), timeout=5) + + assert isinstance(first, RunStartedEvent) + assert (first.thread_id, first.run_id) == ("thread_1", "run_1") + + gate.set() + rest = [event async for event in events] + assert not any(isinstance(event, RunStartedEvent) for event in rest) + assert isinstance(rest[-1], RunFinishedEvent) + + +async def test_run_started_waits_for_the_service_run_id_when_none_is_supplied(stub_agent): + """Without a supplied run ID, RunStarted still carries the service response ID.""" + from agent_framework.ag_ui import AgentFrameworkAgent + + updates = [AgentResponseUpdate(contents=[Content.from_text(text="Hello")], response_id="resp_1")] + wrapper = AgentFrameworkAgent(agent=stub_agent(updates=updates)) + input_data = {"messages": [{"role": "user", "content": "Hi"}], "threadId": "thread_1"} + + events = [event async for event in wrapper.run(input_data)] + + assert isinstance(events[0], RunStartedEvent) + assert (events[0].thread_id, events[0].run_id) == ("thread_1", "resp_1")