fix(multiagent): preserve shared context and cumulative accounting across serialize/deserialize (#3396)

This commit is contained in:
opieter-aws
2026-07-23 14:25:14 -04:00
committed by GitHub
parent 767b8019b0
commit ba37a272c9
5 changed files with 324 additions and 9 deletions
+42
View File
@@ -4,6 +4,7 @@ Provides minimal foundation for multi-agent patterns (Swarm, Graph).
"""
import logging
import time
import warnings
from abc import ABC, abstractmethod
from collections.abc import AsyncIterator, Mapping
@@ -187,6 +188,47 @@ class MultiAgentBase(ABC):
"""
id: str
# Wall-clock start of the active invocation, or None when no invocation is running. Set at
# invocation start; folded into the orchestrator's committed execution-time total exactly once
# at finalization.
_invocation_start_time: float | None
def __init__(self) -> None:
"""Initialize base multi-agent state."""
self._invocation_start_time = None
def _execution_time_with_active_interval(self, committed_time: int) -> int:
"""Committed execution time (ms) plus the active invocation's in-flight interval.
The active interval is folded into the committed total only once, at finalization, so any
read that must reflect elapsed time — checkpoints, result building — adds it on here.
Args:
committed_time: Execution time in milliseconds already committed by prior invocations.
Returns:
committed_time plus the current invocation's elapsed milliseconds, or committed_time
unchanged when no invocation is running.
"""
if self._invocation_start_time is None:
return committed_time
return committed_time + round((time.time() - self._invocation_start_time) * 1000)
def _commit_active_interval(self, committed_time: int) -> int:
"""Fold the active invocation's interval into committed_time and end the interval.
Idempotent: with no active interval this returns committed_time unchanged, so calling it at
finalization never double-counts.
Args:
committed_time: Execution time in milliseconds already committed by prior invocations.
Returns:
The new committed total including the interval that just ended.
"""
total = self._execution_time_with_active_interval(committed_time)
self._invocation_start_time = None
return total
@abstractmethod
async def invoke_async(
+16 -3
View File
@@ -56,7 +56,7 @@ from ..types.event_loop import Metrics, Usage
from ..types.multiagent import MultiAgentInput
from ..types.session import decode_bytes_values, encode_bytes_values
from ..types.traces import AttributeValue
from .base import MultiAgentBase, MultiAgentResult, NodeResult, Status
from .base import MultiAgentBase, MultiAgentResult, NodeResult, Status, _parse_metrics, _parse_usage
logger = logging.getLogger(__name__)
@@ -664,6 +664,8 @@ class Graph(MultiAgentBase):
with trace_api.use_span(span, end_on_exit=True):
interrupts = []
self._invocation_start_time = start_time
try:
logger.debug(
"max_node_executions=<%s>, execution_timeout=<%s>s, node_timeout=<%s>s | graph execution config",
@@ -697,7 +699,7 @@ class Graph(MultiAgentBase):
self.state.status = Status.FAILED
raise
finally:
self.state.execution_time += round((time.time() - start_time) * 1000)
self.state.execution_time = self._commit_active_interval(self.state.execution_time)
await self.hooks.invoke_callbacks_async(AfterMultiAgentInvocationEvent(self))
self._resume_from_session = False
self._resume_next_nodes.clear()
@@ -1249,7 +1251,7 @@ class Graph(MultiAgentBase):
accumulated_usage=self.state.accumulated_usage,
accumulated_metrics=self.state.accumulated_metrics,
execution_count=self.state.execution_count,
execution_time=self.state.execution_time,
execution_time=self._execution_time_with_active_interval(self.state.execution_time),
total_nodes=self.state.total_nodes,
completed_nodes=len(self.state.completed_nodes),
failed_nodes=len(self.state.failed_nodes),
@@ -1275,6 +1277,10 @@ class Graph(MultiAgentBase):
"next_nodes_to_execute": next_nodes,
"current_task": encode_bytes_values(self.state.task),
"execution_order": [n.node_id for n in self.state.execution_order],
"accumulated_usage": self.state.accumulated_usage,
"accumulated_metrics": self.state.accumulated_metrics,
"execution_count": self.state.execution_count,
"execution_time": self._execution_time_with_active_interval(self.state.execution_time),
"_internal_state": {
"interrupt_state": self._interrupt_state.to_dict(),
},
@@ -1483,6 +1489,13 @@ class Graph(MultiAgentBase):
# Task
self.state.task = decode_bytes_values(payload.get("current_task", self.state.task))
# Cumulative accounting: restore so the timeout budget (should_continue) and the reported
# totals stay correct across a resume, rather than resetting to zero.
self.state.accumulated_usage = _parse_usage(payload.get("accumulated_usage") or {})
self.state.accumulated_metrics = _parse_metrics(payload.get("accumulated_metrics") or {})
self.state.execution_count = int(payload.get("execution_count") or 0)
self.state.execution_time = int(payload.get("execution_time") or 0)
# next nodes to execute
next_nodes = [self.nodes[nid] for nid in (payload.get("next_nodes_to_execute") or []) if nid in self.nodes]
self._resume_next_nodes = next_nodes
+19 -6
View File
@@ -56,7 +56,7 @@ from ..types.event_loop import Metrics, Usage
from ..types.multiagent import MultiAgentInput
from ..types.session import decode_bytes_values, encode_bytes_values
from ..types.traces import AttributeValue
from .base import MultiAgentBase, MultiAgentResult, NodeResult, Status
from .base import MultiAgentBase, MultiAgentResult, NodeResult, Status, _parse_metrics, _parse_usage
logger = logging.getLogger(__name__)
@@ -414,6 +414,8 @@ class Swarm(MultiAgentBase):
with trace_api.use_span(span, end_on_exit=True):
interrupts = []
self._invocation_start_time = self.state.start_time
try:
current_node = cast(SwarmNode, self.state.current_node)
logger.debug("current_node=<%s> | starting swarm execution with node", current_node.node_id)
@@ -435,7 +437,7 @@ class Swarm(MultiAgentBase):
self.state.completion_status = Status.FAILED
raise
finally:
self.state.execution_time += round((time.time() - self.state.start_time) * 1000)
self.state.execution_time = self._commit_active_interval(self.state.execution_time)
await self.hooks.invoke_callbacks_async(AfterMultiAgentInvocationEvent(self, invocation_state))
self._resume_from_session = False
@@ -993,6 +995,9 @@ class Swarm(MultiAgentBase):
"node_results": {k: v.to_dict() for k, v in self.state.results.items()},
"next_nodes_to_execute": next_nodes,
"current_task": encode_bytes_values(self.state.task),
"accumulated_usage": self.state.accumulated_usage,
"accumulated_metrics": self.state.accumulated_metrics,
"execution_time": self._execution_time_with_active_interval(self.state.execution_time),
"context": {
"shared_context": getattr(self.state.shared_context, "context", {}) or {},
"handoff_node": self.state.handoff_node.node_id if self.state.handoff_node else None,
@@ -1007,14 +1012,14 @@ class Swarm(MultiAgentBase):
"""Restore swarm state from a session dict and prepare for execution.
This method handles two scenarios:
1. If the persisted status is COMPLETED, FAILED resets all nodes and graph state
to allow re-execution from the beginning.
1. If the payload omits next_nodes_to_execute (a terminal or fresh state), resets all
nodes and swarm state to allow re-execution from the beginning.
2. Otherwise, restores the persisted state and prepares to resume execution
from the next ready nodes.
from the next node.
Args:
payload: Dictionary containing persisted state data including status,
completed nodes, results, and next nodes to execute.
node history, results, and next nodes to execute.
"""
if "_internal_state" in payload:
internal_state = payload["_internal_state"]
@@ -1036,12 +1041,20 @@ class Swarm(MultiAgentBase):
def _from_dict(self, payload: dict[str, Any]) -> None:
self.state.completion_status = Status(payload["status"])
# Point the state's shared context at the swarm-owned object, matching the identity a fresh run
# establishes in stream_async, so the serialize path (reads self.state.shared_context) and the
# node-input builder (reads self.shared_context) stay in sync after resume.
self.state.shared_context = self.shared_context
# Hydrate completed nodes & results
context = payload["context"] or {}
self.shared_context.context = context.get("shared_context") or {}
self.state.handoff_message = context.get("handoff_message")
self.state.handoff_node = self.nodes[context["handoff_node"]] if context.get("handoff_node") else None
self.state.accumulated_usage = _parse_usage(payload.get("accumulated_usage") or {})
self.state.accumulated_metrics = _parse_metrics(payload.get("accumulated_metrics") or {})
self.state.execution_time = int(payload.get("execution_time") or 0)
self.state.node_history = [self.nodes[nid] for nid in (payload.get("node_history") or []) if nid in self.nodes]
raw_results = payload.get("node_results") or {}
@@ -2046,6 +2046,157 @@ async def test_graph_persisted(mock_strands_tracer, mock_use_span):
assert "test_node" in final_state["node_results"]
def test_graph_serialize_deserialize_serialize_preserves_cumulative_state():
"""serialize -> deserialize -> serialize is value-preserving on the resume path.
Guarantees that a resumed graph re-serializes the same cumulative accounting (accumulated_usage /
accumulated_metrics / execution_count / execution_time) it was restored with, so the timeout
budget (should_continue) and the totals reported in GraphResult reflect the whole run.
"""
builder = GraphBuilder()
builder.add_node(create_mock_agent("test_agent"), "test_node")
builder.set_entry_point("test_node")
graph = builder.build()
payload = {
"type": "graph",
"id": "default_graph",
"status": "executing",
"completed_nodes": [],
"failed_nodes": [],
"interrupted_nodes": [],
"node_results": {},
"next_nodes_to_execute": ["test_node"],
"current_task": "resume me",
"execution_order": [],
"accumulated_usage": {"inputTokens": 11, "outputTokens": 22, "totalTokens": 33},
"accumulated_metrics": {"latencyMs": 44},
"execution_count": 3,
"execution_time": 555,
"_internal_state": {"interrupt_state": {"activated": False, "context": {}, "interrupts": {}}},
}
graph.deserialize_state(payload)
# Cumulative accounting is restored, not reset to zero.
assert graph.state.accumulated_usage == {"inputTokens": 11, "outputTokens": 22, "totalTokens": 33}
assert graph.state.accumulated_metrics == {"latencyMs": 44}
assert graph.state.execution_count == 3
assert graph.state.execution_time == 555
serialize1 = graph.serialize_state()
graph.deserialize_state(serialize1)
serialize2 = graph.serialize_state()
assert serialize2["accumulated_usage"] == serialize1["accumulated_usage"]
assert serialize2["accumulated_metrics"] == serialize1["accumulated_metrics"]
assert serialize2["execution_count"] == serialize1["execution_count"]
assert serialize2["execution_time"] == serialize1["execution_time"]
@pytest.mark.asyncio
async def test_graph_execution_time_reflects_active_invocation(mock_strands_tracer, mock_use_span):
"""The final GraphResult includes the current invocation's interval on top of restored prior time.
execution_time is committed to state once, at finalization. GraphResult is built before that
commit, so it must fold in the in-flight interval itself — and finalization must not double-count.
"""
# Monotonic fake clock advanced explicitly; robust to how many times time.time() is called.
clock = {"now": 1000.0}
with patch("strands.multiagent.graph.time.time", lambda: clock["now"]):
builder = GraphBuilder()
builder.add_node(create_mock_agent("test_agent"), "test_node")
builder.set_entry_point("test_node")
graph = builder.build()
# Resume from a checkpoint that already accrued 555ms in a prior invocation.
graph.deserialize_state(
{
"type": "graph",
"id": "default_graph",
"status": "executing",
"completed_nodes": [],
"failed_nodes": [],
"interrupted_nodes": [],
"node_results": {},
"next_nodes_to_execute": ["test_node"],
"current_task": "resume me",
"execution_order": [],
"accumulated_usage": {"inputTokens": 0, "outputTokens": 0, "totalTokens": 0},
"accumulated_metrics": {"latencyMs": 0},
"execution_count": 0,
"execution_time": 555,
"_internal_state": {"interrupt_state": {"activated": False, "context": {}, "interrupts": {}}},
}
)
# The node advances the clock by 1200ms while it "runs".
async def advancing_stream(*args, **kwargs):
clock["now"] += 1.2
yield {"result": graph.nodes["test_node"].executor.return_value}
graph.nodes["test_node"].executor.stream_async = Mock(side_effect=advancing_stream)
result = await graph.invoke_async("resume me")
# 555 restored + 1200 in-flight = 1755ms; committed exactly once.
tru_result_time = result.execution_time
exp_result_time = 1755
assert tru_result_time == exp_result_time
assert graph.state.execution_time == exp_result_time
@pytest.mark.asyncio
async def test_graph_checkpoint_persists_in_flight_execution_time(mock_strands_tracer, mock_use_span):
"""A mid-run per-node checkpoint persists elapsed time so a resumed run keeps its timeout budget.
Guards the crash-restart path: the AfterNodeCall session sync serializes before the invocation's
finally commits the interval, so serialize_state must fold the in-flight interval into
execution_time rather than persisting the stale pre-invocation value (which would reset the budget).
"""
clock = {"now": 2000.4}
builder = GraphBuilder()
builder.add_node(create_mock_agent("test_agent"), "test_node")
builder.set_entry_point("test_node")
graph = builder.build()
with patch("strands.multiagent.graph.time.time", lambda: clock["now"]):
# Marker set at invocation start; a checkpoint taken 400ms in must reflect that interval.
graph._invocation_start_time = 2000.0
graph.state.status = Status.EXECUTING
checkpoint = graph.serialize_state()
assert checkpoint["execution_time"] == 400
@pytest.mark.asyncio
async def test_graph_tracing_setup_failure_does_not_leak_timer(mock_strands_tracer, mock_use_span):
"""A tracing setup failure must not leave the invocation timer running.
The timer starts inside the span context so its clearing finally is guaranteed to run. If span
setup raises before then, no interval is started, and a later serialize_state must not accrue
wall time against an abandoned invocation.
"""
clock = {"now": 1000.0}
mock_strands_tracer.start_multiagent_span.side_effect = RuntimeError("span setup failed")
builder = GraphBuilder()
builder.add_node(create_mock_agent("test_agent"), "test_node")
builder.set_entry_point("test_node")
graph = builder.build()
with patch("strands.multiagent.graph.time.time", lambda: clock["now"]):
with pytest.raises(RuntimeError, match="span setup failed"):
await graph.invoke_async("go")
clock["now"] = 1000.5 # 500ms later
checkpoint = graph.serialize_state()
assert graph._invocation_start_time is None
assert checkpoint["execution_time"] == 0
@pytest.mark.parametrize(
("cancel_node", "cancel_message"),
[(True, "node cancelled by user"), ("custom cancel message", "custom cancel message")],
@@ -1212,6 +1212,102 @@ async def test_swarm_persistence(mock_strands_tracer, mock_use_span):
assert "test_agent" in final_state["node_results"]
def test_swarm_serialize_deserialize_serialize_preserves_state():
"""serialize -> deserialize -> serialize is value-preserving on the resume path.
Guarantees that a resumed swarm re-serializes the same shared_context and cumulative
accounting (execution_time / accumulated_usage / accumulated_metrics) it was restored with,
and that the restored SwarmState and Swarm share the same shared_context object.
"""
agent = create_mock_agent("first")
swarm = Swarm([agent])
payload = {
"type": "swarm",
"id": "default_swarm",
"status": "executing",
"node_history": [],
"node_results": {},
"next_nodes_to_execute": ["first"],
"current_task": "resume me",
"accumulated_usage": {"inputTokens": 11, "outputTokens": 22, "totalTokens": 33},
"accumulated_metrics": {"latencyMs": 44},
"execution_time": 555,
"context": {
"shared_context": {"first": {"fact": "persist-me"}},
"handoff_node": None,
"handoff_message": None,
},
"_internal_state": {"interrupt_state": {"activated": False, "context": {}, "interrupts": {}}},
}
swarm.deserialize_state(payload)
# The swarm-owned and state-owned shared contexts must be the same restored object.
assert swarm.shared_context.context == {"first": {"fact": "persist-me"}}
assert swarm.state.shared_context.context == {"first": {"fact": "persist-me"}}
assert swarm.state.shared_context is swarm.shared_context
# Cumulative accounting is restored, not reset to zero.
assert swarm.state.accumulated_usage == {"inputTokens": 11, "outputTokens": 22, "totalTokens": 33}
assert swarm.state.accumulated_metrics == {"latencyMs": 44}
assert swarm.state.execution_time == 555
serialize1 = swarm.serialize_state()
swarm.deserialize_state(serialize1)
serialize2 = swarm.serialize_state()
assert serialize2["context"]["shared_context"] == serialize1["context"]["shared_context"]
assert serialize2["context"]["shared_context"] == {"first": {"fact": "persist-me"}}
assert serialize2["accumulated_usage"] == serialize1["accumulated_usage"]
assert serialize2["accumulated_metrics"] == serialize1["accumulated_metrics"]
assert serialize2["execution_time"] == serialize1["execution_time"]
def test_swarm_checkpoint_persists_in_flight_execution_time():
"""A mid-run per-node checkpoint persists elapsed time so a resumed swarm keeps its timeout budget.
Guards the crash-restart path: the AfterNodeCall session sync serializes before the invocation's
finally commits the interval, so serialize_state must fold the in-flight interval into
execution_time rather than persisting the stale pre-invocation value (which would reset the budget).
"""
clock = {"now": 3000.6}
swarm = Swarm([create_mock_agent("first")])
with patch("strands.multiagent.swarm.time.time", lambda: clock["now"]):
# Marker set at invocation start; a checkpoint taken 600ms in must reflect that interval.
swarm._invocation_start_time = 3000.0
swarm.state.completion_status = Status.EXECUTING
swarm.state.current_node = swarm.nodes["first"]
checkpoint = swarm.serialize_state()
assert checkpoint["execution_time"] == 600
@pytest.mark.asyncio
async def test_swarm_tracing_setup_failure_does_not_leak_timer(mock_strands_tracer, mock_use_span):
"""A tracing setup failure must not leave the invocation timer running.
The timer starts inside the span context so its clearing finally is guaranteed to run. If span
setup raises before then, no interval is started, and a later serialize_state must not accrue
wall time against an abandoned invocation.
"""
clock = {"now": 2000.0}
mock_strands_tracer.start_multiagent_span.side_effect = RuntimeError("span setup failed")
swarm = Swarm([create_mock_agent("first")])
with patch("strands.multiagent.swarm.time.time", lambda: clock["now"]):
with pytest.raises(RuntimeError, match="span setup failed"):
await swarm.invoke_async("go")
clock["now"] = 2000.5 # 500ms later
checkpoint = swarm.serialize_state()
assert swarm._invocation_start_time is None
assert checkpoint["execution_time"] == 0
@pytest.mark.asyncio
async def test_swarm_handle_handoff():
first_agent = create_mock_agent("first")