diff --git a/agentops/instrumentation/common/token_counting.py b/agentops/instrumentation/common/token_counting.py index 7d2f19f07..827353a07 100644 --- a/agentops/instrumentation/common/token_counting.py +++ b/agentops/instrumentation/common/token_counting.py @@ -6,6 +6,7 @@ from typing import Dict, Any, Optional from dataclasses import dataclass +from collections.abc import Mapping from agentops.logging import logger from agentops.semconv import SpanAttributes @@ -84,10 +85,20 @@ def extract_from_response(response: Any) -> TokenUsage: @staticmethod def _extract_from_usage_object(usage_data: Any) -> TokenUsage: - """Extract from a usage object with standard attributes.""" + """Extract from a usage object or mapping with standard attribute names.""" if not usage_data: return TokenUsage() + if isinstance(usage_data, Mapping): + return TokenUsage( + prompt_tokens=usage_data.get("prompt_tokens"), + completion_tokens=usage_data.get("completion_tokens"), + total_tokens=usage_data.get("total_tokens"), + cached_prompt_tokens=usage_data.get("cached_prompt_tokens"), + cached_read_tokens=usage_data.get("cache_read_input_tokens"), + reasoning_tokens=usage_data.get("reasoning_tokens"), + ) + return TokenUsage( prompt_tokens=getattr(usage_data, "prompt_tokens", None), completion_tokens=getattr(usage_data, "completion_tokens", None), diff --git a/tests/unit/instrumentation/common/test_token_counting.py b/tests/unit/instrumentation/common/test_token_counting.py index 2f00561b9..c7084a69b 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,31 @@ 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_mapping_usage_metadata(self): + response = SimpleNamespace( + usage_metadata={ + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + } + ) + + usage = TokenUsageExtractor.extract_from_response(response) + + assert usage.prompt_tokens == 10 + assert usage.completion_tokens == 5 + assert usage.total_tokens == 15 + + def test_extracts_attribute_usage(self): + response = SimpleNamespace( + usage=SimpleNamespace(prompt_tokens=7, completion_tokens=3, total_tokens=10) + ) + + usage = TokenUsageExtractor.extract_from_response(response) + + assert usage.prompt_tokens == 7 + assert usage.completion_tokens == 3 + assert usage.total_tokens == 10