mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
[BREAKING] Python: Refactor SharedState to State with sync methods and superstep caching (#3667)
* Refactor SharedState to State with sync methods and superstep caching * Fixes * Address PR feedback * Remove dead links * Fix lab test import
This commit is contained in:
committed by
GitHub
Unverified
parent
4e25917644
commit
10afb86213
@@ -93,8 +93,8 @@ async def test_agent_executor_checkpoint_stores_and_restores_state() -> None:
|
||||
)
|
||||
|
||||
# Verify checkpoint contains executor state with both cache and thread
|
||||
assert "_executor_state" in restore_checkpoint.shared_state
|
||||
executor_states = restore_checkpoint.shared_state["_executor_state"]
|
||||
assert "_executor_state" in restore_checkpoint.state
|
||||
executor_states = restore_checkpoint.state["_executor_state"]
|
||||
assert isinstance(executor_states, dict)
|
||||
assert executor.id in executor_states
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ def test_workflow_checkpoint_default_values():
|
||||
assert checkpoint.workflow_id == ""
|
||||
assert checkpoint.timestamp != ""
|
||||
assert checkpoint.messages == {}
|
||||
assert checkpoint.shared_state == {}
|
||||
assert checkpoint.state == {}
|
||||
assert checkpoint.pending_request_info_events == {}
|
||||
assert checkpoint.iteration_count == 0
|
||||
assert checkpoint.metadata == {}
|
||||
@@ -34,7 +34,7 @@ def test_workflow_checkpoint_custom_values():
|
||||
timestamp=custom_timestamp,
|
||||
messages={"executor1": [{"data": "test"}]},
|
||||
pending_request_info_events={"req123": {"data": "test"}},
|
||||
shared_state={"key": "value"},
|
||||
state={"key": "value"},
|
||||
iteration_count=5,
|
||||
metadata={"test": True},
|
||||
version="2.0",
|
||||
@@ -44,7 +44,7 @@ def test_workflow_checkpoint_custom_values():
|
||||
assert checkpoint.workflow_id == "test-workflow-456"
|
||||
assert checkpoint.timestamp == custom_timestamp
|
||||
assert checkpoint.messages == {"executor1": [{"data": "test"}]}
|
||||
assert checkpoint.shared_state == {"key": "value"}
|
||||
assert checkpoint.state == {"key": "value"}
|
||||
assert checkpoint.pending_request_info_events == {"req123": {"data": "test"}}
|
||||
assert checkpoint.iteration_count == 5
|
||||
assert checkpoint.metadata == {"test": True}
|
||||
@@ -159,7 +159,7 @@ async def test_file_checkpoint_storage_save_and_load():
|
||||
checkpoint = WorkflowCheckpoint(
|
||||
workflow_id="test-workflow",
|
||||
messages={"executor1": [{"data": "hello", "source_id": "test", "target_id": None}]},
|
||||
shared_state={"key": "value"},
|
||||
state={"key": "value"},
|
||||
pending_request_info_events={"req123": {"data": "test"}},
|
||||
)
|
||||
|
||||
@@ -177,7 +177,7 @@ async def test_file_checkpoint_storage_save_and_load():
|
||||
assert loaded_checkpoint.checkpoint_id == checkpoint.checkpoint_id
|
||||
assert loaded_checkpoint.workflow_id == checkpoint.workflow_id
|
||||
assert loaded_checkpoint.messages == checkpoint.messages
|
||||
assert loaded_checkpoint.shared_state == checkpoint.shared_state
|
||||
assert loaded_checkpoint.state == checkpoint.state
|
||||
assert loaded_checkpoint.pending_request_info_events == checkpoint.pending_request_info_events
|
||||
|
||||
|
||||
@@ -293,7 +293,7 @@ async def test_file_checkpoint_storage_json_serialization():
|
||||
checkpoint = WorkflowCheckpoint(
|
||||
workflow_id="complex-workflow",
|
||||
messages={"executor1": [{"data": {"nested": {"value": 42}}, "source_id": "test", "target_id": None}]},
|
||||
shared_state={"list": [1, 2, 3], "dict": {"a": "b", "c": {"d": "e"}}, "bool": True, "null": None},
|
||||
state={"list": [1, 2, 3], "dict": {"a": "b", "c": {"d": "e"}}, "bool": True, "null": None},
|
||||
pending_request_info_events={"req123": {"data": "test"}},
|
||||
)
|
||||
|
||||
@@ -303,7 +303,7 @@ async def test_file_checkpoint_storage_json_serialization():
|
||||
|
||||
assert loaded is not None
|
||||
assert loaded.messages == checkpoint.messages
|
||||
assert loaded.shared_state == checkpoint.shared_state
|
||||
assert loaded.state == checkpoint.state
|
||||
|
||||
# Verify the JSON file is properly formatted
|
||||
file_path = Path(temp_dir) / f"{checkpoint.checkpoint_id}.json"
|
||||
@@ -311,9 +311,9 @@ async def test_file_checkpoint_storage_json_serialization():
|
||||
data = json.load(f)
|
||||
|
||||
assert data["messages"]["executor1"][0]["data"]["nested"]["value"] == 42
|
||||
assert data["shared_state"]["list"] == [1, 2, 3]
|
||||
assert data["shared_state"]["bool"] is True
|
||||
assert data["shared_state"]["null"] is None
|
||||
assert data["state"]["list"] == [1, 2, 3]
|
||||
assert data["state"]["bool"] is True
|
||||
assert data["state"]["null"] is None
|
||||
assert data["pending_request_info_events"]["req123"]["data"] == "test"
|
||||
|
||||
|
||||
|
||||
@@ -23,11 +23,9 @@ from agent_framework._workflows._edge import (
|
||||
SwitchCaseEdgeGroupDefault,
|
||||
)
|
||||
from agent_framework._workflows._edge_runner import create_edge_runner
|
||||
from agent_framework._workflows._shared_state import SharedState
|
||||
from agent_framework._workflows._state import State
|
||||
from agent_framework.observability import EdgeGroupDeliveryStatus
|
||||
|
||||
# Add for test
|
||||
|
||||
|
||||
@dataclass
|
||||
class MockMessage:
|
||||
@@ -191,13 +189,13 @@ async def test_single_edge_group_send_message() -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
|
||||
|
||||
@@ -210,13 +208,13 @@ async def test_single_edge_group_send_message_with_target() -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
|
||||
|
||||
@@ -229,13 +227,13 @@ async def test_single_edge_group_send_message_with_invalid_target() -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -248,13 +246,13 @@ async def test_single_edge_group_send_message_with_invalid_data() -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -268,13 +266,13 @@ async def test_single_edge_group_send_message_with_condition_pass() -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id, condition=lambda x: x.data == "test")
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
assert target.call_count == 1
|
||||
assert target.last_message.data == "test"
|
||||
@@ -290,13 +288,13 @@ async def test_single_edge_group_send_message_with_condition_fail() -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id, condition=lambda x: x.data == "test")
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="different")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
# Should return True because message was processed, but condition failed
|
||||
assert success is True
|
||||
# Target should not be called because condition failed
|
||||
@@ -312,7 +310,7 @@ async def test_single_edge_group_tracing_success(span_exporter) -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
# Create trace context and span IDs to simulate a message with tracing information
|
||||
@@ -325,7 +323,7 @@ async def test_single_edge_group_tracing_success(span_exporter) -> None:
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -361,7 +359,7 @@ async def test_single_edge_group_tracing_condition_failure(span_exporter) -> Non
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id, condition=lambda x: x.data == "pass")
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="fail")
|
||||
@@ -370,7 +368,7 @@ async def test_single_edge_group_tracing_condition_failure(span_exporter) -> Non
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True # Returns True but condition failed
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -395,7 +393,7 @@ async def test_single_edge_group_tracing_type_mismatch(span_exporter) -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
# Send incompatible data type
|
||||
@@ -405,7 +403,7 @@ async def test_single_edge_group_tracing_type_mismatch(span_exporter) -> None:
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -430,7 +428,7 @@ async def test_single_edge_group_tracing_target_mismatch(span_exporter) -> None:
|
||||
edge_group = SingleEdgeGroup(source_id=source.id, target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
@@ -439,7 +437,7 @@ async def test_single_edge_group_tracing_target_mismatch(span_exporter) -> None:
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -498,13 +496,13 @@ async def test_source_edge_group_send_message() -> None:
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert target1.call_count == 1
|
||||
@@ -521,13 +519,13 @@ async def test_source_edge_group_send_message_with_target() -> None:
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert target1.call_count == 1
|
||||
@@ -544,13 +542,13 @@ async def test_source_edge_group_send_message_with_invalid_target() -> None:
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -564,13 +562,13 @@ async def test_source_edge_group_send_message_with_invalid_data() -> None:
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -584,13 +582,13 @@ async def test_source_edge_group_send_message_only_one_successful_send() -> None
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert target1.call_count == 1 # target1 can handle MockMessage
|
||||
@@ -633,14 +631,14 @@ async def test_source_edge_group_with_selection_func_send_message() -> None:
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with patch("agent_framework._workflows._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
|
||||
@@ -661,14 +659,14 @@ async def test_source_edge_group_with_selection_func_send_message_with_invalid_s
|
||||
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
await edge_runner.send_message(message, shared_state, ctx)
|
||||
await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
|
||||
async def test_source_edge_group_with_selection_func_send_message_with_target() -> None:
|
||||
@@ -686,14 +684,14 @@ async def test_source_edge_group_with_selection_func_send_message_with_target()
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
with patch("agent_framework._workflows._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
@@ -715,13 +713,13 @@ async def test_source_edge_group_with_selection_func_send_message_with_target_no
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source.id, target_id=target2.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -740,13 +738,13 @@ async def test_source_edge_group_with_selection_func_send_message_with_invalid_d
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -765,13 +763,13 @@ async def test_source_edge_group_with_selection_func_send_message_with_target_in
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -785,7 +783,7 @@ async def test_fan_out_edge_group_tracing_success(span_exporter) -> None:
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
# Create trace context and span IDs to simulate a message with tracing information
|
||||
@@ -798,7 +796,7 @@ async def test_fan_out_edge_group_tracing_success(span_exporter) -> None:
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -835,7 +833,7 @@ async def test_fan_out_edge_group_tracing_with_target(span_exporter) -> None:
|
||||
edge_group = FanOutEdgeGroup(source_id=source.id, target_ids=[target1.id, target2.id])
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
# Create trace context and span IDs to simulate a message with tracing information
|
||||
@@ -854,7 +852,7 @@ async def test_fan_out_edge_group_tracing_with_target(span_exporter) -> None:
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -922,7 +920,7 @@ async def test_target_edge_group_send_message_buffer() -> None:
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
@@ -930,7 +928,7 @@ async def test_target_edge_group_send_message_buffer() -> None:
|
||||
with patch("agent_framework._workflows._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source1.id),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
|
||||
@@ -940,7 +938,7 @@ async def test_target_edge_group_send_message_buffer() -> None:
|
||||
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source2.id),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
assert success is True
|
||||
@@ -961,13 +959,13 @@ async def test_target_edge_group_send_message_with_invalid_target() -> None:
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
message = Message(data=data, source_id=source1.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -982,13 +980,13 @@ async def test_target_edge_group_send_message_with_invalid_data() -> None:
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source1.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -1002,7 +1000,7 @@ async def test_fan_in_edge_group_tracing_buffered(span_exporter) -> None:
|
||||
edge_group = FanInEdgeGroup(source_ids=[source1.id, source2.id], target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
@@ -1020,7 +1018,7 @@ async def test_fan_in_edge_group_tracing_buffered(span_exporter) -> None:
|
||||
# Send first message (should be buffered)
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source1.id, trace_contexts=trace_contexts1, source_span_ids=source_span_ids1),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
assert success is True
|
||||
@@ -1052,7 +1050,7 @@ async def test_fan_in_edge_group_tracing_buffered(span_exporter) -> None:
|
||||
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source2.id, trace_contexts=trace_contexts2, source_span_ids=source_span_ids2),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
assert success is True
|
||||
@@ -1090,7 +1088,7 @@ async def test_fan_in_edge_group_tracing_type_mismatch(span_exporter) -> None:
|
||||
edge_group = FanInEdgeGroup(source_ids=[source1.id, source2.id], target_id=target.id)
|
||||
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
# Send incompatible data type
|
||||
@@ -1100,7 +1098,7 @@ async def test_fan_in_edge_group_tracing_type_mismatch(span_exporter) -> None:
|
||||
# Clear any build spans
|
||||
span_exporter.clear()
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
@@ -1126,14 +1124,14 @@ async def test_fan_in_edge_group_with_multiple_message_types() -> None:
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source1.id),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
assert success
|
||||
@@ -1141,7 +1139,7 @@ async def test_fan_in_edge_group_with_multiple_message_types() -> None:
|
||||
data2 = MockMessageSecondary(data="test")
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data2, source_id=source2.id),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
assert success
|
||||
@@ -1157,14 +1155,14 @@ async def test_fan_in_edge_group_with_multiple_message_types_failed() -> None:
|
||||
executors: dict[str, Executor] = {source1.id: source1, source2.id: source2, target.id: target}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data="test")
|
||||
|
||||
success = await edge_runner.send_message(
|
||||
Message(data=data, source_id=source1.id),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
assert success
|
||||
@@ -1178,7 +1176,7 @@ async def test_fan_in_edge_group_with_multiple_message_types_failed() -> None:
|
||||
data2 = MockMessageSecondary(data="test")
|
||||
_ = await edge_runner.send_message(
|
||||
Message(data=data2, source_id=source2.id),
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
|
||||
@@ -1273,14 +1271,14 @@ async def test_switch_case_edge_group_send_message() -> None:
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data=-1)
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
with patch("agent_framework._workflows._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
@@ -1289,7 +1287,7 @@ async def test_switch_case_edge_group_send_message() -> None:
|
||||
data = MockMessage(data=1)
|
||||
message = Message(data=data, source_id=source.id)
|
||||
with patch("agent_framework._workflows._edge_runner.EdgeRunner._execute_on_target") as mock_send:
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
|
||||
assert success is True
|
||||
assert mock_send.call_count == 1
|
||||
@@ -1312,13 +1310,13 @@ async def test_switch_case_edge_group_send_message_with_invalid_target() -> None
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data=-1)
|
||||
message = Message(data=data, source_id=source.id, target_id="invalid_target")
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
@@ -1339,18 +1337,18 @@ async def test_switch_case_edge_group_send_message_with_valid_target() -> None:
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = MockMessage(data=1) # Condition will fail
|
||||
message = Message(data=data, source_id=source.id, target_id=target1.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, 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_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is True
|
||||
|
||||
|
||||
@@ -1371,13 +1369,13 @@ async def test_switch_case_edge_group_send_message_with_invalid_data() -> None:
|
||||
executors: dict[str, Executor] = {source.id: source, target1.id: target1, target2.id: target2}
|
||||
edge_runner = create_edge_runner(edge_group, executors)
|
||||
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
data = "invalid_data"
|
||||
message = Message(data=data, source_id=source.id)
|
||||
|
||||
success = await edge_runner.send_message(message, shared_state, ctx)
|
||||
success = await edge_runner.send_message(message, state, ctx)
|
||||
assert success is False
|
||||
|
||||
|
||||
|
||||
@@ -898,7 +898,7 @@ async def test_magentic_checkpoint_restore_no_duplicate_history():
|
||||
|
||||
latest_checkpoint = checkpoints[-1]
|
||||
|
||||
# Load checkpoint and verify no duplicates in shared state
|
||||
# Load checkpoint and verify no duplicates in state
|
||||
checkpoint_data = await storage.load_checkpoint(latest_checkpoint.checkpoint_id)
|
||||
assert checkpoint_data is not None
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ from agent_framework import InMemoryCheckpointStorage, InProcRunnerContext
|
||||
from agent_framework._workflows._checkpoint_encoding import DATACLASS_MARKER, encode_checkpoint_value
|
||||
from agent_framework._workflows._checkpoint_summary import get_checkpoint_summary
|
||||
from agent_framework._workflows._events import RequestInfoEvent
|
||||
from agent_framework._workflows._shared_state import SharedState
|
||||
from agent_framework._workflows._state import State
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -46,7 +46,7 @@ async def test_rehydrate_request_info_event() -> None:
|
||||
runner_context = InProcRunnerContext(InMemoryCheckpointStorage())
|
||||
await runner_context.add_request_info_event(request_info_event)
|
||||
|
||||
checkpoint_id = await runner_context.create_checkpoint(SharedState(), iteration_count=1)
|
||||
checkpoint_id = await runner_context.create_checkpoint(State(), iteration_count=1)
|
||||
checkpoint = await runner_context.load_checkpoint(checkpoint_id)
|
||||
|
||||
assert checkpoint is not None
|
||||
@@ -79,7 +79,7 @@ async def test_rehydrate_fails_when_request_type_missing() -> None:
|
||||
runner_context = InProcRunnerContext(InMemoryCheckpointStorage())
|
||||
await runner_context.add_request_info_event(request_info_event)
|
||||
|
||||
checkpoint_id = await runner_context.create_checkpoint(SharedState(), iteration_count=1)
|
||||
checkpoint_id = await runner_context.create_checkpoint(State(), iteration_count=1)
|
||||
checkpoint = await runner_context.load_checkpoint(checkpoint_id)
|
||||
|
||||
assert checkpoint is not None
|
||||
@@ -107,7 +107,7 @@ async def test_rehydrate_fails_when_request_type_mismatch() -> None:
|
||||
runner_context = InProcRunnerContext(InMemoryCheckpointStorage())
|
||||
await runner_context.add_request_info_event(request_info_event)
|
||||
|
||||
checkpoint_id = await runner_context.create_checkpoint(SharedState(), iteration_count=1)
|
||||
checkpoint_id = await runner_context.create_checkpoint(State(), iteration_count=1)
|
||||
checkpoint = await runner_context.load_checkpoint(checkpoint_id)
|
||||
|
||||
assert checkpoint is not None
|
||||
@@ -137,7 +137,7 @@ async def test_pending_requests_in_summary() -> None:
|
||||
runner_context = InProcRunnerContext(InMemoryCheckpointStorage())
|
||||
await runner_context.add_request_info_event(request_info_event)
|
||||
|
||||
checkpoint_id = await runner_context.create_checkpoint(SharedState(), iteration_count=1)
|
||||
checkpoint_id = await runner_context.create_checkpoint(State(), iteration_count=1)
|
||||
checkpoint = await runner_context.load_checkpoint(checkpoint_id)
|
||||
|
||||
assert checkpoint is not None
|
||||
@@ -175,7 +175,7 @@ async def test_request_info_event_serializes_non_json_payloads() -> None:
|
||||
await runner_context.add_request_info_event(req_1)
|
||||
await runner_context.add_request_info_event(req_2)
|
||||
|
||||
checkpoint_id = await runner_context.create_checkpoint(SharedState(), iteration_count=1)
|
||||
checkpoint_id = await runner_context.create_checkpoint(State(), iteration_count=1)
|
||||
checkpoint = await runner_context.load_checkpoint(checkpoint_id)
|
||||
|
||||
# Should be JSON serializable despite datetime/slots
|
||||
|
||||
@@ -25,7 +25,7 @@ from agent_framework._workflows._runner_context import (
|
||||
Message,
|
||||
RunnerContext,
|
||||
)
|
||||
from agent_framework._workflows._shared_state import SharedState
|
||||
from agent_framework._workflows._state import State
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -48,7 +48,7 @@ class MockExecutor(Executor):
|
||||
|
||||
|
||||
def test_create_runner():
|
||||
"""Test creating a runner with edges and shared state."""
|
||||
"""Test creating a runner with edges and state."""
|
||||
executor_a = MockExecutor(id="executor_a")
|
||||
executor_b = MockExecutor(id="executor_b")
|
||||
|
||||
@@ -63,7 +63,7 @@ def test_create_runner():
|
||||
executor_b.id: executor_b,
|
||||
}
|
||||
|
||||
runner = Runner(edge_groups, executors, shared_state=SharedState(), ctx=InProcRunnerContext())
|
||||
runner = Runner(edge_groups, executors, state=State(), ctx=InProcRunnerContext())
|
||||
|
||||
assert runner.context is not None and isinstance(runner.context, RunnerContext)
|
||||
|
||||
@@ -83,16 +83,16 @@ async def test_runner_run_until_convergence():
|
||||
executor_a.id: executor_a,
|
||||
executor_b.id: executor_b,
|
||||
}
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
runner = Runner(edges, executors, shared_state, ctx)
|
||||
runner = Runner(edges, executors, state, ctx)
|
||||
|
||||
result: int | None = None
|
||||
await executor_a.execute(
|
||||
MockMessage(data=0),
|
||||
["START"], # source_executor_ids
|
||||
shared_state, # shared_state
|
||||
state, # state
|
||||
ctx, # runner_context
|
||||
)
|
||||
async for event in runner.run_until_convergence():
|
||||
@@ -121,15 +121,15 @@ async def test_runner_run_until_convergence_not_completed():
|
||||
executor_a.id: executor_a,
|
||||
executor_b.id: executor_b,
|
||||
}
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
runner = Runner(edges, executors, shared_state, ctx, max_iterations=5)
|
||||
runner = Runner(edges, executors, state, ctx, max_iterations=5)
|
||||
|
||||
await executor_a.execute(
|
||||
MockMessage(data=0),
|
||||
["START"], # source_executor_ids
|
||||
shared_state, # shared_state
|
||||
state, # state
|
||||
ctx, # runner_context
|
||||
)
|
||||
with pytest.raises(
|
||||
@@ -155,15 +155,15 @@ async def test_runner_already_running():
|
||||
executor_a.id: executor_a,
|
||||
executor_b.id: executor_b,
|
||||
}
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
runner = Runner(edges, executors, shared_state, ctx)
|
||||
runner = Runner(edges, executors, state, ctx)
|
||||
|
||||
await executor_a.execute(
|
||||
MockMessage(data=0),
|
||||
["START"], # source_executor_ids
|
||||
shared_state, # shared_state
|
||||
state, # state
|
||||
ctx, # runner_context
|
||||
)
|
||||
|
||||
@@ -178,7 +178,7 @@ async def test_runner_already_running():
|
||||
|
||||
async def test_runner_emits_runner_completion_for_agent_response_without_targets():
|
||||
ctx = InProcRunnerContext()
|
||||
runner = Runner([], {}, SharedState(), ctx)
|
||||
runner = Runner([], {}, State(), ctx)
|
||||
|
||||
await ctx.send_message(
|
||||
Message(
|
||||
@@ -227,7 +227,7 @@ async def test_runner_cancellation_stops_active_executor():
|
||||
executor_a.id: executor_a,
|
||||
executor_b.id: executor_b,
|
||||
}
|
||||
shared_state = SharedState()
|
||||
shared_state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
runner = Runner(edges, executors, shared_state, ctx)
|
||||
|
||||
@@ -623,7 +623,7 @@ class TestSerializationWorkflowClasses:
|
||||
|
||||
# These private runtime fields should not be in the serialized data
|
||||
assert "_runner_context" not in data
|
||||
assert "_shared_state" not in data
|
||||
assert "_state" not in data
|
||||
assert "_runner" not in data
|
||||
|
||||
def test_workflow_name_description_serialization(self) -> None:
|
||||
@@ -760,7 +760,7 @@ def test_comprehensive_edge_groups_workflow_serialization() -> None:
|
||||
|
||||
# Verify that serialization excludes non-serializable fields
|
||||
assert "_runner_context" not in data
|
||||
assert "_shared_state" not in data
|
||||
assert "_state" not in data
|
||||
assert "_runner" not in data
|
||||
|
||||
# Test that we can identify each edge group type by examining their structure
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Unit tests for the State class superstep caching behavior."""
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework._workflows._state import State
|
||||
|
||||
|
||||
class TestStateBasicOperations:
|
||||
"""Tests for basic State get/set/has/delete operations."""
|
||||
|
||||
def test_set_and_get(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
assert state.get("key") == "value"
|
||||
|
||||
def test_get_with_default(self) -> None:
|
||||
state = State()
|
||||
assert state.get("missing") is None
|
||||
assert state.get("missing", "default") == "default"
|
||||
|
||||
def test_has_returns_true_for_existing_key(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
assert state.has("key") is True
|
||||
|
||||
def test_has_returns_false_for_missing_key(self) -> None:
|
||||
state = State()
|
||||
assert state.has("missing") is False
|
||||
|
||||
def test_delete_existing_key(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
state.commit()
|
||||
state.delete("key")
|
||||
state.commit()
|
||||
assert state.has("key") is False
|
||||
assert state.get("key") is None
|
||||
|
||||
def test_delete_missing_key_raises(self) -> None:
|
||||
state = State()
|
||||
with pytest.raises(KeyError, match="Key 'missing' not found"):
|
||||
state.delete("missing")
|
||||
|
||||
def test_clear(self) -> None:
|
||||
state = State()
|
||||
state.set("key1", "value1")
|
||||
state.commit()
|
||||
state.set("key2", "value2")
|
||||
state.clear()
|
||||
assert state.get("key1") is None
|
||||
assert state.get("key2") is None
|
||||
|
||||
|
||||
class TestSuperstepCaching:
|
||||
"""Tests for superstep caching semantics - pending vs committed state."""
|
||||
|
||||
def test_set_writes_to_pending_not_committed(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
|
||||
# Value is in pending
|
||||
assert "key" in state._pending
|
||||
# Value is NOT in committed
|
||||
assert "key" not in state._committed
|
||||
# But get() still returns it
|
||||
assert state.get("key") == "value"
|
||||
|
||||
def test_commit_moves_pending_to_committed(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
|
||||
# Before commit: in pending, not committed
|
||||
assert "key" in state._pending
|
||||
assert "key" not in state._committed
|
||||
|
||||
state.commit()
|
||||
|
||||
# After commit: in committed, pending cleared
|
||||
assert "key" not in state._pending
|
||||
assert "key" in state._committed
|
||||
assert state.get("key") == "value"
|
||||
|
||||
def test_discard_clears_pending_without_committing(self) -> None:
|
||||
state = State()
|
||||
state.set("existing", "original")
|
||||
state.commit()
|
||||
|
||||
# Make a pending change
|
||||
state.set("existing", "modified")
|
||||
state.set("new_key", "new_value")
|
||||
|
||||
# Discard pending changes
|
||||
state.discard()
|
||||
|
||||
# Original value is preserved, new key never committed
|
||||
assert state.get("existing") == "original"
|
||||
assert state.get("new_key") is None
|
||||
|
||||
def test_pending_overrides_committed_on_get(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "committed_value")
|
||||
state.commit()
|
||||
|
||||
state.set("key", "pending_value")
|
||||
|
||||
# get() returns pending value, not committed
|
||||
assert state.get("key") == "pending_value"
|
||||
# But committed still has old value
|
||||
assert state._committed["key"] == "committed_value"
|
||||
|
||||
def test_multiple_sets_before_commit(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value1")
|
||||
state.set("key", "value2")
|
||||
state.set("key", "value3")
|
||||
|
||||
# Only final value is in pending
|
||||
assert state.get("key") == "value3"
|
||||
|
||||
state.commit()
|
||||
assert state.get("key") == "value3"
|
||||
|
||||
|
||||
class TestDeleteWithSuperstepCaching:
|
||||
"""Tests for delete behavior with superstep caching."""
|
||||
|
||||
def test_delete_pending_only_key(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
# Key only in pending, not committed
|
||||
assert "key" in state._pending
|
||||
assert "key" not in state._committed
|
||||
|
||||
state.delete("key")
|
||||
|
||||
# Should be removed from pending
|
||||
assert "key" not in state._pending
|
||||
assert state.get("key") is None
|
||||
assert state.has("key") is False
|
||||
|
||||
def test_delete_committed_key_marks_for_deletion(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
state.commit()
|
||||
|
||||
state.delete("key")
|
||||
|
||||
# Key should be marked for deletion in pending (sentinel)
|
||||
assert "key" in state._pending
|
||||
# get() should return default (not the sentinel!)
|
||||
assert state.get("key") is None
|
||||
assert state.get("key", "default") == "default"
|
||||
# has() should return False
|
||||
assert state.has("key") is False
|
||||
# But committed still has it until commit()
|
||||
assert "key" in state._committed
|
||||
|
||||
def test_delete_committed_key_removed_on_commit(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
state.commit()
|
||||
|
||||
state.delete("key")
|
||||
state.commit()
|
||||
|
||||
# Now it should be gone from committed too
|
||||
assert "key" not in state._committed
|
||||
assert "key" not in state._pending
|
||||
|
||||
def test_delete_key_in_both_pending_and_committed(self) -> None:
|
||||
"""Test delete when key exists in both pending (modified) and committed."""
|
||||
state = State()
|
||||
state.set("key", "original")
|
||||
state.commit()
|
||||
|
||||
# Modify the key (now in both pending and committed)
|
||||
state.set("key", "modified")
|
||||
assert state._pending["key"] == "modified"
|
||||
assert state._committed["key"] == "original"
|
||||
|
||||
# Delete should mark for deletion from committed
|
||||
state.delete("key")
|
||||
|
||||
# Should be marked for deletion
|
||||
assert state.get("key") is None
|
||||
assert state.has("key") is False
|
||||
|
||||
# After commit, key should be fully removed
|
||||
state.commit()
|
||||
assert "key" not in state._committed
|
||||
assert "key" not in state._pending
|
||||
|
||||
def test_discard_after_delete_restores_committed_value(self) -> None:
|
||||
state = State()
|
||||
state.set("key", "value")
|
||||
state.commit()
|
||||
|
||||
state.delete("key")
|
||||
# Key appears deleted
|
||||
assert state.has("key") is False
|
||||
|
||||
state.discard()
|
||||
# After discard, committed value is restored
|
||||
assert state.has("key") is True
|
||||
assert state.get("key") == "value"
|
||||
|
||||
|
||||
class TestFailureScenarios:
|
||||
"""Tests simulating failure scenarios - pending changes should not leak to committed."""
|
||||
|
||||
def test_failure_before_commit_preserves_committed_state(self) -> None:
|
||||
"""Simulate executor failure - pending changes should not affect committed state."""
|
||||
state = State()
|
||||
state.set("key1", "original1")
|
||||
state.set("key2", "original2")
|
||||
state.commit()
|
||||
|
||||
# Superstep starts - make some changes
|
||||
state.set("key1", "modified1")
|
||||
state.set("key3", "new_value")
|
||||
state.delete("key2")
|
||||
|
||||
# Simulate failure - we call discard() instead of commit()
|
||||
state.discard()
|
||||
|
||||
# All original values should be intact
|
||||
assert state.get("key1") == "original1"
|
||||
assert state.get("key2") == "original2"
|
||||
assert state.get("key3") is None
|
||||
|
||||
def test_no_partial_commits(self) -> None:
|
||||
"""Ensure commit is atomic - either all changes apply or none."""
|
||||
state = State()
|
||||
state.set("key1", "value1")
|
||||
state.set("key2", "value2")
|
||||
state.set("key3", "value3")
|
||||
|
||||
# Before commit - nothing in committed
|
||||
assert len(state._committed) == 0
|
||||
|
||||
state.commit()
|
||||
|
||||
# After commit - all three values committed together
|
||||
assert state._committed == {"key1": "value1", "key2": "value2", "key3": "value3"}
|
||||
|
||||
def test_repeated_supersteps_are_isolated(self) -> None:
|
||||
"""Test that each superstep's changes are isolated until committed."""
|
||||
state = State()
|
||||
|
||||
# Superstep 1
|
||||
state.set("counter", 1)
|
||||
state.commit()
|
||||
assert state.get("counter") == 1
|
||||
|
||||
# Superstep 2
|
||||
state.set("counter", 2)
|
||||
state.set("temp", "should_be_discarded")
|
||||
state.discard() # Simulate failure
|
||||
assert state.get("counter") == 1 # Reverted to superstep 1 value
|
||||
assert state.get("temp") is None
|
||||
|
||||
# Superstep 3
|
||||
state.set("counter", 3)
|
||||
state.commit()
|
||||
assert state.get("counter") == 3
|
||||
|
||||
|
||||
class TestExportImport:
|
||||
"""Tests for state serialization (export/import)."""
|
||||
|
||||
def test_export_returns_committed_only(self) -> None:
|
||||
state = State()
|
||||
state.set("committed_key", "committed_value")
|
||||
state.commit()
|
||||
state.set("pending_key", "pending_value")
|
||||
|
||||
exported = state.export_state()
|
||||
|
||||
# Only committed state is exported
|
||||
assert exported == {"committed_key": "committed_value"}
|
||||
assert "pending_key" not in exported
|
||||
|
||||
def test_import_merges_into_committed(self) -> None:
|
||||
state = State()
|
||||
state.set("existing", "original")
|
||||
state.commit()
|
||||
|
||||
state.import_state({"imported": "value", "existing": "overwritten"})
|
||||
|
||||
assert state.get("imported") == "value"
|
||||
assert state.get("existing") == "overwritten"
|
||||
|
||||
def test_import_does_not_affect_pending(self) -> None:
|
||||
state = State()
|
||||
state.set("pending_key", "pending_value")
|
||||
|
||||
state.import_state({"imported": "value"})
|
||||
|
||||
# Pending is still there
|
||||
assert state.get("pending_key") == "pending_value"
|
||||
assert "pending_key" in state._pending
|
||||
@@ -87,7 +87,7 @@ class MockExecutorRequestApproval(Executor):
|
||||
@handler
|
||||
async def mock_handler_a(self, message: NumberMessage, ctx: WorkflowContext) -> None:
|
||||
"""A mock handler that requests approval."""
|
||||
await ctx.set_shared_state(self.id, message.data)
|
||||
ctx.set_state(self.id, message.data)
|
||||
await ctx.request_info(MockRequest(prompt="Mock approval request"), ApprovalMessage)
|
||||
|
||||
@response_handler
|
||||
@@ -98,7 +98,7 @@ class MockExecutorRequestApproval(Executor):
|
||||
ctx: WorkflowContext[NumberMessage, int],
|
||||
) -> None:
|
||||
"""A mock handler that processes the approval response."""
|
||||
data = await ctx.get_shared_state(self.id)
|
||||
data = ctx.get_state(self.id)
|
||||
assert isinstance(data, int)
|
||||
if response.approved:
|
||||
await ctx.yield_output(data)
|
||||
@@ -368,7 +368,7 @@ async def test_workflow_run_stream_from_checkpoint_with_external_storage(
|
||||
test_checkpoint = WorkflowCheckpoint(
|
||||
workflow_id="test-workflow",
|
||||
messages={},
|
||||
shared_state={},
|
||||
state={},
|
||||
iteration_count=0,
|
||||
)
|
||||
checkpoint_id = await storage.save_checkpoint(test_checkpoint)
|
||||
@@ -403,7 +403,7 @@ async def test_workflow_run_from_checkpoint_non_streaming(simple_executor: Execu
|
||||
test_checkpoint = WorkflowCheckpoint(
|
||||
workflow_id="test-workflow",
|
||||
messages={},
|
||||
shared_state={},
|
||||
state={},
|
||||
iteration_count=0,
|
||||
)
|
||||
checkpoint_id = await storage.save_checkpoint(test_checkpoint)
|
||||
@@ -436,7 +436,7 @@ async def test_workflow_run_stream_from_checkpoint_with_responses(
|
||||
test_checkpoint = WorkflowCheckpoint(
|
||||
workflow_id="test-workflow",
|
||||
messages={},
|
||||
shared_state={},
|
||||
state={},
|
||||
pending_request_info_events={
|
||||
"request_123": RequestInfoEvent(
|
||||
request_id="request_123",
|
||||
@@ -480,7 +480,7 @@ class StateTrackingMessage:
|
||||
|
||||
|
||||
class StateTrackingExecutor(Executor):
|
||||
"""An executor that tracks state in shared state to test context reset behavior."""
|
||||
"""An executor that tracks state in workflow state to test context reset behavior."""
|
||||
|
||||
@handler
|
||||
async def handle_message(
|
||||
@@ -488,19 +488,16 @@ class StateTrackingExecutor(Executor):
|
||||
message: StateTrackingMessage,
|
||||
ctx: WorkflowContext[StateTrackingMessage, list[str]],
|
||||
) -> None:
|
||||
"""Handle the message and track it in shared state."""
|
||||
# Get existing messages from shared state
|
||||
try:
|
||||
existing_messages = await ctx.get_shared_state("processed_messages")
|
||||
except KeyError:
|
||||
existing_messages = []
|
||||
"""Handle the message and track it in workflow state."""
|
||||
# Get existing messages from workflow state
|
||||
existing_messages = ctx.get_state("processed_messages") or []
|
||||
|
||||
# Record this message
|
||||
message_record = f"{message.run_id}:{message.data}"
|
||||
existing_messages.append(message_record) # type: ignore
|
||||
|
||||
# Update shared state
|
||||
await ctx.set_shared_state("processed_messages", existing_messages)
|
||||
# Update workflow state
|
||||
ctx.set_state("processed_messages", existing_messages)
|
||||
|
||||
# Yield output
|
||||
await ctx.yield_output(existing_messages.copy()) # type: ignore
|
||||
@@ -511,7 +508,7 @@ async def test_workflow_multiple_runs_no_state_collision():
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
storage = FileCheckpointStorage(temp_dir)
|
||||
|
||||
# Create executor that tracks state in shared state
|
||||
# Create executor that tracks state in workflow state
|
||||
state_executor = StateTrackingExecutor(id="state_executor")
|
||||
|
||||
# Build workflow with checkpointing
|
||||
|
||||
@@ -41,15 +41,15 @@ async def make_context(
|
||||
executor_id: str = "exec",
|
||||
) -> AsyncIterator[tuple[WorkflowContext[object], "InProcRunnerContext"]]:
|
||||
from agent_framework._workflows._runner_context import InProcRunnerContext
|
||||
from agent_framework._workflows._shared_state import SharedState
|
||||
from agent_framework._workflows._state import State
|
||||
|
||||
mock_executor = MockExecutor(executor_id)
|
||||
runner_ctx = InProcRunnerContext()
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
workflow_ctx: WorkflowContext[object] = WorkflowContext(
|
||||
mock_executor,
|
||||
["source"],
|
||||
shared_state,
|
||||
state,
|
||||
runner_ctx,
|
||||
)
|
||||
try:
|
||||
|
||||
@@ -208,48 +208,48 @@ async def test_groupchat_kwargs_flow_to_agents() -> None:
|
||||
# endregion
|
||||
|
||||
|
||||
# region SharedState Verification Tests
|
||||
# region State Verification Tests
|
||||
|
||||
|
||||
async def test_kwargs_stored_in_shared_state() -> None:
|
||||
"""Test that kwargs are stored in SharedState with the correct key."""
|
||||
async def test_kwargs_stored_in_state() -> None:
|
||||
"""Test that kwargs are stored in State with the correct key."""
|
||||
from agent_framework import Executor, WorkflowContext, handler
|
||||
|
||||
stored_kwargs: dict[str, Any] | None = None
|
||||
|
||||
class _SharedStateInspector(Executor):
|
||||
class _StateInspector(Executor):
|
||||
@handler
|
||||
async def inspect(self, msgs: list[ChatMessage], ctx: WorkflowContext[list[ChatMessage]]) -> None:
|
||||
nonlocal stored_kwargs
|
||||
stored_kwargs = await ctx.get_shared_state(WORKFLOW_RUN_KWARGS_KEY)
|
||||
stored_kwargs = ctx.get_state(WORKFLOW_RUN_KWARGS_KEY)
|
||||
await ctx.send_message(msgs)
|
||||
|
||||
inspector = _SharedStateInspector(id="inspector")
|
||||
inspector = _StateInspector(id="inspector")
|
||||
workflow = SequentialBuilder().participants([inspector]).build()
|
||||
|
||||
async for event in workflow.run_stream("test", my_kwarg="my_value", another=123):
|
||||
if isinstance(event, WorkflowStatusEvent) and event.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
assert stored_kwargs is not None, "kwargs should be stored in SharedState"
|
||||
assert stored_kwargs is not None, "kwargs should be stored in State"
|
||||
assert stored_kwargs.get("my_kwarg") == "my_value"
|
||||
assert stored_kwargs.get("another") == 123
|
||||
|
||||
|
||||
async def test_empty_kwargs_stored_as_empty_dict() -> None:
|
||||
"""Test that empty kwargs are stored as empty dict in SharedState."""
|
||||
"""Test that empty kwargs are stored as empty dict in State."""
|
||||
from agent_framework import Executor, WorkflowContext, handler
|
||||
|
||||
stored_kwargs: Any = "NOT_CHECKED"
|
||||
|
||||
class _SharedStateChecker(Executor):
|
||||
class _StateChecker(Executor):
|
||||
@handler
|
||||
async def check(self, msgs: list[ChatMessage], ctx: WorkflowContext[list[ChatMessage]]) -> None:
|
||||
nonlocal stored_kwargs
|
||||
stored_kwargs = await ctx.get_shared_state(WORKFLOW_RUN_KWARGS_KEY)
|
||||
stored_kwargs = ctx.get_state(WORKFLOW_RUN_KWARGS_KEY)
|
||||
await ctx.send_message(msgs)
|
||||
|
||||
checker = _SharedStateChecker(id="checker")
|
||||
checker = _StateChecker(id="checker")
|
||||
workflow = SequentialBuilder().participants([checker]).build()
|
||||
|
||||
# Run without any kwargs
|
||||
@@ -257,7 +257,7 @@ async def test_empty_kwargs_stored_as_empty_dict() -> None:
|
||||
if isinstance(event, WorkflowStatusEvent) and event.state == WorkflowRunState.IDLE:
|
||||
break
|
||||
|
||||
# SharedState should have empty dict when no kwargs provided
|
||||
# State should have empty dict when no kwargs provided
|
||||
assert stored_kwargs == {}, f"Expected empty dict, got: {stored_kwargs}"
|
||||
|
||||
|
||||
@@ -420,8 +420,8 @@ async def test_magentic_kwargs_flow_to_agents() -> None:
|
||||
# A more comprehensive integration test would require the manager to select an agent.
|
||||
|
||||
|
||||
async def test_magentic_kwargs_stored_in_shared_state() -> None:
|
||||
"""Test that kwargs are stored in SharedState when using MagenticWorkflow.run_stream()."""
|
||||
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 (
|
||||
MagenticContext,
|
||||
@@ -639,10 +639,10 @@ async def test_subworkflow_kwargs_propagation() -> None:
|
||||
)
|
||||
|
||||
|
||||
async def test_subworkflow_kwargs_accessible_via_shared_state() -> None:
|
||||
"""Test that kwargs are accessible via SharedState within subworkflow.
|
||||
async def test_subworkflow_kwargs_accessible_via_state() -> None:
|
||||
"""Test that kwargs are accessible via State within subworkflow.
|
||||
|
||||
Verifies that WORKFLOW_RUN_KWARGS_KEY is populated in the subworkflow's SharedState
|
||||
Verifies that WORKFLOW_RUN_KWARGS_KEY is populated in the subworkflow's State
|
||||
with kwargs from the parent workflow.
|
||||
"""
|
||||
from agent_framework import Executor, WorkflowContext, handler
|
||||
@@ -650,17 +650,17 @@ async def test_subworkflow_kwargs_accessible_via_shared_state() -> None:
|
||||
|
||||
captured_kwargs_from_state: list[dict[str, Any]] = []
|
||||
|
||||
class _SharedStateReader(Executor):
|
||||
"""Executor that reads kwargs from SharedState for verification."""
|
||||
class _StateReader(Executor):
|
||||
"""Executor that reads kwargs from State for verification."""
|
||||
|
||||
@handler
|
||||
async def read_kwargs(self, msgs: list[ChatMessage], ctx: WorkflowContext[list[ChatMessage]]) -> None:
|
||||
kwargs_from_state = await ctx.get_shared_state(WORKFLOW_RUN_KWARGS_KEY)
|
||||
kwargs_from_state = ctx.get_state(WORKFLOW_RUN_KWARGS_KEY)
|
||||
captured_kwargs_from_state.append(kwargs_from_state or {})
|
||||
await ctx.send_message(msgs)
|
||||
|
||||
# Build inner workflow with SharedState reader
|
||||
state_reader = _SharedStateReader(id="state_reader")
|
||||
# Build inner workflow with State reader
|
||||
state_reader = _StateReader(id="state_reader")
|
||||
inner_workflow = SequentialBuilder().participants([state_reader]).build()
|
||||
|
||||
# Wrap as subworkflow
|
||||
@@ -679,15 +679,15 @@ async def test_subworkflow_kwargs_accessible_via_shared_state() -> None:
|
||||
break
|
||||
|
||||
# Verify the state reader was invoked
|
||||
assert len(captured_kwargs_from_state) >= 1, "SharedState reader should have been invoked"
|
||||
assert len(captured_kwargs_from_state) >= 1, "State reader should have been invoked"
|
||||
|
||||
kwargs_in_subworkflow = captured_kwargs_from_state[0]
|
||||
|
||||
assert kwargs_in_subworkflow.get("my_custom_kwarg") == "should_be_propagated", (
|
||||
f"Expected 'my_custom_kwarg' in subworkflow SharedState, got: {kwargs_in_subworkflow}"
|
||||
f"Expected 'my_custom_kwarg' in subworkflow got: {kwargs_in_subworkflow}"
|
||||
)
|
||||
assert kwargs_in_subworkflow.get("another_kwarg") == 42, (
|
||||
f"Expected 'another_kwarg'=42 in subworkflow SharedState, got: {kwargs_in_subworkflow}"
|
||||
f"Expected 'another_kwarg'=42 in subworkflow got: {kwargs_in_subworkflow}"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanE
|
||||
from agent_framework import InMemoryCheckpointStorage, WorkflowBuilder
|
||||
from agent_framework._workflows._executor import Executor, handler
|
||||
from agent_framework._workflows._runner_context import InProcRunnerContext, Message, MessageType
|
||||
from agent_framework._workflows._shared_state import SharedState
|
||||
from agent_framework._workflows._state import State
|
||||
from agent_framework._workflows._workflow import Workflow
|
||||
from agent_framework._workflows._workflow_context import WorkflowContext
|
||||
from agent_framework.observability import (
|
||||
@@ -170,7 +170,7 @@ async def test_span_creation_and_attributes(span_exporter: InMemorySpanExporter)
|
||||
|
||||
async def test_trace_context_handling(span_exporter: InMemorySpanExporter) -> None:
|
||||
"""Test trace context propagation and handling in messages and executors."""
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
executor = MockExecutor("test-executor")
|
||||
|
||||
@@ -180,7 +180,7 @@ async def test_trace_context_handling(span_exporter: InMemorySpanExporter) -> No
|
||||
workflow_ctx: WorkflowContext[str] = WorkflowContext(
|
||||
executor,
|
||||
["source"],
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
trace_contexts=[{"traceparent": "00-12345678901234567890123456789012-1234567890123456-01"}],
|
||||
source_span_ids=["1234567890123456"],
|
||||
@@ -202,7 +202,7 @@ async def test_trace_context_handling(span_exporter: InMemorySpanExporter) -> No
|
||||
await executor.execute(
|
||||
"test message",
|
||||
["source"], # source_executor_ids
|
||||
shared_state, # shared_state
|
||||
state, # state
|
||||
ctx, # runner_context
|
||||
trace_contexts=[{"traceparent": "00-12345678901234567890123456789012-1234567890123456-01"}],
|
||||
source_span_ids=["1234567890123456"],
|
||||
@@ -236,13 +236,13 @@ async def test_trace_context_disabled_when_tracing_disabled(
|
||||
"""Test that no trace context is added when tracing is disabled."""
|
||||
# Tracing should be disabled by default
|
||||
executor = MockExecutor("test-executor")
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
ctx = InProcRunnerContext()
|
||||
|
||||
workflow_ctx: WorkflowContext[str] = WorkflowContext(
|
||||
executor,
|
||||
["source"],
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
|
||||
@@ -452,7 +452,7 @@ async def test_message_trace_context_serialization(span_exporter: InMemorySpanEx
|
||||
await ctx.send_message(message)
|
||||
|
||||
# Create a checkpoint that includes the message
|
||||
checkpoint_id = await ctx.create_checkpoint(SharedState(), 0)
|
||||
checkpoint_id = await ctx.create_checkpoint(State(), 0)
|
||||
checkpoint = await ctx.load_checkpoint(checkpoint_id)
|
||||
assert checkpoint is not None
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ from agent_framework import (
|
||||
WorkflowStatusEvent,
|
||||
handler,
|
||||
)
|
||||
from agent_framework._workflows._shared_state import SharedState
|
||||
from agent_framework._workflows._state import State
|
||||
|
||||
|
||||
class FailingExecutor(Executor):
|
||||
@@ -62,12 +62,12 @@ async def test_executor_failed_and_workflow_failed_events_streaming():
|
||||
async def test_executor_failed_event_emitted_on_direct_execute():
|
||||
failing = FailingExecutor(id="f")
|
||||
ctx = InProcRunnerContext()
|
||||
shared_state = SharedState()
|
||||
state = State()
|
||||
with pytest.raises(RuntimeError, match="boom"):
|
||||
await failing.execute(
|
||||
0,
|
||||
["START"],
|
||||
shared_state,
|
||||
state,
|
||||
ctx,
|
||||
)
|
||||
drained = await ctx.drain_events()
|
||||
|
||||
Reference in New Issue
Block a user