refactor(vertex): share single Vertex metadata field tuple across redaction and streaming

This commit is contained in:
mateo-berri 2026-06-08 21:13:14 +00:00
parent c3bcd7630c
commit 514c057772
No known key found for this signature in database
3 changed files with 17 additions and 22 deletions

View file

@ -18,6 +18,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
from litellm.secret_managers.main import str_to_bool
from litellm.types.llms.vertex_ai import VERTEX_AI_PROVIDER_METADATA_FIELDS
from litellm.types.utils import StandardCallbackDynamicParams
if TYPE_CHECKING:
@ -29,14 +30,6 @@ if TYPE_CHECKING:
else:
LiteLLMLoggingObject = Any
VERTEX_PROVIDER_METADATA_FIELDS = (
"vertex_ai_grounding_metadata",
"vertex_ai_url_context_metadata",
"vertex_ai_safety_ratings",
"vertex_ai_safety_results",
"vertex_ai_citation_metadata",
)
def redact_message_input_output_from_custom_logger(
litellm_logging_obj: LiteLLMLoggingObject, result, custom_logger: CustomLogger
@ -110,21 +103,21 @@ def _redact_responses_api_output_dict(output_items, redacted_str: str):
def _redact_vertex_provider_metadata(obj: Any) -> None:
if isinstance(obj, dict):
for field in VERTEX_PROVIDER_METADATA_FIELDS:
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
if field in obj:
obj[field] = []
hidden_params = obj.get("_hidden_params")
if isinstance(hidden_params, dict):
for field in VERTEX_PROVIDER_METADATA_FIELDS:
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
hidden_params.pop(field, None)
return
for field in VERTEX_PROVIDER_METADATA_FIELDS:
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
if hasattr(obj, field):
setattr(obj, field, [])
hidden_params = getattr(obj, "_hidden_params", None)
if isinstance(hidden_params, dict):
for field in VERTEX_PROVIDER_METADATA_FIELDS:
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
hidden_params.pop(field, None)
@ -148,7 +141,7 @@ def _redact_vertex_provider_metadata_from_litellm_params(
hidden_params = metadata.get("hidden_params")
if not isinstance(hidden_params, dict):
continue
for field in VERTEX_PROVIDER_METADATA_FIELDS:
for field in VERTEX_AI_PROVIDER_METADATA_FIELDS:
hidden_params.pop(field, None)

View file

@ -63,6 +63,7 @@ from litellm.types.llms.openai import (
OpenAIChatCompletionFinishReason,
)
from litellm.types.llms.vertex_ai import (
VERTEX_AI_PROVIDER_METADATA_FIELDS,
VERTEX_CREDENTIALS_TYPES,
Candidates,
ContentType,
@ -2253,14 +2254,6 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
citation_metadata,
)
_STREAM_METADATA_FIELDS = (
"vertex_ai_grounding_metadata",
"vertex_ai_url_context_metadata",
"vertex_ai_safety_ratings",
"vertex_ai_safety_results",
"vertex_ai_citation_metadata",
)
@staticmethod
def _get_stream_chunk_attr(chunk: Any, field_name: str) -> Any:
if isinstance(chunk, dict):
@ -2312,7 +2305,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
response: ModelResponse,
chunks: List[Any],
) -> None:
for field_name in VertexGeminiConfig._STREAM_METADATA_FIELDS:
for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS:
merged: List[Any] = []
for chunk in chunks:
value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name)

View file

@ -757,3 +757,12 @@ class VertexPartnerProvider(str, Enum):
llama = "llama"
ai21 = "ai21"
claude = "claude"
VERTEX_AI_PROVIDER_METADATA_FIELDS = (
"vertex_ai_grounding_metadata",
"vertex_ai_url_context_metadata",
"vertex_ai_safety_ratings",
"vertex_ai_safety_results",
"vertex_ai_citation_metadata",
)