fix(models): keep streamed usage metadata when a later chunk reports none

Co-authored-by: George Weale <gweale@google.com>
PiperOrigin-RevId: 959816931
This commit is contained in:
George Weale
2026-08-05 12:49:53 -07:00
committed by Copybara-Service
parent 3bbc8ed2cf
commit aebb2a13b3
2 changed files with 103 additions and 2 deletions
+7 -2
View File
@@ -35,7 +35,9 @@ class StreamingResponseAggregator:
def __init__(self) -> None:
self._text: list[str] = []
self._thought_text: list[str] = []
self._usage_metadata = None
self._usage_metadata: Optional[
types.GenerateContentResponseUsageMetadata
] = None
self._grounding_metadata: Optional[types.GroundingMetadata] = None
self._citation_metadata: Optional[types.CitationMetadata] = None
self._response = None
@@ -264,7 +266,10 @@ class StreamingResponseAggregator:
# results = []
self._response = response
llm_response = LlmResponse.create(response)
self._usage_metadata = llm_response.usage_metadata
# Usage is typically reported on a single chunk; keep the last reported
# value rather than letting a usage-less trailing chunk erase it.
if llm_response.usage_metadata:
self._usage_metadata = llm_response.usage_metadata
if llm_response.grounding_metadata:
self._grounding_metadata = llm_response.grounding_metadata
if llm_response.citation_metadata:
@@ -392,6 +392,102 @@ class TestStreamingResponseAggregator:
else:
await run_test()
@pytest.mark.asyncio
@pytest.mark.parametrize("use_progressive_sse", [False, True])
async def test_close_preserves_usage_metadata_from_earlier_chunk(
self, use_progressive_sse
):
"""A later chunk without usage must not erase an earlier chunk's counts.
Providers typically report token usage on a single chunk; the trailing
chunks of the same turn carry none. The aggregated response is the one
that gets persisted, so it must retain the counts it already saw.
"""
with temporary_feature_override(
FeatureName.PROGRESSIVE_SSE_STREAMING, use_progressive_sse
):
aggregator = streaming_utils.StreamingResponseAggregator()
# First chunk carries the token counts.
response1 = types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(parts=[types.Part(text="Hello ")]),
)
],
usage_metadata=types.GenerateContentResponseUsageMetadata(
prompt_token_count=10,
candidates_token_count=5,
total_token_count=15,
),
)
# Second chunk carries none.
response2 = types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(parts=[types.Part(text="World!")]),
finish_reason=types.FinishReason.STOP,
)
],
)
async for _ in aggregator.process_response(response1):
pass
async for _ in aggregator.process_response(response2):
pass
closed_response = aggregator.close()
assert closed_response is not None
assert closed_response.usage_metadata is not None
assert closed_response.usage_metadata.prompt_token_count == 10
assert closed_response.usage_metadata.candidates_token_count == 5
assert closed_response.usage_metadata.total_token_count == 15
@pytest.mark.asyncio
@pytest.mark.parametrize("use_progressive_sse", [False, True])
async def test_close_uses_latest_reported_usage_metadata(
self, use_progressive_sse
):
"""When several chunks report usage, the most recent one wins."""
with temporary_feature_override(
FeatureName.PROGRESSIVE_SSE_STREAMING, use_progressive_sse
):
aggregator = streaming_utils.StreamingResponseAggregator()
response1 = types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(parts=[types.Part(text="Hello ")]),
)
],
usage_metadata=types.GenerateContentResponseUsageMetadata(
prompt_token_count=10,
candidates_token_count=5,
total_token_count=15,
),
)
response2 = types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(parts=[types.Part(text="World!")]),
finish_reason=types.FinishReason.STOP,
)
],
usage_metadata=types.GenerateContentResponseUsageMetadata(
prompt_token_count=10,
candidates_token_count=9,
total_token_count=19,
),
)
async for _ in aggregator.process_response(response1):
pass
async for _ in aggregator.process_response(response2):
pass
closed_response = aggregator.close()
assert closed_response is not None
assert closed_response.usage_metadata is not None
assert closed_response.usage_metadata.total_token_count == 19
@pytest.mark.asyncio
@pytest.mark.parametrize("use_progressive_sse", [False, True])
async def test_close_propagates_model_version(self, use_progressive_sse):