fix(vertex): propagate Vertex AI metadata in streaming success callbacks

Streaming calls assembled via stream_chunk_builder were missing
vertex_ai_grounding_metadata and vertex_ai_url_context_metadata in
standard_logging_object.response. Merge metadata from chunks into the
assembled response and mirror non-streaming hidden_params on Gemini chunks.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-06-08 09:49:31 +05:30
parent aaf1e2444b
commit d791df59b6
No known key found for this signature in database
5 changed files with 171 additions and 0 deletions

View file

@ -32,6 +32,14 @@ if TYPE_CHECKING:
)
VERTEX_AI_STREAM_METADATA_FIELDS = (
"vertex_ai_grounding_metadata",
"vertex_ai_url_context_metadata",
"vertex_ai_safety_ratings",
"vertex_ai_citation_metadata",
)
class ChunkProcessor:
def __init__(self, chunks: List, messages: Optional[list] = None):
self.chunks = self._sort_chunks(chunks)
@ -79,6 +87,42 @@ class ChunkProcessor:
model_response._hidden_params = chunk.get("_hidden_params", {})
return model_response
@staticmethod
def _get_chunk_attr(chunk: Any, field_name: str) -> Any:
if isinstance(chunk, dict):
value = chunk.get(field_name)
if value is not None:
return value
model_extra = chunk.get("model_extra")
if isinstance(model_extra, dict):
return model_extra.get(field_name)
return None
return getattr(chunk, field_name, None)
@staticmethod
def propagate_vertex_ai_metadata_from_chunks(
response: ModelResponse, chunks: List[Any]
) -> None:
"""
Merge Vertex AI metadata from streaming chunks into the assembled response.
Gemini/Vertex streaming sets these fields on individual chunks but
stream_chunk_builder must propagate them for logging callbacks.
"""
for field_name in VERTEX_AI_STREAM_METADATA_FIELDS:
merged: List[Any] = []
for chunk in chunks:
value = ChunkProcessor._get_chunk_attr(chunk, field_name)
if not value:
continue
if isinstance(value, list):
merged.extend(value)
else:
merged.append(value)
if merged:
setattr(response, field_name, merged)
response._hidden_params[field_name] = merged
@staticmethod
def _get_chunk_id(chunks: List[Dict[str, Any]]) -> str:
"""

View file

@ -3386,9 +3386,17 @@ class ModelResponseIterator:
choice.finish_reason = "tool_calls"
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
model_response._hidden_params["vertex_ai_grounding_metadata"] = (
grounding_metadata
)
setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore
model_response._hidden_params["vertex_ai_url_context_metadata"] = (
url_context_metadata
)
setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore
model_response._hidden_params["vertex_ai_safety_ratings"] = safety_ratings
setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore
model_response._hidden_params["vertex_ai_citation_metadata"] = citation_metadata
return (
grounding_metadata,

View file

@ -7761,6 +7761,7 @@ def stream_chunk_builder( # noqa: PLR0915
"cost",
logging_obj._response_cost_calculator(result=response),
)
processor.propagate_vertex_ai_metadata_from_chunks(response, chunks)
return response
tool_call_chunks = [
@ -7940,6 +7941,7 @@ def stream_chunk_builder( # noqa: PLR0915
usage, "cost", logging_obj._response_cost_calculator(result=response)
)
processor.propagate_vertex_ai_metadata_from_chunks(response, chunks)
return response
except Exception as e:
verbose_logger.exception(

View file

@ -2165,6 +2165,41 @@ def test_get_assembled_streaming_response_returns_result_for_streaming():
assert assembled is result
def test_streaming_success_handler_includes_vertex_ai_metadata_in_standard_logging():
"""Assembled streaming responses should include Vertex AI metadata in logging payload."""
import datetime
from litellm.types.utils import Choices, Message
logging_obj = _make_logging_obj(stream=True)
grounding_metadata = [{"webSearchQueries": ["weather in SF"]}]
url_context_metadata = [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}]
result = ModelResponse(
id="resp-1",
choices=[
Choices(
index=0,
message=Message(role="assistant", content="hello"),
finish_reason="stop",
)
],
model="gemini-2.5-flash",
)
setattr(result, "vertex_ai_grounding_metadata", grounding_metadata)
setattr(result, "vertex_ai_url_context_metadata", url_context_metadata)
result._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata
result._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata
start = datetime.datetime.now()
end = datetime.datetime.now()
logging_obj.success_handler(result=result, start_time=start, end_time=end)
payload = logging_obj.model_call_details.get("standard_logging_object")
assert payload is not None
assert payload["response"]["vertex_ai_grounding_metadata"] == grounding_metadata
assert payload["response"]["vertex_ai_url_context_metadata"] == url_context_metadata
def test_get_assembled_streaming_response_returns_none_for_non_streaming_text_completion():
"""Non-streaming TextCompletionResponse should also return None."""
import datetime

View file

@ -613,3 +613,85 @@ def test_stream_chunk_builder_dict_snapshot_preserves_hidden_provider_fields():
assert (
response._hidden_params["provider_specific_fields"]["traffic_type"] == "default"
)
def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_chunks():
"""Vertex AI metadata on streaming chunks must appear on assembled response."""
grounding_metadata = [{"webSearchQueries": ["weather in SF"]}]
url_context_metadata = [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}]
chunk1 = ModelResponseStream(
id="chatcmpl-vertex-1",
created=1,
model="gemini-2.5-flash",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(content="The weather", role="assistant"),
)
],
)
setattr(chunk1, "vertex_ai_grounding_metadata", grounding_metadata)
chunk1._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata
chunk2 = ModelResponseStream(
id="chatcmpl-vertex-1",
created=1,
model="gemini-2.5-flash",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content=" is sunny.", role="assistant"),
)
],
)
setattr(chunk2, "vertex_ai_url_context_metadata", url_context_metadata)
chunk2._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata
response = stream_chunk_builder(chunks=[chunk1, chunk2])
assert response is not None
assert getattr(response, "vertex_ai_grounding_metadata") == grounding_metadata
assert getattr(response, "vertex_ai_url_context_metadata") == url_context_metadata
assert response._hidden_params["vertex_ai_grounding_metadata"] == grounding_metadata
assert (
response._hidden_params["vertex_ai_url_context_metadata"]
== url_context_metadata
)
dumped = response.model_dump()
assert dumped["vertex_ai_grounding_metadata"] == grounding_metadata
assert dumped["vertex_ai_url_context_metadata"] == url_context_metadata
def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_dict_chunks():
"""Dict snapshot chunks (model_dump) should also propagate Vertex AI metadata."""
chunk_dict = ModelResponseStream(
id="chatcmpl-vertex-2",
created=1,
model="gemini-2.5-flash",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="hello", role="assistant"),
)
],
).model_dump()
chunk_dict["vertex_ai_grounding_metadata"] = [{"webSearchQueries": ["test query"]}]
chunk_dict["_hidden_params"] = {
"vertex_ai_grounding_metadata": [{"webSearchQueries": ["test query"]}]
}
response = stream_chunk_builder(chunks=[chunk_dict])
assert response is not None
assert getattr(response, "vertex_ai_grounding_metadata") == [
{"webSearchQueries": ["test query"]}
]
assert response.model_dump()["vertex_ai_grounding_metadata"] == [
{"webSearchQueries": ["test query"]}
]