Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions strands-py/src/strands/agent/agent_result.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,8 @@ def context_size(self) -> int | None:
"""Most recent context size in tokens from the last LLM call.

Returns:
The input token count from the most recent cycle, or None if no data is available.
The total prompt the model processed on the most recent cycle, including cached tokens, or
None if no data is available.
"""
return self.metrics.latest_context_size

Expand All @@ -54,7 +55,8 @@ def projected_context_size(self) -> int | None:
"""Projected context size for the next model call.

Returns:
The projected token count (inputTokens + outputTokens), or None if no data is available.
The projected token count (total prompt including cached tokens plus generated output), or
None if no data is available.
"""
return self.metrics.projected_context_size

Expand Down
13 changes: 7 additions & 6 deletions strands-py/src/strands/event_loop/event_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from .._middleware.stages import InvokeModelContext, InvokeModelStage
from ..experimental.checkpoint import Checkpoint, CheckpointPosition
from ..hooks import AfterModelCallEvent, AfterToolsEvent, BeforeModelCallEvent, BeforeToolsEvent
from ..telemetry.metrics import Trace
from ..telemetry.metrics import Trace, _total_prompt_tokens
from ..telemetry.tracer import Tracer, get_tracer
from ..tools._validator import validate_and_prepare_tools
from ..tools.structured_output._structured_output_context import StructuredOutputContext
Expand Down Expand Up @@ -123,10 +123,11 @@ def _has_tool_use_in_latest_message(messages: "Messages") -> bool:
async def _estimate_input_tokens(agent: "Agent") -> int:
"""Estimate the input token count for the next model call.

Reads inputTokens + outputTokens from the last assistant message's metadata as a known
baseline, then estimates only new messages added after it. Falls back to full estimation
when no metadata is available (cold start or first call). On cold start, tool specs are
resolved lazily so that the caller does not need to resolve them before BeforeModelCallEvent.
Reads the total prompt the model processed (including cached tokens) plus outputTokens from the
last assistant message's metadata as a known baseline, then estimates only new messages added
after it. Falls back to full estimation when no metadata is available (cold start or first call).
On cold start, tool specs are resolved lazily so that the caller does not need to resolve them
before BeforeModelCallEvent.

Args:
agent: The agent instance with messages and model.
Expand All @@ -145,7 +146,7 @@ async def _estimate_input_tokens(agent: "Agent") -> int:

if last_assistant_idx >= 0:
usage = messages[last_assistant_idx]["metadata"]["usage"]
known_baseline = usage["inputTokens"] + usage["outputTokens"]
known_baseline = _total_prompt_tokens(usage) + usage["outputTokens"]
new_messages = messages[last_assistant_idx + 1 :]
if not new_messages:
return known_baseline
Expand Down
37 changes: 31 additions & 6 deletions strands-py/src/strands/telemetry/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,28 @@ class AgentInvocation:
usage: Usage = field(default_factory=lambda: Usage(inputTokens=0, outputTokens=0, totalTokens=0))


def _total_prompt_tokens(usage: Usage) -> int:
"""Return the total prompt the model processed, including cached tokens.

Providers either fold cache tokens into ``inputTokens`` or report them separately; the totals tell
us which. ``inputTokens + outputTokens == totalTokens`` means cache is already inside
``inputTokens``; otherwise the cache counters are added on top. Uncached usage returns
Comment thread
opieter-aws marked this conversation as resolved.
``inputTokens`` either way.

Args:
usage: Token usage from a model invocation.

Returns:
The total prompt token count.
"""
input_tokens = usage["inputTokens"]
output_tokens = usage["outputTokens"]
total_tokens = usage["totalTokens"]
if input_tokens + output_tokens == total_tokens:
Comment thread
opieter-aws marked this conversation as resolved.
return input_tokens
return input_tokens + usage.get("cacheReadInputTokens", 0) + usage.get("cacheWriteInputTokens", 0)


@dataclass
class EventLoopMetrics:
"""Aggregated metrics for an event loop's execution.
Expand Down Expand Up @@ -210,19 +232,22 @@ def latest_context_size(self) -> int | None:
This represents the current context size as reported by the model.

Returns:
The input token count from the most recent cycle, or None if no data is available.
The total prompt the model processed on the most recent cycle, including cached tokens, or
None if no data is available.
"""
if self.agent_invocations and self.agent_invocations[-1].cycles:
return self.agent_invocations[-1].cycles[-1].usage.get("inputTokens")
usage = self.agent_invocations[-1].cycles[-1].usage
if usage.get("inputTokens") is not None:
return _total_prompt_tokens(usage)
return None

@property
def projected_context_size(self) -> int | None:
"""Projected context size for the next model call.

