Skip to content
Open
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
18 changes: 12 additions & 6 deletions agentops/instrumentation/common/token_counting.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
information from various response formats.
"""

from collections.abc import Mapping
from typing import Dict, Any, Optional
from dataclasses import dataclass

Expand Down Expand Up @@ -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
Expand Down
15 changes: 3 additions & 12 deletions tests/unit/instrumentation/common/test_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
55 changes: 54 additions & 1 deletion tests/unit/instrumentation/common/test_token_counting.py
Original file line number Diff line number Diff line change
@@ -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


Expand All @@ -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)