diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index dbaf6b59465..a2130c50706 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -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) 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 2d996567058..c919a367fbf 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 @@ -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) diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 51429d0769e..ea5d4471e82 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -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", +)