Skip to content

Commit 7ae8c29

Browse files
committed
fix(live): propagate output token count in live API usage metadata
1 parent abcaa08 commit 7ae8c29

4 files changed

Lines changed: 102 additions & 7 deletions

File tree

src/google/adk/flows/llm_flows/base_llm_flow.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,8 @@
2828
from websockets.exceptions import ConnectionClosed
2929
from websockets.exceptions import ConnectionClosedOK
3030

31+
from . import _output_schema_processor
32+
from . import functions
3133
from ...agents.base_agent import BaseAgent
3234
from ...agents.callback_context import CallbackContext
3335
from ...agents.invocation_context import InvocationContext
@@ -50,8 +52,6 @@
5052
from ...tools.tool_context import ToolContext
5153
from ...utils import model_name_utils
5254
from ...utils.context_utils import Aclosing
53-
from . import _output_schema_processor
54-
from . import functions
5555
from .audio_cache_manager import AudioCacheManager
5656
from .functions import build_auth_request_event
5757

src/google/adk/models/gemini_llm_connection.py

Lines changed: 33 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -217,6 +217,35 @@ def __build_full_text_response(
217217
live_session_id=self._gemini_session.session_id,
218218
)
219219

220+
def _to_generate_content_usage_metadata(
221+
self, usage_metadata: types.UsageMetadata
222+
) -> types.GenerateContentResponseUsageMetadata:
223+
"""Converts live API usage metadata to GenerateContentResponse usage metadata.
224+
225+
The live API names output tokens `response_token_count`/
226+
`response_tokens_details`, whereas `GenerateContentResponseUsageMetadata`
227+
names them `candidates_token_count`/`candidates_tokens_details`.
228+
229+
Args:
230+
usage_metadata: The live API usage metadata.
231+
232+
Returns:
233+
The converted usage metadata.
234+
"""
235+
return types.GenerateContentResponseUsageMetadata(
236+
prompt_token_count=usage_metadata.prompt_token_count,
237+
cached_content_token_count=usage_metadata.cached_content_token_count,
238+
candidates_token_count=usage_metadata.response_token_count,
239+
total_token_count=usage_metadata.total_token_count,
240+
thoughts_token_count=usage_metadata.thoughts_token_count,
241+
tool_use_prompt_token_count=usage_metadata.tool_use_prompt_token_count,
242+
prompt_tokens_details=usage_metadata.prompt_tokens_details,
243+
cache_tokens_details=usage_metadata.cache_tokens_details,
244+
candidates_tokens_details=usage_metadata.response_tokens_details,
245+
tool_use_prompt_tokens_details=usage_metadata.tool_use_prompt_tokens_details,
246+
traffic_type=usage_metadata.traffic_type,
247+
)
248+
220249
async def receive(self) -> AsyncGenerator[LlmResponse, None]:
221250
"""Receives the model response using the llm server connection.
222251
@@ -235,9 +264,11 @@ async def receive(self) -> AsyncGenerator[LlmResponse, None]:
235264
logger.debug('Got LLM Live message: %s', message)
236265
live_session_id = self._gemini_session.session_id
237266
if message.usage_metadata:
238-
# Tracks token usage data per model.
267+
# Remap live token usage to GenerateContentResponse usage metadata.
239268
yield LlmResponse(
240-
usage_metadata=message.usage_metadata,
269+
usage_metadata=self._to_generate_content_usage_metadata(
270+
message.usage_metadata
271+
),
241272
model_version=self._model_version,
242273
live_session_id=live_session_id,
243274
)

src/google/adk/telemetry/_instrumentation.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,9 +26,9 @@
2626
from opentelemetry import trace
2727
import opentelemetry.context as context_api
2828

29-
from ..events import event as event_lib
3029
from . import _metrics
3130
from . import tracing
31+
from ..events import event as event_lib
3232

3333
if TYPE_CHECKING:
3434
from ..agents.base_agent import BaseAgent

tests/unittests/models/test_gemini_llm_connection.py

Lines changed: 66 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -231,10 +231,12 @@ async def mock_receive_generator():
231231
content_response = next((r for r in responses if r.content), None)
232232
assert content_response is not None
233233

234+
# The live API's `response_token_count`/`response_tokens_details` are remapped
235+
# to `candidates_token_count`/`candidates_tokens_details`.
234236
expected_usage = types.GenerateContentResponseUsageMetadata(
235237
prompt_token_count=10,
236238
cached_content_token_count=5,
237-
candidates_token_count=None,
239+
candidates_token_count=20,
238240
total_token_count=35,
239241
thoughts_token_count=2,
240242
prompt_tokens_details=[
@@ -243,12 +245,74 @@ async def mock_receive_generator():
243245
cache_tokens_details=[
244246
types.ModalityTokenCount(modality='text', token_count=5)
245247
],
246-
candidates_tokens_details=None,
248+
candidates_tokens_details=[
249+
types.ModalityTokenCount(modality='text', token_count=20)
250+
],
247251
)
248252
assert usage_response.usage_metadata == expected_usage
249253
assert content_response.content == mock_content
250254

251255

256+
async def test_receive_usage_metadata_remaps_output_tokens(
257+
gemini_connection, mock_gemini_session
258+
):
259+
"""Test that live API output tokens are remapped to candidates_token_count."""
260+
usage_metadata = types.UsageMetadata(
261+
prompt_token_count=10,
262+
cached_content_token_count=5,
263+
response_token_count=20,
264+
total_token_count=35,
265+
thoughts_token_count=2,
266+
tool_use_prompt_token_count=3,
267+
prompt_tokens_details=[
268+
types.ModalityTokenCount(modality='text', token_count=10)
269+
],
270+
cache_tokens_details=[
271+
types.ModalityTokenCount(modality='text', token_count=5)
272+
],
273+
response_tokens_details=[
274+
types.ModalityTokenCount(modality='text', token_count=20)
275+
],
276+
)
277+
278+
mock_message = mock.AsyncMock()
279+
mock_message.usage_metadata = usage_metadata
280+
mock_message.server_content = None
281+
mock_message.tool_call = None
282+
mock_message.session_resumption_update = None
283+
mock_message.go_away = None
284+
285+
async def mock_receive_generator():
286+
yield mock_message
287+
288+
receive_mock = mock.Mock(return_value=mock_receive_generator())
289+
mock_gemini_session.receive = receive_mock
290+
291+
responses = [resp async for resp in gemini_connection.receive()]
292+
293+
usage_response = next((r for r in responses if r.usage_metadata), None)
294+
assert usage_response is not None
295+
result = usage_response.usage_metadata
296+
assert isinstance(result, types.GenerateContentResponseUsageMetadata)
297+
# Output tokens are remapped from response_* to candidates_*.
298+
assert result.candidates_token_count == 20
299+
assert result.candidates_tokens_details == [
300+
types.ModalityTokenCount(modality='text', token_count=20)
301+
]
302+
# Shared fields are carried over unchanged.
303+
assert result.prompt_token_count == 10
304+
assert result.cached_content_token_count == 5
305+
assert result.total_token_count == 35
306+
assert result.thoughts_token_count == 2
307+
assert result.tool_use_prompt_token_count == 3
308+
assert result.prompt_tokens_details == [
309+
types.ModalityTokenCount(modality='text', token_count=10)
310+
]
311+
assert result.cache_tokens_details == [
312+
types.ModalityTokenCount(modality='text', token_count=5)
313+
]
314+
315+
252316
async def test_receive_populates_live_session_id(
253317
gemini_connection, mock_gemini_session
254318
):

0 commit comments

Comments
 (0)