Computed as inputTokens + outputTokens from the most recent cycle's usage,
representing the approximate input token count for the next model call
(prior input + generated output that is now part of the conversation).
Computed from the most recent cycle's usage as the total prompt the model processed (including
cached tokens) plus the generated output that is now part of the conversation, approximating the
input token count for the next model call.

Returns:
The projected token count, or None if no data is available.
Expand All @@ -232,7 +257,7 @@ def projected_context_size(self) -> int | None:
input_tokens = usage.get("inputTokens")
output_tokens = usage.get("outputTokens")
if input_tokens is not None and output_tokens is not None:
return input_tokens + output_tokens
return _total_prompt_tokens(usage) + output_tokens
return None

@property
Expand Down
11 changes: 7 additions & 4 deletions strands-py/src/strands/telemetry/tracer.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
from ..types.streaming import Metrics, StopReason, Usage
from ..types.tools import ToolResult, ToolUse
from ..types.traces import Attributes, AttributeValue
from .metrics import _total_prompt_tokens

if TYPE_CHECKING:
from ..memory.types import MemoryEntry
Expand Down Expand Up @@ -443,9 +444,10 @@ def end_model_invoke_span(
if not span or not span.is_recording():
return

prompt_tokens = _total_prompt_tokens(usage)
attributes: dict[str, AttributeValue] = {
"gen_ai.usage.prompt_tokens": usage["inputTokens"],
"gen_ai.usage.input_tokens": usage["inputTokens"],
"gen_ai.usage.prompt_tokens": prompt_tokens,
"gen_ai.usage.input_tokens": prompt_tokens,
"gen_ai.usage.completion_tokens": usage["outputTokens"],
"gen_ai.usage.output_tokens": usage["outputTokens"],
"gen_ai.usage.total_tokens": usage["totalTokens"],
Expand Down Expand Up @@ -821,11 +823,12 @@ def end_agent_span(
usage = latest_invocation.usage
else:
usage = response.metrics.accumulated_usage
prompt_tokens = _total_prompt_tokens(usage)
attributes.update(
{
"gen_ai.usage.prompt_tokens": usage["inputTokens"],
"gen_ai.usage.prompt_tokens": prompt_tokens,
"gen_ai.usage.completion_tokens": usage["outputTokens"],
"gen_ai.usage.input_tokens": usage["inputTokens"],
"gen_ai.usage.input_tokens": prompt_tokens,
"gen_ai.usage.output_tokens": usage["outputTokens"],
"gen_ai.usage.total_tokens": usage["totalTokens"],
"gen_ai.usage.cache_read.input_tokens": usage.get("cacheReadInputTokens", 0),
Expand Down
32 changes: 32 additions & 0 deletions strands-py/tests/strands/event_loop/test_event_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -2132,6 +2132,38 @@ async def test_baseline_plus_delta(self):
assert result == 180
agent.model.count_tokens.assert_called_once()

@pytest.mark.asyncio
async def test_baseline_counts_disjoint_cache_tokens(self):
"""Baseline includes cache reads on disjoint providers (#3546).

Without counting the cache read, a large cached prompt reads as a handful of tokens and
proactive compaction never fires. Here inputTokens + outputTokens != totalTokens, so the
cache read is additional to inputTokens and must be included in the baseline.
"""
agent = unittest.mock.AsyncMock()
agent.messages = [
{"role": "user", "content": [{"text": "Hi"}]},
{
"role": "assistant",
"content": [{"text": "Hello"}],
"metadata": {
"usage": {
"inputTokens": 10,
"outputTokens": 4,
"totalTokens": 5862,
"cacheReadInputTokens": 5848,
}
},
},
]
agent.system_prompt = "You are helpful"

result = await strands.event_loop.event_loop._estimate_input_tokens(agent)

# total prompt (10 + 5848 cache read) + output (4) = 5862, not 14
assert result == 5862
agent.model.count_tokens.assert_not_called()

@pytest.mark.asyncio
async def test_error_fallback_returns_none_at_call_site(self):
"""When count_tokens raises, the caller catches and sets projected_input_tokens to None."""
Expand Down
85 changes: 85 additions & 0 deletions strands-py/tests/strands/telemetry/test_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -653,6 +653,71 @@ def test_latest_context_size_missing_input_tokens_key(event_loop_metrics):
assert event_loop_metrics.latest_context_size is None


def test_latest_context_size_counts_disjoint_cache_tokens(event_loop_metrics, mock_get_meter_provider):
"""Counts cache reads on disjoint providers (Bedrock/Anthropic) where they add to inputTokens.

Regression for #3546: a large cache read must not read as a handful of tokens, or proactive
compaction never fires. Here inputTokens + outputTokens != totalTokens, so the cache is additional.
"""
event_loop_metrics.reset_usage_metrics()
event_loop_metrics.start_cycle(attributes={"event_loop_cycle_id": "c1"})
event_loop_metrics.update_usage(Usage(inputTokens=10, outputTokens=4, totalTokens=5862, cacheReadInputTokens=5848))

assert event_loop_metrics.latest_context_size == 5858


def test_latest_context_size_does_not_double_count_subset_cache_tokens(event_loop_metrics, mock_get_meter_provider):
"""Does not double-count cache reads on subset providers (OpenAI/Gemini) where they sit inside input.

Regression for #3546: inputTokens + outputTokens == totalTokens signals the cache is already inside
inputTokens, so the total prompt is inputTokens (not inputTokens + cache).
"""
event_loop_metrics.reset_usage_metrics()
event_loop_metrics.start_cycle(attributes={"event_loop_cycle_id": "c1"})
event_loop_metrics.update_usage(
Usage(inputTokens=12936, outputTokens=10, totalTokens=12946, cacheReadInputTokens=6457)
)

assert event_loop_metrics.latest_context_size == 12936


def test_latest_context_size_undercounts_anthropic_direct_cache(event_loop_metrics, mock_get_meter_provider):
"""Documents the #3546 known limitation: Anthropic-direct cache is not counted.
Comment thread
opieter-aws marked this conversation as resolved.

Anthropic-direct reports cache as a separate counter yet computes totalTokens as
inputTokens + outputTokens, so it is arithmetically indistinguishable from a subset provider and
the cache read is dropped -- an undercount that matches the prior baseline. Adapter-side
normalization to the disjoint convention (#3546) will make totalTokens include the cache and flip
this to the total prompt (5858); this test guards the boundary so that flip is intentional.
"""
event_loop_metrics.reset_usage_metrics()
event_loop_metrics.start_cycle(attributes={"event_loop_cycle_id": "c1"})
event_loop_metrics.update_usage(Usage(inputTokens=10, outputTokens=4, totalTokens=14, cacheReadInputTokens=5848))

assert event_loop_metrics.latest_context_size == 10


@pytest.mark.parametrize(
("usage", "exp_total_prompt"),
[
# Regression for #3546: cache tokens count toward the total prompt under both provider conventions.
# disjoint (Bedrock/Anthropic): cache reads add to inputTokens (inputTokens + outputTokens != totalTokens).
(Usage(inputTokens=10, outputTokens=4, totalTokens=5862, cacheReadInputTokens=5848), 5858),
# subset (OpenAI/Gemini): cache reads sit inside inputTokens (inputTokens + outputTokens == totalTokens).
(Usage(inputTokens=12936, outputTokens=10, totalTokens=12946, cacheReadInputTokens=6457), 12936),
# disjoint with both cache reads and writes added on top.
(
Usage(inputTokens=10, outputTokens=5, totalTokens=100, cacheReadInputTokens=60, cacheWriteInputTokens=25),
95,
),
# No cache tokens: both branches collapse to inputTokens (no behavior change for non-caching providers).
(Usage(inputTokens=100, outputTokens=50, totalTokens=150), 100),
],
)
def test_total_prompt_tokens(usage, exp_total_prompt):
assert strands.telemetry.metrics._total_prompt_tokens(usage) == exp_total_prompt


def test_projected_context_size_no_invocations(event_loop_metrics):
assert event_loop_metrics.projected_context_size is None

Expand Down Expand Up @@ -692,3 +757,23 @@ def test_projected_context_size_missing_tokens_key(event_loop_metrics):
)
)
assert event_loop_metrics.projected_context_size is None


def test_projected_context_size_counts_disjoint_cache_tokens(event_loop_metrics, mock_get_meter_provider):
"""Projects total prompt + output on disjoint providers where cache adds to inputTokens (#3546)."""
event_loop_metrics.reset_usage_metrics()
event_loop_metrics.start_cycle(attributes={"event_loop_cycle_id": "c1"})
event_loop_metrics.update_usage(Usage(inputTokens=10, outputTokens=4, totalTokens=5862, cacheReadInputTokens=5848))

assert event_loop_metrics.projected_context_size == 5862


def test_projected_context_size_does_not_double_count_subset_cache_tokens(event_loop_metrics, mock_get_meter_provider):
"""Projects total prompt + output on subset providers without double-counting cache (#3546)."""
event_loop_metrics.reset_usage_metrics()
event_loop_metrics.start_cycle(attributes={"event_loop_cycle_id": "c1"})
event_loop_metrics.update_usage(
Usage(inputTokens=12936, outputTokens=10, totalTokens=12946, cacheReadInputTokens=6457)
)

assert event_loop_metrics.projected_context_size == 12946
82 changes: 82 additions & 0 deletions strands-py/tests/strands/telemetry/test_tracer.py
Original file line number Diff line number Diff line change
Expand Up @@ -1178,6 +1178,46 @@ def test_end_model_invoke_span_with_cache_metrics(mock_span):
mock_span.end.assert_called_once()


def test_end_model_invoke_span_counts_disjoint_cache_tokens(mock_span):
"""Regression for #3546: input_tokens is the total prompt when cache is additional to inputTokens.

On disjoint providers (Bedrock/Anthropic) inputTokens + outputTokens != totalTokens, so the cache
reads/writes are additional and count toward the prompt the model processed. gen_ai.usage.input_tokens
(and its prompt_tokens alias) report 38, not the bare inputTokens of 10.
"""
tracer = Tracer()
message = {"role": "assistant", "content": [{"text": "Response"}]}
usage = Usage(
inputTokens=10,
outputTokens=20,
totalTokens=58,
cacheReadInputTokens=25,
cacheWriteInputTokens=3,
)
stop_reason: StopReason = "end_turn"
metrics = Metrics(latencyMs=10, timeToFirstByteMs=5)

tracer.end_model_invoke_span(mock_span, message, usage, metrics, stop_reason)

mock_span.set_attributes.assert_called_once_with(
{
"gen_ai.usage.prompt_tokens": 38,
"gen_ai.usage.input_tokens": 38,
"gen_ai.usage.completion_tokens": 20,
"gen_ai.usage.output_tokens": 20,
"gen_ai.usage.total_tokens": 58,
"gen_ai.usage.cache_read.input_tokens": 25,
"gen_ai.usage.cache_creation.input_tokens": 3,
"gen_ai.usage.cache_read_input_tokens": 25,
"gen_ai.usage.cache_write_input_tokens": 3,
"gen_ai.server.request.duration": 10,
"gen_ai.server.time_to_first_token": 5,
}
)
mock_span.set_status.assert_called_once_with(StatusCode.OK)
mock_span.end.assert_called_once()


def test_end_agent_span_with_cache_metrics(mock_span):
"""Test ending an agent span with cache metrics."""
tracer = Tracer()
Expand Down Expand Up @@ -1216,6 +1256,48 @@ def test_end_agent_span_with_cache_metrics(mock_span):
mock_span.end.assert_called_once()


def test_end_agent_span_counts_disjoint_cache_tokens(mock_span):
"""Regression for #3546: input_tokens is the total prompt when cache is additional to inputTokens.

On disjoint providers (Bedrock/Anthropic) inputTokens + outputTokens != totalTokens, so the cache
reads/writes are additional and count toward the prompt the model processed. gen_ai.usage.input_tokens
(and its prompt_tokens alias) report 85, not the bare inputTokens of 50.
"""
tracer = Tracer()

mock_metrics = mock.MagicMock()
mock_metrics.accumulated_usage = {
"inputTokens": 50,
"outputTokens": 100,
"totalTokens": 185,
"cacheReadInputTokens": 25,
"cacheWriteInputTokens": 10,
}

mock_response = mock.MagicMock()
mock_response.metrics = mock_metrics
mock_response.stop_reason = "end_turn"
mock_response.__str__ = mock.MagicMock(return_value="Agent response")

tracer.end_agent_span(mock_span, mock_response)

mock_span.set_attributes.assert_called_once_with(
{
"gen_ai.usage.prompt_tokens": 85,
"gen_ai.usage.input_tokens": 85,
"gen_ai.usage.completion_tokens": 100,
"gen_ai.usage.output_tokens": 100,
"gen_ai.usage.total_tokens": 185,
"gen_ai.usage.cache_read.input_tokens": 25,
"gen_ai.usage.cache_creation.input_tokens": 10,
"gen_ai.usage.cache_read_input_tokens": 25,
"gen_ai.usage.cache_write_input_tokens": 10,
}
)
mock_span.set_status.assert_called_once_with(StatusCode.OK)
mock_span.end.assert_called_once()


def test_end_model_invoke_span_dual_emits_semconv_and_deprecated_cache_names(mock_span):
"""Cache usage is emitted under both the semconv names and the deprecated pre-semconv aliases.

Expand Down
Loading
Loading