fix(vertex): mirror vertex_ai_safety_results on assembled streaming responses

The non-streaming transform_response stores safety data under
vertex_ai_safety_results, but the streaming path only wrote
vertex_ai_safety_ratings. Assembled streaming responses therefore never
carried vertex_ai_safety_results, so any consumer reading that field saw
a silent difference between streaming and non-streaming calls.

Set vertex_ai_safety_results alongside vertex_ai_safety_ratings in the
shared stream metadata setter and add it to the assembled metadata field
list so it propagates through stream_chunk_builder.
This commit is contained in:
mateo-berri 2026-06-08 20:45:59 +00:00
parent 15507be90a
commit 8702ccfad2
No known key found for this signature in database
3 changed files with 54 additions and 0 deletions

View file

@ -2257,6 +2257,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"vertex_ai_grounding_metadata",
"vertex_ai_url_context_metadata",
"vertex_ai_safety_ratings",
"vertex_ai_safety_results",
"vertex_ai_citation_metadata",
)
@ -2296,8 +2297,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
url_context_metadata
)
setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore
setattr(model_response, "vertex_ai_safety_results", safety_ratings) # type: ignore
if safety_ratings:
model_response._hidden_params["vertex_ai_safety_ratings"] = safety_ratings
model_response._hidden_params["vertex_ai_safety_results"] = safety_ratings
setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore
if citation_metadata:
model_response._hidden_params["vertex_ai_citation_metadata"] = (

View file

@ -705,6 +705,37 @@ def test_stream_chunk_builder_uses_assembled_model_for_provider_metadata():
assert getattr(response, "vertex_ai_grounding_metadata") == grounding_metadata
def test_stream_chunk_builder_propagates_vertex_ai_safety_results():
"""Assembled response must expose safety data under the non-streaming field name."""
safety_ratings = [
[{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}]
]
chunk = ModelResponseStream(
id="chatcmpl-vertex-safety",
created=1,
model="gemini-2.5-flash",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="hello", role="assistant"),
)
],
)
setattr(chunk, "vertex_ai_safety_ratings", safety_ratings)
setattr(chunk, "vertex_ai_safety_results", safety_ratings)
chunk._hidden_params["vertex_ai_safety_ratings"] = safety_ratings
chunk._hidden_params["vertex_ai_safety_results"] = safety_ratings
response = stream_chunk_builder(chunks=[chunk])
assert response is not None
assert getattr(response, "vertex_ai_safety_results") == safety_ratings
assert response._hidden_params["vertex_ai_safety_results"] == safety_ratings
assert response.model_dump()["vertex_ai_safety_results"] == safety_ratings
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(

View file

@ -1459,6 +1459,26 @@ def test_vertex_ai_process_candidates_with_grounding_metadata():
assert len(result[0]) == 1
def test_set_stream_metadata_mirrors_non_streaming_safety_field_names():
safety_ratings = [
[{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}]
]
model_response = ModelResponse()
VertexGeminiConfig._set_stream_metadata_on_response(
model_response=model_response,
grounding_metadata=[],
url_context_metadata=[],
safety_ratings=safety_ratings,
citation_metadata=[],
)
assert getattr(model_response, "vertex_ai_safety_ratings") == safety_ratings
assert getattr(model_response, "vertex_ai_safety_results") == safety_ratings
assert model_response._hidden_params["vertex_ai_safety_ratings"] == safety_ratings
assert model_response._hidden_params["vertex_ai_safety_results"] == safety_ratings
def test_vertex_ai_tool_call_id_format():
"""
Test that tool call IDs have the correct format and length.