mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
aaf1e2444b
commit
d791df59b6
5 changed files with 171 additions and 0 deletions
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]}
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue