diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index 99e610d69ad..f1de30fa72d 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -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: diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 5f35a58ce1f..8f9d5cad7c4 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -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]: diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index d6436f744c4..0a7ba70dacc 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -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, diff --git a/litellm/main.py b/litellm/main.py index c7f0b4a2091..1a0d0312d73 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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(