diff --git a/python/packages/anthropic/agent_framework_anthropic/_chat_client.py b/python/packages/anthropic/agent_framework_anthropic/_chat_client.py index 367362c703..abda651e31 100644 --- a/python/packages/anthropic/agent_framework_anthropic/_chat_client.py +++ b/python/packages/anthropic/agent_framework_anthropic/_chat_client.py @@ -24,6 +24,7 @@ ResponseStream, TextSpanRegion, UsageDetails, + add_usage_details, tool, ) from agent_framework._settings import SecretString, load_settings @@ -545,8 +546,12 @@ def _inner_get_response( if stream: # Streaming mode async def _stream() -> AsyncIterable[ChatResponseUpdate]: + # Anthropic streams cumulative usage snapshots (message_start seeds it, + # each message_delta carries the running total), so thread a per-stream + # accumulator to _process_stream_event to emit increments instead. + emitted_usage: dict[str, int] = {} async for chunk in await self.anthropic_client.beta.messages.create(**run_options, stream=True): - parsed_chunk = self._process_stream_event(chunk) + parsed_chunk = self._process_stream_event(chunk, emitted_usage) if parsed_chunk: yield parsed_chunk @@ -1076,11 +1081,17 @@ def _process_message(self, message: BetaMessage, options: Mapping[str, Any]) -> raw_representation=message, ) - def _process_stream_event(self, event: BetaRawMessageStreamEvent) -> ChatResponseUpdate | None: + def _process_stream_event( + self, event: BetaRawMessageStreamEvent, emitted_usage: dict[str, int] | None = None + ) -> ChatResponseUpdate | None: """Process a streaming event from the Anthropic client. Args: event: The streaming event returned by the Anthropic client. + emitted_usage: Per-stream accumulator of the cumulative usage already + emitted, used to convert Anthropic's cumulative usage snapshots into + increments (see ``_incremental_usage``). Pass ``None`` for a one-off + event to keep the snapshot unchanged. Returns: A ChatResponseUpdate object containing the processed update. @@ -1089,7 +1100,9 @@ def _process_stream_event(self, event: BetaRawMessageStreamEvent) -> ChatRespons case "message_start": usage_details: list[Content] = [] if event.message.usage and (details := self._parse_usage_from_anthropic(event.message.usage)): - usage_details.append(Content.from_usage(usage_details=details)) + usage_details.append( + Content.from_usage(usage_details=self._incremental_usage(details, emitted_usage)) + ) return ChatResponseUpdate( role="assistant", @@ -1107,7 +1120,14 @@ def _process_stream_event(self, event: BetaRawMessageStreamEvent) -> ChatRespons case "message_delta": usage = self._parse_usage_from_anthropic(event.usage) return ChatResponseUpdate( - contents=[Content.from_usage(usage_details=usage, raw_representation=event.usage)] if usage else [], + contents=[ + Content.from_usage( + usage_details=self._incremental_usage(usage, emitted_usage), + raw_representation=event.usage, + ) + ] + if usage + else [], finish_reason=FINISH_REASON_MAP.get(event.delta.stop_reason) if event.delta.stop_reason else None, raw_representation=event, ) @@ -1146,6 +1166,40 @@ def _parse_usage_from_anthropic(self, usage: BetaUsage | BetaMessageDeltaUsage | usage_details["cache_read_input_token_count"] = usage.cache_read_input_tokens return usage_details + @staticmethod + def _incremental_usage(cumulative: UsageDetails, emitted: dict[str, int] | None) -> UsageDetails: + """Convert a cumulative Anthropic usage snapshot into the increment since the last one. + + Anthropic streams cumulative usage: ``message_start`` seeds it (with an + ``output_tokens`` placeholder of 1) and each ``message_delta`` reports the + running total for the message, not a per-delta increment. ``ChatResponse.from_updates`` + sums every usage ``Content``, so emitting the raw snapshots inflates the total by + the earlier ones (e.g. the seed makes ``output_token_count`` land at 26 when the + API reported 25). Emitting the increment over what has already been emitted makes + that summation reconstruct the final cumulative usage instead. + + ``emitted`` accumulates the last-seen cumulative value per key and is mutated in + place; it is threaded across a single stream. When it is ``None`` (e.g. a direct, + single-event call) the snapshot is returned unchanged. + """ + if emitted is None: + return cumulative + # The increment is cumulative minus what was already emitted, restricted to + # the keys this snapshot reports (a delta event may carry only a subset). + # add_usage_details sums int values and skips anything else, so negating the + # emitted totals and adding them yields exactly that increment. + negated = UsageDetails() + for key in cumulative: + emitted_value = emitted.get(key) + if isinstance(emitted_value, int): + negated[key] = -emitted_value + delta = add_usage_details(cumulative, negated) + emitted.clear() + for key, value in cumulative.items(): + if isinstance(value, int): + emitted[key] = value + return delta + def _parse_contents_from_anthropic( self, content: Sequence[BetaContentBlock | BetaRawContentBlockDelta | BetaTextBlock], diff --git a/python/packages/anthropic/tests/test_anthropic_client.py b/python/packages/anthropic/tests/test_anthropic_client.py index 8876823cc6..3c41754140 100644 --- a/python/packages/anthropic/tests/test_anthropic_client.py +++ b/python/packages/anthropic/tests/test_anthropic_client.py @@ -10,11 +10,13 @@ Agent, ChatMiddlewareLayer, ChatOptions, + ChatResponse, ChatResponseUpdate, Content, FunctionInvocationLayer, Message, SupportsChatGetResponse, + UsageDetails, tool, ) from agent_framework._settings import load_settings @@ -22,6 +24,7 @@ from agent_framework.observability import ChatTelemetryLayer from anthropic.types.beta import ( BetaMessage, + BetaMessageDeltaUsage, BetaTextBlock, BetaToolUseBlock, BetaUsage, @@ -1716,6 +1719,82 @@ def test_process_stream_event_message_start_sets_assistant_role(mock_anthropic_c assert result.role == "assistant" +def _usage_message_start_event(*, input_tokens: int, output_tokens: int) -> MagicMock: + event = MagicMock() + event.type = "message_start" + event.message.id = "msg_usage" + event.message.role = "assistant" + event.message.model = "claude-3-5-sonnet-20241022" + event.message.content = [] + event.message.stop_reason = None + event.message.usage = BetaUsage(input_tokens=input_tokens, output_tokens=output_tokens) + return event + + +def _usage_message_delta_event(*, output_tokens: int, input_tokens: int | None = None) -> MagicMock: + event = MagicMock() + event.type = "message_delta" + event.usage = BetaMessageDeltaUsage( + input_tokens=input_tokens, + output_tokens=output_tokens, + cache_creation_input_tokens=None, + cache_read_input_tokens=None, + ) + event.delta.stop_reason = "end_turn" + return event + + +def test_streaming_usage_not_double_counted(mock_anthropic_client: MagicMock) -> None: + """message_start's seed usage must not be summed onto message_delta's cumulative total. + + Anthropic reports cumulative usage on message_delta (per their streaming docs), while + message_start carries an output_tokens=1 seed. ChatResponse.from_updates sums every + usage Content, which used to inflate output_token_count by the seed — 26 when the API + reported 25. + """ + client = create_test_anthropic_client(mock_anthropic_client) + emitted: dict[str, int] = {} + updates = [ + u + for u in ( + client._process_stream_event(_usage_message_start_event(input_tokens=10, output_tokens=1), emitted), + client._process_stream_event(_usage_message_delta_event(output_tokens=25), emitted), + ) + if u is not None + ] + + response = ChatResponse.from_updates(updates) + + assert response.usage_details is not None + assert response.usage_details["output_token_count"] == 25 + assert response.usage_details["input_token_count"] == 10 + + +def test_streaming_usage_delta_input_not_double_counted(mock_anthropic_client: MagicMock) -> None: + """When message_delta also reports cumulative input tokens, the input must not double. + + Server-tool turns report cumulative input_tokens on message_delta; summing them onto + message_start's input snapshot double-counted the prompt. The final input_token_count + should equal the last cumulative value message_delta reports. + """ + client = create_test_anthropic_client(mock_anthropic_client) + emitted: dict[str, int] = {} + updates = [ + u + for u in ( + client._process_stream_event(_usage_message_start_event(input_tokens=10, output_tokens=1), emitted), + client._process_stream_event(_usage_message_delta_event(input_tokens=12, output_tokens=25), emitted), + ) + if u is not None + ] + + response = ChatResponse.from_updates(updates) + + assert response.usage_details is not None + assert response.usage_details["output_token_count"] == 25 + assert response.usage_details["input_token_count"] == 12 + + def test_process_stream_event_message_start_role_prevents_tool_use_collapse() -> None: """Regression test: tool_use blocks must not end up in a user-role message.