mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
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:
committed by
GitHub
Unverified
parent
b1c7c7c844
commit
d8b9409e96
@@ -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")
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user