From 8f0533b8e967a29d3a2714cfd869dd916fdf2312 Mon Sep 17 00:00:00 2001 From: Gabriel Santana Date: Mon, 28 Sep 2026 18:15:00 -0300 Subject: [PATCH] Python: Emit AG-UI RUN_STARTED before the agent runs when IDs are supplied When an AG-UI request supplies both threadId and runId, emit RUN_STARTED, PredictState, the initial state snapshot, and approval results before invoking the agent instead of waiting for its first update. Requests that omit either ID keep waiting for service-generated IDs. The provider conversation ID is still captured from the first update in both cases. The run-start events are built by a helper that does not mutate stream state; the caller keeps ownership of latest_state_snapshot. Fixes #8540 --- .../ag-ui/agent_framework_ag_ui/_agent_run.py | 118 +++++++++++------- .../tests/ag_ui/test_service_thread_id.py | 94 +++++++++++++- 2 files changed, 165 insertions(+), 47 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 c70800c5b8e..cece3ed33dd 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 @@ -743,6 +743,32 @@ def _make_approval_tool_result_events(resolved_approval_results: list[Content]) return events +def _run_start_events( + *, + run_id: str, + thread_id: str, + predict_state_config: dict[str, dict[str, str]], + state_snapshot: dict[str, Any] | None, + resolved_approval_results: list[Content], +) -> list[BaseEvent]: + """Build the events that open an agent run: RUN_STARTED, PredictState, initial state, and approval results.""" + events: list[BaseEvent] = [RunStartedEvent(run_id=run_id, thread_id=thread_id)] + 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)) + if state_snapshot is not None: + events.append(StateSnapshotEvent(snapshot=state_snapshot)) + events.extend(_make_approval_tool_result_events(resolved_approval_results)) + return events + + def _pending_approval_name(entry: ApprovalOccurrence) -> str: return entry.name @@ -3346,13 +3372,32 @@ 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 + # Stream from agent. RunStarted waits for the first update when the request omits + # thread or run IDs, so service-generated IDs can be used instead. run_started_emitted = False + first_update_received = 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 ) + initial_state_snapshot = flow.current_state if state_schema and flow.current_state else None + + # With both IDs supplied there is nothing to wait for, so start the run before + # context providers and the first model call. + if supplied_thread_id and supplied_run_id: + if initial_state_snapshot is not None: + latest_state_snapshot = cast(dict[str, Any], make_json_safe(initial_state_snapshot)) + for event in _run_start_events( + run_id=run_id, + thread_id=thread_id, + predict_state_config=predict_state_config, + state_snapshot=initial_state_snapshot, + resolved_approval_results=resolved_approval_results, + ): + yield event + run_started_emitted = True + # 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 @@ -3376,38 +3421,31 @@ async def _run_agent_stream( if response_format is not None: all_updates.append(update) - # 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_received: + first_update_received = 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 + # 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 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 + if initial_state_snapshot is not None: + latest_state_snapshot = cast(dict[str, Any], make_json_safe(initial_state_snapshot)) + for event in _run_start_events( + run_id=run_id, + thread_id=thread_id, + predict_state_config=predict_state_config, + state_snapshot=initial_state_snapshot, + resolved_approval_results=resolved_approval_results, + ): + yield event + run_started_emitted = True # Feature #4: Detect tool-only messages (no text content) # Emit TextMessageStartEvent to create message context for tool calls @@ -3563,21 +3601,13 @@ async def _run_agent_stream( # If no updates at all, still emit RunStarted if not run_started_emitted: - yield RunStartedEvent(run_id=run_id, thread_id=thread_id) - 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) - if state_schema and flow.current_state: - yield StateSnapshotEvent(snapshot=flow.current_state) - - for event in _make_approval_tool_result_events(resolved_approval_results): + for event in _run_start_events( + run_id=run_id, + thread_id=thread_id, + predict_state_config=predict_state_config, + state_snapshot=initial_state_snapshot, + resolved_approval_results=resolved_approval_results, + ): yield event if response_format is not None and all_updates: from agent_framework import AgentResponse 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..63350ad7ec1 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 @@ -2,11 +2,35 @@ """Tests for service-managed thread IDs, and service-generated response ids.""" +import asyncio +from collections.abc import AsyncIterator from typing import Any -from ag_ui.core import RunFinishedEvent, RunStartedEvent -from agent_framework import Content -from agent_framework._types import AgentResponseUpdate, ChatResponseUpdate +import pytest +from ag_ui.core import CustomEvent, RunFinishedEvent, RunStartedEvent, StateSnapshotEvent +from agent_framework import AgentResponse, Content +from agent_framework._types import AgentResponseUpdate, ChatResponseUpdate, ResponseStream + + +def _gate_agent_stream(agent: Any, monkeypatch: pytest.MonkeyPatch) -> asyncio.Event: + """Hold the agent's streamed updates until the returned event is set.""" + release = asyncio.Event() + original_run = agent.run + + def gated_run(messages: Any = None, *, stream: bool = False, **kwargs: Any) -> Any: + inner = original_run(messages, stream=stream, **kwargs) + if not stream: + return inner + + async def _gated() -> AsyncIterator[AgentResponseUpdate]: + await release.wait() + async for update in inner: + yield update + + return ResponseStream(_gated(), finalizer=AgentResponse.from_updates) + + monkeypatch.setattr(agent, "run", gated_run) + return release async def test_service_thread_id_when_there_are_updates(stub_agent): @@ -80,3 +104,67 @@ 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_agent_updates_when_ids_are_supplied(stub_agent, monkeypatch): + """Supplied thread and run IDs let the run start before the agent produces its first update.""" + from agent_framework.ag_ui import AgentFrameworkAgent + + agent = stub_agent() + release = _gate_agent_stream(agent, monkeypatch) + wrapper = AgentFrameworkAgent( + agent=agent, + state_schema={"document": {"type": "string"}}, + predict_state_config={"document": {"tool": "write_doc", "tool_argument": "content"}}, + ) + + input_data: dict[str, Any] = { + "messages": [{"role": "user", "content": "Hi"}], + "state": {"document": "draft"}, + "threadId": "thread-1", + "runId": "run-1", + } + + events = wrapper.run(input_data) + opening_events = [await asyncio.wait_for(anext(events), timeout=5) for _ in range(3)] + release.set() + remaining_events = [event async for event in events] + + run_started, predict_state, state_snapshot = opening_events + assert isinstance(run_started, RunStartedEvent) + assert (run_started.thread_id, run_started.run_id) == ("thread-1", "run-1") + assert isinstance(predict_state, CustomEvent) + assert predict_state.name == "PredictState" + assert isinstance(state_snapshot, StateSnapshotEvent) + assert state_snapshot.snapshot == {"document": "draft"} + assert not any(isinstance(event, RunStartedEvent) for event in remaining_events) + assert any(event.type == "TEXT_MESSAGE_CONTENT" for event in remaining_events) + assert isinstance(remaining_events[-1], RunFinishedEvent) + assert (remaining_events[-1].thread_id, remaining_events[-1].run_id) == ("thread-1", "run-1") + + +async def test_run_started_waits_for_service_run_id_when_only_thread_id_is_supplied(stub_agent): + """A missing run ID is still taken from the first update before the run starts.""" + from agent_framework.ag_ui import AgentFrameworkAgent + + updates: list[AgentResponseUpdate] = [ + AgentResponseUpdate( + contents=[Content.from_text(text="Hello, user!")], + response_id="resp_67890", + raw_representation=ChatResponseUpdate( + contents=[Content.from_text(text="Hello, user!")], + conversation_id="conv_12345", + response_id="resp_67890", + ), + ) + ] + wrapper = AgentFrameworkAgent(agent=stub_agent(updates=updates)) + + input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}], "threadId": "thread-1"} + + events: list[Any] = [event async for event in wrapper.run(input_data)] + + run_started_events = [event for event in events if isinstance(event, RunStartedEvent)] + assert run_started_events == [events[0]] + assert (events[0].thread_id, events[0].run_id) == ("thread-1", "resp_67890") + assert isinstance(events[-1], RunFinishedEvent)