Skip to content

Commit 587d9fe

Browse files
authored
feat(tracing): report cached input token usage (#1148)
1 parent 6f5e444 commit 587d9fe

4 files changed

Lines changed: 130 additions & 17 deletions

File tree

‎tests/test_llm_attributes_extractors.py‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,12 +16,17 @@
1616
from types import SimpleNamespace
1717
from unittest.mock import Mock, call
1818

19+
import pytest
1920
from google.adk.models.llm_request import LlmRequest
21+
from google.adk.models.llm_response import LlmResponse
2022
from google.adk.tools.function_tool import FunctionTool
23+
from google.genai import types
2124
from opentelemetry.sdk.trace import TracerProvider
2225

2326
from veadk.tracing.telemetry.attributes.extractors.llm_attributes_extractors import (
2427
llm_gen_ai_request_functions,
28+
llm_gen_ai_usage_cache_creation_input_tokens,
29+
llm_gen_ai_usage_cache_read_input_tokens,
2530
llm_gen_ai_usage_output_tokens,
2631
)
2732
from veadk.tracing.telemetry.attributes.extractors.types import ExtractorResponse
@@ -118,6 +123,45 @@ def test_missing_output_token_count_is_not_written_to_span():
118123
span.set_attribute.assert_not_called()
119124

120125

126+
@pytest.mark.parametrize("cached_tokens", [600, 0, None])
127+
def test_cached_tokens_are_only_reported_as_cache_reads(cached_tokens):
128+
params = SimpleNamespace(
129+
llm_response=LlmResponse(
130+
usage_metadata=types.GenerateContentResponseUsageMetadata(
131+
cached_content_token_count=cached_tokens,
132+
)
133+
)
134+
)
135+
provider = TracerProvider()
136+
try:
137+
with provider.get_tracer(__name__).start_as_current_span("cache-usage") as span:
138+
ExtractorResponse.update_span(
139+
span,
140+
"gen_ai.usage.cache_read_input_tokens",
141+
llm_gen_ai_usage_cache_read_input_tokens(params),
142+
)
143+
ExtractorResponse.update_span(
144+
span,
145+
"gen_ai.usage.cache_creation_input_tokens",
146+
llm_gen_ai_usage_cache_creation_input_tokens(params),
147+
)
148+
expected = (
149+
{"gen_ai.usage.cache_read_input_tokens": cached_tokens}
150+
if cached_tokens is not None
151+
else {}
152+
)
153+
assert dict(span.attributes) == expected
154+
finally:
155+
provider.shutdown()
156+
157+
158+
def test_missing_usage_omits_cache_attributes():
159+
params = SimpleNamespace(llm_response=LlmResponse())
160+
161+
assert llm_gen_ai_usage_cache_read_input_tokens(params).content is None
162+
assert llm_gen_ai_usage_cache_creation_input_tokens(params).content is None
163+
164+
121165
def test_falsy_attribute_values_are_written_to_span():
122166
span = Mock()
123167

‎tests/test_portal_metrics.py‎

Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,10 @@
1616
from types import SimpleNamespace
1717

1818
import pytest
19+
from google.adk.agents.run_config import RunConfig, StreamingMode
20+
from google.adk.models.llm_request import LlmRequest
21+
from google.adk.models.llm_response import LlmResponse
22+
from google.genai import types
1923
from opentelemetry import metrics as metrics_api
2024
from opentelemetry.metrics import _internal as metrics_internal
2125
from opentelemetry.sdk import metrics as metrics_sdk
@@ -138,6 +142,79 @@ def test_proxy_instruments_follow_provider_installed_later(
138142
]
139143

140144

145+
@pytest.mark.parametrize(
146+
"server_address",
147+
[
148+
"https://ark.cn-beijing.volces.com/api/v3/",
149+
"https://ark.ap-southeast.bytepluses.com/api/v3",
150+
],
151+
ids=["volcengine", "byteplus"],
152+
)
153+
@pytest.mark.parametrize("streaming_mode", [StreamingMode.NONE, StreamingMode.SSE])
154+
@pytest.mark.parametrize("cached_tokens", [600, 0, None])
155+
def test_cache_read_token_metrics_preserve_input_usage(
156+
fresh_global_meter_provider, server_address, streaming_mode, cached_tokens
157+
):
158+
reader = InMemoryMetricReader()
159+
metrics_api.set_meter_provider(metrics_sdk.MeterProvider(metric_readers=[reader]))
160+
recorder = PortalMetricRecorder(name="test-cache-read")
161+
context = SimpleNamespace(
162+
agent=SimpleNamespace(model_api_base=server_address),
163+
run_config=RunConfig(streaming_mode=streaming_mode),
164+
)
165+
request = LlmRequest(model="test-model")
166+
167+
if streaming_mode == StreamingMode.SSE:
168+
recorder.record_call_llm(context, "partial", request, LlmResponse(partial=True))
169+
170+
recorder.record_call_llm(
171+
context,
172+
"final",
173+
request,
174+
LlmResponse(
175+
usage_metadata=types.GenerateContentResponseUsageMetadata(
176+
prompt_token_count=1000,
177+
candidates_token_count=200,
178+
total_token_count=1200,
179+
cached_content_token_count=cached_tokens,
180+
),
181+
),
182+
)
183+
184+
metrics_data = reader.get_metrics_data()
185+
recorded_metrics = {
186+
metric.name: list(metric.data.data_points)
187+
for resource_metrics in metrics_data.resource_metrics
188+
for scope_metrics in resource_metrics.scope_metrics
189+
for metric in scope_metrics.metrics
190+
}
191+
points = recorded_metrics["gen_ai.client.token.usage"]
192+
expected = {"input": 1000, "output": 200}
193+
if cached_tokens is not None:
194+
expected["cache_read"] = cached_tokens
195+
assert {point.attributes["gen_ai_token_type"]: point.sum for point in points} == (
196+
expected
197+
)
198+
assert len(points) == len(expected)
199+
for point in points:
200+
assert point.count == 1
201+
assert point.attributes["server_address"] == server_address
202+
assert point.attributes["stream"] == (streaming_mode == StreamingMode.SSE)
203+
assert point.attributes["gen_ai_response_model"] == "test-model"
204+
assert [point.value for point in recorded_metrics["gen_ai.chat.count"]] == [1]
205+
206+
207+
def test_missing_usage_does_not_record_token_metrics(fresh_global_meter_provider):
208+
reader = InMemoryMetricReader()
209+
metrics_api.set_meter_provider(metrics_sdk.MeterProvider(metric_readers=[reader]))
210+
recorder = PortalMetricRecorder(name="test-missing-usage")
211+
context = SimpleNamespace(agent=SimpleNamespace(), run_config=None)
212+
213+
recorder.record_call_llm(context, "event", LlmRequest(), LlmResponse())
214+
215+
assert reader.get_metrics_data() is None
216+
217+
141218
def test_apmplus_reuses_preconfigured_global_provider(
142219
fresh_global_meter_provider,
143220
monkeypatch,

‎veadk/tracing/telemetry/attributes/extractors/llm_attributes_extractors.py‎

Lines changed: 2 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -205,29 +205,16 @@ def llm_gen_ai_usage_total_tokens(params: LLMAttributesParams) -> ExtractorRespo
205205
return ExtractorResponse(content=None)
206206

207207

208-
# FIXME
209208
def llm_gen_ai_usage_cache_creation_input_tokens(
210209
params: LLMAttributesParams,
211210
) -> ExtractorResponse:
212-
"""Extract the number of tokens used for cache creation.
211+
"""Omit cache creation usage until response metadata provides it
213212
214-
Provides the count of tokens used for creating cached content,
215-
which affects cost calculation in caching-enabled models.
216-
217-
Args:
218-
params: LLM execution parameters containing response metadata
219-
220-
Returns:
221-
ExtractorResponse: Response containing cache creation token count or None
213+
cached_content_token_count measures cache reads, not cache creation
222214
"""
223-
if params.llm_response.usage_metadata:
224-
return ExtractorResponse(
225-
content=params.llm_response.usage_metadata.cached_content_token_count,
226-
)
227215
return ExtractorResponse(content=None)
228216

229217

230-
# FIXME
231218
def llm_gen_ai_usage_cache_read_input_tokens(
232219
params: LLMAttributesParams,
233220
) -> ExtractorResponse:

‎veadk/tracing/telemetry/portal_metrics.py‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -162,7 +162,7 @@ class PortalMetricRecorder:
162162
163163
Metrics Collected:
164164
- LLM invocation counts and frequencies
165-
- Token consumption (input/output) with histogram distribution
165+
- Token consumption (input/output/cache_read) with histogram distribution
166166
- Operation latency with performance bucket analysis
167167
- Error rates and exception details
168168
- Span-level performance metrics for APMPlus dashboards
@@ -258,7 +258,7 @@ def record_call_llm(
258258
259259
Metrics Recorded:
260260
- Invocation count with model and operation attributes
261-
- Input/output token usage with separate tracking
261+
- Input/output token usage and the cached subset of input tokens
262262
- Operation duration from span timing data
263263
- Error counts and exception details
264264
- Span latency for performance analysis
@@ -290,13 +290,18 @@ def record_call_llm(
290290
# upload token usage
291291
input_token = llm_response.usage_metadata.prompt_token_count
292292
output_token = llm_response.usage_metadata.candidates_token_count
293+
cache_read_token = llm_response.usage_metadata.cached_content_token_count
293294

294295
if input_token:
295296
token_attributes = {**attributes, "gen_ai_token_type": "input"}
296297
self.token_usage.record(input_token, attributes=token_attributes)
297298
if output_token:
298299
token_attributes = {**attributes, "gen_ai_token_type": "output"}
299300
self.token_usage.record(output_token, attributes=token_attributes)
301+
# Cache reads are already included in input tokens, not extra usage
302+
if cache_read_token is not None:
303+
token_attributes = {**attributes, "gen_ai_token_type": "cache_read"}
304+
self.token_usage.record(cache_read_token, attributes=token_attributes)
300305

301306
# Get llm duration
302307
span = trace.get_current_span()

0 commit comments

Comments
 (0)