mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
fix(multiagent): preserve shared context and cumulative accounting across serialize/deserialize (#3396)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user