mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Move Workflow, Edge, EdgeGroup and Executor to AFBaseModel (#472)
* Refactor workflow to introduce EdgeRunner for edge execution. * Fix edge cases * Convert Workflow, Edge, EdgeGroup, and Executor into AFBaseModel to support object model serialization * format * remove accidental file * fix typing * Add type information to EdgeGroup and Executor subclasses * fix format * Add condition_name field to Edge * Add new fields * remove Optional * Update
This commit is contained in:
committed by
GitHub
Unverified
parent
88c51013c6
commit
6e9d35830f
@@ -8,14 +8,17 @@ import pytest
|
||||
from agent_framework.workflow import Executor, WorkflowContext, handler
|
||||
|
||||
from agent_framework_workflow._edge import (
|
||||
Case,
|
||||
Default,
|
||||
Edge,
|
||||
FanInEdgeGroup,
|
||||
FanOutEdgeGroup,
|
||||
SingleEdgeGroup,
|
||||
SwitchCaseEdgeGroup,
|
||||
SwitchCaseEdgeGroupCase,
|
||||
SwitchCaseEdgeGroupDefault,
|
||||
)
|
||||
from agent_framework_workflow._edge_runner import create_edge_runner
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -35,28 +38,40 @@ class MockMessageSecondary:
|
||||
class MockExecutor(Executor):
|
||||
"""A mock executor for testing purposes."""
|
||||
|
||||
call_count: int = 0
|
||||
last_message: Any = None
|
||||
|
||||
@handler
|
||||
async def mock_handler(self, message: MockMessage, ctx: WorkflowContext) -> None:
|
||||
"""A mock handler that does nothing."""
|
||||
pass
|
||||
self.call_count += 1
|
||||
self.last_message = message
|
||||
|
||||
|
||||
class MockExecutorSecondary(Executor):
|
||||
"""A secondary mock executor for testing purposes."""
|
||||
|
||||
call_count: int = 0
|
||||
last_message: Any = None
|
||||
|
||||
@handler
|
||||
async def mock_handler_secondary(self, message: MockMessageSecondary, ctx: WorkflowContext) -> None:
|
||||
"""A secondary mock handler that does nothing."""
|
||||
pass
|
||||
self.call_count += 1
|
||||
self.last_message = message
|
||||
|
||||
|
||||
class MockAggregator(Executor):
|
||||
"""A mock aggregator for testing purposes."""
|
||||
|
||||
call_count: int = 0
|
||||
last_message: Any = None
|
||||
|
||||
@handler
|
||||
async def mock_aggregator_handler(self, message: list[MockMessage], ctx: WorkflowContext) -> None:
|
||||
"""A mock aggregator handler that does nothing."""
|
||||
pass
|
||||
self.call_count += 1
|
||||
self.last_message = message
|
||||
|
||||
|
||||
# region Edge
|
||||
@@ -67,7 +82,7 @@ def test_create_edge():
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge = Edge(source=source, target=target)
|
||||
edge = Edge(source_id=source.id, target_id=target.id)
|
||||
|
||||
assert edge.source_id == "source_executor"
|
||||
assert edge.target_id == "target_executor"
|
||||
@@ -79,9 +94,9 @@ def test_edge_can_handle():
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge = Edge(source=source, target=target)
|
||||
edge = Edge(source_id=source.id, target_id=target.id)
|
||||
|
||||
assert edge.can_handle(MockMessage(data="test"))
|
||||
assert edge.should_route(MockMessage(data="test"))
|
||||
|
||||
|
||||
# endregion Edge
|
||||
@@ -94,10 +109,10 @@ def test_single_edge_group():
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge_group = SingleEdgeGroup(source=source, target=target)
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
assert edge_group.source_executors == [source]
|
||||
assert edge_group.target_executors == [target]
|
||||
assert edge_group.source_executor_ids == [source.id]
|
||||
assert edge_group.target_executor_ids == [target.id]
|
||||
assert edge_group.edges[0].source_id == "source_executor"
|
||||
assert edge_group.edges[0].target_id == "target_executor"
|
||||
|
||||
@@ -107,92 +122,88 @@ def test_single_edge_group_with_condition():
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge_group = SingleEdgeGroup(source=source, target=target, condition=lambda x: x.data == "test")
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id, condition=lambda x: x.data == "test")
|
||||
|
||||
assert edge_group.source_executors == [source]
|
||||
assert edge_group.target_executors == [target]
|
||||
assert edge_group.source_executor_ids == [source.id]
|
||||
assert edge_group.target_executor_ids == [target.id]
|
||||
assert edge_group.edges[0].source_id == "source_executor"
|
||||
assert edge_group.edges[0].target_id == "target_executor"
|
||||
assert edge_group.edges[0]._condition is not None # type: ignore
|
||||
|
||||
|
||||
async def test_single_edge_group_send_message():
|
||||
"""Test sending a message through a single edge group."""
|
||||
async def test_single_edge_group_send_message() -> None:
|
||||
"""Test sending a message through a single edge runner."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge_group = SingleEdgeGroup(source=source, target=target)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target.id: target}
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is True
|
||||
|
||||
|
||||
async def test_single_edge_group_send_message_with_target():
|
||||
"""Test sending a message through a single edge group."""
|
||||
async def test_single_edge_group_send_message_with_target() -> None:
|
||||
"""Test sending a message through a single edge runner."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge_group = SingleEdgeGroup(source=source, target=target)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target.id: target}
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is True
|
||||
|
||||
|
||||
async def test_single_edge_group_send_message_with_invalid_target():
|
||||
"""Test sending a message through a single edge group."""
|
||||
async def test_single_edge_group_send_message_with_invalid_target() -> None:
|
||||
"""Test sending a message through a single edge runner."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge_group = SingleEdgeGroup(source=source, target=target)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target.id: target}
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_single_edge_group_send_message_with_invalid_data():
|
||||
"""Test sending a message through a single edge group."""
|
||||
async def test_single_edge_group_send_message_with_invalid_data() -> None:
|
||||
"""Test sending a message through a single edge runner with invalid data."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
edge_group = SingleEdgeGroup(source=source, target=target)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target.id: target}
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -208,10 +219,10 @@ def test_source_edge_group():
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(source=source, targets=[target1, target2])
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
assert edge_group.source_executors == [source]
|
||||
assert edge_group.target_executors == [target1, target2]
|
||||
assert edge_group.source_executor_ids == [source.id]
|
||||
assert edge_group.target_executor_ids == [target1.id, target2.id]
|
||||
assert len(edge_group.edges) == 2
|
||||
assert edge_group.edges[0].source_id == "source_executor"
|
||||
assert edge_group.edges[0].target_id == "target_executor_1"
|
||||
@@ -219,128 +230,122 @@ def test_source_edge_group():
|
||||
assert edge_group.edges[1].target_id == "target_executor_2"
|
||||
|
||||
|
||||
def test_source_edge_group_invalid_number_of_targets():
|
||||
def test_source_edge_group_invalid_number_of_targets() -> None:
|
||||
"""Test creating a fan-out group with an invalid number of targets."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target = MockExecutor(id="target_executor")
|
||||
|
||||
with pytest.raises(ValueError, match="FanOutEdgeGroup must contain at least two targets"):
|
||||
FanOutEdgeGroup(source=source, targets=[target])
|
||||
FanOutEdgeGroup(source_id=source.id, target_ids=[target.id])
|
||||
|
||||
|
||||
async def test_source_edge_group_send_message():
|
||||
"""Test sending a message through a fan-out group."""
|
||||
async def test_source_edge_group_send_message() -> None:
|
||||
"""Test sending a message through a fan-out edge runner."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(source=source, targets=[target1, target2])
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 2
|
||||
assert success is True
|
||||
assert target1.call_count == 1
|
||||
assert target2.call_count == 1
|
||||
|
||||
|
||||
async def test_source_edge_group_send_message_with_target():
|
||||
async def test_source_edge_group_send_message_with_target() -> None:
|
||||
"""Test sending a message through a fan-out group with a target."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(source=source, targets=[target1, target2])
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
assert mock_send.call_args[0][0].target_id == target1.id
|
||||
assert success is True
|
||||
assert target1.call_count == 1
|
||||
assert target2.call_count == 0 # target2 should not be called since message targets target1
|
||||
|
||||
|
||||
async def test_source_edge_group_send_message_with_invalid_target():
|
||||
async def test_source_edge_group_send_message_with_invalid_target() -> None:
|
||||
"""Test sending a message through a fan-out group with an invalid target."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(source=source, targets=[target1, target2])
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_source_edge_group_send_message_with_invalid_data():
|
||||
async def test_source_edge_group_send_message_with_invalid_data() -> None:
|
||||
"""Test sending a message through a fan-out group with invalid data."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(source=source, targets=[target1, target2])
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_source_edge_group_send_message_only_one_successful_send():
|
||||
async def test_source_edge_group_send_message_only_one_successful_send() -> None:
|
||||
"""Test sending a message through a fan-out group where only one edge can handle the message."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutorSecondary(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(source=source, targets=[target1, target2])
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
assert success is True
|
||||
assert target1.call_count == 1 # target1 can handle MockMessage
|
||||
assert target2.call_count == 0 # target2 (MockExecutorSecondary) cannot handle MockMessage
|
||||
|
||||
|
||||
def test_source_edge_group_with_selection_func():
|
||||
@@ -350,13 +355,13 @@ def test_source_edge_group_with_selection_func():
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source,
|
||||
targets=[target1, target2],
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id],
|
||||
)
|
||||
|
||||
assert edge_group.source_executors == [source]
|
||||
assert edge_group.target_executors == [target1, target2]
|
||||
assert edge_group.source_executor_ids == [source.id]
|
||||
assert edge_group.target_executor_ids == [target1.id, target2.id]
|
||||
assert len(edge_group.edges) == 2
|
||||
assert edge_group.edges[0].source_id == "source_executor"
|
||||
assert edge_group.edges[0].target_id == "target_executor_1"
|
||||
@@ -364,20 +369,20 @@ def test_source_edge_group_with_selection_func():
|
||||
assert edge_group.edges[1].target_id == "target_executor_2"
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message():
|
||||
async def test_source_edge_group_with_selection_func_send_message() -> None:
|
||||
"""Test sending a message through a fan-out group with a selection function."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source,
|
||||
targets=[target1, target2],
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id, target2.id],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -385,28 +390,28 @@ async def test_source_edge_group_with_selection_func_send_message():
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
with patch("agent_framework_workflow._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
|
||||
assert mock_send.call_count == 2
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_invalid_selection_result():
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_invalid_selection_result() -> None:
|
||||
"""Test sending a message through a fan-out group with a selection func with an invalid selection result."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source,
|
||||
targets=[target1, target2],
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id, "invalid_target"],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
@@ -414,23 +419,23 @@ async def test_source_edge_group_with_selection_func_send_message_with_invalid_s
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
await edge_group.send_message(message, shared_state, ctx)
|
||||
await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target():
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target() -> None:
|
||||
"""Test sending a message through a fan-out group with a selection func with a target."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source,
|
||||
targets=[target1, target2],
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id, target2.id],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -438,28 +443,28 @@ async def test_source_edge_group_with_selection_func_send_message_with_target():
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
with patch("agent_framework_workflow._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
assert mock_send.call_args[0][0].target_id == target1.id
|
||||
assert mock_send.call_args[0][0] == target1.id
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target_not_in_selection():
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target_not_in_selection() -> None:
|
||||
"""Test sending a message through a fan-out group with a selection func with a target not in the selection."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source,
|
||||
targets=[target1, target2],
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id], # Only target1 will receive the message
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -467,22 +472,24 @@ async def test_source_edge_group_with_selection_func_send_message_with_target_no
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target2.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_invalid_data():
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_invalid_data() -> None:
|
||||
"""Test sending a message through a fan-out group with a selection func with invalid data."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source, targets=[target1, target2], selection_func=lambda data, target_ids: [target1.id, target2.id]
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id, target2.id],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -490,22 +497,24 @@ async def test_source_edge_group_with_selection_func_send_message_with_invalid_d
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target_invalid_data():
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target_invalid_data() -> None:
|
||||
"""Test sending a message through a fan-out group with a selection func with a target and invalid data."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = FanOutEdgeGroup(
|
||||
source=source, targets=[target1, target2], selection_func=lambda data, target_ids: [target1.id, target2.id]
|
||||
source_id=source.id,
|
||||
target_ids=[target1.id, target2.id],
|
||||
selection_func=lambda data, target_ids: [target1.id, target2.id],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -513,7 +522,7 @@ async def test_source_edge_group_with_selection_func_send_message_with_target_in
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -528,10 +537,10 @@ def test_target_edge_group():
|
||||
source2 = MockExecutor(id="source_executor_2")
|
||||
target = MockAggregator(id="target_executor")
|
||||
|
||||
edge_group = FanInEdgeGroup(sources=[source1, source2], target=target)
|
||||
edge_group = FanInEdgeGroup(source_ids=[source1.id, source2.id], target_id=target.id)
|
||||
|
||||
assert edge_group.source_executors == [source1, source2]
|
||||
assert edge_group.target_executors == [target]
|
||||
assert edge_group.source_executor_ids == [source1.id, source2.id]
|
||||
assert edge_group.target_executor_ids == [target.id]
|
||||
assert len(edge_group.edges) == 2
|
||||
assert edge_group.edges[0].source_id == "source_executor_1"
|
||||
assert edge_group.edges[0].target_id == "target_executor"
|
||||
@@ -545,27 +554,27 @@ def test_target_edge_group_invalid_number_of_sources():
|
||||
target = MockAggregator(id="target_executor")
|
||||
|
||||
with pytest.raises(ValueError, match="FanInEdgeGroup must contain at least two sources"):
|
||||
FanInEdgeGroup(sources=[source], target=target)
|
||||
FanInEdgeGroup(source_ids=[source.id], target_id=target.id)
|
||||
|
||||
|
||||
async def test_target_edge_group_send_message_buffer():
|
||||
async def test_target_edge_group_send_message_buffer() -> None:
|
||||
"""Test sending a message through a fan-in edge group with buffering."""
|
||||
source1 = MockExecutor(id="source_executor_1")
|
||||
source2 = MockExecutor(id="source_executor_2")
|
||||
target = MockAggregator(id="target_executor")
|
||||
|
||||
edge_group = FanInEdgeGroup(sources=[source1, source2], target=target)
|
||||
edge_group = FanInEdgeGroup(source_ids=[source1.id, source2.id], target_id=target.id)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(
|
||||
with patch("agent_framework_workflow._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source1.id),
|
||||
shared_state,
|
||||
ctx,
|
||||
@@ -573,9 +582,9 @@ async def test_target_edge_group_send_message_buffer():
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 0 # The message should be buffered and wait for the second source
|
||||
assert len(edge_group._buffer[source1.id]) == 1 # type: ignore
|
||||
assert len(edge_runner._buffer[source1.id]) == 1 # type: ignore
|
||||
|
||||
success = await edge_group.send_message(
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source2.id),
|
||||
shared_state,
|
||||
ctx,
|
||||
@@ -584,19 +593,19 @@ async def test_target_edge_group_send_message_buffer():
|
||||
assert mock_send.call_count == 1 # The message should be sent now that both sources have sent their messages
|
||||
|
||||
# Buffer should be cleared after sending
|
||||
assert not edge_group._buffer # type: ignore
|
||||
assert not edge_runner._buffer # type: ignore
|
||||
|
||||
|
||||
async def test_target_edge_group_send_message_with_invalid_target():
|
||||
async def test_target_edge_group_send_message_with_invalid_target() -> None:
|
||||
"""Test sending a message through a fan-in edge group with an invalid target."""
|
||||
source1 = MockExecutor(id="source_executor_1")
|
||||
source2 = MockExecutor(id="source_executor_2")
|
||||
target = MockAggregator(id="target_executor")
|
||||
|
||||
edge_group = FanInEdgeGroup(sources=[source1, source2], target=target)
|
||||
edge_group = FanInEdgeGroup(source_ids=[source1.id, source2.id], target_id=target.id)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -604,20 +613,20 @@ async def test_target_edge_group_send_message_with_invalid_target():
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source1.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_target_edge_group_send_message_with_invalid_data():
|
||||
async def test_target_edge_group_send_message_with_invalid_data() -> None:
|
||||
"""Test sending a message through a fan-in edge group with invalid data."""
|
||||
source1 = MockExecutor(id="source_executor_1")
|
||||
source2 = MockExecutor(id="source_executor_2")
|
||||
target = MockAggregator(id="target_executor")
|
||||
|
||||
edge_group = FanInEdgeGroup(sources=[source1, source2], target=target)
|
||||
edge_group = FanInEdgeGroup(source_ids=[source1.id, source2.id], target_id=target.id)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -625,7 +634,7 @@ async def test_target_edge_group_send_message_with_invalid_data():
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source1.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -634,22 +643,22 @@ async def test_target_edge_group_send_message_with_invalid_data():
|
||||
# region SwitchCaseEdgeGroup
|
||||
|
||||
|
||||
def test_switch_case_edge_group():
|
||||
def test_switch_case_edge_group() -> None:
|
||||
"""Test creating a switch case edge group."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target1),
|
||||
Default(target=target2),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target1.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
],
|
||||
)
|
||||
|
||||
assert edge_group.source_executors == [source]
|
||||
assert edge_group.target_executors == [target1, target2]
|
||||
assert edge_group.source_executor_ids == [source.id]
|
||||
assert edge_group.target_executor_ids == [target1.id, target2.id]
|
||||
assert len(edge_group.edges) == 2
|
||||
assert edge_group.edges[0].source_id == "source_executor"
|
||||
assert edge_group.edges[0].target_id == "target_executor_1"
|
||||
@@ -670,18 +679,18 @@ def test_switch_case_edge_group_invalid_number_of_cases():
|
||||
ValueError, match=r"SwitchCaseEdgeGroup must contain at least two cases \(including the default case\)."
|
||||
):
|
||||
SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target.id),
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="SwitchCaseEdgeGroup must contain exactly one default case."):
|
||||
SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target),
|
||||
Case(condition=lambda x: x.data >= 0, target=target),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target.id),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data >= 0, target_id=target.id),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -694,31 +703,30 @@ def test_switch_case_edge_group_invalid_number_of_default_cases():
|
||||
|
||||
with pytest.raises(ValueError, match="SwitchCaseEdgeGroup must contain exactly one default case."):
|
||||
SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target1),
|
||||
Default(target=target2),
|
||||
Default(target=target2),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target1.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
async def test_switch_case_edge_group_send_message():
|
||||
async def test_switch_case_edge_group_send_message() -> None:
|
||||
"""Test sending a message through a switch case edge group."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target1),
|
||||
Default(target=target2),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target1.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -726,8 +734,8 @@ async def test_switch_case_edge_group_send_message():
|
||||
data = MockMessage(data=-1)
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
with patch("agent_framework_workflow._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
@@ -735,29 +743,29 @@ async def test_switch_case_edge_group_send_message():
|
||||
# Default condition should
|
||||
data = MockMessage(data=1)
|
||||
message = Message(data=data, source_id=source.id)
|
||||
with patch("agent_framework_workflow._edge.Edge.send_message") as mock_send:
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
with patch("agent_framework_workflow._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
|
||||
|
||||
async def test_switch_case_edge_group_send_message_with_invalid_target():
|
||||
async def test_switch_case_edge_group_send_message_with_invalid_target() -> None:
|
||||
"""Test sending a message through a switch case edge group with an invalid target."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target1),
|
||||
Default(target=target2),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target1.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -765,26 +773,26 @@ async def test_switch_case_edge_group_send_message_with_invalid_target():
|
||||
data = MockMessage(data=-1)
|
||||
message = Message(data=data, source_id=source.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
async def test_switch_case_edge_group_send_message_with_valid_target():
|
||||
async def test_switch_case_edge_group_send_message_with_valid_target() -> None:
|
||||
"""Test sending a message through a switch case edge group with a target."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target1),
|
||||
Default(target=target2),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target1.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -792,31 +800,31 @@ async def test_switch_case_edge_group_send_message_with_valid_target():
|
||||
data = MockMessage(data=1) # Condition will fail
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
data = MockMessage(data=-1) # Condition will pass
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is True
|
||||
|
||||
|
||||
async def test_switch_case_edge_group_send_message_with_invalid_data():
|
||||
async def test_switch_case_edge_group_send_message_with_invalid_data() -> None:
|
||||
"""Test sending a message through a switch case edge group with invalid data."""
|
||||
source = MockExecutor(id="source_executor")
|
||||
target1 = MockExecutor(id="target_executor_1")
|
||||
target2 = MockExecutor(id="target_executor_2")
|
||||
|
||||
edge_group = SwitchCaseEdgeGroup(
|
||||
source=source,
|
||||
source_id=source.id,
|
||||
cases=[
|
||||
Case(condition=lambda x: x.data < 0, target=target1),
|
||||
Default(target=target2),
|
||||
SwitchCaseEdgeGroupCase(condition=lambda x: x.data < 0, target_id=target1.id),
|
||||
SwitchCaseEdgeGroupDefault(target_id=target2.id),
|
||||
],
|
||||
)
|
||||
|
||||
from agent_framework_workflow._runner_context import InProcRunnerContext, Message
|
||||
from agent_framework_workflow._shared_state import SharedState
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
ctx = InProcRunnerContext()
|
||||
@@ -824,7 +832,7 @@ async def test_switch_case_edge_group_send_message_with_invalid_data():
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_group.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user