From 29706d8ddbc4b200a51c90c3d683f1de15b555ab Mon Sep 17 00:00:00 2001 From: KXH Date: Sat, 25 Jul 2026 21:46:13 +0800 Subject: [PATCH] fix: extract token usage from mappings --- .../instrumentation/common/token_counting.py | 18 ++++-- .../instrumentation/common/test_streaming.py | 15 +---- .../common/test_token_counting.py | 55 ++++++++++++++++++- 3 files changed, 69 insertions(+), 19 deletions(-) diff --git a/agentops/instrumentation/common/token_counting.py b/agentops/instrumentation/common/token_counting.py index 7d2f19f07..9721b12f1 100644 --- a/agentops/instrumentation/common/token_counting.py +++ b/agentops/instrumentation/common/token_counting.py @@ -4,6 +4,7 @@ information from various response formats. """ +from collections.abc import Mapping from typing import Dict, Any, Optional from dataclasses import dataclass @@ -88,13 +89,18 @@ def _extract_from_usage_object(usage_data: Any) -> TokenUsage: if not usage_data: return TokenUsage() + def get_value(field: str) -> Optional[int]: + if isinstance(usage_data, Mapping): + return usage_data.get(field) + return getattr(usage_data, field, None) + return TokenUsage( - prompt_tokens=getattr(usage_data, "prompt_tokens", None), - completion_tokens=getattr(usage_data, "completion_tokens", None), - total_tokens=getattr(usage_data, "total_tokens", None), - cached_prompt_tokens=getattr(usage_data, "cached_prompt_tokens", None), - cached_read_tokens=getattr(usage_data, "cache_read_input_tokens", None), - reasoning_tokens=getattr(usage_data, "reasoning_tokens", None), + prompt_tokens=get_value("prompt_tokens"), + completion_tokens=get_value("completion_tokens"), + total_tokens=get_value("total_tokens"), + cached_prompt_tokens=get_value("cached_prompt_tokens"), + cached_read_tokens=get_value("cache_read_input_tokens"), + reasoning_tokens=get_value("reasoning_tokens"), ) @staticmethod diff --git a/tests/unit/instrumentation/common/test_streaming.py b/tests/unit/instrumentation/common/test_streaming.py index 31fabc683..624db1e30 100644 --- a/tests/unit/instrumentation/common/test_streaming.py +++ b/tests/unit/instrumentation/common/test_streaming.py @@ -133,19 +133,10 @@ def test_process_chunk_with_usage_metadata(self): extract_content = lambda x: "test" wrapper = BaseStreamWrapper(mock_stream, mock_span, extract_content) - # Mock chunk with usage_metadata - mock_chunk = Mock() - mock_chunk.usage_metadata = {"prompt_tokens": 10, "completion_tokens": 5} - - with patch( - "agentops.instrumentation.common.token_counting.TokenUsageExtractor.extract_from_response" - ) as mock_extract: - mock_usage = Mock() - mock_usage.prompt_tokens = 10 - mock_usage.completion_tokens = 5 - mock_extract.return_value = mock_usage + # Use the mapping shape emitted by some streaming providers. + mock_chunk = SimpleNamespace(usage_metadata={"prompt_tokens": 10, "completion_tokens": 5}) - wrapper._process_chunk(mock_chunk) + wrapper._process_chunk(mock_chunk) assert wrapper.token_usage.prompt_tokens == 10 assert wrapper.token_usage.completion_tokens == 5 diff --git a/tests/unit/instrumentation/common/test_token_counting.py b/tests/unit/instrumentation/common/test_token_counting.py index 2f00561b9..779117b83 100644 --- a/tests/unit/instrumentation/common/test_token_counting.py +++ b/tests/unit/instrumentation/common/test_token_counting.py @@ -1,4 +1,6 @@ -from agentops.instrumentation.common.token_counting import TokenUsage +from types import SimpleNamespace + +from agentops.instrumentation.common.token_counting import TokenUsage, TokenUsageExtractor from agentops.semconv import SpanAttributes @@ -24,3 +26,54 @@ def test_includes_positive_values_only(self): assert SpanAttributes.LLM_USAGE_COMPLETION_TOKENS not in attrs assert SpanAttributes.LLM_USAGE_TOTAL_TOKENS in attrs assert attrs[SpanAttributes.LLM_USAGE_TOTAL_TOKENS] == 5 + + +class TestTokenUsageExtractor: + def test_extracts_object_usage(self): + response = SimpleNamespace( + usage=SimpleNamespace( + prompt_tokens=6, + completion_tokens=3, + total_tokens=9, + ) + ) + + usage = TokenUsageExtractor.extract_from_response(response) + + assert usage == TokenUsage(prompt_tokens=6, completion_tokens=3, total_tokens=9) + + def test_extracts_mapping_usage(self): + response = SimpleNamespace( + usage={ + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "cached_prompt_tokens": 3, + "cache_read_input_tokens": 2, + "reasoning_tokens": 1, + } + ) + + usage = TokenUsageExtractor.extract_from_response(response) + + assert usage == TokenUsage( + prompt_tokens=10, + completion_tokens=5, + total_tokens=15, + cached_prompt_tokens=3, + cached_read_tokens=2, + reasoning_tokens=1, + ) + + def test_extracts_mapping_usage_metadata(self): + response = SimpleNamespace( + usage_metadata={ + "prompt_tokens": 8, + "completion_tokens": 4, + "total_tokens": 12, + } + ) + + usage = TokenUsageExtractor.extract_from_response(response) + + assert usage == TokenUsage(prompt_tokens=8, completion_tokens=4, total_tokens=12)