Python: (ag-ui): Add Workflow Support, Harden Streaming Semantics, and add Dynamic Handoff Demo (#3911)

* fix Workflow.as_agent() streaming regression in ag-ui

* Address PR feedback

* workflows wip

* wip

* wip

* Workflow AG-UI demo

* Fixes for handoff workflow demo

* Fixes to workflows support in AG-UI

* Fixes

* Add headers to some demo files

* Fix comment

* Fixes for store

* Make _input_schema lazy-loaded

* fix mypy

* revert session change to handoff only for now

---------

Co-authored-by: Eduard van Valkenburg <eavanvalkenburg@users.noreply.github.com>
This commit is contained in:
Evan Mattson
2026-02-23 20:59:56 +09:00
committed by GitHub
Unverified
parent b1c7c7c844
commit d8b9409e96
60 changed files with 8349 additions and 512 deletions
@@ -356,3 +356,31 @@ class TestAGUIChatClient:
response = await client.inner_get_response(messages=messages, options=chat_options)
assert response is not None
async def test_interrupt_options_transmission(self, monkeypatch: MonkeyPatch) -> None:
"""Interrupt option fields are forwarded to the HTTP service."""
available_interrupts = [{"id": "req_1", "type": "request_info"}]
resume_payload = {"interrupts": [{"id": "req_1", "value": "approved"}]}
mock_events = [
{"type": "RUN_STARTED", "threadId": "thread_1", "runId": "run_1"},
{"type": "RUN_FINISHED", "threadId": "thread_1", "runId": "run_1"},
]
async def mock_post_run(*args: object, **kwargs: Any) -> AsyncGenerator[dict[str, Any], None]:
assert kwargs.get("available_interrupts") == available_interrupts
assert kwargs.get("resume") == resume_payload
for event in mock_events:
yield event
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
monkeypatch.setattr(client.http_service, "post_run", mock_post_run)
messages = [Message(role="user", text="continue")]
options = {
"available_interrupts": available_interrupts,
"resume": resume_payload,
}
response = await client.inner_get_response(messages=messages, options=options)
assert response is not None
@@ -103,7 +103,7 @@ async def test_run_started_event_emission(streaming_chat_client_stub):
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# First event should be RunStartedEvent
@@ -131,7 +131,7 @@ async def test_predict_state_custom_event_emission(streaming_chat_client_stub):
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Find PredictState event
@@ -144,6 +144,83 @@ async def test_predict_state_custom_event_emission(streaming_chat_client_stub):
assert {"state_key": "summary", "tool": "summarize", "tool_argument": "text"} in predict_value
async def test_usage_content_emits_custom_usage_event(streaming_chat_client_stub):
"""Usage content from the wrapped agent should be surfaced as a custom usage event."""
from agent_framework.ag_ui import AgentFrameworkAgent
async def stream_fn(
messages: MutableSequence[Message], options: dict[str, Any], **kwargs: Any
) -> AsyncIterator[ChatResponseUpdate]:
del messages, options, kwargs
yield ChatResponseUpdate(
contents=[
Content.from_usage(
{
"input_token_count": 10,
"output_token_count": 4,
"total_token_count": 14,
}
)
]
)
agent = Agent(name="usage_agent", instructions="Usage test", client=streaming_chat_client_stub(stream_fn))
wrapper = AgentFrameworkAgent(agent=agent)
events: list[Any] = []
async for event in wrapper.run({"messages": [{"role": "user", "content": "Hi"}]}):
events.append(event)
usage_events = [event for event in events if event.type == "CUSTOM" and event.name == "usage"]
assert len(usage_events) == 1
assert usage_events[0].value["input_token_count"] == 10
assert usage_events[0].value["output_token_count"] == 4
assert usage_events[0].value["total_token_count"] == 14
async def test_multimodal_input_is_forwarded_to_agent_run(streaming_chat_client_stub):
"""Multimodal AG-UI input should be converted and passed through to agent.run."""
from agent_framework.ag_ui import AgentFrameworkAgent
captured_messages: list[Message] = []
async def stream_fn(
messages: MutableSequence[Message], options: dict[str, Any], **kwargs: Any
) -> AsyncIterator[ChatResponseUpdate]:
del options, kwargs
captured_messages[:] = list(messages)
yield ChatResponseUpdate(contents=[Content.from_text(text="Processed multimodal input")])
agent = Agent(name="multimodal_agent", instructions="Multimodal test", client=streaming_chat_client_stub(stream_fn))
wrapper = AgentFrameworkAgent(agent=agent)
input_data = {
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "What is in this image?"},
{
"type": "image",
"source": {"type": "url", "url": "https://example.com/cat.png", "mimeType": "image/png"},
},
],
}
]
}
_ = [event async for event in wrapper.run(input_data)]
assert len(captured_messages) == 1
message = captured_messages[0]
assert message.role == "user"
assert len(message.contents) == 2
assert message.contents[0].type == "text"
assert message.contents[0].text == "What is in this image?"
assert message.contents[1].type == "uri"
assert message.contents[1].uri == "https://example.com/cat.png"
async def test_initial_state_snapshot_with_schema(streaming_chat_client_stub):
"""Test initial StateSnapshotEvent emission when state_schema present."""
from agent_framework.ag_ui import AgentFrameworkAgent
@@ -163,7 +240,7 @@ async def test_initial_state_snapshot_with_schema(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Find StateSnapshotEvent
@@ -190,7 +267,7 @@ async def test_state_initialization_object_type(streaming_chat_client_stub):
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Find StateSnapshotEvent
@@ -217,7 +294,7 @@ async def test_state_initialization_array_type(streaming_chat_client_stub):
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Find StateSnapshotEvent
@@ -243,7 +320,7 @@ async def test_run_finished_event_emission(streaming_chat_client_stub):
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Last event should be RunFinishedEvent
@@ -280,7 +357,7 @@ async def test_tool_result_confirm_changes_accepted(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit text message confirming acceptance
@@ -322,7 +399,7 @@ async def test_tool_result_confirm_changes_rejected(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit text message asking what to change
@@ -362,7 +439,7 @@ async def test_tool_result_function_approval_accepted(streaming_chat_client_stub
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should list enabled steps
@@ -405,7 +482,7 @@ async def test_tool_result_function_approval_rejected(streaming_chat_client_stub
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should ask what to change about the plan
@@ -441,7 +518,7 @@ async def test_thread_metadata_tracking(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# AG-UI internal metadata should NOT be passed to chat client options
@@ -479,7 +556,7 @@ async def test_state_context_injection(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Current state should NOT be passed to chat client options
@@ -502,7 +579,7 @@ async def test_no_messages_provided(streaming_chat_client_stub):
input_data: dict[str, Any] = {"messages": []}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit RunStartedEvent and RunFinishedEvent only
@@ -526,7 +603,7 @@ async def test_message_end_event_emission(streaming_chat_client_stub):
input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should have TextMessageEndEvent before RunFinishedEvent
@@ -556,7 +633,7 @@ async def test_error_handling_with_exception(streaming_chat_client_stub):
input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}]}
with pytest.raises(RuntimeError, match="Simulated failure"):
async for _ in wrapper.run_agent(input_data):
async for _ in wrapper.run(input_data):
pass
@@ -586,7 +663,7 @@ async def test_json_decode_error_in_tool_result(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Orphaned tool result should be sanitized out
@@ -616,7 +693,7 @@ async def test_agent_with_use_service_session_is_false(streaming_chat_client_stu
input_data = {"messages": [{"role": "user", "content": "Hi"}], "thread_id": "conv_123456"}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
assert request_service_session_id is None # type: ignore[attr-defined] (service_session_id should be set)
@@ -643,7 +720,7 @@ async def test_agent_with_use_service_session_is_true(streaming_chat_client_stub
input_data = {"messages": [{"role": "user", "content": "Hi"}], "thread_id": "conv_123456"}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
request_service_session_id = agent.client.last_service_session_id
assert request_service_session_id == "conv_123456" # type: ignore[attr-defined] (service_session_id should be set)
@@ -714,7 +791,7 @@ async def test_function_approval_mode_executes_tool(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Verify the run completed successfully
@@ -802,7 +879,7 @@ async def test_function_approval_mode_rejection(streaming_chat_client_stub):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Verify the run completed
@@ -3,9 +3,18 @@
"""Tests for FastAPI endpoint creation (_endpoint.py)."""
import json
from typing import Any
import pytest
from agent_framework import Agent, ChatResponseUpdate, Content
from ag_ui.core import RunStartedEvent
from agent_framework import (
Agent,
ChatResponseUpdate,
Content,
WorkflowBuilder,
WorkflowContext,
executor,
)
from agent_framework.orchestrations import SequentialBuilder
from fastapi import FastAPI, Header, HTTPException
from fastapi.params import Depends
@@ -13,6 +22,7 @@ from fastapi.testclient import TestClient
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui._agent import AgentFrameworkAgent
from agent_framework_ag_ui._workflow import AgentFrameworkWorkflow
@pytest.fixture
@@ -55,6 +65,32 @@ async def test_add_endpoint_with_wrapped_agent(build_chat_client):
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
async def test_add_endpoint_with_workflow_protocol():
"""Test adding endpoint with native Workflow support."""
@executor(id="start")
async def start(message: Any, ctx: WorkflowContext) -> None:
await ctx.yield_output("Workflow response")
app = FastAPI()
workflow = WorkflowBuilder(start_executor=start).build()
add_agent_framework_fastapi_endpoint(app, workflow, path="/workflow")
client = TestClient(app)
response = client.post("/workflow", json={"messages": [{"role": "user", "content": "Hello"}]})
assert response.status_code == 200
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
content = response.content.decode("utf-8")
lines = [line for line in content.split("\n") if line.startswith("data: ")]
event_types = [json.loads(line[6:]).get("type") for line in lines]
assert "RUN_STARTED" in event_types
assert "TEXT_MESSAGE_CONTENT" in event_types
assert "RUN_FINISHED" in event_types
async def test_endpoint_with_state_schema(build_chat_client):
"""Test endpoint with state_schema parameter."""
app = FastAPI()
@@ -403,8 +439,32 @@ async def test_endpoint_internal_error_handling(build_chat_client):
mock_deepcopy.side_effect = Exception("Simulated internal error")
response = client.post("/error-test", json={"messages": [{"role": "user", "content": "Hello"}]})
assert response.status_code == 500
assert response.json() == {"detail": "An internal error has occurred."}
async def test_endpoint_streaming_error_emits_run_error_event():
"""Streaming exceptions should emit RUN_ERROR instead of terminating silently."""
class FailingStreamWorkflow(AgentFrameworkWorkflow):
async def run(self, input_data: dict[str, Any]):
del input_data
yield RunStartedEvent(run_id="run-1", thread_id="thread-1")
raise RuntimeError("stream exploded")
app = FastAPI()
add_agent_framework_fastapi_endpoint(app, FailingStreamWorkflow(), path="/stream-error")
client = TestClient(app)
response = client.post("/stream-error", json={"messages": [{"role": "user", "content": "Hello"}]})
assert response.status_code == 200
assert response.json() == {"error": "An internal error has occurred."}
content = response.content.decode("utf-8")
lines = [line for line in content.split("\n") if line.startswith("data: ")]
event_types = [json.loads(line[6:]).get("type") for line in lines]
assert "RUN_STARTED" in event_types
assert "RUN_ERROR" in event_types
async def test_endpoint_with_dependencies_blocks_unauthorized(build_chat_client):
@@ -207,6 +207,26 @@ class TestAGUIEventConverter:
assert update.additional_properties["thread_id"] == "thread_123"
assert update.additional_properties["run_id"] == "run_456"
def test_run_finished_event_with_interrupt(self) -> None:
"""RUN_FINISHED interrupt metadata is preserved in additional_properties."""
converter = AGUIEventConverter()
converter.thread_id = "thread_123"
converter.run_id = "run_456"
event = {
"type": "RUN_FINISHED",
"threadId": "thread_123",
"runId": "run_456",
"interrupt": [{"id": "req_1", "value": {"question": "Continue?"}}],
"result": {"status": "paused"},
}
update = converter.convert_event(event)
assert update is not None
assert update.additional_properties["interrupt"] == [{"id": "req_1", "value": {"question": "Continue?"}}]
assert update.additional_properties["result"] == {"status": "paused"}
def test_run_error_event(self) -> None:
"""Test conversion of RUN_ERROR event."""
converter = AGUIEventConverter()
@@ -239,6 +259,37 @@ class TestAGUIEventConverter:
assert update is None
def test_custom_event_conversion(self) -> None:
"""CUSTOM events are converted to update metadata."""
converter = AGUIEventConverter()
event = {
"type": "CUSTOM",
"name": "progress",
"value": {"percent": 10},
}
update = converter.convert_event(event)
assert update is not None
assert update.additional_properties["ag_ui_custom_event"]["name"] == "progress"
assert update.additional_properties["ag_ui_custom_event"]["value"] == {"percent": 10}
assert update.additional_properties["ag_ui_custom_event"]["raw_type"] == "CUSTOM"
def test_custom_event_alias_conversion(self) -> None:
"""CUSTOM_EVENT/custom_event aliases map to CUSTOM behavior."""
converter = AGUIEventConverter()
events = [
{"type": "CUSTOM_EVENT", "name": "alias_upper", "value": {"v": 1}},
{"type": "custom_event", "name": "alias_lower", "value": {"v": 2}},
]
updates = [converter.convert_event(event) for event in events]
assert updates[0] is not None
assert updates[1] is not None
assert updates[0].additional_properties["ag_ui_custom_event"]["raw_type"] == "CUSTOM_EVENT"
assert updates[1].additional_properties["ag_ui_custom_event"]["raw_type"] == "custom_event"
def test_full_conversation_flow(self) -> None:
"""Test complete conversation flow with multiple event types."""
converter = AGUIEventConverter()
@@ -107,8 +107,8 @@ async def test_post_run_successful_streaming(mock_http_client, sample_events):
assert call_args.kwargs["headers"] == {"Accept": "text/event-stream"}
async def test_post_run_with_state_and_tools(mock_http_client):
"""Test posting run with state and tools."""
async def test_post_run_with_state_tools_and_interrupts(mock_http_client):
"""Test posting run with state, tools, and interrupt metadata."""
async def mock_aiter_lines():
return
@@ -127,8 +127,18 @@ async def test_post_run_with_state_and_tools(mock_http_client):
state = {"user_context": {"name": "Alice"}}
tools = [{"type": "function", "function": {"name": "test_tool"}}]
available_interrupts = [{"id": "req_1", "type": "request_info"}]
resume = {"interrupts": [{"id": "req_1", "value": "approved"}]}
async for _ in service.post_run(thread_id="thread_123", run_id="run_456", messages=[], state=state, tools=tools):
async for _ in service.post_run(
thread_id="thread_123",
run_id="run_456",
messages=[],
state=state,
tools=tools,
available_interrupts=available_interrupts,
resume=resume,
):
pass
# Verify state and tools were included in request
@@ -136,6 +146,8 @@ async def test_post_run_with_state_and_tools(mock_http_client):
request_data = call_args.kwargs["json"]
assert request_data["state"] == state
assert request_data["tools"] == tools
assert request_data["availableInterrupts"] == available_interrupts
assert request_data["resume"] == resume
async def test_post_run_http_error(mock_http_client):
@@ -2,7 +2,9 @@
"""Tests for message adapters."""
import base64
import json
import logging
import pytest
from agent_framework import Content, Message
@@ -406,6 +408,101 @@ def test_agui_non_string_content():
assert "nested" in messages[0].contents[0].text
def test_agui_multimodal_legacy_binary_to_agent_framework():
"""Legacy text/binary multimodal content converts to text + media Content."""
messages = agui_messages_to_agent_framework(
[
{
"role": "user",
"content": [
{"type": "text", "text": "See this image"},
{"type": "binary", "mimeType": "image/png", "url": "https://example.com/image.png"},
],
}
]
)
assert len(messages) == 1
assert len(messages[0].contents) == 2
assert messages[0].contents[0].type == "text"
assert messages[0].contents[0].text == "See this image"
assert messages[0].contents[1].type == "uri"
assert messages[0].contents[1].uri == "https://example.com/image.png"
assert messages[0].contents[1].media_type == "image/png"
def test_agui_multimodal_draft_source_base64_to_agent_framework():
"""Draft-style media source payload converts into data Content."""
payload = base64.b64encode(b"abc").decode("utf-8")
messages = agui_messages_to_agent_framework(
[
{
"role": "user",
"content": [
{
"type": "audio",
"source": {"type": "base64", "data": payload, "mimeType": "audio/wav"},
}
],
}
]
)
assert len(messages) == 1
assert len(messages[0].contents) == 1
assert messages[0].contents[0].type == "data"
assert messages[0].contents[0].media_type == "audio/wav"
assert isinstance(messages[0].contents[0].uri, str)
assert messages[0].contents[0].uri.startswith("data:audio/wav;base64,")
def test_agui_multimodal_invalid_base64_logs_warning(caplog):
"""Malformed base64 payloads should log and fall back to data URI."""
with caplog.at_level(logging.WARNING):
messages = agui_messages_to_agent_framework(
[
{
"role": "user",
"content": [
{
"type": "image",
"source": {"type": "base64", "data": "abc", "mimeType": "image/png"},
}
],
}
]
)
assert len(messages) == 1
assert len(messages[0].contents) == 1
assert messages[0].contents[0].type in {"data", "uri"}
assert messages[0].contents[0].uri == "data:image/png;base64,abc"
assert any("Failed to decode AG-UI media payload as base64" in record.message for record in caplog.records)
def test_agui_multimodal_mixed_order_preserved():
"""Mixed text/media multimodal input keeps content ordering."""
messages = agui_messages_to_agent_framework(
[
{
"role": "user",
"content": [
{"type": "text", "text": "First"},
{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}},
{"type": "text", "text": "Last"},
],
}
]
)
assert len(messages[0].contents) == 3
assert messages[0].contents[0].type == "text"
assert messages[0].contents[0].text == "First"
assert messages[0].contents[1].type == "uri"
assert messages[0].contents[2].type == "text"
assert messages[0].contents[2].text == "Last"
def test_agui_message_without_id():
"""Test message without ID field."""
messages = agui_messages_to_agent_framework([{"role": "user", "content": "No ID"}])
@@ -414,6 +511,31 @@ def test_agui_message_without_id():
assert messages[0].message_id is None
def test_agui_snapshot_format_preserves_multimodal_content():
"""Snapshot normalization emits legacy binary parts for multimodal content."""
normalized = agui_messages_to_snapshot_format(
[
{
"role": "user",
"content": [
{"type": "input_text", "text": "Caption"},
{
"type": "image",
"source": {"type": "url", "url": "https://example.com/image.png", "mime_type": "image/png"},
},
],
}
]
)
assert isinstance(normalized[0]["content"], list)
content_parts = normalized[0]["content"]
assert content_parts[0]["type"] == "text"
assert content_parts[1]["type"] == "binary"
assert content_parts[1]["mimeType"] == "image/png"
assert content_parts[1]["url"] == "https://example.com/image.png"
def test_agui_with_tool_calls_to_agent_framework():
"""Assistant message with tool_calls is converted to FunctionCallContent."""
agui_msg = {
@@ -66,7 +66,7 @@ def test_convert_approval_results_to_tool_messages() -> None:
results ended up in user messages instead of tool messages, causing OpenAI to
reject the request with 'tool_call_ids did not have response messages'.
"""
from agent_framework_ag_ui._run import _convert_approval_results_to_tool_messages
from agent_framework_ag_ui._agent_run import _convert_approval_results_to_tool_messages
# Simulate what happens after _resolve_approval_responses:
# A user message contains function_result content (the executed tool result)
@@ -106,7 +106,7 @@ def test_convert_approval_results_preserves_other_user_content() -> None:
the function_result content should be extracted to a tool message while the
remaining content stays in the user message.
"""
from agent_framework_ag_ui._run import _convert_approval_results_to_tool_messages
from agent_framework_ag_ui._agent_run import _convert_approval_results_to_tool_messages
messages = [
Message(
@@ -0,0 +1,24 @@
# Copyright (c) Microsoft. All rights reserved.
"""Public export coverage for AG-UI package surfaces."""
def test_agent_framework_ag_ui_exports_workflow() -> None:
"""Runtime package should export AgentFrameworkWorkflow."""
from agent_framework_ag_ui import AgentFrameworkWorkflow
assert AgentFrameworkWorkflow.__name__ == "AgentFrameworkWorkflow"
def test_core_ag_ui_lazy_exports_include_only_stable_api() -> None:
"""Core facade should expose only the stable high-level AG-UI API."""
from agent_framework import ag_ui
assert hasattr(ag_ui, "AgentFrameworkWorkflow")
assert hasattr(ag_ui, "AgentFrameworkAgent")
assert hasattr(ag_ui, "AGUIChatClient")
assert hasattr(ag_ui, "add_agent_framework_fastapi_endpoint")
assert not hasattr(ag_ui, "WorkflowFactory")
assert not hasattr(ag_ui, "AGUIRequest")
assert not hasattr(ag_ui, "RunMetadata")
+126 -27
View File
@@ -1,6 +1,6 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for _run.py helper functions and FlowState."""
"""Tests for _agent_run.py helper functions and FlowState."""
import pytest
from ag_ui.core import (
@@ -10,17 +10,25 @@ from ag_ui.core import (
from agent_framework import AgentResponseUpdate, Content, Message, ResponseStream
from agent_framework.exceptions import AgentInvalidResponseException
from agent_framework_ag_ui._run import (
FlowState,
from agent_framework_ag_ui._agent_run import (
_build_safe_metadata,
_create_state_context_message,
_emit_content,
_emit_tool_result,
_has_only_tool_calls,
_inject_state_context,
_normalize_response_stream,
_resume_to_tool_messages,
_should_suppress_intermediate_snapshot,
)
from agent_framework_ag_ui._run_common import (
FlowState,
_build_run_finished_event,
_emit_approval_request,
_emit_content,
_emit_text,
_emit_tool_call,
_emit_tool_result,
_extract_resume_payload,
_has_only_tool_calls,
)
class TestBuildSafeMetadata:
@@ -150,6 +158,7 @@ class TestFlowState:
assert flow.tool_calls_by_id == {}
assert flow.tool_results == []
assert flow.tool_calls_ended == set()
assert flow.interrupts == []
def test_get_tool_name(self):
"""Tests get_tool_name method."""
@@ -308,13 +317,11 @@ class TestInjectStateContext:
assert "Hello" in result[2].contents[0].text
# Additional tests for _run.py functions
# Additional tests for _agent_run.py functions
def test_emit_text_basic():
"""Test _emit_text emits correct events."""
from agent_framework_ag_ui._run import _emit_text
flow = FlowState()
content = Content.from_text("Hello world")
@@ -327,8 +334,6 @@ def test_emit_text_basic():
def test_emit_text_skip_empty():
"""Test _emit_text skips empty text."""
from agent_framework_ag_ui._run import _emit_text
flow = FlowState()
content = Content.from_text("")
@@ -339,8 +344,6 @@ def test_emit_text_skip_empty():
def test_emit_text_continues_existing_message():
"""Test _emit_text continues existing message."""
from agent_framework_ag_ui._run import _emit_text
flow = FlowState()
flow.message_id = "existing-id"
content = Content.from_text("more text")
@@ -351,10 +354,21 @@ def test_emit_text_continues_existing_message():
assert flow.message_id == "existing-id"
def test_emit_text_skips_duplicate_full_message_delta():
"""Test _emit_text skips replayed full-message chunks on an open message."""
flow = FlowState()
flow.message_id = "existing-id"
flow.accumulated_text = "Case complete."
content = Content.from_text("Case complete.")
events = _emit_text(content, flow)
assert events == []
assert flow.accumulated_text == "Case complete."
def test_emit_text_skips_when_waiting_for_approval():
"""Test _emit_text skips when waiting for approval."""
from agent_framework_ag_ui._run import _emit_text
flow = FlowState()
flow.waiting_for_approval = True
content = Content.from_text("should skip")
@@ -366,8 +380,6 @@ def test_emit_text_skips_when_waiting_for_approval():
def test_emit_text_skips_when_skip_text_flag():
"""Test _emit_text skips with skip_text flag."""
from agent_framework_ag_ui._run import _emit_text
flow = FlowState()
content = Content.from_text("should skip")
@@ -378,8 +390,6 @@ def test_emit_text_skips_when_skip_text_flag():
def test_emit_tool_call_basic():
"""Test _emit_tool_call emits correct events."""
from agent_framework_ag_ui._run import _emit_tool_call
flow = FlowState()
content = Content.from_function_call(
call_id="call_123",
@@ -396,8 +406,6 @@ def test_emit_tool_call_basic():
def test_emit_tool_call_generates_id():
"""Test _emit_tool_call generates ID when not provided."""
from agent_framework_ag_ui._run import _emit_tool_call
flow = FlowState()
# Create content without call_id
content = Content(type="function_call", name="test_tool", arguments="{}")
@@ -452,9 +460,100 @@ def test_emit_tool_result_no_open_message():
assert len(text_end_events) == 0
def test_emit_tool_result_serializes_non_string_result():
"""Non-string tool results should be serialized before emitting TOOL_CALL_RESULT."""
flow = FlowState()
content = Content.from_function_result(call_id="call_789", result={"ok": True, "items": [1, 2]})
events = _emit_tool_result(content, flow, predictive_handler=None)
result_event = next(event for event in events if getattr(event, "type", None) == "TOOL_CALL_RESULT")
assert isinstance(result_event.content, str)
assert '"ok": true' in result_event.content
assert flow.tool_results[0]["content"] == result_event.content
def test_emit_content_usage_emits_custom_usage_event():
"""Usage content should be emitted as a custom usage event."""
flow = FlowState()
content = Content.from_usage({"input_token_count": 3, "output_token_count": 2, "total_token_count": 5})
events = _emit_content(content, flow)
assert len(events) == 1
assert events[0].type == "CUSTOM"
assert events[0].name == "usage"
assert events[0].value["total_token_count"] == 5
def test_emit_approval_request_populates_interrupt_metadata():
"""Approval requests should populate FlowState interrupts for RUN_FINISHED metadata."""
flow = FlowState(message_id="msg-1")
function_call = Content.from_function_call(call_id="call_123", name="write_doc", arguments={"content": "x"})
approval_content = Content.from_function_approval_request(id="approval_1", function_call=function_call)
_emit_approval_request(approval_content, flow)
assert flow.waiting_for_approval is True
assert len(flow.interrupts) == 1
assert flow.interrupts[0]["id"] == "call_123"
assert flow.interrupts[0]["value"]["type"] == "function_approval_request"
def test_resume_to_tool_messages_from_interrupts_payload():
"""Resume payload interrupt responses map to tool messages."""
resume = {
"interrupts": [
{"id": "req_1", "value": {"accepted": True, "steps": []}},
{"id": "req_2", "value": "plain value"},
]
}
messages = _resume_to_tool_messages(resume)
assert len(messages) == 2
assert messages[0]["role"] == "tool"
assert messages[0]["toolCallId"] == "req_1"
assert '"accepted": true' in messages[0]["content"]
assert messages[1]["content"] == "plain value"
def test_extract_resume_payload_prefers_top_level_resume():
"""Top-level resume should take precedence over forwarded props."""
payload = {
"resume": {"interrupts": [{"id": "req_1", "value": "approved"}]},
"forwarded_props": {"command": {"resume": "ignored"}},
}
result = _extract_resume_payload(payload)
assert result == {"interrupts": [{"id": "req_1", "value": "approved"}]}
def test_extract_resume_payload_reads_forwarded_command_resume():
"""Forwarded command.resume should be treated as a resume payload."""
payload = {
"forwarded_props": {
"command": {"resume": '{"airline":"KLM","departure":"Amsterdam (AMS)","arrival":"San Francisco (SFO)"}'}
}
}
result = _extract_resume_payload(payload)
assert isinstance(result, str)
assert "KLM" in result
def test_build_run_finished_event_with_interrupt():
"""RUN_FINISHED helper should preserve interrupt payloads."""
event = _build_run_finished_event("run-1", "thread-1", interrupts=[{"id": "req_1", "value": {"x": 1}}])
dumped = event.model_dump()
assert dumped["run_id"] == "run-1"
assert dumped["thread_id"] == "thread-1"
assert dumped["interrupt"] == [{"id": "req_1", "value": {"x": 1}}]
def test_extract_approved_state_updates_no_handler():
"""Test _extract_approved_state_updates returns empty with no handler."""
from agent_framework_ag_ui._run import _extract_approved_state_updates
from agent_framework_ag_ui._agent_run import _extract_approved_state_updates
messages = [Message(role="user", contents=[Content.from_text("Hello")])]
result = _extract_approved_state_updates(messages, None)
@@ -463,8 +562,8 @@ def test_extract_approved_state_updates_no_handler():
def test_extract_approved_state_updates_no_approval():
"""Test _extract_approved_state_updates returns empty when no approval content."""
from agent_framework_ag_ui._agent_run import _extract_approved_state_updates
from agent_framework_ag_ui._orchestration._predictive_state import PredictiveStateHandler
from agent_framework_ag_ui._run import _extract_approved_state_updates
handler = PredictiveStateHandler(predict_state_config={"doc": {"tool": "write", "tool_argument": "content"}})
messages = [Message(role="user", contents=[Content.from_text("Hello")])]
@@ -481,7 +580,7 @@ class TestBuildMessagesSnapshot:
This is a regression test for issue #3619 where tool calls and content
were incorrectly merged into a single assistant message.
"""
from agent_framework_ag_ui._run import FlowState, _build_messages_snapshot
from agent_framework_ag_ui._agent_run import FlowState, _build_messages_snapshot
flow = FlowState()
flow.message_id = "msg-123"
@@ -518,7 +617,7 @@ class TestBuildMessagesSnapshot:
def test_only_tool_calls_no_text(self):
"""Test snapshot with only tool calls and no accumulated text."""
from agent_framework_ag_ui._run import FlowState, _build_messages_snapshot
from agent_framework_ag_ui._agent_run import FlowState, _build_messages_snapshot
flow = FlowState()
flow.message_id = "msg-123"
@@ -538,7 +637,7 @@ class TestBuildMessagesSnapshot:
def test_only_text_no_tool_calls(self):
"""Test snapshot with only text and no tool calls."""
from agent_framework_ag_ui._run import FlowState, _build_messages_snapshot
from agent_framework_ag_ui._agent_run import FlowState, _build_messages_snapshot
flow = FlowState()
flow.message_id = "msg-123"
@@ -558,7 +657,7 @@ class TestBuildMessagesSnapshot:
def test_preserves_snapshot_messages(self):
"""Test that existing snapshot messages are preserved."""
from agent_framework_ag_ui._run import FlowState, _build_messages_snapshot
from agent_framework_ag_ui._agent_run import FlowState, _build_messages_snapshot
flow = FlowState()
flow.pending_tool_calls = []
@@ -32,7 +32,7 @@ async def test_service_thread_id_when_there_are_updates(stub_agent):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
assert isinstance(events[0], RunStartedEvent)
@@ -54,7 +54,7 @@ async def test_service_thread_id_when_no_user_message(stub_agent):
}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
assert len(events) == 2
@@ -74,7 +74,7 @@ async def test_service_thread_id_when_user_supplied_thread_id(stub_agent):
input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}], "threadId": "conv_12345"}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
assert isinstance(events[0], RunStartedEvent)
@@ -52,7 +52,7 @@ async def test_structured_output_with_recipe(streaming_chat_client_stub, stream_
input_data = {"messages": [{"role": "user", "content": "Make pasta"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit StateSnapshotEvent with recipe
@@ -94,7 +94,7 @@ async def test_structured_output_with_steps(streaming_chat_client_stub, stream_f
input_data = {"messages": [{"role": "user", "content": "Do steps"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit StateSnapshotEvent with steps
@@ -129,7 +129,7 @@ async def test_structured_output_with_no_schema_match(streaming_chat_client_stub
input_data = {"messages": [{"role": "user", "content": "Generate data"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit StateSnapshotEvent but with no state updates since no schema fields match
@@ -164,7 +164,7 @@ async def test_structured_output_without_schema(streaming_chat_client_stub, stre
input_data = {"messages": [{"role": "user", "content": "Generate data"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit StateSnapshotEvent with both data and info fields
@@ -194,7 +194,7 @@ async def test_no_structured_output_when_no_response_format(streaming_chat_clien
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit text content normally
@@ -224,7 +224,7 @@ async def test_structured_output_with_message_field(streaming_chat_client_stub,
input_data = {"messages": [{"role": "user", "content": "Make salad"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should emit the message as text
@@ -256,7 +256,7 @@ async def test_empty_updates_no_structured_processing(streaming_chat_client_stub
input_data = {"messages": [{"role": "user", "content": "Test"}]}
events: list[Any] = []
async for event in wrapper.run_agent(input_data):
async for event in wrapper.run(input_data):
events.append(event)
# Should only have start and end events
@@ -0,0 +1,238 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for the subgraphs example agent used by Dojo."""
from __future__ import annotations
import json
from typing import Any
from agent_framework_ag_ui_examples.agents.subgraphs_agent import subgraphs_agent
async def _run(agent: Any, payload: dict[str, Any]) -> list[Any]:
return [event async for event in agent.run(payload)]
async def test_subgraphs_example_initial_run_emits_flight_interrupt() -> None:
"""Initial run should publish flight options and pause with an interrupt."""
agent = subgraphs_agent()
events = await _run(
agent,
{
"thread_id": "thread-subgraphs-initial",
"run_id": "run-initial",
"messages": [{"role": "user", "content": "Help me plan a trip to San Francisco"}],
},
)
event_types = [event.type for event in events]
assert event_types[0] == "RUN_STARTED"
assert "STATE_SNAPSHOT" in event_types
assert "STEP_STARTED" in event_types
assert "STEP_FINISHED" in event_types
assert "TEXT_MESSAGE_CONTENT" in event_types
assert "RUN_FINISHED" in event_types
started_steps = [event.step_name for event in events if event.type == "STEP_STARTED"]
finished_steps = [event.step_name for event in events if event.type == "STEP_FINISHED"]
assert "supervisor_agent" in started_steps
assert "flights_agent" in started_steps
assert "supervisor_agent" in finished_steps
assert "flights_agent" in finished_steps
finished = [event for event in events if event.type == "RUN_FINISHED"][0]
interrupt_payload = finished.model_dump().get("interrupt")
assert isinstance(interrupt_payload, list)
assert interrupt_payload
assert interrupt_payload[0]["value"]["agent"] == "flights"
assert len(interrupt_payload[0]["value"]["options"]) == 2
assert interrupt_payload[0]["value"]["options"][0]["airline"] == "KLM"
custom_event_names = [event.name for event in events if event.type == "CUSTOM"]
assert "WorkflowInterruptEvent" in custom_event_names
async def test_subgraphs_example_resume_flow_reaches_completion() -> None:
"""Flight + hotel resume payloads should complete the itinerary state."""
agent = subgraphs_agent()
thread_id = "thread-subgraphs-complete"
first_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-1",
"messages": [{"role": "user", "content": "I want to visit San Francisco from Amsterdam"}],
},
)
first_interrupt = [event for event in first_events if event.type == "RUN_FINISHED"][0].model_dump()["interrupt"][0]
second_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-2",
"resume": {
"interrupts": [
{
"id": first_interrupt["id"],
"value": json.dumps(
{
"airline": "United",
"departure": "Amsterdam (AMS)",
"arrival": "San Francisco (SFO)",
"price": "$720",
"duration": "12h 15m",
}
),
}
]
},
},
)
second_finished = [event for event in second_events if event.type == "RUN_FINISHED"][0].model_dump()
second_interrupt = second_finished.get("interrupt")
assert isinstance(second_interrupt, list)
assert second_interrupt[0]["value"]["agent"] == "hotels"
third_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-3",
"resume": {
"interrupts": [
{
"id": second_interrupt[0]["id"],
"value": json.dumps(
{
"name": "The Ritz-Carlton",
"location": "Nob Hill",
"price_per_night": "$550/night",
"rating": "4.8 stars",
}
),
}
]
},
},
)
third_finished = [event for event in third_events if event.type == "RUN_FINISHED"][0].model_dump()
assert "interrupt" not in third_finished
snapshots = [event.snapshot for event in third_events if event.type == "STATE_SNAPSHOT"]
assert snapshots
final_snapshot = snapshots[-1]
assert final_snapshot["planning_step"] == "complete"
assert final_snapshot["active_agent"] == "supervisor"
assert final_snapshot["itinerary"]["flight"]["airline"] == "United"
assert final_snapshot["itinerary"]["hotel"]["name"] == "The Ritz-Carlton"
assert len(final_snapshot["experiences"]) == 4
async def test_subgraphs_example_requires_structured_resume_for_selection() -> None:
"""Agent should re-issue interrupts when user sends plain text instead of resume payload."""
agent = subgraphs_agent()
thread_id = "thread-subgraphs-text"
first_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-a",
"messages": [{"role": "user", "content": "Plan a trip for me"}],
},
)
first_finished = [event for event in first_events if event.type == "RUN_FINISHED"][0].model_dump()
assert isinstance(first_finished.get("interrupt"), list)
assert first_finished["interrupt"][0]["value"]["agent"] == "flights"
second_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-b",
"messages": [{"role": "user", "content": "Let's do the United flight"}],
},
)
second_finished = [event for event in second_events if event.type == "RUN_FINISHED"][0].model_dump()
assert isinstance(second_finished.get("interrupt"), list)
assert second_finished["interrupt"][0]["value"]["agent"] == "flights"
assert "TOOL_CALL_START" in [event.type for event in second_events]
assert "TEXT_MESSAGE_CONTENT" not in [event.type for event in second_events]
third_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-c",
"resume": {
"interrupts": [
{
"id": second_finished["interrupt"][0]["id"],
"value": json.dumps(
{
"airline": "United",
"departure": "Amsterdam (AMS)",
"arrival": "San Francisco (SFO)",
"price": "$720",
"duration": "12h 15m",
}
),
}
]
},
},
)
third_finished = [event for event in third_events if event.type == "RUN_FINISHED"][0].model_dump()
assert isinstance(third_finished.get("interrupt"), list)
assert third_finished["interrupt"][0]["value"]["agent"] == "hotels"
third_snapshots = [event.snapshot for event in third_events if event.type == "STATE_SNAPSHOT"]
assert third_snapshots[-1]["itinerary"]["flight"]["airline"] == "United"
async def test_subgraphs_example_forwarded_command_resume_reaches_hotels_interrupt() -> None:
"""CopilotKit-style forwarded command.resume should continue workflow interrupts."""
agent = subgraphs_agent()
thread_id = "thread-subgraphs-forwarded-resume"
first_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-forwarded-1",
"messages": [{"role": "user", "content": "Plan my trip"}],
},
)
first_interrupt = [event for event in first_events if event.type == "RUN_FINISHED"][0].model_dump()["interrupt"][0]
second_events = await _run(
agent,
{
"thread_id": thread_id,
"run_id": "run-forwarded-2",
"messages": [],
"forwarded_props": {
"command": {
"resume": json.dumps(
{
"airline": "KLM",
"departure": "Amsterdam (AMS)",
"arrival": "San Francisco (SFO)",
"price": "$650",
"duration": "11h 30m",
}
)
}
},
},
)
second_finished = [event for event in second_events if event.type == "RUN_FINISHED"][0].model_dump()
second_interrupt = second_finished.get("interrupt")
assert isinstance(second_interrupt, list)
assert second_interrupt[0]["value"]["agent"] == "hotels"
assert second_interrupt[0]["id"] != first_interrupt["id"]
@@ -183,6 +183,21 @@ class TestAGUIRequest:
assert request.forwarded_props == {"custom_key": "custom_value"}
assert request.parent_run_id == "parent-run-789"
def test_agui_request_camel_case_aliases(self) -> None:
"""Test AGUIRequest accepts camelCase aliases from AG-UI HTTP clients."""
request = AGUIRequest(
messages=[{"role": "user", "content": "Hello"}],
runId="run-camel-1",
threadId="thread-camel-1",
forwardedProps={"k": "v"},
parentRunId="parent-camel-1",
)
assert request.run_id == "run-camel-1"
assert request.thread_id == "thread-camel-1"
assert request.forwarded_props == {"k": "v"}
assert request.parent_run_id == "parent-camel-1"
def test_agui_request_model_dump_excludes_none(self) -> None:
"""Test that model_dump(exclude_none=True) excludes None fields."""
request = AGUIRequest(
@@ -223,3 +238,15 @@ class TestAGUIRequest:
assert dumped["context"] == [{"type": "snippet", "content": "code here"}]
assert dumped["forwarded_props"] == {"auth_token": "secret", "user_id": "user-1"}
assert dumped["parent_run_id"] == "parent-456"
def test_agui_request_available_interrupts_alias_round_trip(self) -> None:
"""availableInterrupts should deserialize, while dumps remain snake_case."""
request = AGUIRequest(
messages=[{"role": "user", "content": "Hello"}],
availableInterrupts=[{"id": "req_1", "value": {"choice": "A"}}],
)
assert request.available_interrupts == [{"id": "req_1", "value": {"choice": "A"}}]
dumped = request.model_dump(exclude_none=True)
assert dumped["available_interrupts"] == [{"id": "req_1", "value": {"choice": "A"}}]
assert "availableInterrupts" not in dumped
@@ -0,0 +1,112 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for AgentFrameworkWorkflow wrapper behavior."""
from __future__ import annotations
from typing import Any, cast
import pytest
from agent_framework import Workflow, WorkflowBuilder, WorkflowContext, executor
from agent_framework_ag_ui import AgentFrameworkWorkflow
async def _run(agent: AgentFrameworkWorkflow, payload: dict[str, Any]) -> list[Any]:
return [event async for event in agent.run(payload)]
async def test_workflow_wrapper_rejects_workflow_and_factory_at_once() -> None:
"""Workflow wrapper should reject ambiguous workflow source configuration."""
@executor(id="start")
async def start(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.yield_output("ok")
workflow = WorkflowBuilder(start_executor=start).build()
with pytest.raises(ValueError, match="workflow_factory"):
AgentFrameworkWorkflow(workflow=workflow, workflow_factory=lambda _thread_id: workflow)
async def test_workflow_wrapper_factory_is_thread_scoped() -> None:
"""Thread-scoped workflow factories should isolate workflow instances by thread id."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.request_info({"message": "Choose an option", "options": ["a", "b"]}, dict, request_id="choice")
factory_calls: dict[str, int] = {}
def workflow_factory(thread_id: str) -> Workflow:
factory_calls[thread_id] = factory_calls.get(thread_id, 0) + 1
return WorkflowBuilder(start_executor=requester).build()
agent = AgentFrameworkWorkflow(workflow_factory=workflow_factory)
first_events = await _run(
agent,
{
"thread_id": "thread-a",
"messages": [{"role": "user", "content": "start"}],
},
)
first_finished = [event for event in first_events if event.type == "RUN_FINISHED"][0].model_dump()
first_interrupt = first_finished.get("interrupt")
assert isinstance(first_interrupt, list)
assert first_interrupt[0]["id"] == "choice"
assert factory_calls["thread-a"] == 1
second_events = await _run(
agent,
{
"thread_id": "thread-a",
"messages": [],
"resume": {"interrupts": [{"id": "choice", "value": {"selection": "a"}}]},
},
)
second_types = [event.type for event in second_events]
assert "RUN_ERROR" not in second_types
second_finished = [event for event in second_events if event.type == "RUN_FINISHED"][0].model_dump()
assert "interrupt" not in second_finished
assert factory_calls["thread-a"] == 1
third_events = await _run(
agent,
{
"thread_id": "thread-b",
"messages": [{"role": "user", "content": "start"}],
},
)
third_finished = [event for event in third_events if event.type == "RUN_FINISHED"][0].model_dump()
third_interrupt = third_finished.get("interrupt")
assert isinstance(third_interrupt, list)
assert third_interrupt[0]["id"] == "choice"
assert factory_calls["thread-b"] == 1
agent.clear_thread_workflow("thread-a")
await _run(
agent,
{
"thread_id": "thread-a",
"messages": [{"role": "user", "content": "restart"}],
},
)
assert factory_calls["thread-a"] == 2
async def test_workflow_wrapper_without_workflow_raises_not_implemented() -> None:
"""Without workflow/workflow_factory, run should raise NotImplementedError."""
agent = AgentFrameworkWorkflow()
with pytest.raises(NotImplementedError, match="No workflow is attached"):
_ = [event async for event in agent.run({"messages": [{"role": "user", "content": "start"}]})]
async def test_workflow_wrapper_factory_return_type_is_validated() -> None:
"""Factory outputs must be Workflow instances."""
agent = AgentFrameworkWorkflow(workflow_factory=lambda _thread_id: cast(Any, object()))
with pytest.raises(TypeError, match="workflow_factory must return a Workflow instance"):
_ = [event async for event in agent.run({"thread_id": "thread-a", "messages": []})]
@@ -0,0 +1,679 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for native workflow AG-UI runner."""
import json
from types import SimpleNamespace
from typing import Any, cast
from ag_ui.core import EventType, StateSnapshotEvent
from agent_framework import (
AgentResponse,
Content,
Executor,
Message,
WorkflowBuilder,
WorkflowContext,
WorkflowEvent,
executor,
handler,
response_handler,
)
from typing_extensions import Never
from agent_framework_ag_ui._workflow_run import (
_coerce_message,
_coerce_response_for_request,
run_workflow_stream,
)
class ProgressEvent(WorkflowEvent):
"""Custom workflow event used to validate CUSTOM mapping."""
def __init__(self, progress: int) -> None:
super().__init__("custom_progress", data={"progress": progress})
async def test_workflow_run_maps_custom_and_text_events():
"""Custom workflow events and yielded text are mapped to AG-UI events."""
@executor(id="start")
async def start(message: Any, ctx: WorkflowContext[Never, str]) -> None:
await ctx.add_event(ProgressEvent(10))
await ctx.yield_output("Hello workflow")
workflow = WorkflowBuilder(start_executor=start).build()
input_data = {"messages": [{"role": "user", "content": "go"}]}
events = [event async for event in run_workflow_stream(input_data, workflow)]
event_types = [event.type for event in events]
assert "RUN_STARTED" in event_types
assert "CUSTOM" in event_types
assert "TEXT_MESSAGE_CONTENT" in event_types
assert "STEP_STARTED" in event_types
assert "STEP_FINISHED" in event_types
assert "RUN_FINISHED" in event_types
custom_events = [event for event in events if event.type == "CUSTOM" and event.name == "custom_progress"]
assert len(custom_events) == 1
assert custom_events[0].value == {"progress": 10}
async def test_workflow_run_request_info_emits_interrupt_and_resume_works():
"""request_info should emit interrupt metadata and resume should continue run."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
await ctx.request_info("Need approval", str)
workflow = WorkflowBuilder(start_executor=requester).build()
first_run_events = [
event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)
]
run_finished_events = [event for event in first_run_events if event.type == "RUN_FINISHED"]
assert len(run_finished_events) == 1
interrupt_payload = run_finished_events[0].model_dump().get("interrupt")
assert isinstance(interrupt_payload, list)
assert len(interrupt_payload) == 1
request_id = str(interrupt_payload[0]["id"])
assert request_id
resumed_events = [
event
async for event in run_workflow_stream(
{"messages": [], "resume": {"interrupts": [{"id": request_id, "value": "approved"}]}},
workflow,
)
]
resumed_types = [event.type for event in resumed_events]
assert "RUN_STARTED" in resumed_types
assert "RUN_FINISHED" in resumed_types
assert "RUN_ERROR" not in resumed_types
async def test_workflow_run_request_info_closes_open_text_message() -> None:
"""Text output should end before request_info interrupt events begin."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.yield_output("Please confirm this action.")
await ctx.request_info("Need approval", str, request_id="approval-1")
workflow = WorkflowBuilder(start_executor=requester).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
content_index = next(i for i, event in enumerate(events) if event.type == "TEXT_MESSAGE_CONTENT")
end_index = next(i for i, event in enumerate(events) if event.type == "TEXT_MESSAGE_END")
request_start_index = next(
i
for i, event in enumerate(events)
if event.type == "TOOL_CALL_START" and getattr(event, "tool_call_id", None) == "approval-1"
)
assert content_index < end_index < request_start_index
async def test_workflow_run_request_info_interrupt_uses_raw_dict_value():
"""Dict request payloads should be surfaced directly in RUN_FINISHED.interrupt.value."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
await ctx.request_info(
{
"message": "Choose a flight",
"options": [{"airline": "KLM"}],
"recommendation": {"airline": "KLM"},
"agent": "flights",
},
dict,
request_id="flights-choice",
)
workflow = WorkflowBuilder(start_executor=requester).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
run_finished = [event for event in events if event.type == "RUN_FINISHED"][0].model_dump()
interrupt_payload = run_finished.get("interrupt")
assert isinstance(interrupt_payload, list)
assert interrupt_payload[0]["id"] == "flights-choice"
assert interrupt_payload[0]["value"]["agent"] == "flights"
assert interrupt_payload[0]["value"]["message"] == "Choose a flight"
async def test_workflow_run_resume_from_forwarded_command_payload() -> None:
"""forwarded_props.command.resume should resume a pending dict request."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.request_info({"options": [{"airline": "KLM"}]}, dict, request_id="flights-choice")
workflow = WorkflowBuilder(start_executor=requester).build()
_ = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
resumed_events = [
event
async for event in run_workflow_stream(
{
"messages": [],
"forwarded_props": {
"command": {"resume": json.dumps({"airline": "KLM", "departure": "AMS", "arrival": "SFO"})}
},
},
workflow,
)
]
resumed_types = [event.type for event in resumed_events]
assert "RUN_ERROR" not in resumed_types
finished = [event for event in resumed_events if event.type == "RUN_FINISHED"][0].model_dump()
assert "interrupt" not in finished
async def test_workflow_run_structured_user_json_resumes_single_pending_request() -> None:
"""A JSON user reply should resume a single pending dict request without heuristics."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.request_info({"options": [{"name": "Hotel Zoe"}]}, dict, request_id="hotel-choice")
workflow = WorkflowBuilder(start_executor=requester).build()
_ = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
resumed_events = [
event
async for event in run_workflow_stream(
{
"messages": [{"role": "user", "content": json.dumps({"name": "Hotel Zoe"})}],
},
workflow,
)
]
resumed_types = [event.type for event in resumed_events]
assert "RUN_ERROR" not in resumed_types
finished = [event for event in resumed_events if event.type == "RUN_FINISHED"][0].model_dump()
assert "interrupt" not in finished
async def test_workflow_run_resume_content_response_from_json_payload() -> None:
"""JSON resume payloads should coerce into Content responses for approval requests."""
class ApprovalExecutor(Executor):
def __init__(self) -> None:
super().__init__(id="approval_executor")
@handler
async def start(self, message: Any, ctx: WorkflowContext) -> None:
del message
function_call = Content.from_function_call(
call_id="refund-call",
name="submit_refund",
arguments={"order_id": "12345", "amount": "$89.99"},
)
approval_request = Content.from_function_approval_request(id="approval-1", function_call=function_call)
await ctx.request_info(approval_request, Content, request_id="approval-1")
@response_handler
async def handle_approval(self, original_request: Content, response: Content, ctx: WorkflowContext) -> None:
del original_request
status = "approved" if bool(response.approved) else "rejected"
await ctx.yield_output(f"Refund tool call {status}.")
workflow = WorkflowBuilder(start_executor=ApprovalExecutor()).build()
first_events = [
event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)
]
first_finished = [event for event in first_events if event.type == "RUN_FINISHED"][0].model_dump()
interrupt_payload = cast(list[dict[str, Any]], first_finished.get("interrupt"))
interrupt_value = cast(dict[str, Any], interrupt_payload[0]["value"])
resumed_events = [
event
async for event in run_workflow_stream(
{
"messages": [],
"resume": {
"interrupts": [
{
"id": "approval-1",
"value": {
"type": "function_approval_response",
"approved": True,
"id": interrupt_value.get("id", "approval-1"),
"function_call": interrupt_value.get("function_call"),
},
}
]
},
},
workflow,
)
]
resumed_types = [event.type for event in resumed_events]
assert "RUN_ERROR" not in resumed_types
assert "TEXT_MESSAGE_CONTENT" in resumed_types
resumed_finished = [event for event in resumed_events if event.type == "RUN_FINISHED"][0].model_dump()
assert "interrupt" not in resumed_finished
text_deltas = [event.delta for event in resumed_events if event.type == "TEXT_MESSAGE_CONTENT"]
assert any("approved" in delta for delta in text_deltas)
async def test_workflow_run_resume_message_list_from_json_payload() -> None:
"""Resume payloads should coerce AG-UI message dictionaries into list[Message] responses."""
class MessageRequestExecutor(Executor):
def __init__(self) -> None:
super().__init__(id="message_request_executor")
@handler
async def start(self, message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.request_info({"prompt": "Need user follow-up"}, list[Message], request_id="handoff-user-input")
@response_handler
async def handle_user_input(
self, original_request: dict, response: list[Message], ctx: WorkflowContext
) -> None:
del original_request
user_text = response[0].text if response else ""
await ctx.yield_output(f"Captured response: {user_text}")
workflow = WorkflowBuilder(start_executor=MessageRequestExecutor()).build()
_ = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "start"}]}, workflow)]
resumed_events = [
event
async for event in run_workflow_stream(
{
"messages": [],
"resume": {
"interrupts": [
{
"id": "handoff-user-input",
"value": [
{
"role": "user",
"contents": [{"type": "text", "text": "Please ship a replacement instead."}],
}
],
}
]
},
},
workflow,
)
]
resumed_types = [event.type for event in resumed_events]
assert "RUN_ERROR" not in resumed_types
assert "TEXT_MESSAGE_CONTENT" in resumed_types
resumed_finished = [event for event in resumed_events if event.type == "RUN_FINISHED"][0].model_dump()
assert "interrupt" not in resumed_finished
text_deltas = [event.delta for event in resumed_events if event.type == "TEXT_MESSAGE_CONTENT"]
assert any("replacement" in delta for delta in text_deltas)
async def test_workflow_run_non_chat_output_maps_to_custom_output_event():
"""Non-chat workflow outputs are emitted as CUSTOM workflow_output events."""
@executor(id="structured")
async def structured(message: Any, ctx: WorkflowContext[Never, dict[str, int]]) -> None:
await ctx.yield_output({"count": 3})
workflow = WorkflowBuilder(start_executor=structured).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
output_custom = [event for event in events if event.type == "CUSTOM" and event.name == "workflow_output"]
assert len(output_custom) == 1
assert output_custom[0].value == {"count": 3}
async def test_workflow_run_passthroughs_ag_ui_base_events():
"""Workflow outputs that are AG-UI BaseEvent instances should be emitted directly."""
@executor(id="stateful")
async def stateful(message: Any, ctx: WorkflowContext[Never, StateSnapshotEvent]) -> None:
await ctx.yield_output(StateSnapshotEvent(type=EventType.STATE_SNAPSHOT, snapshot={"active_agent": "flights"}))
workflow = WorkflowBuilder(start_executor=stateful).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
snapshots = [event for event in events if event.type == "STATE_SNAPSHOT"]
assert len(snapshots) == 1
assert snapshots[0].snapshot["active_agent"] == "flights"
async def test_workflow_run_plain_text_follow_up_does_not_infer_interrupt_response():
"""User follow-up text should not be coerced into request_info responses for workflows."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.request_info(
{
"message": "Choose a flight",
"options": [{"airline": "KLM"}, {"airline": "United"}],
"agent": "flights",
},
dict,
request_id="flights-choice",
)
workflow = WorkflowBuilder(start_executor=requester).build()
_ = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
follow_up_events = [
event
async for event in run_workflow_stream(
{
"messages": [
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "flights-choice",
"type": "function",
"function": {"name": "request_info", "arguments": "{}"},
}
],
},
{"role": "user", "content": "I prefer KLM please"},
]
},
workflow,
)
]
follow_up_types = [event.type for event in follow_up_events]
assert "RUN_ERROR" not in follow_up_types
assert "TOOL_CALL_START" in follow_up_types
run_finished = [event for event in follow_up_events if event.type == "RUN_FINISHED"][0].model_dump()
interrupt_payload = run_finished.get("interrupt")
assert isinstance(interrupt_payload, list)
assert interrupt_payload[0]["id"] == "flights-choice"
assert interrupt_payload[0]["value"]["agent"] == "flights"
async def test_workflow_run_empty_turn_with_pending_request_preserves_interrupts():
"""An empty turn should still return pending workflow interrupts without errors."""
@executor(id="requester")
async def requester(message: Any, ctx: WorkflowContext) -> None:
del message
await ctx.request_info({"prompt": "choose"}, dict, request_id="pick-one")
workflow = WorkflowBuilder(start_executor=requester).build()
_ = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
events = [event async for event in run_workflow_stream({"messages": []}, workflow)]
types = [event.type for event in events]
assert types[0] == "RUN_STARTED"
assert "RUN_FINISHED" in types
assert "RUN_ERROR" not in types
finished = [event for event in events if event.type == "RUN_FINISHED"][0].model_dump()
interrupts = finished.get("interrupt")
assert isinstance(interrupts, list)
assert interrupts[0]["id"] == "pick-one"
async def test_workflow_run_agent_response_output_uses_latest_assistant_message_only() -> None:
"""Conversation payload outputs should not flatten full history into one assistant message."""
@executor(id="responder")
async def responder(message: Any, ctx: WorkflowContext[Never, AgentResponse]) -> None:
del message
response = AgentResponse(
messages=[
Message(role="user", contents=[Content.from_text("My order arrived damaged")]),
Message(
role="assistant",
contents=[Content.from_text("Order Agent: Got it. I submitted the replacement request.")],
),
]
)
await ctx.yield_output(response)
workflow = WorkflowBuilder(start_executor=responder).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
text_deltas = [event.delta for event in events if event.type == "TEXT_MESSAGE_CONTENT"]
assert text_deltas == ["Order Agent: Got it. I submitted the replacement request."]
async def test_workflow_run_skips_duplicate_text_from_conversation_snapshot() -> None:
"""Do not emit duplicate assistant text when a snapshot repeats the latest output."""
@executor(id="responder")
async def responder(message: Any, ctx: WorkflowContext[Never, Any]) -> None:
del message
duplicate_text = "Order Agent: Got it. I submitted the replacement request."
await ctx.yield_output(duplicate_text)
await ctx.yield_output(
AgentResponse(
messages=[
Message(role="user", contents=[Content.from_text("standard")]),
Message(role="assistant", contents=[Content.from_text(duplicate_text)]),
]
)
)
workflow = WorkflowBuilder(start_executor=responder).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
text_deltas = [event.delta for event in events if event.type == "TEXT_MESSAGE_CONTENT"]
assert text_deltas == ["Order Agent: Got it. I submitted the replacement request."]
async def test_workflow_run_skips_consecutive_duplicate_text_outputs() -> None:
"""Do not emit duplicate assistant text when consecutive outputs are identical."""
@executor(id="responder")
async def responder(message: Any, ctx: WorkflowContext[Never, Any]) -> None:
del message
duplicate_text = "Order Agent: Replacement processed. Case complete."
await ctx.yield_output(duplicate_text)
await ctx.yield_output(duplicate_text)
workflow = WorkflowBuilder(start_executor=responder).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
text_deltas = [event.delta for event in events if event.type == "TEXT_MESSAGE_CONTENT"]
assert text_deltas == ["Order Agent: Replacement processed. Case complete."]
async def test_workflow_run_skips_final_snapshot_when_streamed_chunks_already_match() -> None:
"""Do not append full snapshot text when prior chunk outputs already formed the same message."""
@executor(id="responder")
async def responder(message: Any, ctx: WorkflowContext[Never, Any]) -> None:
del message
full_text = (
"Your replacement request for order 28939393 has been submitted with expedited shipping, "
"as you requested.\n\nCase complete."
)
await ctx.yield_output(
"Your replacement request for order 28939393 has been submitted with expedited shipping, "
)
await ctx.yield_output("as you requested.\n\nCase complete.")
await ctx.yield_output(
AgentResponse(
messages=[
Message(role="user", contents=[Content.from_text("My order is 28939393.")]),
Message(role="assistant", contents=[Content.from_text(full_text)]),
]
)
)
workflow = WorkflowBuilder(start_executor=responder).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
text_deltas = [event.delta for event in events if event.type == "TEXT_MESSAGE_CONTENT"]
assert text_deltas == [
"Your replacement request for order 28939393 has been submitted with expedited shipping, ",
"as you requested.\n\nCase complete.",
]
async def test_workflow_run_usage_content_emits_custom_usage_event() -> None:
"""Usage output from workflows should be surfaced as a custom usage event."""
@executor(id="usage")
async def usage(message: Any, ctx: WorkflowContext[Never, Content]) -> None:
del message
await ctx.yield_output(
Content.from_usage(
{
"input_token_count": 12,
"output_token_count": 6,
"total_token_count": 18,
}
)
)
workflow = WorkflowBuilder(start_executor=usage).build()
events = [event async for event in run_workflow_stream({"messages": [{"role": "user", "content": "go"}]}, workflow)]
usage_events = [event for event in events if event.type == "CUSTOM" and event.name == "usage"]
assert len(usage_events) == 1
assert usage_events[0].value["input_token_count"] == 12
assert usage_events[0].value["output_token_count"] == 6
assert usage_events[0].value["total_token_count"] == 18
async def test_workflow_run_accepts_multimodal_input_messages() -> None:
"""Workflow runner should normalize multimodal input into workflow Message content."""
class CapturingWorkflow:
def __init__(self) -> None:
self.captured_message: list[Message] | None = None
def run(self, **kwargs: Any):
self.captured_message = cast(list[Message] | None, kwargs.get("message"))
async def _stream():
yield SimpleNamespace(type="started")
return _stream()
workflow = CapturingWorkflow()
events = [
event
async for event in run_workflow_stream(
{
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": "Please analyze this image"},
{
"type": "image",
"source": {
"type": "url",
"url": "https://example.com/diagram.png",
"mimeType": "image/png",
},
},
],
}
]
},
cast(Any, workflow),
)
]
event_types = [event.type for event in events]
assert "RUN_STARTED" in event_types
assert "RUN_FINISHED" in event_types
assert "RUN_ERROR" not in event_types
assert workflow.captured_message is not None
assert len(workflow.captured_message) == 1
user_message = workflow.captured_message[0]
assert user_message.role == "user"
assert len(user_message.contents) == 2
assert user_message.contents[0].type == "text"
assert user_message.contents[0].text == "Please analyze this image"
assert user_message.contents[1].type == "uri"
assert user_message.contents[1].uri == "https://example.com/diagram.png"
def test_coerce_message_accepts_string_payload() -> None:
"""String values should coerce into a user Message with one text content."""
message = _coerce_message("Please continue")
assert message is not None
assert message.role == "user"
assert len(message.contents) == 1
assert message.contents[0].type == "text"
assert message.contents[0].text == "Please continue"
def test_coerce_message_accepts_content_key_variant() -> None:
"""The 'content' key variant should map into Message.contents."""
message = _coerce_message({"role": "assistant", "content": {"type": "text", "content": "Done"}})
assert message is not None
assert message.role == "assistant"
assert len(message.contents) == 1
assert message.contents[0].type == "text"
assert message.contents[0].text == "Done"
def test_coerce_response_for_request_bool_int_float_and_mismatch() -> None:
"""Scalar coercion should enforce bool/int/float rules and return None on mismatches."""
bool_request = SimpleNamespace(response_type=bool)
assert _coerce_response_for_request(bool_request, True) is True
assert _coerce_response_for_request(bool_request, "true") is True
assert _coerce_response_for_request(bool_request, 1) is None
int_request = SimpleNamespace(response_type=int)
assert _coerce_response_for_request(int_request, 7) == 7
assert _coerce_response_for_request(int_request, "7") == 7
assert _coerce_response_for_request(int_request, True) is None
float_request = SimpleNamespace(response_type=float)
assert _coerce_response_for_request(float_request, 2) == 2
assert _coerce_response_for_request(float_request, "2.5") == 2.5
assert _coerce_response_for_request(float_request, True) is None
dict_request = SimpleNamespace(response_type=dict)
assert _coerce_response_for_request(dict_request, "[1,2,3]") is None
async def test_workflow_run_emits_run_error_when_stream_raises() -> None:
"""Unexpected stream exceptions should be converted into RUN_ERROR events."""
class FailingWorkflow:
def run(self, **kwargs: Any):
del kwargs
async def _stream():
raise RuntimeError("workflow stream exploded")
yield # pragma: no cover
return _stream()
events = [
event
async for event in run_workflow_stream(
{"messages": [{"role": "user", "content": "go"}]},
cast(Any, FailingWorkflow()),
)
]
event_types = [event.type for event in events]
assert event_types[0] == "RUN_STARTED"
assert "RUN_ERROR" in event_types
run_error = next(event for event in events if event.type == "RUN_ERROR")
assert "workflow stream exploded" in run_error.message