fix(multiagent): accumulate cache token counters in Graph and Swarm (#3884)

This commit is contained in:
opieter-aws
2026-08-20 07:28:26 -04:00
committed by GitHub
parent f79ce81011
commit 606e1b84dc
5 changed files with 177 additions and 3 deletions
@@ -385,3 +385,11 @@ def _parse_usage(usage_data: dict[str, Any]) -> Usage:
def _parse_metrics(metrics_data: dict[str, Any]) -> Metrics:
"""Parse Metrics from dict data."""
return Metrics(latencyMs=metrics_data.get("latencyMs", 0))
def _accumulate_cache_usage(target: Usage, source: Usage) -> None:
"""Add source's optional cache token counters into target."""
if "cacheReadInputTokens" in source:
target["cacheReadInputTokens"] = target.get("cacheReadInputTokens", 0) + source["cacheReadInputTokens"]
if "cacheWriteInputTokens" in source:
target["cacheWriteInputTokens"] = target.get("cacheWriteInputTokens", 0) + source["cacheWriteInputTokens"]
+10 -1
View File
@@ -56,7 +56,15 @@ 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, _parse_metrics, _parse_usage
from .base import (
MultiAgentBase,
MultiAgentResult,
NodeResult,
Status,
_accumulate_cache_usage,
_parse_metrics,
_parse_usage,
)
logger = logging.getLogger(__name__)
@@ -1147,6 +1155,7 @@ class Graph(MultiAgentBase):
self.state.accumulated_usage["inputTokens"] += node_result.accumulated_usage.get("inputTokens", 0)
self.state.accumulated_usage["outputTokens"] += node_result.accumulated_usage.get("outputTokens", 0)
self.state.accumulated_usage["totalTokens"] += node_result.accumulated_usage.get("totalTokens", 0)
_accumulate_cache_usage(self.state.accumulated_usage, node_result.accumulated_usage)
self.state.accumulated_metrics["latencyMs"] += node_result.accumulated_metrics.get("latencyMs", 0)
self.state.execution_count += node_result.execution_count
+10 -1
View File
@@ -56,7 +56,15 @@ 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, _parse_metrics, _parse_usage
from .base import (
MultiAgentBase,
MultiAgentResult,
NodeResult,
Status,
_accumulate_cache_usage,
_parse_metrics,
_parse_usage,
)
logger = logging.getLogger(__name__)
@@ -1025,6 +1033,7 @@ class Swarm(MultiAgentBase):
self.state.accumulated_usage["inputTokens"] += node_result.accumulated_usage.get("inputTokens", 0)
self.state.accumulated_usage["outputTokens"] += node_result.accumulated_usage.get("outputTokens", 0)
self.state.accumulated_usage["totalTokens"] += node_result.accumulated_usage.get("totalTokens", 0)
_accumulate_cache_usage(self.state.accumulated_usage, node_result.accumulated_usage)
self.state.accumulated_metrics["latencyMs"] += node_result.accumulated_metrics.get("latencyMs", 0)
def _build_result(self, interrupts: list[Interrupt]) -> SwarmResult:
@@ -814,6 +814,83 @@ async def test_node_reset_executor_state():
assert multi_agent_node.result is None
def test_accumulate_metrics_sums_cache_token_counters():
"""Cache token counters are summed across nodes, not dropped (#797).
Each node's usage is a disjoint Bedrock-style payload where
totalTokens = inputTokens + outputTokens + cacheReadInputTokens + cacheWriteInputTokens,
so the aggregated totals must reconcile once the cache counters are accumulated too.
"""
graph = _make_graph(nodes={})
result_text = AgentResult(
message={"role": "assistant", "content": [{"text": "ok"}]}, stop_reason="end_turn", state={}, metrics={}
)
node_results = [
NodeResult(
result=result_text,
accumulated_usage={
"inputTokens": 10,
"outputTokens": 20,
"totalTokens": 38,
"cacheReadInputTokens": 5,
"cacheWriteInputTokens": 3,
},
execution_count=1,
),
NodeResult(
result=result_text,
accumulated_usage={
"inputTokens": 15,
"outputTokens": 25,
"totalTokens": 51,
"cacheReadInputTokens": 7,
"cacheWriteInputTokens": 4,
},
execution_count=1,
),
]
for node_result in node_results:
graph._accumulate_metrics(node_result)
tru_usage = graph.state.accumulated_usage
exp_usage = {
"inputTokens": 25,
"outputTokens": 45,
"totalTokens": 89,
"cacheReadInputTokens": 12,
"cacheWriteInputTokens": 7,
}
assert tru_usage == exp_usage
assert (
tru_usage["inputTokens"]
+ tru_usage["outputTokens"]
+ tru_usage["cacheReadInputTokens"]
+ tru_usage["cacheWriteInputTokens"]
== tru_usage["totalTokens"]
)
def test_accumulate_metrics_without_cache_counters_omits_them():
"""Nodes without cache counters must not materialize cache keys as 0 (#797)."""
graph = _make_graph(nodes={})
result_text = AgentResult(
message={"role": "assistant", "content": [{"text": "ok"}]}, stop_reason="end_turn", state={}, metrics={}
)
graph._accumulate_metrics(
NodeResult(
result=result_text,
accumulated_usage={"inputTokens": 10, "outputTokens": 20, "totalTokens": 30},
execution_count=1,
)
)
tru_usage = graph.state.accumulated_usage
exp_usage = {"inputTokens": 10, "outputTokens": 20, "totalTokens": 30}
assert tru_usage == exp_usage
def test_graph_dataclasses_and_enums():
"""Test dataclass initialization, properties, and enum behavior."""
# Test Status enum
@@ -11,7 +11,7 @@ from strands.agent.state import AgentState
from strands.hooks import AfterMultiAgentInvocationEvent, AfterNodeCallEvent, BeforeNodeCallEvent
from strands.hooks.registry import HookRegistry
from strands.interrupt import Interrupt, _InterruptState
from strands.multiagent.base import Status
from strands.multiagent.base import NodeResult, Status
from strands.multiagent.swarm import SharedContext, Swarm, SwarmNode, SwarmResult, SwarmState, _InflightTurn
from strands.session.file_session_manager import FileSessionManager
from strands.session.session_manager import SessionManager
@@ -452,6 +452,77 @@ def test_swarm_metrics_handling():
assert result.status == Status.COMPLETED
def test_accumulate_metrics_sums_cache_token_counters():
"""Cache token counters are summed across nodes, not dropped (#797).
Each node's usage is a disjoint Bedrock-style payload where
totalTokens = inputTokens + outputTokens + cacheReadInputTokens + cacheWriteInputTokens,
so the aggregated totals must reconcile once the cache counters are accumulated too.
"""
swarm = Swarm(nodes=[create_mock_agent("agent1"), create_mock_agent("agent2")])
result_text = AgentResult(
message={"role": "assistant", "content": [{"text": "ok"}]}, stop_reason="end_turn", state={}, metrics={}
)
node_results = [
NodeResult(
result=result_text,
accumulated_usage={
"inputTokens": 10,
"outputTokens": 20,
"totalTokens": 38,
"cacheReadInputTokens": 5,
"cacheWriteInputTokens": 3,
},
),
NodeResult(
result=result_text,
accumulated_usage={
"inputTokens": 15,
"outputTokens": 25,
"totalTokens": 51,
"cacheReadInputTokens": 7,
"cacheWriteInputTokens": 4,
},
),
]
for node_result in node_results:
swarm._accumulate_metrics(node_result)
tru_usage = swarm.state.accumulated_usage
exp_usage = {
"inputTokens": 25,
"outputTokens": 45,
"totalTokens": 89,
"cacheReadInputTokens": 12,
"cacheWriteInputTokens": 7,
}
assert tru_usage == exp_usage
assert (
tru_usage["inputTokens"]
+ tru_usage["outputTokens"]
+ tru_usage["cacheReadInputTokens"]
+ tru_usage["cacheWriteInputTokens"]
== tru_usage["totalTokens"]
)
def test_accumulate_metrics_without_cache_counters_omits_them():
"""Nodes without cache counters must not materialize cache keys as 0 (#797)."""
swarm = Swarm(nodes=[create_mock_agent("agent1"), create_mock_agent("agent2")])
result_text = AgentResult(
message={"role": "assistant", "content": [{"text": "ok"}]}, stop_reason="end_turn", state={}, metrics={}
)
swarm._accumulate_metrics(
NodeResult(result=result_text, accumulated_usage={"inputTokens": 10, "outputTokens": 20, "totalTokens": 30})
)
tru_usage = swarm.state.accumulated_usage
exp_usage = {"inputTokens": 10, "outputTokens": 20, "totalTokens": 30}
assert tru_usage == exp_usage
def test_swarm_auto_completion_without_handoff():
"""Test swarm auto-completion when no handoff occurs."""
# Create a simple agent that doesn't hand off