mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Add executor I/O data to ExecutorInvokedEvent and ExecutorCompletedEvent (#2591)
* Add executor I/O data to ExecutorInvokedEvent and ExecutorCompletedEvent * Sample cleanup
This commit is contained in:
committed by
GitHub
Unverified
parent
b86c130411
commit
f7a9005235
@@ -2,7 +2,15 @@
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import Executor, Message, WorkflowContext, handler
|
||||
from agent_framework import (
|
||||
Executor,
|
||||
ExecutorCompletedEvent,
|
||||
ExecutorInvokedEvent,
|
||||
Message,
|
||||
WorkflowBuilder,
|
||||
WorkflowContext,
|
||||
handler,
|
||||
)
|
||||
|
||||
|
||||
def test_executor_without_id():
|
||||
@@ -101,3 +109,155 @@ def test_executor_handlers_with_output_types():
|
||||
assert int_handler._handler_spec["name"] == "handle_integer" # type: ignore
|
||||
assert int_handler._handler_spec["message_type"] is int # type: ignore
|
||||
assert int_handler._handler_spec["output_types"] == [int] # type: ignore
|
||||
|
||||
|
||||
async def test_executor_invoked_event_contains_input_data():
|
||||
"""Test that ExecutorInvokedEvent contains the input message data."""
|
||||
|
||||
class UpperCaseExecutor(Executor):
|
||||
@handler
|
||||
async def handle(self, text: str, ctx: WorkflowContext[str]) -> None:
|
||||
await ctx.send_message(text.upper())
|
||||
|
||||
class CollectorExecutor(Executor):
|
||||
@handler
|
||||
async def handle(self, text: str, ctx: WorkflowContext) -> None:
|
||||
pass
|
||||
|
||||
upper = UpperCaseExecutor(id="upper")
|
||||
collector = CollectorExecutor(id="collector")
|
||||
|
||||
workflow = WorkflowBuilder().add_edge(upper, collector).set_start_executor(upper).build()
|
||||
|
||||
events = await workflow.run("hello world")
|
||||
invoked_events = [e for e in events if isinstance(e, ExecutorInvokedEvent)]
|
||||
|
||||
assert len(invoked_events) == 2
|
||||
|
||||
# First invoked event should be for 'upper' executor with input "hello world"
|
||||
upper_invoked = next(e for e in invoked_events if e.executor_id == "upper")
|
||||
assert upper_invoked.data == "hello world"
|
||||
|
||||
# Second invoked event should be for 'collector' executor with input "HELLO WORLD"
|
||||
collector_invoked = next(e for e in invoked_events if e.executor_id == "collector")
|
||||
assert collector_invoked.data == "HELLO WORLD"
|
||||
|
||||
|
||||
async def test_executor_completed_event_contains_sent_messages():
|
||||
"""Test that ExecutorCompletedEvent contains the messages sent via ctx.send_message()."""
|
||||
|
||||
class MultiSenderExecutor(Executor):
|
||||
@handler
|
||||
async def handle(self, text: str, ctx: WorkflowContext[str]) -> None:
|
||||
await ctx.send_message(f"{text}-first")
|
||||
await ctx.send_message(f"{text}-second")
|
||||
|
||||
class CollectorExecutor(Executor):
|
||||
def __init__(self, id: str) -> None:
|
||||
super().__init__(id=id)
|
||||
self.received: list[str] = []
|
||||
|
||||
@handler
|
||||
async def handle(self, text: str, ctx: WorkflowContext) -> None:
|
||||
self.received.append(text)
|
||||
|
||||
sender = MultiSenderExecutor(id="sender")
|
||||
collector = CollectorExecutor(id="collector")
|
||||
|
||||
workflow = WorkflowBuilder().add_edge(sender, collector).set_start_executor(sender).build()
|
||||
|
||||
events = await workflow.run("hello")
|
||||
completed_events = [e for e in events if isinstance(e, ExecutorCompletedEvent)]
|
||||
|
||||
# Sender should have completed with the sent messages
|
||||
sender_completed = next(e for e in completed_events if e.executor_id == "sender")
|
||||
assert sender_completed.data is not None
|
||||
assert sender_completed.data == ["hello-first", "hello-second"]
|
||||
|
||||
# Collector should have completed with no sent messages (None)
|
||||
collector_completed_events = [e for e in completed_events if e.executor_id == "collector"]
|
||||
# Collector is called twice (once per message from sender)
|
||||
assert len(collector_completed_events) == 2
|
||||
for collector_completed in collector_completed_events:
|
||||
assert collector_completed.data is None
|
||||
|
||||
|
||||
async def test_executor_completed_event_none_when_no_messages_sent():
|
||||
"""Test that ExecutorCompletedEvent.data is None when no messages are sent."""
|
||||
from typing_extensions import Never
|
||||
|
||||
from agent_framework import WorkflowOutputEvent
|
||||
|
||||
class YieldOnlyExecutor(Executor):
|
||||
@handler
|
||||
async def handle(self, text: str, ctx: WorkflowContext[Never, str]) -> None:
|
||||
await ctx.yield_output(text.upper())
|
||||
|
||||
executor = YieldOnlyExecutor(id="yielder")
|
||||
workflow = WorkflowBuilder().set_start_executor(executor).build()
|
||||
|
||||
events = await workflow.run("test")
|
||||
completed_events = [e for e in events if isinstance(e, ExecutorCompletedEvent)]
|
||||
|
||||
assert len(completed_events) == 1
|
||||
assert completed_events[0].executor_id == "yielder"
|
||||
assert completed_events[0].data is None
|
||||
|
||||
# Verify the output was still yielded correctly
|
||||
output_events = [e for e in events if isinstance(e, WorkflowOutputEvent)]
|
||||
assert len(output_events) == 1
|
||||
assert output_events[0].data == "TEST"
|
||||
|
||||
|
||||
async def test_executor_events_with_complex_message_types():
|
||||
"""Test that executor events correctly capture complex message types."""
|
||||
from dataclasses import dataclass
|
||||
|
||||
@dataclass
|
||||
class Request:
|
||||
query: str
|
||||
limit: int
|
||||
|
||||
@dataclass
|
||||
class Response:
|
||||
results: list[str]
|
||||
|
||||
class ProcessorExecutor(Executor):
|
||||
@handler
|
||||
async def handle(self, request: Request, ctx: WorkflowContext[Response]) -> None:
|
||||
response = Response(results=[request.query.upper()] * request.limit)
|
||||
await ctx.send_message(response)
|
||||
|
||||
class CollectorExecutor(Executor):
|
||||
@handler
|
||||
async def handle(self, response: Response, ctx: WorkflowContext) -> None:
|
||||
pass
|
||||
|
||||
processor = ProcessorExecutor(id="processor")
|
||||
collector = CollectorExecutor(id="collector")
|
||||
|
||||
workflow = WorkflowBuilder().add_edge(processor, collector).set_start_executor(processor).build()
|
||||
|
||||
input_request = Request(query="hello", limit=3)
|
||||
events = await workflow.run(input_request)
|
||||
|
||||
invoked_events = [e for e in events if isinstance(e, ExecutorInvokedEvent)]
|
||||
completed_events = [e for e in events if isinstance(e, ExecutorCompletedEvent)]
|
||||
|
||||
# Check processor invoked event has the Request object
|
||||
processor_invoked = next(e for e in invoked_events if e.executor_id == "processor")
|
||||
assert isinstance(processor_invoked.data, Request)
|
||||
assert processor_invoked.data.query == "hello"
|
||||
assert processor_invoked.data.limit == 3
|
||||
|
||||
# Check processor completed event has the Response object
|
||||
processor_completed = next(e for e in completed_events if e.executor_id == "processor")
|
||||
assert processor_completed.data is not None
|
||||
assert len(processor_completed.data) == 1
|
||||
assert isinstance(processor_completed.data[0], Response)
|
||||
assert processor_completed.data[0].results == ["HELLO", "HELLO", "HELLO"]
|
||||
|
||||
# Check collector invoked event has the Response object
|
||||
collector_invoked = next(e for e in invoked_events if e.executor_id == "collector")
|
||||
assert isinstance(collector_invoked.data, Response)
|
||||
assert collector_invoked.data.results == ["HELLO", "HELLO", "HELLO"]
|
||||
|
||||
@@ -23,6 +23,7 @@ from agent_framework import (
|
||||
WorkflowOutputEvent,
|
||||
)
|
||||
from agent_framework._mcp import MCPTool
|
||||
from agent_framework._workflows import AgentRunEvent
|
||||
from agent_framework._workflows import _handoff as handoff_module # type: ignore
|
||||
from agent_framework._workflows._handoff import _clone_chat_agent # type: ignore[reportPrivateUsage]
|
||||
from agent_framework._workflows._workflow_builder import WorkflowBuilder
|
||||
@@ -224,12 +225,12 @@ async def test_handoff_preserves_complex_additional_properties(complex_metadata:
|
||||
|
||||
# Initial run should preserve complex metadata in the triage response
|
||||
events = await _drain(workflow.run_stream("Need help with a return"))
|
||||
agent_events = [ev for ev in events if hasattr(ev, "data") and hasattr(ev.data, "messages")]
|
||||
agent_events = [ev for ev in events if isinstance(ev, AgentRunEvent)]
|
||||
if agent_events:
|
||||
first_agent_event = agent_events[0]
|
||||
first_agent_event_data = first_agent_event.data
|
||||
if first_agent_event_data and hasattr(first_agent_event_data, "messages"):
|
||||
first_agent_message = first_agent_event_data.messages[0] # type: ignore[attr-defined]
|
||||
if first_agent_event_data and first_agent_event_data.messages:
|
||||
first_agent_message = first_agent_event_data.messages[0]
|
||||
assert "complex" in first_agent_message.additional_properties, "Agent event lost complex metadata"
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests, "Workflow should request additional user input"
|
||||
|
||||
Reference in New Issue
Block a user