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