[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:
Evan Mattson
2026-02-05 10:42:52 +09:00
committed by GitHub
Unverified
parent 4e25917644
commit 10afb86213
48 changed files with 1971 additions and 1724 deletions
@@ -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()