Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
118 changes: 74 additions & 44 deletions python/packages/ag-ui/agent_framework_ag_ui/_agent_run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
94 changes: 91 additions & 3 deletions python/packages/ag-ui/tests/ag_ui/test_service_thread_id.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Loading