refactor(vertex): move streaming metadata merge into provider config hook

Address review feedback by delegating assembled-stream metadata propagation
to VertexGeminiConfig via BaseConfig.apply_assembled_streaming_response_metadata,
and only write chunk hidden_params when metadata is non-empty.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-06-08 10:06:59 +05:30
parent d791df59b6
commit 1add616c3e
No known key found for this signature in database
4 changed files with 130 additions and 53 deletions

View file

@ -32,14 +32,6 @@ 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)
@ -88,40 +80,53 @@ class ChunkProcessor:
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]
def apply_provider_assembled_streaming_metadata(
response: ModelResponse,
chunks: List[Any],
logging_obj: Optional[Any] = None,
) -> None:
"""
Merge Vertex AI metadata from streaming chunks into the assembled response.
if not chunks:
return
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
first_chunk = chunks[0]
model = (
first_chunk.get("model")
if isinstance(first_chunk, dict)
else getattr(first_chunk, "model", None)
)
if not model:
return
custom_llm_provider = None
if logging_obj is not None:
custom_llm_provider = logging_obj.model_call_details.get(
"custom_llm_provider"
)
try:
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
if custom_llm_provider:
provider = LlmProviders(custom_llm_provider)
else:
_, provider_str, _, _ = get_llm_provider(model)
provider = LlmProviders(provider_str)
provider_config = ProviderConfigManager.get_provider_chat_config(
model=model,
provider=provider,
)
if provider_config is not None:
provider_config.apply_assembled_streaming_response_metadata(
response=response,
chunks=chunks,
)
except Exception:
return
@staticmethod
def _get_chunk_id(chunks: List[Dict[str, Any]]) -> str:

View file

@ -442,6 +442,14 @@ class BaseConfig(ABC):
"""Hook for providers to post-process streaming responses. Default: pass-through."""
return stream
def apply_assembled_streaming_response_metadata(
self,
response: "ModelResponse",
chunks: List[Any],
) -> None:
"""Hook for providers to merge chunk metadata into assembled streaming responses."""
return None
def calculate_additional_costs(
self, model: str, prompt_tokens: int, completion_tokens: int
) -> Optional[dict]:

View file

@ -2253,6 +2253,71 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
citation_metadata,
)
_STREAM_METADATA_FIELDS = (
"vertex_ai_grounding_metadata",
"vertex_ai_url_context_metadata",
"vertex_ai_safety_ratings",
"vertex_ai_citation_metadata",
)
@staticmethod
def _get_stream_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 _set_stream_metadata_on_response(
model_response: Any,
grounding_metadata: List[dict],
url_context_metadata: List[dict],
safety_ratings: List[dict],
citation_metadata: List[dict],
) -> None:
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
if grounding_metadata:
model_response._hidden_params["vertex_ai_grounding_metadata"] = (
grounding_metadata
)
setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore
if url_context_metadata:
model_response._hidden_params["vertex_ai_url_context_metadata"] = (
url_context_metadata
)
setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore
if safety_ratings:
model_response._hidden_params["vertex_ai_safety_ratings"] = safety_ratings
setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore
if citation_metadata:
model_response._hidden_params["vertex_ai_citation_metadata"] = (
citation_metadata
)
def apply_assembled_streaming_response_metadata(
self,
response: ModelResponse,
chunks: List[Any],
) -> None:
for field_name in VertexGeminiConfig._STREAM_METADATA_FIELDS:
merged: List[Any] = []
for chunk in chunks:
value = VertexGeminiConfig._get_stream_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 _convert_grounding_metadata_to_annotations(
grounding_metadata: List[dict],
@ -3385,18 +3450,13 @@ class ModelResponseIterator:
if choice.finish_reason == "stop":
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
VertexGeminiConfig._set_stream_metadata_on_response(
model_response,
grounding_metadata,
url_context_metadata,
safety_ratings,
citation_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,7 +7761,9 @@ def stream_chunk_builder( # noqa: PLR0915
"cost",
logging_obj._response_cost_calculator(result=response),
)
processor.propagate_vertex_ai_metadata_from_chunks(response, chunks)
processor.apply_provider_assembled_streaming_metadata(
response, chunks, logging_obj
)
return response
tool_call_chunks = [
@ -7941,7 +7943,9 @@ def stream_chunk_builder( # noqa: PLR0915
usage, "cost", logging_obj._response_cost_calculator(result=response)
)
processor.propagate_vertex_ai_metadata_from_chunks(response, chunks)
processor.apply_provider_assembled_streaming_metadata(
response, chunks, logging_obj
)
return response
except Exception as e:
verbose_logger.exception(