mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
parent
15507be90a
commit
8702ccfad2
3 changed files with 54 additions and 0 deletions
|
|
@ -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"] = (
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue