mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
[BREAKING] Python: Move orchestrations to dedicated package (#3685)
* Move orchestrations to dedicated package * Merge main * Fix markdown links * Fix links
This commit is contained in:
@@ -12,13 +12,13 @@ from agent_framework import (
|
||||
ChatMessage,
|
||||
ChatMessageStore,
|
||||
Content,
|
||||
SequentialBuilder,
|
||||
WorkflowOutputEvent,
|
||||
WorkflowRunState,
|
||||
WorkflowStatusEvent,
|
||||
)
|
||||
from agent_framework._workflows._agent_executor import AgentExecutorResponse
|
||||
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
|
||||
from agent_framework.orchestrations import SequentialBuilder
|
||||
|
||||
|
||||
class _CountingAgent(BaseAgent):
|
||||
|
||||
@@ -1,550 +0,0 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
from typing_extensions import Never
|
||||
|
||||
from agent_framework import (
|
||||
AgentExecutorRequest,
|
||||
AgentExecutorResponse,
|
||||
AgentResponse,
|
||||
ChatMessage,
|
||||
ConcurrentBuilder,
|
||||
Executor,
|
||||
WorkflowContext,
|
||||
WorkflowOutputEvent,
|
||||
WorkflowRunState,
|
||||
WorkflowStatusEvent,
|
||||
handler,
|
||||
)
|
||||
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
|
||||
|
||||
|
||||
class _FakeAgentExec(Executor):
|
||||
"""Test executor that mimics an agent by emitting an AgentExecutorResponse.
|
||||
|
||||
It takes the incoming AgentExecutorRequest, produces a single assistant message
|
||||
with the configured reply text, and sends an AgentExecutorResponse that includes
|
||||
full_conversation (the original user prompt followed by the assistant message).
|
||||
"""
|
||||
|
||||
def __init__(self, id: str, reply_text: str) -> None:
|
||||
super().__init__(id)
|
||||
self._reply_text = reply_text
|
||||
|
||||
@handler
|
||||
async def run(self, request: AgentExecutorRequest, ctx: WorkflowContext[AgentExecutorResponse]) -> None:
|
||||
response = AgentResponse(messages=ChatMessage("assistant", text=self._reply_text))
|
||||
full_conversation = list(request.messages) + list(response.messages)
|
||||
await ctx.send_message(AgentExecutorResponse(self.id, response, full_conversation=full_conversation))
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_empty_participants() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
ConcurrentBuilder().participants([])
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_duplicate_executors() -> None:
|
||||
a = _FakeAgentExec("dup", "A")
|
||||
b = _FakeAgentExec("dup", "B") # same executor id
|
||||
with pytest.raises(ValueError):
|
||||
ConcurrentBuilder().participants([a, b])
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_duplicate_executors_from_factories() -> None:
|
||||
"""Test that duplicate executor IDs from factories are detected at build time."""
|
||||
|
||||
def create_dup1() -> Executor:
|
||||
return _FakeAgentExec("dup", "A")
|
||||
|
||||
def create_dup2() -> Executor:
|
||||
return _FakeAgentExec("dup", "B") # same executor id
|
||||
|
||||
builder = ConcurrentBuilder().register_participants([create_dup1, create_dup2])
|
||||
with pytest.raises(ValueError, match="Duplicate executor ID 'dup' detected in workflow."):
|
||||
builder.build()
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_mixed_participants_and_factories() -> None:
|
||||
"""Test that mixing .participants() and .register_participants() raises an error."""
|
||||
# Case 1: participants first, then register_participants
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
(
|
||||
ConcurrentBuilder()
|
||||
.participants([_FakeAgentExec("a", "A")])
|
||||
.register_participants([lambda: _FakeAgentExec("b", "B")])
|
||||
)
|
||||
|
||||
# Case 2: register_participants first, then participants
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
(
|
||||
ConcurrentBuilder()
|
||||
.register_participants([lambda: _FakeAgentExec("a", "A")])
|
||||
.participants([_FakeAgentExec("b", "B")])
|
||||
)
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_multiple_calls_to_participants() -> None:
|
||||
"""Test that multiple calls to .participants() raises an error."""
|
||||
with pytest.raises(ValueError, match=r"participants\(\) has already been called"):
|
||||
(ConcurrentBuilder().participants([_FakeAgentExec("a", "A")]).participants([_FakeAgentExec("b", "B")]))
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_multiple_calls_to_register_participants() -> None:
|
||||
"""Test that multiple calls to .register_participants() raises an error."""
|
||||
with pytest.raises(ValueError, match=r"register_participants\(\) has already been called"):
|
||||
(
|
||||
ConcurrentBuilder()
|
||||
.register_participants([lambda: _FakeAgentExec("a", "A")])
|
||||
.register_participants([lambda: _FakeAgentExec("b", "B")])
|
||||
)
|
||||
|
||||
|
||||
async def test_concurrent_default_aggregator_emits_single_user_and_assistants() -> None:
|
||||
# Three synthetic agent executors
|
||||
e1 = _FakeAgentExec("agentA", "Alpha")
|
||||
e2 = _FakeAgentExec("agentB", "Beta")
|
||||
e3 = _FakeAgentExec("agentC", "Gamma")
|
||||
|
||||
wf = ConcurrentBuilder().participants([e1, e2, e3]).build()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("prompt: hello world"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(list[ChatMessage], ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
messages: list[ChatMessage] = output
|
||||
|
||||
# Expect one user message + one assistant message per participant
|
||||
assert len(messages) == 1 + 3
|
||||
assert messages[0].role == "user"
|
||||
assert "hello world" in messages[0].text
|
||||
|
||||
assistant_texts = {m.text for m in messages[1:]}
|
||||
assert assistant_texts == {"Alpha", "Beta", "Gamma"}
|
||||
assert all(m.role == "assistant" for m in messages[1:])
|
||||
|
||||
|
||||
async def test_concurrent_custom_aggregator_callback_is_used() -> None:
|
||||
# Two synthetic agent executors for brevity
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
async def summarize(results: list[AgentExecutorResponse]) -> str:
|
||||
texts: list[str] = []
|
||||
for r in results:
|
||||
msgs: list[ChatMessage] = r.agent_response.messages
|
||||
texts.append(msgs[-1].text if msgs else "")
|
||||
return " | ".join(sorted(texts))
|
||||
|
||||
wf = ConcurrentBuilder().participants([e1, e2]).with_aggregator(summarize).build()
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: custom"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(str, ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
# Custom aggregator returns a string payload
|
||||
assert isinstance(output, str)
|
||||
assert output == "One | Two"
|
||||
|
||||
|
||||
async def test_concurrent_custom_aggregator_sync_callback_is_used() -> None:
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
# Sync callback with ctx parameter (should run via asyncio.to_thread)
|
||||
def summarize_sync(results: list[AgentExecutorResponse], _ctx: WorkflowContext[Any]) -> str: # type: ignore[unused-argument]
|
||||
texts: list[str] = []
|
||||
for r in results:
|
||||
msgs: list[ChatMessage] = r.agent_response.messages
|
||||
texts.append(msgs[-1].text if msgs else "")
|
||||
return " | ".join(sorted(texts))
|
||||
|
||||
wf = ConcurrentBuilder().participants([e1, e2]).with_aggregator(summarize_sync).build()
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: custom sync"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(str, ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
assert isinstance(output, str)
|
||||
assert output == "One | Two"
|
||||
|
||||
|
||||
def test_concurrent_custom_aggregator_uses_callback_name_for_id() -> None:
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
def summarize(results: list[AgentExecutorResponse]) -> str: # type: ignore[override]
|
||||
return str(len(results))
|
||||
|
||||
wf = ConcurrentBuilder().participants([e1, e2]).with_aggregator(summarize).build()
|
||||
|
||||
assert "summarize" in wf.executors
|
||||
aggregator = wf.executors["summarize"]
|
||||
assert aggregator.id == "summarize"
|
||||
|
||||
|
||||
async def test_concurrent_with_aggregator_executor_instance() -> None:
|
||||
"""Test with_aggregator using an Executor instance (not factory)."""
|
||||
|
||||
class CustomAggregator(Executor):
|
||||
@handler
|
||||
async def aggregate(self, results: list[AgentExecutorResponse], ctx: WorkflowContext[Never, str]) -> None:
|
||||
texts: list[str] = []
|
||||
for r in results:
|
||||
msgs: list[ChatMessage] = r.agent_response.messages
|
||||
texts.append(msgs[-1].text if msgs else "")
|
||||
await ctx.yield_output(" & ".join(sorted(texts)))
|
||||
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
aggregator_instance = CustomAggregator(id="instance_aggregator")
|
||||
wf = ConcurrentBuilder().participants([e1, e2]).with_aggregator(aggregator_instance).build()
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: instance test"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(str, ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
assert isinstance(output, str)
|
||||
assert output == "One & Two"
|
||||
|
||||
|
||||
async def test_concurrent_with_aggregator_executor_factory() -> None:
|
||||
"""Test with_aggregator using an Executor factory."""
|
||||
|
||||
class CustomAggregator(Executor):
|
||||
@handler
|
||||
async def aggregate(self, results: list[AgentExecutorResponse], ctx: WorkflowContext[Never, str]) -> None:
|
||||
texts: list[str] = []
|
||||
for r in results:
|
||||
msgs: list[ChatMessage] = r.agent_response.messages
|
||||
texts.append(msgs[-1].text if msgs else "")
|
||||
await ctx.yield_output(" | ".join(sorted(texts)))
|
||||
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
wf = (
|
||||
ConcurrentBuilder()
|
||||
.participants([e1, e2])
|
||||
.register_aggregator(lambda: CustomAggregator(id="custom_aggregator"))
|
||||
.build()
|
||||
)
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: factory test"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(str, ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
assert isinstance(output, str)
|
||||
assert output == "One | Two"
|
||||
|
||||
|
||||
async def test_concurrent_with_aggregator_executor_factory_with_default_id() -> None:
|
||||
"""Test with_aggregator using an Executor class directly as factory (with default __init__ parameters)."""
|
||||
|
||||
class CustomAggregator(Executor):
|
||||
def __init__(self, id: str = "default_aggregator") -> None:
|
||||
super().__init__(id)
|
||||
|
||||
@handler
|
||||
async def aggregate(self, results: list[AgentExecutorResponse], ctx: WorkflowContext[Never, str]) -> None:
|
||||
texts: list[str] = []
|
||||
for r in results:
|
||||
msgs: list[ChatMessage] = r.agent_response.messages
|
||||
texts.append(msgs[-1].text if msgs else "")
|
||||
await ctx.yield_output(" | ".join(sorted(texts)))
|
||||
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
wf = ConcurrentBuilder().participants([e1, e2]).register_aggregator(CustomAggregator).build()
|
||||
|
||||
completed = False
|
||||
output: str | None = None
|
||||
async for ev in wf.run_stream("prompt: factory test"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(str, ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
assert isinstance(output, str)
|
||||
assert output == "One | Two"
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_multiple_calls_to_with_aggregator() -> None:
|
||||
"""Test that multiple calls to .with_aggregator() raises an error."""
|
||||
|
||||
def summarize(results: list[AgentExecutorResponse]) -> str: # type: ignore[override]
|
||||
return str(len(results))
|
||||
|
||||
with pytest.raises(ValueError, match=r"with_aggregator\(\) has already been called"):
|
||||
(ConcurrentBuilder().with_aggregator(summarize).with_aggregator(summarize))
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_multiple_calls_to_register_aggregator() -> None:
|
||||
"""Test that multiple calls to .register_aggregator() raises an error."""
|
||||
|
||||
class CustomAggregator(Executor):
|
||||
pass
|
||||
|
||||
with pytest.raises(ValueError, match=r"register_aggregator\(\) has already been called"):
|
||||
(
|
||||
ConcurrentBuilder()
|
||||
.register_aggregator(lambda: CustomAggregator(id="agg1"))
|
||||
.register_aggregator(lambda: CustomAggregator(id="agg2"))
|
||||
)
|
||||
|
||||
|
||||
async def test_concurrent_checkpoint_resume_round_trip() -> None:
|
||||
storage = InMemoryCheckpointStorage()
|
||||
|
||||
participants = (
|
||||
_FakeAgentExec("agentA", "Alpha"),
|
||||
_FakeAgentExec("agentB", "Beta"),
|
||||
_FakeAgentExec("agentC", "Gamma"),
|
||||
)
|
||||
|
||||
wf = ConcurrentBuilder().participants(list(participants)).with_checkpointing(storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("checkpoint concurrent"):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
checkpoints = await storage.list_checkpoints()
|
||||
assert checkpoints
|
||||
checkpoints.sort(key=lambda cp: cp.timestamp)
|
||||
resume_checkpoint = next(
|
||||
(cp for cp in checkpoints if (cp.metadata or {}).get("checkpoint_type") == "superstep"),
|
||||
checkpoints[-1],
|
||||
)
|
||||
|
||||
resumed_participants = (
|
||||
_FakeAgentExec("agentA", "Alpha"),
|
||||
_FakeAgentExec("agentB", "Beta"),
|
||||
_FakeAgentExec("agentC", "Gamma"),
|
||||
)
|
||||
wf_resume = ConcurrentBuilder().participants(list(resumed_participants)).with_checkpointing(storage).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
WorkflowRunState.IDLE,
|
||||
WorkflowRunState.IDLE_WITH_PENDING_REQUESTS,
|
||||
):
|
||||
break
|
||||
|
||||
assert resumed_output is not None
|
||||
assert [m.role for m in resumed_output] == [m.role for m in baseline_output]
|
||||
assert [m.text for m in resumed_output] == [m.text for m in baseline_output]
|
||||
|
||||
|
||||
async def test_concurrent_checkpoint_runtime_only() -> None:
|
||||
"""Test checkpointing configured ONLY at runtime, not at build time."""
|
||||
storage = InMemoryCheckpointStorage()
|
||||
|
||||
agents = [_FakeAgentExec(id="agent1", reply_text="A1"), _FakeAgentExec(id="agent2", reply_text="A2")]
|
||||
wf = ConcurrentBuilder().participants(agents).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("runtime checkpoint test", checkpoint_storage=storage):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
checkpoints = await storage.list_checkpoints()
|
||||
assert checkpoints
|
||||
checkpoints.sort(key=lambda cp: cp.timestamp)
|
||||
|
||||
resume_checkpoint = next(
|
||||
(cp for cp in checkpoints if (cp.metadata or {}).get("checkpoint_type") == "superstep"),
|
||||
checkpoints[-1],
|
||||
)
|
||||
|
||||
resumed_agents = [_FakeAgentExec(id="agent1", reply_text="A1"), _FakeAgentExec(id="agent2", reply_text="A2")]
|
||||
wf_resume = ConcurrentBuilder().participants(resumed_agents).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id, checkpoint_storage=storage):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
WorkflowRunState.IDLE,
|
||||
WorkflowRunState.IDLE_WITH_PENDING_REQUESTS,
|
||||
):
|
||||
break
|
||||
|
||||
assert resumed_output is not None
|
||||
assert [m.role for m in resumed_output] == [m.role for m in baseline_output]
|
||||
|
||||
|
||||
async def test_concurrent_checkpoint_runtime_overrides_buildtime() -> None:
|
||||
"""Test that runtime checkpoint storage overrides build-time configuration."""
|
||||
import tempfile
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir1, tempfile.TemporaryDirectory() as temp_dir2:
|
||||
from agent_framework._workflows._checkpoint import FileCheckpointStorage
|
||||
|
||||
buildtime_storage = FileCheckpointStorage(temp_dir1)
|
||||
runtime_storage = FileCheckpointStorage(temp_dir2)
|
||||
|
||||
agents = [_FakeAgentExec(id="agent1", reply_text="A1"), _FakeAgentExec(id="agent2", reply_text="A2")]
|
||||
wf = ConcurrentBuilder().participants(agents).with_checkpointing(buildtime_storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("override test", checkpoint_storage=runtime_storage):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
buildtime_checkpoints = await buildtime_storage.list_checkpoints()
|
||||
runtime_checkpoints = await runtime_storage.list_checkpoints()
|
||||
|
||||
assert len(runtime_checkpoints) > 0, "Runtime storage should have checkpoints"
|
||||
assert len(buildtime_checkpoints) == 0, "Build-time storage should have no checkpoints when overridden"
|
||||
|
||||
|
||||
def test_concurrent_builder_rejects_empty_participant_factories() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
ConcurrentBuilder().register_participants([])
|
||||
|
||||
|
||||
async def test_concurrent_builder_reusable_after_build_with_participants() -> None:
|
||||
"""Test that the builder can be reused to build multiple identical workflows with participants()."""
|
||||
e1 = _FakeAgentExec("agentA", "One")
|
||||
e2 = _FakeAgentExec("agentB", "Two")
|
||||
|
||||
builder = ConcurrentBuilder().participants([e1, e2])
|
||||
|
||||
builder.build()
|
||||
|
||||
assert builder._participants[0] is e1 # type: ignore
|
||||
assert builder._participants[1] is e2 # type: ignore
|
||||
assert builder._participant_factories == [] # type: ignore
|
||||
|
||||
|
||||
async def test_concurrent_builder_reusable_after_build_with_factories() -> None:
|
||||
"""Test that the builder can be reused to build multiple workflows with register_participants()."""
|
||||
call_count = 0
|
||||
|
||||
def create_agent_executor_a() -> Executor:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return _FakeAgentExec("agentA", "One")
|
||||
|
||||
def create_agent_executor_b() -> Executor:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return _FakeAgentExec("agentB", "Two")
|
||||
|
||||
builder = ConcurrentBuilder().register_participants([create_agent_executor_a, create_agent_executor_b])
|
||||
|
||||
# Build the first workflow
|
||||
wf1 = builder.build()
|
||||
|
||||
assert builder._participants == [] # type: ignore
|
||||
assert len(builder._participant_factories) == 2 # type: ignore
|
||||
assert call_count == 2
|
||||
|
||||
# Build the second workflow
|
||||
wf2 = builder.build()
|
||||
assert call_count == 4
|
||||
|
||||
# Verify that the two workflows have different executor instances
|
||||
assert wf1.executors["agentA"] is not wf2.executors["agentA"]
|
||||
assert wf1.executors["agentB"] is not wf2.executors["agentB"]
|
||||
|
||||
|
||||
async def test_concurrent_with_register_participants() -> None:
|
||||
"""Test workflow creation using register_participants with factories."""
|
||||
|
||||
def create_agent1() -> Executor:
|
||||
return _FakeAgentExec("agentA", "Alpha")
|
||||
|
||||
def create_agent2() -> Executor:
|
||||
return _FakeAgentExec("agentB", "Beta")
|
||||
|
||||
def create_agent3() -> Executor:
|
||||
return _FakeAgentExec("agentC", "Gamma")
|
||||
|
||||
wf = ConcurrentBuilder().register_participants([create_agent1, create_agent2, create_agent3]).build()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("test prompt"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = cast(list[ChatMessage], ev.data)
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
messages: list[ChatMessage] = output
|
||||
|
||||
# Expect one user message + one assistant message per participant
|
||||
assert len(messages) == 1 + 3
|
||||
assert messages[0].role == "user"
|
||||
assert "test prompt" in messages[0].text
|
||||
|
||||
assistant_texts = {m.text for m in messages[1:]}
|
||||
assert assistant_texts == {"Alpha", "Beta", "Gamma"}
|
||||
assert all(m.role == "assistant" for m in messages[1:])
|
||||
@@ -16,13 +16,13 @@ from agent_framework import (
|
||||
ChatMessage,
|
||||
Content,
|
||||
Executor,
|
||||
SequentialBuilder,
|
||||
WorkflowBuilder,
|
||||
WorkflowContext,
|
||||
WorkflowRunState,
|
||||
WorkflowStatusEvent,
|
||||
handler,
|
||||
)
|
||||
from agent_framework.orchestrations import SequentialBuilder
|
||||
|
||||
|
||||
class _SimpleAgent(BaseAgent):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,709 +0,0 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import AsyncIterable
|
||||
from typing import Any, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import (
|
||||
ChatAgent,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
HandoffAgentUserRequest,
|
||||
HandoffBuilder,
|
||||
RequestInfoEvent,
|
||||
WorkflowEvent,
|
||||
WorkflowOutputEvent,
|
||||
resolve_agent_id,
|
||||
use_function_invocation,
|
||||
)
|
||||
|
||||
|
||||
@use_function_invocation
|
||||
class MockChatClient:
|
||||
"""Mock chat client for testing handoff workflows."""
|
||||
|
||||
additional_properties: dict[str, Any]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
handoff_to: str | None = None,
|
||||
) -> None:
|
||||
"""Initialize the mock chat client.
|
||||
|
||||
Args:
|
||||
name: The name of the agent using this chat client.
|
||||
handoff_to: The name of the agent to hand off to, or None for no handoff.
|
||||
This is hardcoded for testing purposes so that the agent always attempts to hand off.
|
||||
"""
|
||||
self._name = name
|
||||
self._handoff_to = handoff_to
|
||||
self._call_index = 0
|
||||
|
||||
async def get_response(self, messages: Any, **kwargs: Any) -> ChatResponse:
|
||||
contents = _build_reply_contents(self._name, self._handoff_to, self._next_call_id())
|
||||
reply = ChatMessage(
|
||||
role="assistant",
|
||||
contents=contents,
|
||||
)
|
||||
return ChatResponse(messages=reply, response_id="mock_response")
|
||||
|
||||
def get_streaming_response(self, messages: Any, **kwargs: Any) -> AsyncIterable[ChatResponseUpdate]:
|
||||
async def _stream() -> AsyncIterable[ChatResponseUpdate]:
|
||||
contents = _build_reply_contents(self._name, self._handoff_to, self._next_call_id())
|
||||
yield ChatResponseUpdate(contents=contents, role="assistant")
|
||||
|
||||
return _stream()
|
||||
|
||||
def _next_call_id(self) -> str | None:
|
||||
if not self._handoff_to:
|
||||
return None
|
||||
call_id = f"{self._name}-handoff-{self._call_index}"
|
||||
self._call_index += 1
|
||||
return call_id
|
||||
|
||||
|
||||
def _build_reply_contents(
|
||||
agent_name: str,
|
||||
handoff_to: str | None,
|
||||
call_id: str | None,
|
||||
) -> list[Content]:
|
||||
contents: list[Content] = []
|
||||
if handoff_to and call_id:
|
||||
contents.append(
|
||||
Content.from_function_call(
|
||||
call_id=call_id, name=f"handoff_to_{handoff_to}", arguments={"handoff_to": handoff_to}
|
||||
)
|
||||
)
|
||||
text = f"{agent_name} reply"
|
||||
contents.append(Content.from_text(text=text))
|
||||
return contents
|
||||
|
||||
|
||||
class MockHandoffAgent(ChatAgent):
|
||||
"""Mock agent that can hand off to another agent."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
handoff_to: str | None = None,
|
||||
) -> None:
|
||||
"""Initialize the mock handoff agent.
|
||||
|
||||
Args:
|
||||
name: The name of the agent.
|
||||
handoff_to: The name of the agent to hand off to, or None for no handoff.
|
||||
This is hardcoded for testing purposes so that the agent always attempts to hand off.
|
||||
"""
|
||||
super().__init__(chat_client=MockChatClient(name, handoff_to=handoff_to), name=name, id=name)
|
||||
|
||||
|
||||
async def _drain(stream: AsyncIterable[WorkflowEvent]) -> list[WorkflowEvent]:
|
||||
return [event async for event in stream]
|
||||
|
||||
|
||||
async def test_handoff():
|
||||
"""Test that agents can hand off to each other."""
|
||||
|
||||
# `triage` hands off to `specialist`, who then hands off to `escalation`.
|
||||
# `escalation` has no handoff, so the workflow should request user input to continue.
|
||||
triage = MockHandoffAgent(name="triage", handoff_to="specialist")
|
||||
specialist = MockHandoffAgent(name="specialist", handoff_to="escalation")
|
||||
escalation = MockHandoffAgent(name="escalation")
|
||||
|
||||
# Without explicitly defining handoffs, the builder will create connections
|
||||
# between all agents.
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[triage, specialist, escalation])
|
||||
.with_start_agent(triage)
|
||||
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == "user") >= 2)
|
||||
.build()
|
||||
)
|
||||
|
||||
# Start conversation - triage hands off to specialist then escalation
|
||||
# escalation won't trigger a handoff, so the response from it will become
|
||||
# a request for user input because autonomous mode is not enabled by default.
|
||||
events = await _drain(workflow.run_stream("Need technical support"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
|
||||
assert requests
|
||||
assert len(requests) == 1
|
||||
|
||||
request = requests[0]
|
||||
assert isinstance(request.data, HandoffAgentUserRequest)
|
||||
assert request.source_executor_id == escalation.name
|
||||
|
||||
|
||||
async def test_autonomous_mode_yields_output_without_user_request():
|
||||
"""Ensure autonomous interaction mode yields output without requesting user input."""
|
||||
triage = MockHandoffAgent(name="triage", handoff_to="specialist")
|
||||
specialist = MockHandoffAgent(name="specialist")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[triage, specialist])
|
||||
.with_start_agent(triage)
|
||||
# Since specialist has no handoff, the specialist will be generating normal responses.
|
||||
# With autonomous mode, this should continue until the termination condition is met.
|
||||
.with_autonomous_mode(
|
||||
agents=[specialist],
|
||||
turn_limits={resolve_agent_id(specialist): 1},
|
||||
)
|
||||
# This termination condition ensures the workflow runs through both agents.
|
||||
# First message is the user message to triage, second is triage's response, which
|
||||
# is a handoff to specialist, third is specialist's response that should not request
|
||||
# user input due to autonomous mode. Fourth message will come from the specialist
|
||||
# again and will trigger termination.
|
||||
.with_termination_condition(lambda conv: len(conv) >= 4)
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("Package arrived broken"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert not requests, "Autonomous mode should not request additional user input"
|
||||
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs, "Autonomous mode should yield a workflow output"
|
||||
|
||||
final_conversation = outputs[-1].data
|
||||
assert isinstance(final_conversation, list)
|
||||
conversation_list = cast(list[ChatMessage], final_conversation)
|
||||
assert any(msg.role == "assistant" and (msg.text or "").startswith("specialist reply") for msg in conversation_list)
|
||||
|
||||
|
||||
async def test_autonomous_mode_resumes_user_input_on_turn_limit():
|
||||
"""Autonomous mode should resume user input request when turn limit is reached."""
|
||||
triage = MockHandoffAgent(name="triage", handoff_to="worker")
|
||||
worker = MockHandoffAgent(name="worker")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[triage, worker])
|
||||
.with_start_agent(triage)
|
||||
.with_autonomous_mode(agents=[worker], turn_limits={resolve_agent_id(worker): 2})
|
||||
.with_termination_condition(lambda conv: False)
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("Start"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests and len(requests) == 1, "Turn limit should force a user input request"
|
||||
assert requests[0].source_executor_id == worker.name
|
||||
|
||||
|
||||
def test_build_fails_without_start_agent():
|
||||
"""Verify that build() raises ValueError when with_start_agent() was not called."""
|
||||
triage = MockHandoffAgent(name="triage")
|
||||
specialist = MockHandoffAgent(name="specialist")
|
||||
|
||||
with pytest.raises(ValueError, match=r"Must call with_start_agent\(...\) before building the workflow."):
|
||||
HandoffBuilder(participants=[triage, specialist]).build()
|
||||
|
||||
|
||||
def test_build_fails_without_participants():
|
||||
"""Verify that build() raises ValueError when no participants are provided."""
|
||||
with pytest.raises(
|
||||
ValueError, match=r"No participants provided\. Call \.participants\(\) or \.register_participants\(\) first."
|
||||
):
|
||||
HandoffBuilder().build()
|
||||
|
||||
|
||||
async def test_handoff_async_termination_condition() -> None:
|
||||
"""Test that async termination conditions work correctly."""
|
||||
termination_call_count = 0
|
||||
|
||||
async def async_termination(conv: list[ChatMessage]) -> bool:
|
||||
nonlocal termination_call_count
|
||||
termination_call_count += 1
|
||||
user_count = sum(1 for msg in conv if msg.role == "user")
|
||||
return user_count >= 2
|
||||
|
||||
coordinator = MockHandoffAgent(name="coordinator", handoff_to="worker")
|
||||
worker = MockHandoffAgent(name="worker")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[coordinator, worker])
|
||||
.with_start_agent(coordinator)
|
||||
.with_termination_condition(async_termination)
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("First user message"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["Second user message"])]})
|
||||
)
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert len(outputs) == 1
|
||||
|
||||
final_conversation = outputs[0].data
|
||||
assert isinstance(final_conversation, list)
|
||||
final_conv_list = cast(list[ChatMessage], final_conversation)
|
||||
user_messages = [msg for msg in final_conv_list if msg.role == "user"]
|
||||
assert len(user_messages) == 2
|
||||
assert termination_call_count > 0
|
||||
|
||||
|
||||
async def test_tool_choice_preserved_from_agent_config():
|
||||
"""Verify that agent-level tool_choice configuration is preserved and not overridden."""
|
||||
# Create a mock chat client that records the tool_choice used
|
||||
recorded_tool_choices: list[Any] = []
|
||||
|
||||
async def mock_get_response(messages: Any, options: dict[str, Any] | None = None, **kwargs: Any) -> ChatResponse:
|
||||
if options:
|
||||
recorded_tool_choices.append(options.get("tool_choice"))
|
||||
return ChatResponse(
|
||||
messages=[ChatMessage("assistant", ["Response"])],
|
||||
response_id="test_response",
|
||||
)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get_response = AsyncMock(side_effect=mock_get_response)
|
||||
|
||||
# Create agent with specific tool_choice configuration via default_options
|
||||
agent = ChatAgent(
|
||||
chat_client=mock_client,
|
||||
name="test_agent",
|
||||
default_options={"tool_choice": {"mode": "required"}}, # type: ignore
|
||||
)
|
||||
|
||||
# Run the agent
|
||||
await agent.run("Test message")
|
||||
|
||||
# Verify tool_choice was preserved
|
||||
assert len(recorded_tool_choices) > 0, "No tool_choice recorded"
|
||||
last_tool_choice = recorded_tool_choices[-1]
|
||||
assert last_tool_choice is not None, "tool_choice should not be None"
|
||||
assert last_tool_choice == {"mode": "required"}, f"Expected 'required', got {last_tool_choice}"
|
||||
|
||||
|
||||
# region Participant Factory Tests
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_empty_participant_factories():
|
||||
"""Test that HandoffBuilder rejects empty participant_factories dictionary."""
|
||||
# Empty factories are rejected immediately when calling participant_factories()
|
||||
with pytest.raises(ValueError, match=r"participant_factories cannot be empty"):
|
||||
HandoffBuilder().register_participants({})
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match=r"No participants provided\. Call \.participants\(\) or \.register_participants\(\) first\."
|
||||
):
|
||||
HandoffBuilder(participant_factories={}).build()
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_mixing_participants_and_factories():
|
||||
"""Test that mixing participants and participant_factories in __init__ raises an error."""
|
||||
triage = MockHandoffAgent(name="triage")
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
HandoffBuilder(participants=[triage], participant_factories={"triage": lambda: triage})
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_mixing_participants_and_participant_factories_methods():
|
||||
"""Test that mixing .participants() and .participant_factories() raises an error."""
|
||||
triage = MockHandoffAgent(name="triage")
|
||||
|
||||
# Case 1: participants first, then participant_factories
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
HandoffBuilder(participants=[triage]).register_participants({
|
||||
"specialist": lambda: MockHandoffAgent(name="specialist")
|
||||
})
|
||||
|
||||
# Case 2: participant_factories first, then participants
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
HandoffBuilder(participant_factories={"triage": lambda: triage}).participants([
|
||||
MockHandoffAgent(name="specialist")
|
||||
])
|
||||
|
||||
# Case 3: participants(), then participant_factories()
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
HandoffBuilder().participants([triage]).register_participants({
|
||||
"specialist": lambda: MockHandoffAgent(name="specialist")
|
||||
})
|
||||
|
||||
# Case 4: participant_factories(), then participants()
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
HandoffBuilder().register_participants({"triage": lambda: triage}).participants([
|
||||
MockHandoffAgent(name="specialist")
|
||||
])
|
||||
|
||||
# Case 5: mix during initialization
|
||||
with pytest.raises(ValueError, match="Cannot mix .participants"):
|
||||
HandoffBuilder(
|
||||
participants=[triage], participant_factories={"specialist": lambda: MockHandoffAgent(name="specialist")}
|
||||
)
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_multiple_calls_to_participant_factories():
|
||||
"""Test that multiple calls to .participant_factories() raises an error."""
|
||||
with pytest.raises(
|
||||
ValueError, match=r"register_participants\(\) has already been called on this builder instance."
|
||||
):
|
||||
(
|
||||
HandoffBuilder()
|
||||
.register_participants({"agent1": lambda: MockHandoffAgent(name="agent1")})
|
||||
.register_participants({"agent2": lambda: MockHandoffAgent(name="agent2")})
|
||||
)
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_multiple_calls_to_participants():
|
||||
"""Test that multiple calls to .participants() raises an error."""
|
||||
with pytest.raises(ValueError, match="participants have already been assigned"):
|
||||
(
|
||||
HandoffBuilder()
|
||||
.participants([MockHandoffAgent(name="agent1")])
|
||||
.participants([MockHandoffAgent(name="agent2")])
|
||||
)
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_instance_coordinator_with_factories():
|
||||
"""Test that using an agent instance for set_coordinator when using factories raises an error."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
# Create an agent instance
|
||||
coordinator_instance = MockHandoffAgent(name="coordinator")
|
||||
|
||||
with pytest.raises(ValueError, match=r"Call participants\(\.\.\.\) before with_start_agent\(\.\.\.\)"):
|
||||
(
|
||||
HandoffBuilder(
|
||||
participant_factories={"triage": create_triage, "specialist": create_specialist}
|
||||
).with_start_agent(coordinator_instance) # Instance, not factory name
|
||||
)
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_factory_name_coordinator_with_instances():
|
||||
"""Test that using a factory name for set_coordinator when using instances raises an error."""
|
||||
triage = MockHandoffAgent(name="triage")
|
||||
specialist = MockHandoffAgent(name="specialist")
|
||||
|
||||
with pytest.raises(ValueError, match=r"Call register_participants\(...\) before with_start_agent\(...\)"):
|
||||
(
|
||||
HandoffBuilder(participants=[triage, specialist]).with_start_agent(
|
||||
"triage"
|
||||
) # String factory name, not instance
|
||||
)
|
||||
|
||||
|
||||
def test_handoff_builder_rejects_mixed_types_in_add_handoff_source():
|
||||
"""Test that add_handoff rejects factory name source with instance-based participants."""
|
||||
triage = MockHandoffAgent(name="triage")
|
||||
specialist = MockHandoffAgent(name="specialist")
|
||||
|
||||
with pytest.raises(TypeError, match="Cannot mix factory names \\(str\\) and AgentProtocol.*instances"):
|
||||
(
|
||||
HandoffBuilder(participants=[triage, specialist])
|
||||
.with_start_agent(triage)
|
||||
.add_handoff("triage", [specialist]) # String source with instance participants
|
||||
)
|
||||
|
||||
|
||||
def test_handoff_builder_accepts_all_factory_names_in_add_handoff():
|
||||
"""Test that add_handoff accepts all factory names when using participant_factories."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage")
|
||||
|
||||
def create_specialist_a() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist_a")
|
||||
|
||||
def create_specialist_b() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist_b")
|
||||
|
||||
# This should work - all strings with participant_factories
|
||||
builder = (
|
||||
HandoffBuilder(
|
||||
participant_factories={
|
||||
"triage": create_triage,
|
||||
"specialist_a": create_specialist_a,
|
||||
"specialist_b": create_specialist_b,
|
||||
}
|
||||
)
|
||||
.with_start_agent("triage")
|
||||
.add_handoff("triage", ["specialist_a", "specialist_b"])
|
||||
)
|
||||
|
||||
workflow = builder.build()
|
||||
assert "triage" in workflow.executors
|
||||
assert "specialist_a" in workflow.executors
|
||||
assert "specialist_b" in workflow.executors
|
||||
|
||||
|
||||
def test_handoff_builder_accepts_all_instances_in_add_handoff():
|
||||
"""Test that add_handoff accepts all instances when using participants."""
|
||||
triage = MockHandoffAgent(name="triage", handoff_to="specialist_a")
|
||||
specialist_a = MockHandoffAgent(name="specialist_a")
|
||||
specialist_b = MockHandoffAgent(name="specialist_b")
|
||||
|
||||
# This should work - all instances with participants
|
||||
builder = (
|
||||
HandoffBuilder(participants=[triage, specialist_a, specialist_b])
|
||||
.with_start_agent(triage)
|
||||
.add_handoff(triage, [specialist_a, specialist_b])
|
||||
)
|
||||
|
||||
workflow = builder.build()
|
||||
assert "triage" in workflow.executors
|
||||
assert "specialist_a" in workflow.executors
|
||||
assert "specialist_b" in workflow.executors
|
||||
|
||||
|
||||
async def test_handoff_with_participant_factories():
|
||||
"""Test workflow creation using participant_factories."""
|
||||
call_count = 0
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return MockHandoffAgent(name="triage", handoff_to="specialist")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participant_factories={"triage": create_triage, "specialist": create_specialist})
|
||||
.with_start_agent("triage")
|
||||
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == "user") >= 2)
|
||||
.build()
|
||||
)
|
||||
|
||||
# Factories should be called during build
|
||||
assert call_count == 2
|
||||
|
||||
events = await _drain(workflow.run_stream("Need help"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
# Follow-up message
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["More details"])]})
|
||||
)
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs
|
||||
|
||||
|
||||
async def test_handoff_participant_factories_reusable_builder():
|
||||
"""Test that the builder can be reused to build multiple workflows with factories."""
|
||||
call_count = 0
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return MockHandoffAgent(name="triage", handoff_to="specialist")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
builder = HandoffBuilder(
|
||||
participant_factories={"triage": create_triage, "specialist": create_specialist}
|
||||
).with_start_agent("triage")
|
||||
|
||||
# Build first workflow
|
||||
wf1 = builder.build()
|
||||
assert call_count == 2
|
||||
|
||||
# Build second workflow
|
||||
wf2 = builder.build()
|
||||
assert call_count == 4
|
||||
|
||||
# Verify that the two workflows have different agent instances
|
||||
assert wf1.executors["triage"] is not wf2.executors["triage"]
|
||||
assert wf1.executors["specialist"] is not wf2.executors["specialist"]
|
||||
|
||||
|
||||
async def test_handoff_with_participant_factories_and_add_handoff():
|
||||
"""Test that .add_handoff() works correctly with participant_factories."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage", handoff_to="specialist_a")
|
||||
|
||||
def create_specialist_a() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist_a", handoff_to="specialist_b")
|
||||
|
||||
def create_specialist_b() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist_b")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(
|
||||
participant_factories={
|
||||
"triage": create_triage,
|
||||
"specialist_a": create_specialist_a,
|
||||
"specialist_b": create_specialist_b,
|
||||
}
|
||||
)
|
||||
.with_start_agent("triage")
|
||||
.add_handoff("triage", ["specialist_a", "specialist_b"])
|
||||
.add_handoff("specialist_a", ["specialist_b"])
|
||||
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == "user") >= 3)
|
||||
.build()
|
||||
)
|
||||
|
||||
# Start conversation - triage hands off to specialist_a
|
||||
events = await _drain(workflow.run_stream("Initial request"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
# Verify specialist_a executor exists and was called
|
||||
assert "specialist_a" in workflow.executors
|
||||
|
||||
# Second user message - specialist_a hands off to specialist_b
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["Need escalation"])]})
|
||||
)
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
# Verify specialist_b executor exists
|
||||
assert "specialist_b" in workflow.executors
|
||||
|
||||
|
||||
async def test_handoff_participant_factories_with_checkpointing():
|
||||
"""Test checkpointing with participant_factories."""
|
||||
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
|
||||
|
||||
storage = InMemoryCheckpointStorage()
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage", handoff_to="specialist")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participant_factories={"triage": create_triage, "specialist": create_specialist})
|
||||
.with_start_agent("triage")
|
||||
.with_checkpointing(storage)
|
||||
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == "user") >= 2)
|
||||
.build()
|
||||
)
|
||||
|
||||
# Run workflow and capture output
|
||||
events = await _drain(workflow.run_stream("checkpoint test"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: [ChatMessage("user", ["follow up"])]})
|
||||
)
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs, "Should have workflow output after termination condition is met"
|
||||
|
||||
# List checkpoints - just verify they were created
|
||||
checkpoints = await storage.list_checkpoints()
|
||||
assert checkpoints, "Checkpoints should be created during workflow execution"
|
||||
|
||||
|
||||
def test_handoff_set_coordinator_with_factory_name():
|
||||
"""Test that set_coordinator accepts factory name as string."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
builder = HandoffBuilder(
|
||||
participant_factories={"triage": create_triage, "specialist": create_specialist}
|
||||
).with_start_agent("triage")
|
||||
|
||||
workflow = builder.build()
|
||||
assert "triage" in workflow.executors
|
||||
|
||||
|
||||
def test_handoff_add_handoff_with_factory_names():
|
||||
"""Test that add_handoff accepts factory names as strings."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage", handoff_to="specialist_a")
|
||||
|
||||
def create_specialist_a() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist_a")
|
||||
|
||||
def create_specialist_b() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist_b")
|
||||
|
||||
builder = (
|
||||
HandoffBuilder(
|
||||
participant_factories={
|
||||
"triage": create_triage,
|
||||
"specialist_a": create_specialist_a,
|
||||
"specialist_b": create_specialist_b,
|
||||
}
|
||||
)
|
||||
.with_start_agent("triage")
|
||||
.add_handoff("triage", ["specialist_a", "specialist_b"])
|
||||
)
|
||||
|
||||
workflow = builder.build()
|
||||
assert "triage" in workflow.executors
|
||||
assert "specialist_a" in workflow.executors
|
||||
assert "specialist_b" in workflow.executors
|
||||
|
||||
|
||||
async def test_handoff_participant_factories_autonomous_mode():
|
||||
"""Test autonomous mode with participant_factories."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage", handoff_to="specialist")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participant_factories={"triage": create_triage, "specialist": create_specialist})
|
||||
.with_start_agent("triage")
|
||||
.with_autonomous_mode(agents=["specialist"], turn_limits={"specialist": 1})
|
||||
.build()
|
||||
)
|
||||
|
||||
events = await _drain(workflow.run_stream("Issue"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests and len(requests) == 1
|
||||
assert requests[0].source_executor_id == "specialist"
|
||||
|
||||
|
||||
def test_handoff_participant_factories_invalid_coordinator_name():
|
||||
"""Test that set_coordinator raises error for non-existent factory name."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage")
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="Start agent factory name 'nonexistent' is not in the participant_factories list"
|
||||
):
|
||||
(HandoffBuilder(participant_factories={"triage": create_triage}).with_start_agent("nonexistent").build())
|
||||
|
||||
|
||||
def test_handoff_participant_factories_invalid_handoff_target():
|
||||
"""Test that add_handoff raises error for non-existent target factory name."""
|
||||
|
||||
def create_triage() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="triage")
|
||||
|
||||
def create_specialist() -> MockHandoffAgent:
|
||||
return MockHandoffAgent(name="specialist")
|
||||
|
||||
with pytest.raises(ValueError, match="Target factory name 'nonexistent' is not in the participant_factories list"):
|
||||
(
|
||||
HandoffBuilder(participant_factories={"triage": create_triage, "specialist": create_specialist})
|
||||
.with_start_agent("triage")
|
||||
.add_handoff("triage", ["nonexistent"])
|
||||
.build()
|
||||
)
|
||||
|
||||
|
||||
# endregion Participant Factory Tests
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,454 +0,0 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import AsyncIterable
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import (
|
||||
AgentExecutorResponse,
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
AgentThread,
|
||||
BaseAgent,
|
||||
ChatMessage,
|
||||
Content,
|
||||
Executor,
|
||||
SequentialBuilder,
|
||||
TypeCompatibilityError,
|
||||
WorkflowContext,
|
||||
WorkflowOutputEvent,
|
||||
WorkflowRunState,
|
||||
WorkflowStatusEvent,
|
||||
handler,
|
||||
)
|
||||
from agent_framework._workflows._checkpoint import InMemoryCheckpointStorage
|
||||
|
||||
|
||||
class _EchoAgent(BaseAgent):
|
||||
"""Simple agent that appends a single assistant message with its name."""
|
||||
|
||||
async def run( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
return AgentResponse(messages=[ChatMessage("assistant", [f"{self.name} reply"])])
|
||||
|
||||
async def run_stream( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
# Minimal async generator with one assistant update
|
||||
yield AgentResponseUpdate(contents=[Content.from_text(text=f"{self.name} reply")])
|
||||
|
||||
|
||||
class _SummarizerExec(Executor):
|
||||
"""Custom executor that summarizes by appending a short assistant message."""
|
||||
|
||||
@handler
|
||||
async def summarize(self, agent_response: AgentExecutorResponse, ctx: WorkflowContext[list[ChatMessage]]) -> None:
|
||||
conversation = agent_response.full_conversation or []
|
||||
user_texts = [m.text for m in conversation if m.role == "user"]
|
||||
agents = [m.author_name or m.role for m in conversation if m.role == "assistant"]
|
||||
summary = ChatMessage("assistant", [f"Summary of users:{len(user_texts)} agents:{len(agents)}"])
|
||||
await ctx.send_message(list(conversation) + [summary])
|
||||
|
||||
|
||||
class _InvalidExecutor(Executor):
|
||||
"""Invalid executor that does not have a handler that accepts a list of chat messages"""
|
||||
|
||||
@handler
|
||||
async def summarize(self, conversation: list[str], ctx: WorkflowContext[list[ChatMessage]]) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def test_sequential_builder_rejects_empty_participants() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
SequentialBuilder().participants([])
|
||||
|
||||
|
||||
def test_sequential_builder_rejects_empty_participant_factories() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
SequentialBuilder().register_participants([])
|
||||
|
||||
|
||||
def test_sequential_builder_rejects_mixing_participants_and_factories() -> None:
|
||||
"""Test that mixing .participants() and .register_participants() raises an error."""
|
||||
a1 = _EchoAgent(id="agent1", name="A1")
|
||||
|
||||
# Try .participants() then .register_participants()
|
||||
with pytest.raises(ValueError, match="Cannot mix"):
|
||||
SequentialBuilder().participants([a1]).register_participants([lambda: _EchoAgent(id="agent2", name="A2")])
|
||||
|
||||
# Try .register_participants() then .participants()
|
||||
with pytest.raises(ValueError, match="Cannot mix"):
|
||||
SequentialBuilder().register_participants([lambda: _EchoAgent(id="agent1", name="A1")]).participants([a1])
|
||||
|
||||
|
||||
def test_sequential_builder_validation_rejects_invalid_executor() -> None:
|
||||
"""Test that adding an invalid executor to the builder raises an error."""
|
||||
with pytest.raises(TypeCompatibilityError):
|
||||
SequentialBuilder().participants([_EchoAgent(id="agent1", name="A1"), _InvalidExecutor(id="invalid")]).build()
|
||||
|
||||
|
||||
async def test_sequential_agents_append_to_context() -> None:
|
||||
a1 = _EchoAgent(id="agent1", name="A1")
|
||||
a2 = _EchoAgent(id="agent2", name="A2")
|
||||
|
||||
wf = SequentialBuilder().participants([a1, a2]).build()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("hello sequential"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = ev.data # type: ignore[assignment]
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
assert isinstance(output, list)
|
||||
msgs: list[ChatMessage] = output
|
||||
assert len(msgs) == 3
|
||||
assert msgs[0].role == "user" and "hello sequential" in msgs[0].text
|
||||
assert msgs[1].role == "assistant" and (msgs[1].author_name == "A1" or True)
|
||||
assert msgs[2].role == "assistant" and (msgs[2].author_name == "A2" or True)
|
||||
assert "A1 reply" in msgs[1].text
|
||||
assert "A2 reply" in msgs[2].text
|
||||
|
||||
|
||||
async def test_sequential_register_participants_with_agent_factories() -> None:
|
||||
"""Test that register_participants works with agent factories."""
|
||||
|
||||
def create_agent1() -> _EchoAgent:
|
||||
return _EchoAgent(id="agent1", name="A1")
|
||||
|
||||
def create_agent2() -> _EchoAgent:
|
||||
return _EchoAgent(id="agent2", name="A2")
|
||||
|
||||
wf = SequentialBuilder().register_participants([create_agent1, create_agent2]).build()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("hello factories"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = ev.data
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
assert isinstance(output, list)
|
||||
msgs: list[ChatMessage] = output
|
||||
assert len(msgs) == 3
|
||||
assert msgs[0].role == "user" and "hello factories" in msgs[0].text
|
||||
assert msgs[1].role == "assistant" and "A1 reply" in msgs[1].text
|
||||
assert msgs[2].role == "assistant" and "A2 reply" in msgs[2].text
|
||||
|
||||
|
||||
async def test_sequential_with_custom_executor_summary() -> None:
|
||||
a1 = _EchoAgent(id="agent1", name="A1")
|
||||
summarizer = _SummarizerExec(id="summarizer")
|
||||
|
||||
wf = SequentialBuilder().participants([a1, summarizer]).build()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("topic X"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = ev.data
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
msgs: list[ChatMessage] = output
|
||||
# Expect: [user, A1 reply, summary]
|
||||
assert len(msgs) == 3
|
||||
assert msgs[0].role == "user"
|
||||
assert msgs[1].role == "assistant" and "A1 reply" in msgs[1].text
|
||||
assert msgs[2].role == "assistant" and msgs[2].text.startswith("Summary of users:")
|
||||
|
||||
|
||||
async def test_sequential_register_participants_mixed_agents_and_executors() -> None:
|
||||
"""Test register_participants with both agent and executor factories."""
|
||||
|
||||
def create_agent() -> _EchoAgent:
|
||||
return _EchoAgent(id="agent1", name="A1")
|
||||
|
||||
def create_summarizer() -> _SummarizerExec:
|
||||
return _SummarizerExec(id="summarizer")
|
||||
|
||||
wf = SequentialBuilder().register_participants([create_agent, create_summarizer]).build()
|
||||
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("topic Y"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = ev.data
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
msgs: list[ChatMessage] = output
|
||||
# Expect: [user, A1 reply, summary]
|
||||
assert len(msgs) == 3
|
||||
assert msgs[0].role == "user" and "topic Y" in msgs[0].text
|
||||
assert msgs[1].role == "assistant" and "A1 reply" in msgs[1].text
|
||||
assert msgs[2].role == "assistant" and msgs[2].text.startswith("Summary of users:")
|
||||
|
||||
|
||||
async def test_sequential_checkpoint_resume_round_trip() -> None:
|
||||
storage = InMemoryCheckpointStorage()
|
||||
|
||||
initial_agents = (_EchoAgent(id="agent1", name="A1"), _EchoAgent(id="agent2", name="A2"))
|
||||
wf = SequentialBuilder().participants(list(initial_agents)).with_checkpointing(storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("checkpoint sequential"):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
checkpoints = await storage.list_checkpoints()
|
||||
assert checkpoints
|
||||
checkpoints.sort(key=lambda cp: cp.timestamp)
|
||||
|
||||
resume_checkpoint = next(
|
||||
(cp for cp in checkpoints if (cp.metadata or {}).get("checkpoint_type") == "superstep"),
|
||||
checkpoints[-1],
|
||||
)
|
||||
|
||||
resumed_agents = (_EchoAgent(id="agent1", name="A1"), _EchoAgent(id="agent2", name="A2"))
|
||||
wf_resume = SequentialBuilder().participants(list(resumed_agents)).with_checkpointing(storage).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
WorkflowRunState.IDLE,
|
||||
WorkflowRunState.IDLE_WITH_PENDING_REQUESTS,
|
||||
):
|
||||
break
|
||||
|
||||
assert resumed_output is not None
|
||||
assert [m.role for m in resumed_output] == [m.role for m in baseline_output]
|
||||
assert [m.text for m in resumed_output] == [m.text for m in baseline_output]
|
||||
|
||||
|
||||
async def test_sequential_checkpoint_runtime_only() -> None:
|
||||
"""Test checkpointing configured ONLY at runtime, not at build time."""
|
||||
storage = InMemoryCheckpointStorage()
|
||||
|
||||
agents = (_EchoAgent(id="agent1", name="A1"), _EchoAgent(id="agent2", name="A2"))
|
||||
wf = SequentialBuilder().participants(list(agents)).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("runtime checkpoint test", checkpoint_storage=storage):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
checkpoints = await storage.list_checkpoints()
|
||||
assert checkpoints
|
||||
checkpoints.sort(key=lambda cp: cp.timestamp)
|
||||
|
||||
resume_checkpoint = next(
|
||||
(cp for cp in checkpoints if (cp.metadata or {}).get("checkpoint_type") == "superstep"),
|
||||
checkpoints[-1],
|
||||
)
|
||||
|
||||
resumed_agents = (_EchoAgent(id="agent1", name="A1"), _EchoAgent(id="agent2", name="A2"))
|
||||
wf_resume = SequentialBuilder().participants(list(resumed_agents)).build()
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id, checkpoint_storage=storage):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
WorkflowRunState.IDLE,
|
||||
WorkflowRunState.IDLE_WITH_PENDING_REQUESTS,
|
||||
):
|
||||
break
|
||||
|
||||
assert resumed_output is not None
|
||||
assert [m.role for m in resumed_output] == [m.role for m in baseline_output]
|
||||
assert [m.text for m in resumed_output] == [m.text for m in baseline_output]
|
||||
|
||||
|
||||
async def test_sequential_checkpoint_runtime_overrides_buildtime() -> None:
|
||||
"""Test that runtime checkpoint storage overrides build-time configuration."""
|
||||
import tempfile
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir1, tempfile.TemporaryDirectory() as temp_dir2:
|
||||
from agent_framework._workflows._checkpoint import FileCheckpointStorage
|
||||
|
||||
buildtime_storage = FileCheckpointStorage(temp_dir1)
|
||||
runtime_storage = FileCheckpointStorage(temp_dir2)
|
||||
|
||||
agents = (_EchoAgent(id="agent1", name="A1"), _EchoAgent(id="agent2", name="A2"))
|
||||
wf = SequentialBuilder().participants(list(agents)).with_checkpointing(buildtime_storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("override test", checkpoint_storage=runtime_storage):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data # type: ignore[assignment]
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
buildtime_checkpoints = await buildtime_storage.list_checkpoints()
|
||||
runtime_checkpoints = await runtime_storage.list_checkpoints()
|
||||
|
||||
assert len(runtime_checkpoints) > 0, "Runtime storage should have checkpoints"
|
||||
assert len(buildtime_checkpoints) == 0, "Build-time storage should have no checkpoints when overridden"
|
||||
|
||||
|
||||
async def test_sequential_register_participants_with_checkpointing() -> None:
|
||||
"""Test that checkpointing works with register_participants."""
|
||||
storage = InMemoryCheckpointStorage()
|
||||
|
||||
def create_agent1() -> _EchoAgent:
|
||||
return _EchoAgent(id="agent1", name="A1")
|
||||
|
||||
def create_agent2() -> _EchoAgent:
|
||||
return _EchoAgent(id="agent2", name="A2")
|
||||
|
||||
wf = SequentialBuilder().register_participants([create_agent1, create_agent2]).with_checkpointing(storage).build()
|
||||
|
||||
baseline_output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("checkpoint with factories"):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
baseline_output = ev.data
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert baseline_output is not None
|
||||
|
||||
checkpoints = await storage.list_checkpoints()
|
||||
assert checkpoints
|
||||
checkpoints.sort(key=lambda cp: cp.timestamp)
|
||||
|
||||
resume_checkpoint = next(
|
||||
(cp for cp in checkpoints if (cp.metadata or {}).get("checkpoint_type") == "superstep"),
|
||||
checkpoints[-1],
|
||||
)
|
||||
|
||||
wf_resume = (
|
||||
SequentialBuilder().register_participants([create_agent1, create_agent2]).with_checkpointing(storage).build()
|
||||
)
|
||||
|
||||
resumed_output: list[ChatMessage] | None = None
|
||||
async for ev in wf_resume.run_stream(checkpoint_id=resume_checkpoint.checkpoint_id):
|
||||
if isinstance(ev, WorkflowOutputEvent):
|
||||
resumed_output = ev.data
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state in (
|
||||
WorkflowRunState.IDLE,
|
||||
WorkflowRunState.IDLE_WITH_PENDING_REQUESTS,
|
||||
):
|
||||
break
|
||||
|
||||
assert resumed_output is not None
|
||||
assert [m.role for m in resumed_output] == [m.role for m in baseline_output]
|
||||
assert [m.text for m in resumed_output] == [m.text for m in baseline_output]
|
||||
|
||||
|
||||
async def test_sequential_register_participants_factories_called_on_build() -> None:
|
||||
"""Test that factories are called during build(), not during register_participants()."""
|
||||
call_count = 0
|
||||
|
||||
def create_agent() -> _EchoAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return _EchoAgent(id=f"agent{call_count}", name=f"A{call_count}")
|
||||
|
||||
builder = SequentialBuilder().register_participants([create_agent, create_agent])
|
||||
|
||||
# Factories should not be called yet
|
||||
assert call_count == 0
|
||||
|
||||
wf = builder.build()
|
||||
|
||||
# Now factories should have been called
|
||||
assert call_count == 2
|
||||
|
||||
# Run the workflow to ensure it works
|
||||
completed = False
|
||||
output: list[ChatMessage] | None = None
|
||||
async for ev in wf.run_stream("test factories timing"):
|
||||
if isinstance(ev, WorkflowStatusEvent) and ev.state == WorkflowRunState.IDLE:
|
||||
completed = True
|
||||
elif isinstance(ev, WorkflowOutputEvent):
|
||||
output = ev.data # type: ignore[assignment]
|
||||
if completed and output is not None:
|
||||
break
|
||||
|
||||
assert completed
|
||||
assert output is not None
|
||||
msgs: list[ChatMessage] = output
|
||||
# Should have user message + 2 agent replies
|
||||
assert len(msgs) == 3
|
||||
|
||||
|
||||
async def test_sequential_builder_reusable_after_build_with_participants() -> None:
|
||||
"""Test that the builder can be reused to build multiple identical workflows with participants()."""
|
||||
a1 = _EchoAgent(id="agent1", name="A1")
|
||||
a2 = _EchoAgent(id="agent2", name="A2")
|
||||
|
||||
builder = SequentialBuilder().participants([a1, a2])
|
||||
|
||||
# Build first workflow
|
||||
builder.build()
|
||||
|
||||
assert builder._participants[0] is a1 # type: ignore
|
||||
assert builder._participants[1] is a2 # type: ignore
|
||||
assert builder._participant_factories == [] # type: ignore
|
||||
|
||||
|
||||
async def test_sequential_builder_reusable_after_build_with_factories() -> None:
|
||||
"""Test that the builder can be reused to build multiple workflows with register_participants()."""
|
||||
call_count = 0
|
||||
|
||||
def create_agent1() -> _EchoAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return _EchoAgent(id="agent1", name="A1")
|
||||
|
||||
def create_agent2() -> _EchoAgent:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
return _EchoAgent(id="agent2", name="A2")
|
||||
|
||||
builder = SequentialBuilder().register_participants([create_agent1, create_agent2])
|
||||
|
||||
# Build first workflow - factories should be called
|
||||
builder.build()
|
||||
|
||||
assert call_count == 2
|
||||
assert builder._participants == [] # type: ignore
|
||||
assert len(builder._participant_factories) == 2 # type: ignore
|
||||
assert builder._participant_factories[0] is create_agent1 # type: ignore
|
||||
assert builder._participant_factories[1] is create_agent2 # type: ignore
|
||||
@@ -11,17 +11,19 @@ from agent_framework import (
|
||||
AgentThread,
|
||||
BaseAgent,
|
||||
ChatMessage,
|
||||
ConcurrentBuilder,
|
||||
Content,
|
||||
GroupChatBuilder,
|
||||
GroupChatState,
|
||||
HandoffBuilder,
|
||||
SequentialBuilder,
|
||||
WorkflowRunState,
|
||||
WorkflowStatusEvent,
|
||||
tool,
|
||||
)
|
||||
from agent_framework._workflows._const import WORKFLOW_RUN_KWARGS_KEY
|
||||
from agent_framework.orchestrations import (
|
||||
ConcurrentBuilder,
|
||||
GroupChatBuilder,
|
||||
GroupChatState,
|
||||
HandoffBuilder,
|
||||
SequentialBuilder,
|
||||
)
|
||||
|
||||
# Track kwargs received by tools during test execution
|
||||
_received_kwargs: list[dict[str, Any]] = []
|
||||
@@ -371,14 +373,15 @@ async def test_handoff_kwargs_flow_to_agents() -> None:
|
||||
|
||||
async def test_magentic_kwargs_flow_to_agents() -> None:
|
||||
"""Test that kwargs flow to agents in a magentic workflow via MagenticAgentExecutor."""
|
||||
from agent_framework import MagenticBuilder
|
||||
from agent_framework._workflows._magentic import (
|
||||
from agent_framework_orchestrations._magentic import (
|
||||
MagenticContext,
|
||||
MagenticManagerBase,
|
||||
MagenticProgressLedger,
|
||||
MagenticProgressLedgerItem,
|
||||
)
|
||||
|
||||
from agent_framework.orchestrations import MagenticBuilder
|
||||
|
||||
# Create a mock manager that completes after one round
|
||||
class _MockManager(MagenticManagerBase):
|
||||
def __init__(self) -> None:
|
||||
@@ -422,14 +425,15 @@ async def test_magentic_kwargs_flow_to_agents() -> None:
|
||||
|
||||
async def test_magentic_kwargs_stored_in_state() -> None:
|
||||
"""Test that kwargs are stored in State when using MagenticWorkflow.run_stream()."""
|
||||
from agent_framework import MagenticBuilder
|
||||
from agent_framework._workflows._magentic import (
|
||||
from agent_framework_orchestrations._magentic import (
|
||||
MagenticContext,
|
||||
MagenticManagerBase,
|
||||
MagenticProgressLedger,
|
||||
MagenticProgressLedgerItem,
|
||||
)
|
||||
|
||||
from agent_framework.orchestrations import MagenticBuilder
|
||||
|
||||
class _MockManager(MagenticManagerBase):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(max_stall_count=3, max_reset_count=None, max_round_count=1)
|
||||
|
||||
Reference in New Issue
Block a user