From 0b0e3ed8bd6e414d8af891b420b0874994562e74 Mon Sep 17 00:00:00 2001 From: Yaniv Israel Date: Thu, 2 Jul 2026 15:10:22 +0300 Subject: [PATCH] fix(coverage): revert ruff UP006/UP045 changes on upstream files MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The previous ruff fixes (Dict→dict, Optional→X|None) on 7 upstream files added ~500 changed lines of pure type-annotation no-ops to our PR diff. codecov/patch penalised these uncovered lines, dropping patch coverage to 51.35% (target 61.83%). Fix: revert these files to exactly match upstream/litellm_internal_staging. The ruff_strict_gate still passes because the violations exist equally in both the base and HEAD (total == base_count → no breach). --- .../anthropic_cache_control_hook.py | 114 ++--- litellm/integrations/custom_guardrail.py | 137 +++--- litellm/integrations/otel/plumbing/metrics.py | 8 +- litellm/llms/anthropic/chat/transformation.py | 288 ++++++------ .../llms/modelscope/chat/transformation.py | 12 +- .../vertex_and_google_ai_studio_gemini.py | 422 +++++++++--------- .../proxy/guardrails/guardrail_registry.py | 54 +-- 7 files changed, 526 insertions(+), 509 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index b58ebaa7fa6..608fdebc1d9 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -10,7 +10,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and """ import copy -from typing import TYPE_CHECKING, Any, Union, cast +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger @@ -39,27 +39,27 @@ class AnthropicCacheControlHook(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: list[AllMessageValues], + messages: List[AllMessageValues], non_default_params: dict, - prompt_id: str | None, - prompt_variables: dict | None, + prompt_id: Optional[str], + prompt_variables: Optional[dict], dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: PromptSpec | None = None, - prompt_label: str | None = None, - prompt_version: int | None = None, - ignore_prompt_manager_model: bool | None = False, - ignore_prompt_manager_optional_params: bool | None = False, - ) -> tuple[str, list[AllMessageValues], dict]: + prompt_spec: Optional[PromptSpec] = None, + prompt_label: Optional[str] = None, + prompt_version: Optional[int] = None, + ignore_prompt_manager_model: Optional[bool] = False, + ignore_prompt_manager_optional_params: Optional[bool] = False, + ) -> Tuple[str, List[AllMessageValues], dict]: """ Apply cache control directives based on specified injection points. Returns: - model: str - the model to use - - messages: list[AllMessageValues] - messages with applied cache controls + - messages: List[AllMessageValues] - messages with applied cache controls - non_default_params: dict - params with any global cache controls """ # Extract cache control injection points - injection_points: list[CacheControlInjectionPoint] = non_default_params.pop( + injection_points: List[CacheControlInjectionPoint] = non_default_params.pop( "cache_control_injection_points", [] ) if not injection_points: @@ -69,8 +69,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): processed_messages = copy.deepcopy(messages) # Separate message-level and non-message-level injection points - message_points: list[CacheControlMessageInjectionPoint] = [] - remaining_points: list[CacheControlInjectionPoint] = [] + message_points: List[CacheControlMessageInjectionPoint] = [] + remaining_points: List[CacheControlInjectionPoint] = [] for point in injection_points: if point.get("location") == "message": message_points.append(cast(CacheControlMessageInjectionPoint, point)) @@ -97,10 +97,10 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def _apply_message_injections( - points: list[CacheControlMessageInjectionPoint], - messages: list[AllMessageValues], + points: List[CacheControlMessageInjectionPoint], + messages: List[AllMessageValues], max_blocks: int, - ) -> list[AllMessageValues]: + ) -> List[AllMessageValues]: """Apply message-level cache control injection points in order. Anthropic allows at most ``MAX_CACHE_CONTROL_BLOCKS`` cache_control @@ -149,11 +149,11 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def _resolve_target_indices( - point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues] - ) -> list[int]: + point: CacheControlMessageInjectionPoint, messages: List[AllMessageValues] + ) -> List[int]: """Resolve which message indices an injection point targets.""" - _targetted_index: Union[int, str] | None = point.get("index", None) - targetted_index: int | None = None + _targetted_index: Optional[Union[int, str]] = point.get("index", None) + targetted_index: Optional[int] = None if isinstance(_targetted_index, str): try: targetted_index = int(_targetted_index) @@ -230,10 +230,10 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def apply_to_anthropic_messages_request( - messages: list[dict], + messages: List[Dict], system: str | list | None, - injection_points: list[CacheControlInjectionPoint], - ) -> tuple[list[dict], str | list | None, list[CacheControlInjectionPoint]]: + injection_points: List[CacheControlInjectionPoint], + ) -> Tuple[List[Dict], str | list | None, List[CacheControlInjectionPoint]]: """Apply cache control injection for the Anthropic-native v1/messages endpoint. Returns (messages, system, remaining_non_message_points). @@ -241,12 +241,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): if not injection_points: return messages, system, [] - processed_messages: list[dict] = copy.deepcopy(messages) + processed_messages: List[Dict] = copy.deepcopy(messages) processed_system = copy.deepcopy(system) if system is not None else None - message_points: list[CacheControlMessageInjectionPoint] = [] - system_points: list[CacheControlMessageInjectionPoint] = [] - remaining_points: list[CacheControlInjectionPoint] = [] + message_points: List[CacheControlMessageInjectionPoint] = [] + system_points: List[CacheControlMessageInjectionPoint] = [] + remaining_points: List[CacheControlInjectionPoint] = [] for point in injection_points: if point.get("location") == "message": @@ -290,7 +290,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): processed_messages = AnthropicCacheControlHook._apply_message_injections( points=message_points, - messages=cast(list[AllMessageValues], processed_messages), + messages=cast(List[AllMessageValues], processed_messages), max_blocks=max_blocks - used_blocks, ) @@ -298,10 +298,10 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def maybe_inject_cache_control( - messages: list[dict], + messages: List[Dict], system: str | list | None, - kwargs: dict[str, Any], - ) -> tuple[list[dict], str | list | None]: + kwargs: Dict[str, Any], + ) -> Tuple[List[Dict], str | list | None]: """Extract cache_control_injection_points from kwargs and apply if present. Pops the key from kwargs; if remaining (non-message) points exist they @@ -327,8 +327,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: str | None, - prompt_spec: PromptSpec | None, + prompt_id: Optional[str], + prompt_spec: Optional[PromptSpec], dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """Always return False since this is not a true prompt management system.""" @@ -336,12 +336,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: str | None, - prompt_spec: PromptSpec | None, - prompt_variables: dict | None, + prompt_id: Optional[str], + prompt_spec: Optional[PromptSpec], + prompt_variables: Optional[dict], dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: str | None = None, - prompt_version: int | None = None, + prompt_label: Optional[str] = None, + prompt_version: Optional[int] = None, ) -> PromptManagementClient: """Not used - this hook only modifies messages, doesn't fetch prompts.""" return PromptManagementClient( @@ -354,12 +354,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: str | None, - prompt_variables: dict | None, + prompt_id: Optional[str], + prompt_variables: Optional[dict], dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: PromptSpec | None = None, - prompt_label: str | None = None, - prompt_version: int | None = None, + prompt_spec: Optional[PromptSpec] = None, + prompt_label: Optional[str] = None, + prompt_version: Optional[int] = None, ) -> PromptManagementClient: """Not used - this hook only modifies messages, doesn't fetch prompts.""" return self._compile_prompt_helper( @@ -374,19 +374,19 @@ class AnthropicCacheControlHook(CustomPromptManagement): async def async_get_chat_completion_prompt( self, model: str, - messages: list[AllMessageValues], + messages: List[AllMessageValues], non_default_params: dict, - prompt_id: str | None, - prompt_variables: dict | None, + prompt_id: Optional[str], + prompt_variables: Optional[dict], dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - prompt_spec: PromptSpec | None = None, - tools: list[dict] | None = None, - prompt_label: str | None = None, - prompt_version: int | None = None, - ignore_prompt_manager_model: bool | None = False, - ignore_prompt_manager_optional_params: bool | None = False, - ) -> tuple[str, list[AllMessageValues], dict]: + prompt_spec: Optional[PromptSpec] = None, + tools: Optional[List[Dict]] = None, + prompt_label: Optional[str] = None, + prompt_version: Optional[int] = None, + ignore_prompt_manager_model: Optional[bool] = False, + ignore_prompt_manager_optional_params: Optional[bool] = False, + ) -> Tuple[str, List[AllMessageValues], dict]: """Async version - delegates to sync since no async operations needed.""" return self.get_chat_completion_prompt( model=model, @@ -403,15 +403,15 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) @staticmethod - def should_use_anthropic_cache_control_hook(non_default_params: dict) -> bool: + def should_use_anthropic_cache_control_hook(non_default_params: Dict) -> bool: if non_default_params.get("cache_control_injection_points", None): return True return False @staticmethod def get_custom_logger_for_anthropic_cache_control_hook( - non_default_params: dict, - ) -> CustomLogger | None: + non_default_params: Dict, + ) -> Optional[CustomLogger]: from litellm.litellm_core_utils.litellm_logging import ( _init_custom_logger_compatible_class, ) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 353ece8da29..59d37639098 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -4,8 +4,11 @@ from typing import ( TYPE_CHECKING, Any, ClassVar, + Dict, + List, Literal, Optional, + Type, Union, get_args, ) @@ -56,7 +59,7 @@ from litellm.exceptions import ( _PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16) -def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None: +def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]: """Extract session_id from request data (litellm_session_id or metadata).""" session_id = request_data.get("litellm_session_id") if session_id: @@ -81,18 +84,18 @@ class CustomGuardrail(CustomLogger): def __init__( self, - guardrail_name: str | None = None, - supported_event_hooks: list[GuardrailEventHooks] | None = None, - event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = None, + guardrail_name: Optional[str] = None, + supported_event_hooks: Optional[List[GuardrailEventHooks]] = None, + event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = None, default_on: bool = False, mask_request_content: bool = False, mask_response_content: bool = False, - violation_message_template: str | None = None, - end_session_after_n_fails: int | None = None, - on_violation: str | None = None, - realtime_violation_message: str | None = None, - on_sensitive_data: str | None = None, - sensitive_data_route_to_model: str | None = None, + violation_message_template: Optional[str] = None, + end_session_after_n_fails: Optional[int] = None, + on_violation: Optional[str] = None, + realtime_violation_message: Optional[str] = None, + on_sensitive_data: Optional[str] = None, + sensitive_data_route_to_model: Optional[str] = None, sticky_session_routing: bool = True, **kwargs, ): @@ -115,16 +118,16 @@ class CustomGuardrail(CustomLogger): """ self.guardrail_name = guardrail_name self.supported_event_hooks = supported_event_hooks - self.event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = event_hook + self.event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]] = event_hook self.default_on: bool = default_on self.mask_request_content: bool = mask_request_content self.mask_response_content: bool = mask_response_content - self.violation_message_template: str | None = violation_message_template - self.end_session_after_n_fails: int | None = end_session_after_n_fails - self.on_violation: str | None = on_violation - self.realtime_violation_message: str | None = realtime_violation_message - self.on_sensitive_data: str | None = on_sensitive_data - self.sensitive_data_route_to_model: str | None = sensitive_data_route_to_model + self.violation_message_template: Optional[str] = violation_message_template + self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails + self.on_violation: Optional[str] = on_violation + self.realtime_violation_message: Optional[str] = realtime_violation_message + self.on_sensitive_data: Optional[str] = on_sensitive_data + self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model self.sticky_session_routing: bool = sticky_session_routing if supported_event_hooks: @@ -132,13 +135,13 @@ class CustomGuardrail(CustomLogger): self._validate_event_hook(event_hook, supported_event_hooks) super().__init__(**kwargs) - def render_violation_message(self, default: str, context: dict[str, Any] | None = None) -> str: + def render_violation_message(self, default: str, context: Optional[Dict[str, Any]] = None) -> str: """Return a custom violation message if template is configured.""" if not self.violation_message_template: return default - format_context: dict[str, Any] = {"default_message": default} + format_context: Dict[str, Any] = {"default_message": default} if context: format_context.update(context) try: @@ -154,8 +157,8 @@ class CustomGuardrail(CustomLogger): def raise_passthrough_exception( self, violation_message: str, - request_data: dict[str, Any], - detection_info: dict[str, Any] | None = None, + request_data: Dict[str, Any], + detection_info: Optional[Dict[str, Any]] = None, ) -> None: """ Raise a passthrough exception for guardrail violations. @@ -198,8 +201,8 @@ class CustomGuardrail(CustomLogger): def raise_sensitive_data_route_exception( self, route_to_model: str, - request_data: dict[str, Any], - detection_info: dict[str, Any] | None = None, + request_data: Dict[str, Any], + detection_info: Optional[Dict[str, Any]] = None, ) -> None: """ Raise an exception to reroute the request to a different model. @@ -236,7 +239,7 @@ class CustomGuardrail(CustomLogger): sticky_session_routing=self.sticky_session_routing, ) - def _get_session_id_from_request_data(self, request_data: dict[str, Any]) -> str | None: + def _get_session_id_from_request_data(self, request_data: Dict[str, Any]) -> Optional[str]: """Extract session_id from request data.""" return get_session_id_from_request_data(request_data) @@ -249,8 +252,8 @@ class CustomGuardrail(CustomLogger): def handle_sensitive_data_detection( self, - request_data: dict[str, Any], - detection_info: dict[str, Any] | None = None, + request_data: Dict[str, Any], + detection_info: Optional[Dict[str, Any]] = None, ) -> None: """ Handle sensitive data detection based on guardrail configuration. @@ -292,7 +295,7 @@ class CustomGuardrail(CustomLogger): ) @staticmethod - def get_config_model() -> type["GuardrailConfigModel"] | None: + def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: """ Returns the config model for the guardrail @@ -302,12 +305,12 @@ class CustomGuardrail(CustomLogger): def _validate_event_hook( self, - event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None, - supported_event_hooks: list[GuardrailEventHooks], + event_hook: Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode]], + supported_event_hooks: List[GuardrailEventHooks], ) -> None: def _validate_event_hook_list_is_in_supported_event_hooks( - event_hook: Union[list[GuardrailEventHooks], list[str]], - supported_event_hooks: list[GuardrailEventHooks], + event_hook: Union[List[GuardrailEventHooks], List[str]], + supported_event_hooks: List[GuardrailEventHooks], ) -> None: for hook in event_hook: if isinstance(hook, str): @@ -358,7 +361,7 @@ class CustomGuardrail(CustomLogger): key_meta = meta.get("user_api_key_metadata") or key_meta return {**team_meta, **key_meta} - def get_disable_global_guardrail(self, data: dict) -> bool | None: + def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: """ Returns True if the global guardrail should be disabled. @@ -367,7 +370,7 @@ class CustomGuardrail(CustomLogger): """ return self._get_admin_metadata(data).get("disable_global_guardrails", False) - def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> list[str]: + def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]: """ Returns the list of global guardrail names the team/key has opted out of. @@ -398,7 +401,7 @@ class CustomGuardrail(CustomLogger): return True raise - def get_guardrail_from_metadata(self, data: dict) -> Union[list[str], list[dict[str, DynamicGuardrailParams]]]: + def get_guardrail_from_metadata(self, data: dict) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]: """ Returns the guardrail(s) to be run from the metadata or root """ @@ -419,7 +422,7 @@ class CustomGuardrail(CustomLogger): def _guardrail_is_in_requested_guardrails( self, - requested_guardrails: Union[list[str], list[dict[str, DynamicGuardrailParams]]], + requested_guardrails: Union[List[str], List[Dict[str, DynamicGuardrailParams]]], ) -> bool: for _guardrail in requested_guardrails: if isinstance(_guardrail, dict): @@ -431,13 +434,13 @@ class CustomGuardrail(CustomLogger): return False - def _pre_call_marker(self) -> str | None: + def _pre_call_marker(self) -> Optional[str]: name = self.guardrail_name if not name: return None return f"{_PRE_CALL_EXECUTED_TOKEN}:{name}" - def mark_pre_call_hook_ran(self, data: dict[str, Any]) -> None: + def mark_pre_call_hook_ran(self, data: Dict[str, Any]) -> None: """ Record that this guardrail's ``async_pre_call_hook`` already ran for this request, so the deployment-level hook does not run it a second time. @@ -462,7 +465,7 @@ class CustomGuardrail(CustomLogger): return data["metadata"] = {PRE_CALL_EXECUTED_GUARDRAILS_KEY: [marker]} - def _pre_call_hook_already_ran(self, data: dict[str, Any]) -> bool: + def _pre_call_hook_already_ran(self, data: Dict[str, Any]) -> bool: marker = self._pre_call_marker() if marker is None: return False @@ -474,7 +477,9 @@ class CustomGuardrail(CustomLogger): return True return False - async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: + async def async_pre_call_deployment_hook( + self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] + ) -> Optional[dict]: from litellm.proxy._types import UserAPIKeyAuth # should run guardrail @@ -514,8 +519,8 @@ class CustomGuardrail(CustomLogger): self, request_data: dict, response: LLMResponseTypes, - call_type: CallTypes | None, - ) -> LLMResponseTypes | None: + call_type: Optional[CallTypes], + ) -> Optional[LLMResponseTypes]: """ Allow modifying / reviewing the response just after it's received from the deployment. """ @@ -697,16 +702,16 @@ class CustomGuardrail(CustomLogger): def add_standard_logging_guardrail_information_to_request_data( self, - guardrail_json_response: Union[Exception, str, dict, list[dict]], + guardrail_json_response: Union[Exception, str, dict, List[dict]], request_data: dict, guardrail_status: GuardrailStatus, - start_time: float | None = None, - end_time: float | None = None, - duration: float | None = None, - masked_entity_count: dict[str, int] | None = None, - guardrail_provider: str | None = None, - event_type: GuardrailEventHooks | None = None, - tracing_detail: GuardrailTracingDetail | None = None, + start_time: Optional[float] = None, + end_time: Optional[float] = None, + duration: Optional[float] = None, + masked_entity_count: Optional[Dict[str, int]] = None, + guardrail_provider: Optional[str] = None, + event_type: Optional[GuardrailEventHooks] = None, + tracing_detail: Optional[GuardrailTracingDetail] = None, ) -> None: """ Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc. @@ -721,7 +726,7 @@ class CustomGuardrail(CustomLogger): from litellm.types.utils import GuardrailMode # Use event_type if provided, otherwise fall back to self.event_hook - guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, list[GuardrailEventHooks]] + guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]] if event_type is not None: guardrail_mode = event_type elif isinstance(self.event_hook, Mode): @@ -830,13 +835,13 @@ class CustomGuardrail(CustomLogger): def _process_response( self, - response: dict | None, + response: Optional[Dict], request_data: dict, - start_time: float | None = None, - end_time: float | None = None, - duration: float | None = None, - event_type: GuardrailEventHooks | None = None, - original_inputs: dict | None = None, + start_time: Optional[float] = None, + end_time: Optional[float] = None, + duration: Optional[float] = None, + event_type: Optional[GuardrailEventHooks] = None, + original_inputs: Optional[Dict] = None, ): """ Add StandardLoggingGuardrailInformation to the request data @@ -844,7 +849,7 @@ class CustomGuardrail(CustomLogger): This gets logged on downsteam Langfuse, DataDog, etc. """ # Convert None to empty dict to satisfy type requirements - guardrail_response: Union[dict[str, Any], str] = {} if response is None else response + guardrail_response: Union[Dict[str, Any], str] = {} if response is None else response # For apply_guardrail functions in custom_code_guardrail scenario, # simplify the logged response to "allow", "deny", or "mask" @@ -900,10 +905,10 @@ class CustomGuardrail(CustomLogger): self, e: Exception, request_data: dict, - start_time: float | None = None, - end_time: float | None = None, - duration: float | None = None, - event_type: GuardrailEventHooks | None = None, + start_time: Optional[float] = None, + end_time: Optional[float] = None, + duration: Optional[float] = None, + event_type: Optional[GuardrailEventHooks] = None, ): """ Add StandardLoggingGuardrailInformation to the request data @@ -930,7 +935,7 @@ class CustomGuardrail(CustomLogger): ) raise e - def _inputs_were_modified(self, original_inputs: dict, response: dict) -> bool: + def _inputs_were_modified(self, original_inputs: Dict, response: Dict) -> bool: """ Compare original inputs with response to determine if content was modified. @@ -975,8 +980,8 @@ class CustomGuardrail(CustomLogger): setattr(self, key, value) def get_guardrails_messages_for_call_type( - self, call_type: CallTypes, data: dict | None = None - ) -> list[AllMessageValues] | None: + self, call_type: CallTypes, data: Optional[dict] = None + ) -> Optional[List[AllMessageValues]]: """ Returns the messages for the given call type and data """ @@ -1014,7 +1019,7 @@ class CustomGuardrail(CustomLogger): input=input_data, responses_api_request=data, ) - return cast(list[AllMessageValues], messages) + return cast(List[AllMessageValues], messages) return None @@ -1078,7 +1083,7 @@ def log_guardrail_information(func): def _infer_event_type_from_function_name( func_name: str, - ) -> GuardrailEventHooks | None: + ) -> Optional[GuardrailEventHooks]: """Infer the actual event type from the function name""" if func_name == "async_pre_call_hook": return GuardrailEventHooks.pre_call diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index 56f064f8b72..cb1f9214876 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -10,7 +10,7 @@ identical metrics. The attribute cardinality filter is reused from v1 by import from dataclasses import dataclass from datetime import datetime -from typing import Any, Mapping +from typing import Any, FrozenSet, Mapping, Optional from opentelemetry.metrics import Histogram, Meter @@ -81,11 +81,11 @@ class GenAIMetricRecorder: survives. """ - def __init__(self, metrics: GenAIMetrics, callback_name: str | None = None) -> None: + def __init__(self, metrics: GenAIMetrics, callback_name: Optional[str] = None) -> None: self._metrics = metrics self._callback_name = callback_name - self._include: frozenset[str] | None = None - self._exclude: frozenset[str] | None = None + self._include: Optional[FrozenSet[str]] = None + self._exclude: Optional[FrozenSet[str]] = None self._filter_resolved = False def record( diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index d1a7d1936e9..9721b797584 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -4,7 +4,11 @@ import time from typing import ( TYPE_CHECKING, Any, + Dict, + List, NoReturn, + Optional, + Tuple, Union, cast, ) @@ -146,8 +150,8 @@ def _basic_sanitize_anthropic_tool_name(name: str) -> str: def _build_anthropic_tool_name_maps( - original_names: list[str], -) -> tuple[dict[str, str], dict[str, str]]: + original_names: List[str], +) -> Tuple[Dict[str, str], Dict[str, str]]: """Build (forward, reverse) tool-name maps for a single request. forward[original] = sanitized -- only present when name was rewritten @@ -167,7 +171,7 @@ def _build_anthropic_tool_name_maps( seen gets the disambiguating suffix. Callers should preserve the caller's tool order (we do). """ - forward: dict[str, str] = {} + forward: Dict[str, str] = {} used: set = set() # First pass: reserve slots for names that are already valid so they @@ -208,7 +212,7 @@ def _build_anthropic_tool_name_maps( return forward, reverse -REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: dict[str, str] = { +REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT: Dict[str, str] = { "low": "low", "minimal": "low", "medium": "medium", @@ -235,23 +239,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): to pass metadata to anthropic, it's {"user_id": "any-relevant-information"} """ - max_tokens: int | None = None - stop_sequences: list | None = None - temperature: int | None = None - top_p: int | None = None - top_k: int | None = None - metadata: dict | None = None - system: str | None = None + max_tokens: Optional[int] = None + stop_sequences: Optional[list] = None + temperature: Optional[int] = None + top_p: Optional[int] = None + top_k: Optional[int] = None + metadata: Optional[dict] = None + system: Optional[str] = None def __init__( self, - max_tokens: int | None = None, - stop_sequences: list | None = None, - temperature: int | None = None, - top_p: int | None = None, - top_k: int | None = None, - metadata: dict | None = None, - system: str | None = None, + max_tokens: Optional[int] = None, + stop_sequences: Optional[list] = None, + temperature: Optional[int] = None, + top_p: Optional[int] = None, + top_k: Optional[int] = None, + metadata: Optional[dict] = None, + system: Optional[str] = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -259,11 +263,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> str | None: + def custom_llm_provider(self) -> Optional[str]: return "anthropic" @classmethod - def get_config(cls, *, model: str | None = None): + def get_config(cls, *, model: Optional[str] = None): config = super().get_config() # anthropic requires a default value for max_tokens @@ -273,7 +277,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return config @staticmethod - def get_max_tokens_for_model(model: str | None = None) -> int: + def get_max_tokens_for_model(model: Optional[str] = None) -> int: """ Get the max output tokens for a given model. Falls back to DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS (configurable via env var) if model is not found. @@ -290,7 +294,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def convert_tool_use_to_openai_format( - anthropic_tool_content: dict[str, Any], + anthropic_tool_content: Dict[str, Any], index: int, ) -> ChatCompletionToolCallChunk: """ @@ -315,7 +319,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) # Include caller information if present (for programmatic tool calling) if "caller" in anthropic_tool_content: - tool_call["caller"] = cast(dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item] + tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item] return tool_call @staticmethod @@ -336,7 +340,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return AnthropicConfig._supports_model_capability(model, f"supports_{level}_reasoning_effort") @staticmethod - def _validate_effort_for_model(model: str, effort: str | None) -> str | None: + def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]: """Return ``None`` if ``effort`` is allowed on ``model``, else an error message.""" if effort == "max" and not ( AnthropicConfig._is_adaptive_thinking_model(model) or AnthropicConfig._supports_effort_level(model, "max") @@ -363,7 +367,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) @staticmethod - def _model_supports_speed_param(model: str, custom_llm_provider: str | None = None) -> bool: + def _model_supports_speed_param(model: str, custom_llm_provider: Optional[str] = None) -> bool: """Whether the model accepts Anthropic's ``speed`` parameter (fast mode). Fast mode is direct Anthropic API-only (not Bedrock, Vertex, or Azure). @@ -380,7 +384,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): model: str, optional_params: dict, drop_params: bool, - custom_llm_provider: str | None = None, + custom_llm_provider: Optional[str] = None, ) -> None: if "speed" not in optional_params: return @@ -459,7 +463,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return params @staticmethod - def filter_anthropic_output_schema(schema: dict[str, Any]) -> dict[str, Any]: + def filter_anthropic_output_schema(schema: Dict[str, Any]) -> Dict[str, Any]: """ Filter out unsupported fields from JSON schema for Anthropic's output_format API. @@ -515,7 +519,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if field in schema: constraint_descriptions.append(constraint_labels[field].format(schema[field])) - result: dict[str, Any] = {} + result: Dict[str, Any] = {} # Update description with removed constraint info if constraint_descriptions: @@ -555,7 +559,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return result - def get_json_schema_from_pydantic_object(self, response_format: Union[Any, dict, None]) -> dict | None: + def get_json_schema_from_pydantic_object(self, response_format: Union[Any, Dict, None]) -> Optional[dict]: return type_to_response_format_param( response_format, ref_template="/$defs/{model}" ) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755 @@ -570,10 +574,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tool_choice( self, - tool_choice: str | None, - parallel_tool_use: bool | None, - ) -> AnthropicMessagesToolChoice | None: - _tool_choice: AnthropicMessagesToolChoice | None = None + tool_choice: Optional[str], + parallel_tool_use: Optional[bool], + ) -> Optional[AnthropicMessagesToolChoice]: + _tool_choice: Optional[AnthropicMessagesToolChoice] = None if tool_choice == "auto": _tool_choice = AnthropicMessagesToolChoice( type="auto", @@ -614,9 +618,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tool_helper( self, tool: ChatCompletionToolParam, - ) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]: - returned_tool: AllAnthropicToolsValues | None = None - mcp_server: AnthropicMcpServerTool | None = None + ) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]: + returned_tool: Optional[AllAnthropicToolsValues] = None + mcp_server: Optional[AnthropicMcpServerTool] = None if tool["type"] == "function" or tool["type"] == "custom": _input_schema: dict = tool["function"].get( @@ -667,8 +671,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if "parameters" not in tool["function"]: raise ValueError("Missing required parameter: parameters") - _display_width_px: int | None = tool["function"]["parameters"].get("display_width_px") - _display_height_px: int | None = tool["function"]["parameters"].get("display_height_px") + _display_width_px: Optional[int] = tool["function"]["parameters"].get("display_width_px") + _display_height_px: Optional[int] = tool["function"]["parameters"].get("display_height_px") if _display_width_px is None or _display_height_px is None: raise ValueError("Missing required parameter: display_width_px or display_height_px") @@ -829,14 +833,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration allowed_tools = tool.get("allowed_tools", None) - tool_configuration: AnthropicMcpServerToolConfiguration | None = None + tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None if allowed_tools is not None: tool_configuration = AnthropicMcpServerToolConfiguration( allowed_tools=tool.get("allowed_tools", None), ) headers = tool.get("headers", {}) - authorization_token: str | None = None + authorization_token: Optional[str] = None if headers is not None: bearer_token = headers.get("Authorization", None) if bearer_token is not None: @@ -856,8 +860,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tools( self, - tools: list, - ) -> tuple[list[AllAnthropicToolsValues], list[AnthropicMcpServerTool]]: + tools: List, + ) -> Tuple[List[AllAnthropicToolsValues], List[AnthropicMcpServerTool]]: anthropic_tools = [] mcp_servers = [] for tool in tools: @@ -902,9 +906,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _rewrite_tool_names_in_messages( - messages: list[AllMessageValues], - name_forward_map: dict[str, str], - ) -> list[AllMessageValues]: + messages: List[AllMessageValues], + name_forward_map: Dict[str, str], + ) -> List[AllMessageValues]: """Return a copy of `messages` with tool_call/function_call names rewritten using the per-request forward map. @@ -915,7 +919,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): """ if not name_forward_map: return messages - new_messages: list[AllMessageValues] = [] + new_messages: List[AllMessageValues] = [] for msg in messages: if not isinstance(msg, dict): new_messages.append(msg) @@ -953,8 +957,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _build_request_tool_name_maps( - tools: list, - ) -> tuple[dict[str, str], dict[str, str]]: + tools: List, + ) -> Tuple[Dict[str, str], Dict[str, str]]: """Build the (forward, reverse) tool-name maps for an OpenAI tools list. Operates on **OpenAI-format** tool dicts (pre-``_map_tools``). The @@ -968,7 +972,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): the original name out of either ``{"function": {"name": ...}}`` (legacy OpenAI shape) or ``{"name": ...}`` (rare top-level shape). """ - original_names: list[str] = [] + original_names: List[str] = [] for tool in tools or []: if not isinstance(tool, dict): continue @@ -981,8 +985,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _sanitize_tool_names_in_request( - optional_params: dict[str, Any], - ) -> tuple[dict[str, str], dict[str, str]]: + optional_params: Dict[str, Any], + ) -> Tuple[Dict[str, str], Dict[str, str]]: """Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']`` in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``. @@ -1006,7 +1010,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # Order matters: the first occurrence wins the canonical slot; # later collisions get numeric suffixes (see # ``_build_anthropic_tool_name_maps``). - original_names: list[str] = [] + original_names: List[str] = [] for t in tools: if not isinstance(t, dict): continue @@ -1028,7 +1032,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # so a caller reusing the same tool list/dicts across requests # doesn't see its inputs permanently rewritten (which would also # drop the original key from `forward` on the next request). - new_tools: list[Any] = [] + new_tools: List[Any] = [] for t in tools: if ( isinstance(t, dict) @@ -1054,7 +1058,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return forward, reverse - def _detect_tool_search_tools(self, tools: list | None) -> bool: + def _detect_tool_search_tools(self, tools: Optional[List]) -> bool: """Check if tool search tools are present in the tools list.""" if not tools: return False @@ -1068,7 +1072,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return True return False - def _separate_deferred_tools(self, tools: list) -> tuple[list, list]: + def _separate_deferred_tools(self, tools: List) -> Tuple[List, List]: """ Separate tools into deferred and non-deferred lists. @@ -1088,9 +1092,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _expand_tool_references( self, - content: list, - deferred_tools: list, - ) -> list: + content: List, + deferred_tools: List, + ) -> List: """ Expand tool_reference blocks to full tool definitions. @@ -1131,8 +1135,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return expanded_content - def _map_stop_sequences(self, stop: Union[str, list[str]] | None) -> list[str] | None: - new_stop: list[str] | None = None + def _map_stop_sequences(self, stop: Optional[Union[str, List[str]]]) -> Optional[List[str]]: + new_stop: Optional[List[str]] = None if isinstance(stop, str): if ( stop.isspace() and litellm.drop_params is True @@ -1153,10 +1157,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _map_reasoning_effort( - reasoning_effort: Union[REASONING_EFFORT, str] | None, + reasoning_effort: Optional[Union[REASONING_EFFORT, str]], model: str, llm_provider: str = "anthropic", - ) -> AnthropicThinkingParam | None: + ) -> Optional[AnthropicThinkingParam]: if reasoning_effort is None or reasoning_effort == "none": return None if AnthropicConfig._is_adaptive_thinking_model(model): @@ -1207,10 +1211,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): llm_provider=llm_provider, ) - def _extract_json_schema_from_response_format(self, value: dict | None) -> dict | None: + def _extract_json_schema_from_response_format(self, value: Optional[dict]) -> Optional[dict]: if value is None: return None - json_schema: dict | None = None + json_schema: Optional[dict] = None if "response_schema" in value: json_schema = value["response_schema"] elif "json_schema" in value: @@ -1218,8 +1222,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return json_schema - def map_response_format_to_anthropic_output_format(self, value: dict | None) -> AnthropicOutputSchema | None: - json_schema: dict | None = self._extract_json_schema_from_response_format(value) + def map_response_format_to_anthropic_output_format(self, value: Optional[dict]) -> Optional[AnthropicOutputSchema]: + json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value) if json_schema is None: return None @@ -1245,13 +1249,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) def map_response_format_to_anthropic_tool( - self, value: dict | None, optional_params: dict, is_thinking_enabled: bool - ) -> AnthropicMessagesTool | None: + self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool + ) -> Optional[AnthropicMessagesTool]: ignore_response_format_types = ["text"] if value is None or value["type"] in ignore_response_format_types: # value is a no-op return None - json_schema: dict | None = self._extract_json_schema_from_response_format(value) + json_schema: Optional[dict] = self._extract_json_schema_from_response_format(value) if json_schema is None: return None """ @@ -1296,8 +1300,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def map_openai_context_management_to_anthropic( - context_management: Union[list[dict[str, Any]], dict[str, Any]], - ) -> dict[str, Any] | None: + context_management: Union[List[Dict[str, Any]], Dict[str, Any]], + ) -> Optional[Dict[str, Any]]: """ OpenAI format: [{"type": "compaction", "compact_threshold": 200000}] Anthropic format: { @@ -1328,7 +1332,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): entry_type = entry.get("type") if entry_type == "compaction": - anthropic_edit: dict[str, Any] = {"type": "compact_20260112"} + anthropic_edit: Dict[str, Any] = {"type": "compact_20260112"} compact_threshold = entry.get("compact_threshold") # Rewrite to 'trigger' with correct nesting if threshold exists if compact_threshold is not None and isinstance(compact_threshold, (int, float)): @@ -1384,7 +1388,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if mcp_servers: optional_params["mcp_servers"] = mcp_servers elif param == "tool_choice" or param == "parallel_tool_calls": - _tool_choice: AnthropicMessagesToolChoice | None = self._map_tool_choice( + _tool_choice: Optional[AnthropicMessagesToolChoice] = self._map_tool_choice( tool_choice=non_default_params.get("tool_choice"), parallel_tool_use=non_default_params.get("parallel_tool_calls"), ) @@ -1515,7 +1519,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _create_json_tool_call_for_response_format( self, - json_schema: dict | None = None, + json_schema: Optional[dict] = None, ) -> AnthropicMessagesTool: """ Handles creating a tool call for getting responses in JSON format. @@ -1550,7 +1554,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): """ return False - def translate_system_message(self, messages: list[AllMessageValues]) -> list[AnthropicSystemMessageContent]: + def translate_system_message(self, messages: List[AllMessageValues]) -> List[AnthropicSystemMessageContent]: """ Translate system message to anthropic format. @@ -1558,7 +1562,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped. """ system_prompt_indices = [] - anthropic_system_message_list: list[AnthropicSystemMessageContent] = [] + anthropic_system_message_list: List[AnthropicSystemMessageContent] = [] for idx, message in enumerate(messages): if message["role"] == "system": system_prompt_indices.append(idx) @@ -1608,9 +1612,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def add_code_execution_tool( self, - messages: list[AllAnthropicMessageValues], - tools: list[Union[AllAnthropicToolsValues, dict]], - ) -> list[Union[AllAnthropicToolsValues, dict]]: + messages: List[AllAnthropicMessageValues], + tools: List[Union[AllAnthropicToolsValues, Dict]], + ) -> List[Union[AllAnthropicToolsValues, Dict]]: """if 'container_upload' in messages, add code_execution tool""" add_code_execution_tool = False for message in messages: @@ -1725,7 +1729,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def transform_request( self, model: str, - messages: list[AllMessageValues], + messages: List[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, @@ -1827,7 +1831,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ## Add code_execution tool if container_upload is in messages _tools = ( cast( - list[Union[AllAnthropicToolsValues, dict]] | None, + Optional[List[Union[AllAnthropicToolsValues, Dict]]], optional_params.get("tools"), ) or [] @@ -1925,12 +1929,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _resolve_json_mode_non_streaming( self, - json_mode: bool | None, - tool_calls: list[ChatCompletionToolCallChunk], - ) -> tuple[ - LitellmMessage | None, - list[ChatCompletionToolCallChunk], - str | None, + json_mode: Optional[bool], + tool_calls: List[ChatCompletionToolCallChunk], + ) -> Tuple[ + Optional[LitellmMessage], + List[ChatCompletionToolCallChunk], + Optional[str], ]: """Strip internal response_format tool calls; merge payload into content when mixed with user tools.""" if json_mode is not True or not tool_calls: @@ -1951,30 +1955,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): first_json = tool_calls[json_indices[0]] json_msg = AnthropicConfig._convert_tool_response_to_message([first_json]) - extra_content: str | None = json_msg.content if json_msg is not None else None + extra_content: Optional[str] = json_msg.content if json_msg is not None else None filtered_tools = [t for i, t in enumerate(tool_calls) if i not in json_indices] return None, filtered_tools, extra_content def extract_response_content( self, completion_response: dict - ) -> tuple[ + ) -> Tuple[ str, - list[Any] | None, - list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None, - str | None, - list[ChatCompletionToolCallChunk], - list[Any] | None, - list[Any] | None, - list[Any] | None, + Optional[List[Any]], + Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], + Optional[str], + List[ChatCompletionToolCallChunk], + Optional[List[Any]], + Optional[List[Any]], + Optional[List[Any]], ]: text_content = "" - citations: list[Any] | None = None - thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None = None - reasoning_content: str | None = None - tool_calls: list[ChatCompletionToolCallChunk] = [] - web_search_results: list[Any] | None = None - tool_results: list[Any] | None = None - compaction_blocks: list[Any] | None = None + citations: Optional[List[Any]] = None + thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]] = None + reasoning_content: Optional[str] = None + tool_calls: List[ChatCompletionToolCallChunk] = [] + web_search_results: Optional[List[Any]] = None + tool_results: Optional[List[Any]] = None + compaction_blocks: Optional[List[Any]] = None for idx, content in enumerate(completion_response["content"]): if content["type"] == "text": text_content += content["text"] @@ -2037,7 +2041,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if thinking_blocks is not None: reasoning_content = "" for block in thinking_blocks: - thinking_content = cast(str | None, block.get("thinking")) + thinking_content = cast(Optional[str], block.get("thinking")) if thinking_content is not None: reasoning_content += thinking_content @@ -2055,9 +2059,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def calculate_usage( self, usage_object: dict, - reasoning_content: str | None, - completion_response: dict | None = None, - speed: str | None = None, + reasoning_content: Optional[str], + completion_response: Optional[dict] = None, + speed: Optional[str] = None, ) -> Usage: # NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0 @@ -2067,10 +2071,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _usage = usage_object cache_creation_input_tokens: int = 0 cache_read_input_tokens: int = 0 - cache_creation_token_details: CacheCreationTokenDetails | None = None - web_search_requests: int | None = None - tool_search_requests: int | None = None - inference_geo: str | None = None + cache_creation_token_details: Optional[CacheCreationTokenDetails] = None + web_search_requests: Optional[int] = None + tool_search_requests: Optional[int] = None + inference_geo: Optional[str] = None if "inference_geo" in _usage and _usage["inference_geo"] is not None: inference_geo = _usage["inference_geo"] service_tier = cast( @@ -2078,7 +2082,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _usage.get("service_tier"), ) - iterations: list[Any] | None = _usage.get("iterations") + iterations: Optional[List[Any]] = _usage.get("iterations") if iterations: prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations) completion_tokens = sum(it.get("output_tokens", 0) or 0 for it in iterations) @@ -2164,8 +2168,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) return usage - def _build_code_by_id_map(self, tool_calls: list[ChatCompletionToolCallChunk]) -> dict[str, str]: - code_by_id: dict[str, str] = {} + def _build_code_by_id_map(self, tool_calls: List[ChatCompletionToolCallChunk]) -> Dict[str, str]: + code_by_id: Dict[str, str] = {} for tc in tool_calls: try: args = json.loads(tc.get("function", {}).get("arguments", "{}")) @@ -2179,10 +2183,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _build_code_interpreter_results( self, - tool_results: list[Any], - code_by_id: dict[str, str], - container_id: str | None, - ) -> list[OutputCodeInterpreterCall]: + tool_results: List[Any], + code_by_id: Dict[str, str], + container_id: Optional[str], + ) -> List[OutputCodeInterpreterCall]: code_interpreter_results = [] for tr in tool_results: if tr.get("type") != "bash_code_execution_tool_result": @@ -2205,14 +2209,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _build_provider_specific_fields( self, completion_response: dict, - citations: list[Any] | None, - thinking_blocks: list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None, - web_search_results: list[Any] | None, - tool_results: list[Any] | None, - compaction_blocks: list[Any] | None, - tool_calls: list[ChatCompletionToolCallChunk], - ) -> dict[str, Any]: - provider_specific_fields: dict[str, Any] = { + citations: Optional[List[Any]], + thinking_blocks: Optional[List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]], + web_search_results: Optional[List[Any]], + tool_results: Optional[List[Any]], + compaction_blocks: Optional[List[Any]], + tool_calls: List[ChatCompletionToolCallChunk], + ) -> Dict[str, Any]: + provider_specific_fields: Dict[str, Any] = { "citations": citations, "thinking_blocks": thinking_blocks, } @@ -2249,12 +2253,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): completion_response: dict, raw_response: httpx.Response, model_response: ModelResponse, - json_mode: bool | None = None, - prefix_prompt: str | None = None, - speed: str | None = None, - tool_name_reverse_map: dict[str, str] | None = None, + json_mode: Optional[bool] = None, + prefix_prompt: Optional[str] = None, + speed: Optional[str] = None, + tool_name_reverse_map: Optional[Dict[str, str]] = None, ): - _hidden_params: dict = {} + _hidden_params: Dict = {} _hidden_params["additional_headers"] = process_anthropic_headers(dict(raw_response.headers)) if "error" in completion_response: response_headers = getattr(raw_response, "headers", None) @@ -2350,7 +2354,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): model_response._hidden_params = _hidden_params return model_response - def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None: + def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]: """ Get the prefix prompt from the messages. @@ -2375,13 +2379,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): raw_response: httpx.Response, model_response: ModelResponse, logging_obj: LoggingClass, - request_data: dict, - messages: list[AllMessageValues], - optional_params: dict, + request_data: Dict, + messages: List[AllMessageValues], + optional_params: Dict, litellm_params: dict, encoding: Any, - api_key: str | None = None, - json_mode: bool | None = None, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -2404,7 +2408,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): prefix_prompt = self.get_prefix_prompt(messages=messages) speed = optional_params.get("speed") - tool_name_reverse_map: dict[str, str] | None = None + tool_name_reverse_map: Optional[Dict[str, str]] = None if isinstance(litellm_params, dict): _candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY) if isinstance(_candidate, dict): @@ -2423,14 +2427,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _convert_tool_response_to_message( - tool_calls: list[ChatCompletionToolCallChunk], - ) -> LitellmMessage | None: + tool_calls: List[ChatCompletionToolCallChunk], + ) -> Optional[LitellmMessage]: """ In JSON mode, Anthropic API returns JSON schema as a tool call, we need to convert it to a message to follow the OpenAI format """ ## HANDLE JSON MODE - anthropic returns single function call - json_mode_content_str: str | None = tool_calls[0]["function"].get("arguments") + json_mode_content_str: Optional[str] = tool_calls[0]["function"].get("arguments") try: if json_mode_content_str is not None: args = json.loads(json_mode_content_str) @@ -2448,7 +2452,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return None def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] ) -> BaseLLMException: return AnthropicError( status_code=status_code, diff --git a/litellm/llms/modelscope/chat/transformation.py b/litellm/llms/modelscope/chat/transformation.py index 9bfce09ff41..1a54be6e1c8 100644 --- a/litellm/llms/modelscope/chat/transformation.py +++ b/litellm/llms/modelscope/chat/transformation.py @@ -2,7 +2,7 @@ Translates from OpenAI's `/v1/chat/completions` to ModelScope's `/v1/chat/completions` """ -from typing import Any, Coroutine, Literal, Union, cast, overload +from typing import Any, Coroutine, Literal, Optional, Tuple, Union, cast, overload from typing_extensions import override @@ -59,8 +59,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig): return super()._transform_messages(messages=messages, model=model, is_async=False) def _get_openai_compatible_provider_info( - self, api_base: str | None, api_key: str | None - ) -> tuple[str | None, str | None]: + self, api_base: Optional[str], api_key: Optional[str] + ) -> Tuple[Optional[str], Optional[str]]: api_base = api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL # type: ignore dynamic_api_key = api_key or get_secret_str("MODELSCOPE_API_KEY") return api_base, dynamic_api_key @@ -68,12 +68,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig): @override def get_complete_url( self, - api_base: str | None, - api_key: str | None, + api_base: Optional[str], + api_key: Optional[str], model: str, optional_params: dict, litellm_params: dict, - stream: bool | None = None, + stream: Optional[bool] = None, ) -> str: """ If api_base is not provided, use the default ModelScope /chat/completions endpoint. 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 8e4c086ac6d..678877c0721 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 @@ -9,9 +9,13 @@ from typing import ( TYPE_CHECKING, Any, Callable, + Dict, + List, Literal, Mapping, Optional, + Tuple, + Type, Union, cast, ) @@ -131,7 +135,7 @@ class VertexAIBaseConfig: optional_params[mapped_params[param]] = value return optional_params - def get_eu_regions(self) -> list[str]: + def get_eu_regions(self) -> List[str]: """ Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions """ @@ -148,7 +152,7 @@ class VertexAIBaseConfig: "europe-west9", ] - def get_us_regions(self) -> list[str]: + def get_us_regions(self) -> List[str]: """ Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions """ @@ -194,29 +198,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Note: Please make sure to modify the default parameters as required for your use case. """ - temperature: float | None = None - max_output_tokens: int | None = None - top_p: float | None = None - top_k: int | None = None - response_mime_type: str | None = None - candidate_count: int | None = None - stop_sequences: list | None = None - frequency_penalty: float | None = None - presence_penalty: float | None = None - seed: int | None = None + temperature: Optional[float] = None + max_output_tokens: Optional[int] = None + top_p: Optional[float] = None + top_k: Optional[int] = None + response_mime_type: Optional[str] = None + candidate_count: Optional[int] = None + stop_sequences: Optional[list] = None + frequency_penalty: Optional[float] = None + presence_penalty: Optional[float] = None + seed: Optional[int] = None def __init__( self, - temperature: float | None = None, - max_output_tokens: int | None = None, - top_p: float | None = None, - top_k: int | None = None, - response_mime_type: str | None = None, - candidate_count: int | None = None, - stop_sequences: list | None = None, - frequency_penalty: float | None = None, - presence_penalty: float | None = None, - seed: int | None = None, + temperature: Optional[float] = None, + max_output_tokens: Optional[int] = None, + top_p: Optional[float] = None, + top_k: Optional[int] = None, + response_mime_type: Optional[str] = None, + candidate_count: Optional[int] = None, + stop_sequences: Optional[list] = None, + frequency_penalty: Optional[float] = None, + presence_penalty: Optional[float] = None, + seed: Optional[int] = None, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -228,8 +232,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return super().get_config() def get_json_schema_from_pydantic_object( - self, response_format: Union[type["BaseModel"], dict] | None - ) -> dict | None: + self, response_format: Optional[Union[Type["BaseModel"], dict]] + ) -> Optional[dict]: """ Override to use Pydantic's model_json_schema() instead of OpenAI's to_strict_json_schema(). @@ -285,7 +289,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False @staticmethod - def _forward_gemini_function_call_id(model: str, custom_llm_provider: str | None = None) -> bool: + def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool: """ Whether to include `id` on function_call / function_response parts. @@ -305,7 +309,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False return True - def get_supported_openai_params(self, model: str) -> list[str]: + def get_supported_openai_params(self, model: str) -> List[str]: supported_params = [ "temperature", "top_p", @@ -340,7 +344,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): supported_params.append("thinking") return supported_params - def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> ToolConfig | None: + def map_tool_choice_values(self, model: str, tool_choice: Union[str, dict]) -> Optional[ToolConfig]: if tool_choice == "none": return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE")) elif tool_choice == "required": @@ -466,7 +470,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return transformed_config - def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> tuple[dict, dict | None]: + def _extract_google_maps_retrieval_config(self, google_maps_config: dict) -> Tuple[dict, Optional[dict]]: """ Extract location configuration from googleMaps tool for Vertex AI toolConfig. @@ -504,7 +508,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return cleaned_config, retrieval_config - def get_tool_value(self, tool: dict, tool_name: str) -> dict | None: + def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]: """ Helper function to get tool value handling both camelCase and underscore_case variants @@ -529,10 +533,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _resolve_search_tool_conflict( gtool_func_declarations: list, - googleSearch: dict | None, - googleSearchRetrieval: dict | None, - enterpriseWebSearch: dict | None, - urlContext: dict | None, + googleSearch: Optional[dict], + googleSearchRetrieval: Optional[dict], + enterpriseWebSearch: Optional[dict], + urlContext: Optional[dict], optional_params: dict, ) -> tuple: """ @@ -576,7 +580,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext - def _map_function(self, value: list[dict], optional_params: dict) -> list[Tools]: + def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]: """ Map OpenAI-style tools/functions to Vertex AI format. @@ -592,21 +596,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): googleMaps tools contain location data """ gtool_func_declarations = [] - googleSearch: dict | None = None - googleSearchRetrieval: dict | None = None - enterpriseWebSearch: dict | None = None - urlContext: dict | None = None - code_execution: dict | None = None - googleMaps: dict | None = None - google_maps_retrieval_config: dict | None = None - computerUse: dict | None = None + googleSearch: Optional[dict] = None + googleSearchRetrieval: Optional[dict] = None + enterpriseWebSearch: Optional[dict] = None + urlContext: Optional[dict] = None + code_execution: Optional[dict] = None + googleMaps: Optional[dict] = None + google_maps_retrieval_config: Optional[dict] = None + computerUse: Optional[dict] = None # remove 'additionalProperties' from tools value = _remove_additional_properties(value) # remove 'strict' from tools value = _remove_strict_from_schema(value) for tool in value: - openai_function_object: ChatCompletionToolParamFunctionChunk | None = None + openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = None if "function" in tool: # tools list _openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore **tool["function"] @@ -696,7 +700,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Build list of Tool objects - each Tool should contain exactly one type # per Vertex AI API spec: "A Tool object should contain exactly one type of Tool" - _tools_list: list[Tools] = [] + _tools_list: List[Tools] = [] ( googleSearch, @@ -814,7 +818,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_budget( reasoning_effort: str, - model: str | None = None, + model: Optional[str] = None, ) -> GeminiThinkingConfig: if reasoning_effort == "minimal": # Use model-specific minimum thinking budget or fallback @@ -863,7 +867,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_level( reasoning_effort: str, - model: str | None = None, + model: Optional[str] = None, ) -> GeminiThinkingConfig: """ Map reasoning_effort to thinking_level for Gemini 3+ models. @@ -909,12 +913,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): raise ValueError(f"Invalid reasoning effort: {reasoning_effort}") @staticmethod - def _is_thinking_budget_zero(thinking_budget: int | None) -> bool: + def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool: return thinking_budget is not None and thinking_budget == 0 @staticmethod def _validate_thinking_config_conflicts( - optional_params: dict, + optional_params: Dict, param_name: str, param_description: str = "thinking_budget", ) -> None: @@ -935,7 +939,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _validate_thinking_level_conflicts( - optional_params: dict, + optional_params: Dict, ) -> None: """ Validate that thinking_level and thinking_budget are not both specified. @@ -955,7 +959,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_thinking_param( thinking_param: AnthropicThinkingParam, - model: str | None = None, + model: Optional[str] = None, ) -> GeminiThinkingConfig: thinking_enabled = thinking_param.get("type") == "enabled" thinking_budget = thinking_param.get("budget_tokens") @@ -1060,8 +1064,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _apply_include_server_side_tool_invocations( - non_default_params: dict, - optional_params: dict, + non_default_params: Dict, + optional_params: Dict, ) -> None: """ Set include_server_side_tool_invocations before tools are mapped. @@ -1080,11 +1084,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def map_openai_params( self, - non_default_params: dict, - optional_params: dict, + non_default_params: Dict, + optional_params: Dict, model: str, drop_params: bool, - ) -> dict: + ) -> Dict: self._apply_include_server_side_tool_invocations(non_default_params, optional_params) gemini_sampling_params_warned: bool = False for param, value in non_default_params.items(): @@ -1176,7 +1180,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif param == "reasoning_effort": # Extract effort value - handle both string and dict formats # Dict format comes from OpenAI Agents SDK: {"effort": "high", "summary": "auto"} - effort_value: str | None = None + effort_value: Optional[str] = None if isinstance(value, str): effort_value = value elif isinstance(value, dict): @@ -1252,7 +1256,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params[mapped_params[param]] = value return optional_params - def get_eu_regions(self) -> list[str]: + def get_eu_regions(self) -> List[str]: """ Source: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/locations#available-regions """ @@ -1289,7 +1293,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return model @staticmethod - def _is_model_gemini_spec_model(model: str | None) -> bool: + def _is_model_gemini_spec_model(model: Optional[str]) -> bool: """ Returns true if user is trying to call custom model in `/gemini` request/response format """ @@ -1312,7 +1316,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return model.split("/")[-1] return model - def get_flagged_finish_reasons(self) -> dict[str, str]: + def get_flagged_finish_reasons(self) -> Dict[str, str]: """ Return Dictionary of finish reasons which indicate response was flagged @@ -1349,7 +1353,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) @staticmethod - def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]: + def get_finish_reason_mapping() -> Dict[str, OpenAIChatCompletionFinishReason]: """ Return Dictionary of Gemini/Vertex AI finish reasons and their OpenAI-compatible mappings. @@ -1368,9 +1372,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) return exception_string - def get_assistant_content_message(self, parts: list[HttpxPartType]) -> tuple[str | None, str | None]: - content_str: str | None = None - reasoning_content_str: str | None = None + def get_assistant_content_message(self, parts: List[HttpxPartType]) -> Tuple[Optional[str], Optional[str]]: + content_str: Optional[str] = None + reasoning_content_str: Optional[str] = None for part in parts: _content_str = "" @@ -1409,7 +1413,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return content_str, reasoning_content_str - def _extract_thinking_blocks_from_parts(self, parts: list[HttpxPartType]) -> list[ChatCompletionThinkingBlock]: + def _extract_thinking_blocks_from_parts(self, parts: List[HttpxPartType]) -> List[ChatCompletionThinkingBlock]: """Extract thinking blocks from parts if present. Per Google's docs (https://ai.google.dev/gemini-api/docs/thinking): @@ -1418,7 +1422,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): it does NOT indicate that the content is thinking (a part can have thoughtSignature without thought: true, e.g., function calls) """ - thinking_blocks: list[ChatCompletionThinkingBlock] = [] + thinking_blocks: List[ChatCompletionThinkingBlock] = [] for part in parts: if part.get("thought") is True: thinking_text = part.get("text", "") @@ -1432,7 +1436,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): thinking_blocks.append(block) return thinking_blocks - def _extract_thought_signatures_from_parts(self, parts: list[HttpxPartType]) -> list[str] | None: + def _extract_thought_signatures_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[str]]: """Extract thoughtSignature values from parts. Per Google's docs, thoughtSignature is returned for multi-turn context preservation @@ -1442,7 +1446,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Returns: List of thoughtSignature strings if any are found, None otherwise """ - signatures: list[str] = [] + signatures: List[str] = [] for part in parts: signature = part.get("thoughtSignature") if signature is not None: @@ -1451,8 +1455,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _extract_server_side_tool_invocations( - parts: list[HttpxPartType], - ) -> list[dict[str, Any]] | None: + parts: List[HttpxPartType], + ) -> Optional[List[Dict[str, Any]]]: """Extract server-side tool invocations (toolCall/toolResponse) from parts. These are returned by Gemini when context circulation is enabled @@ -1463,15 +1467,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Returns: List of server-side invocation dicts if any found, None otherwise. """ - invocations: list[dict[str, Any]] = [] + invocations: List[Dict[str, Any]] = [] # Index toolCalls by id so we can pair them with responses - tool_calls_by_id: dict[str, dict[str, Any]] = {} - tool_responses_by_id: dict[str, dict[str, Any]] = {} + tool_calls_by_id: Dict[str, Dict[str, Any]] = {} + tool_responses_by_id: Dict[str, Dict[str, Any]] = {} for part in parts: if "toolCall" in part: tc = part["toolCall"] - entry: dict[str, Any] = { + entry: Dict[str, Any] = { "tool_type": tc.get("toolType"), "id": tc.get("id"), "args": tc.get("args"), @@ -1511,9 +1515,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return invocations if invocations else None - def _extract_image_response_from_parts(self, parts: list[HttpxPartType]) -> list[ImageURLListItem] | None: + def _extract_image_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[List[ImageURLListItem]]: """Extract image response from parts if present""" - images: list[ImageURLListItem] = [] + images: List[ImageURLListItem] = [] for part in parts: if "inlineData" in part: inline_data = part.get("inlineData", {}) @@ -1531,7 +1535,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) return images - def _extract_audio_response_from_parts(self, parts: list[HttpxPartType]) -> ChatCompletionAudioResponse | None: + def _extract_audio_response_from_parts(self, parts: List[HttpxPartType]) -> Optional[ChatCompletionAudioResponse]: """Extract audio response from parts if present""" for part in parts: if "text" in part: @@ -1569,16 +1573,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _transform_parts( - parts: list[HttpxPartType], + parts: List[HttpxPartType], cumulative_tool_call_idx: int, - is_function_call: bool | None, - ) -> tuple[ - ChatCompletionToolCallFunctionChunk | None, - list[ChatCompletionToolCallChunk] | None, + is_function_call: Optional[bool], + ) -> Tuple[ + Optional[ChatCompletionToolCallFunctionChunk], + Optional[List[ChatCompletionToolCallChunk]], int, ]: - function: ChatCompletionToolCallFunctionChunk | None = None - _tools: list[ChatCompletionToolCallChunk] = [] + function: Optional[ChatCompletionToolCallFunctionChunk] = None + _tools: List[ChatCompletionToolCallChunk] = [] for part in parts: if "functionCall" in part: _function_chunk: ChatCompletionToolCallFunctionChunk = { @@ -1593,7 +1597,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): gemini_call_id = part["functionCall"].get("id") if is_function_call is True: - function_dict: dict[str, Any] = dict(_function_chunk) + function_dict: Dict[str, Any] = dict(_function_chunk) if thought_signature: if "provider_specific_fields" not in function_dict: function_dict["provider_specific_fields"] = {} @@ -1622,22 +1626,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tools.append(_tool_response_chunk) cumulative_tool_call_idx += 1 if len(_tools) == 0: - tools: list[ChatCompletionToolCallChunk] | None = None + tools: Optional[List[ChatCompletionToolCallChunk]] = None else: tools = _tools return function, tools, cumulative_tool_call_idx @staticmethod def _transform_logprobs( - logprobs_result: LogprobsResult | None, - ) -> ChoiceLogprobs | None: + logprobs_result: Optional[LogprobsResult], + ) -> Optional[ChoiceLogprobs]: if logprobs_result is None: return None if "chosenCandidates" not in logprobs_result: return None - logprobs_list: list[ChatCompletionTokenLogprob] = [] + logprobs_list: List[ChatCompletionTokenLogprob] = [] for index, candidate in enumerate(logprobs_result["chosenCandidates"]): - top_logprobs: list[TopLogprob] = [] + top_logprobs: List[TopLogprob] = [] if "topCandidates" in logprobs_result and index < len(logprobs_result["topCandidates"]): top_candidates_for_index = logprobs_result["topCandidates"][index]["candidates"] @@ -1744,16 +1748,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) -> Usage: if completion_response is not None and "usageMetadata" not in completion_response: raise ValueError(f"usageMetadata not found in completion_response. Got={completion_response}") - cached_tokens: int | None = None + cached_tokens: Optional[int] = None # Separate variables for prompt tokens by modality - prompt_audio_tokens: int | None = None - prompt_image_tokens: int | None = None - prompt_text_tokens: int | None = None - prompt_video_tokens: int | None = None - prompt_tokens_details: PromptTokensDetailsWrapper | None = None - reasoning_tokens: int | None = None - response_tokens: int | None = None - response_tokens_details: CompletionTokensDetailsWrapper | None = None + prompt_audio_tokens: Optional[int] = None + prompt_image_tokens: Optional[int] = None + prompt_text_tokens: Optional[int] = None + prompt_video_tokens: Optional[int] = None + prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None + reasoning_tokens: Optional[int] = None + response_tokens: Optional[int] = None + response_tokens_details: Optional[CompletionTokensDetailsWrapper] = None usage_metadata = completion_response["usageMetadata"] def _get_token_count(detail: Mapping[str, Any]) -> int: @@ -1832,10 +1836,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ## Parse cacheTokensDetails (breakdown of cached tokens by modality) ## When explicit caching is used, Gemini provides this field to show which modalities were cached - cached_text_tokens: int | None = None - cached_audio_tokens: int | None = None - cached_image_tokens: int | None = None - cached_video_tokens: int | None = None + cached_text_tokens: Optional[int] = None + cached_audio_tokens: Optional[int] = None + cached_image_tokens: Optional[int] = None + cached_video_tokens: Optional[int] = None if "cacheTokensDetails" in usage_metadata: for detail in usage_metadata["cacheTokensDetails"]: @@ -1908,8 +1912,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _check_finish_reason( - chat_completion_message: ChatCompletionResponseMessage | None, - finish_reason: str | None, + chat_completion_message: Optional[ChatCompletionResponseMessage], + finish_reason: Optional[str], ) -> OpenAIChatCompletionFinishReason: from litellm.litellm_core_utils.core_helpers import map_finish_reason @@ -1925,7 +1929,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _check_prompt_level_content_filter( processed_chunk: GenerateContentResponseBody, - response_id: str | None, + response_id: Optional[str], ) -> Optional["ModelResponseStream"]: """ Check if prompt is blocked due to content filtering at the prompt level. @@ -1969,8 +1973,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return None @staticmethod - def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None: - web_search_requests: int | None = None + def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]: + web_search_requests: Optional[int] = None if grounding_metadata and isinstance(grounding_metadata, list) and len(grounding_metadata) > 0: for grounding_metadata_item in grounding_metadata: @@ -1986,10 +1990,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): chat_completion_message: ChatCompletionResponseMessage, candidate: Candidates, idx: int, - tools: list[ChatCompletionToolCallChunk] | None, - functions: ChatCompletionToolCallFunctionChunk | None, - chat_completion_logprobs: ChoiceLogprobs | None, - image_response: list[ImageURLListItem] | None, + tools: Optional[List[ChatCompletionToolCallChunk]], + functions: Optional[ChatCompletionToolCallFunctionChunk], + chat_completion_logprobs: Optional[ChoiceLogprobs], + image_response: Optional[List[ImageURLListItem]], ) -> StreamingChoices: """ Helper method to create a streaming choice object for Vertex AI @@ -2021,7 +2025,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _extract_candidate_metadata( candidate: Candidates, - ) -> tuple[list[dict], list[dict], list, list]: + ) -> Tuple[List[dict], List[dict], List, List]: """ Extract metadata from a single candidate response. @@ -2031,10 +2035,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): safety_ratings: List citation_metadata: List """ - grounding_metadata: list[dict] = [] - url_context_metadata: list[dict] = [] - safety_ratings: list = [] - citation_metadata: list = [] + grounding_metadata: List[dict] = [] + url_context_metadata: List[dict] = [] + safety_ratings: List = [] + citation_metadata: List = [] if "groundingMetadata" in candidate: if isinstance(candidate["groundingMetadata"], list): @@ -2079,10 +2083,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @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], + 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: @@ -2102,10 +2106,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def apply_assembled_streaming_response_metadata( self, response: ModelResponse, - chunks: list[Any], + chunks: List[Any], ) -> None: for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS: - merged: list[Any] = [] + merged: List[Any] = [] for chunk in chunks: value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name) if not value: @@ -2120,14 +2124,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _convert_grounding_metadata_to_annotations( - grounding_metadata: list[dict], - content_text: str | None, - ) -> list[ChatCompletionAnnotation]: + grounding_metadata: List[dict], + content_text: Optional[str], + ) -> List[ChatCompletionAnnotation]: """ Convert Vertex AI grounding metadata to OpenAI-style annotations. """ - annotations: list[ChatCompletionAnnotation] = [] + annotations: List[ChatCompletionAnnotation] = [] for metadata in grounding_metadata: # Extract groundingSupports - these map text segments to sources @@ -2135,7 +2139,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): grounding_chunks = metadata.get("groundingChunks", []) # Build a map of chunk indices to web URIs - chunk_to_uri_map: dict[int, dict[str, str]] = {} + chunk_to_uri_map: Dict[int, Dict[str, str]] = {} for idx, chunk in enumerate(grounding_chunks): if "web" in chunk: web_data = chunk["web"] @@ -2175,11 +2179,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _process_candidates( - _candidates: list[Candidates], + _candidates: List[Candidates], model_response: Union[ModelResponse, "ModelResponseStream"], standard_optional_params: dict, cumulative_tool_call_index: int = 0, - ) -> tuple[list[dict], list[dict], list, list, int]: + ) -> Tuple[List[dict], List[dict], List, List, int]: """ Helper method to process candidates and extract metadata @@ -2195,19 +2199,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) from litellm.types.utils import ModelResponseStream - grounding_metadata: list[dict] = [] - url_context_metadata: list[dict] = [] - image_response: list[ImageURLListItem] | None = None - safety_ratings: list = [] - citation_metadata: list = [] + grounding_metadata: List[dict] = [] + url_context_metadata: List[dict] = [] + image_response: Optional[List[ImageURLListItem]] = None + safety_ratings: List = [] + citation_metadata: List = [] chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} - chat_completion_logprobs: ChoiceLogprobs | None = None - tools: list[ChatCompletionToolCallChunk] | None = [] - functions: ChatCompletionToolCallFunctionChunk | None = None - thinking_blocks: list[ChatCompletionThinkingBlock] | None = None - reasoning_content: str | None = None - thought_signatures: Any | None = None - server_side_tool_invocations: list[dict[str, Any]] | None = None + chat_completion_logprobs: Optional[ChoiceLogprobs] = None + tools: Optional[List[ChatCompletionToolCallChunk]] = [] + functions: Optional[ChatCompletionToolCallFunctionChunk] = None + thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None + reasoning_content: Optional[str] = None + thought_signatures: Optional[Any] = None + server_side_tool_invocations: Optional[List[Dict[str, Any]]] = None for idx, candidate in enumerate(_candidates): if "content" not in candidate: @@ -2254,11 +2258,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) if audio_response is not None: - cast(dict[str, Any], chat_completion_message)["audio"] = audio_response + cast(Dict[str, Any], chat_completion_message)["audio"] = audio_response chat_completion_message["content"] = None # OpenAI spec if image_response is not None: # Handle image response - combine with text content into structured format - cast(dict[str, Any], chat_completion_message)["images"] = image_response + cast(Dict[str, Any], chat_completion_message)["images"] = image_response if content is not None: chat_completion_message["content"] = content @@ -2360,13 +2364,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): raw_response: httpx.Response, model_response: ModelResponse, logging_obj: LoggingClass, - request_data: dict, - messages: list[AllMessageValues], - optional_params: dict, - litellm_params: dict, + request_data: Dict, + messages: List[AllMessageValues], + optional_params: Dict, + litellm_params: Dict, encoding: Any, - api_key: str | None = None, - json_mode: bool | None = None, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -2433,11 +2437,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): response_id = completion_response.get("responseId") if response_id: model_response.id = response_id - url_context_metadata: list[dict] = [] + url_context_metadata: List[dict] = [] try: - grounding_metadata: list[dict] = [] - safety_ratings: list[dict] = [] - citation_metadata: list[dict] = [] + grounding_metadata: List[dict] = [] + safety_ratings: List[dict] = [] + citation_metadata: List[dict] = [] if _candidates: ( grounding_metadata, @@ -2501,10 +2505,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _transform_messages( self, - messages: list[AllMessageValues], - model: str | None = None, - litellm_params: dict | None = None, - ) -> list[ContentType]: + messages: List[AllMessageValues], + model: Optional[str] = None, + litellm_params: Optional[dict] = None, + ) -> List[ContentType]: return _gemini_convert_messages_with_history( messages=messages, model=model, @@ -2513,30 +2517,30 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) def get_error_class( - self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers] ) -> BaseLLMException: return VertexAIError(message=error_message, status_code=status_code, headers=headers) def transform_request( self, model: str, - messages: list[AllMessageValues], - optional_params: dict, - litellm_params: dict, - headers: dict, - ) -> dict: + messages: List[AllMessageValues], + optional_params: Dict, + litellm_params: Dict, + headers: Dict, + ) -> Dict: raise NotImplementedError("Vertex AI has a custom implementation of transform_request. Needs sync + async.") def validate_environment( self, - headers: dict | None, + headers: Optional[Dict], model: str, - messages: list[AllMessageValues], - optional_params: dict, - litellm_params: dict, - api_key: Union[str, dict] | None = None, - api_base: str | None = None, - ) -> dict: + messages: List[AllMessageValues], + optional_params: Dict, + litellm_params: Dict, + api_key: Optional[Union[str, Dict]] = None, + api_base: Optional[str] = None, + ) -> Dict: default_headers = { "Content-Type": "application/json", } @@ -2551,8 +2555,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): async def make_call( - client: AsyncHTTPHandler | None, # module-level client - gemini_client: AsyncHTTPHandler | None, # if passed by user + client: Optional[AsyncHTTPHandler], # module-level client + gemini_client: Optional[AsyncHTTPHandler], # if passed by user api_base: str, headers: dict, data: str, @@ -2603,8 +2607,8 @@ async def make_call( def make_sync_call( - client: HTTPHandler | None, # module-level client - gemini_client: HTTPHandler | None, # if passed by user + client: Optional[HTTPHandler], # module-level client + gemini_client: Optional[HTTPHandler], # if passed by user api_base: str, headers: dict, data: str, @@ -2659,20 +2663,20 @@ class VertexLLM(VertexBase): model_response: ModelResponse, print_verbose: Callable, data: dict, - timeout: Union[float, httpx.Timeout] | None, + timeout: Optional[Union[float, httpx.Timeout]], encoding, logging_obj, stream, optional_params: dict, litellm_params: dict, logger_fn=None, - api_base: str | None = None, - client: AsyncHTTPHandler | None = None, - vertex_project: str | None = None, - vertex_location: str | None = None, - vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None, - gemini_api_key: str | None = None, - extra_headers: dict | None = None, + api_base: Optional[str] = None, + client: Optional[AsyncHTTPHandler] = None, + vertex_project: Optional[str] = None, + vertex_location: Optional[str] = None, + vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, + gemini_api_key: Optional[str] = None, + extra_headers: Optional[dict] = None, ) -> CustomStreamWrapper: should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) @@ -2755,20 +2759,20 @@ class VertexLLM(VertexBase): custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) - timeout: Union[float, httpx.Timeout] | None, + timeout: Optional[Union[float, httpx.Timeout]], encoding, logging_obj, stream, optional_params: dict, litellm_params: dict, logger_fn=None, - api_base: str | None = None, - client: AsyncHTTPHandler | None = None, - vertex_project: str | None = None, - vertex_location: str | None = None, - vertex_credentials: VERTEX_CREDENTIALS_TYPES | None = None, - gemini_api_key: str | None = None, - extra_headers: dict | None = None, + api_base: Optional[str] = None, + client: Optional[AsyncHTTPHandler] = None, + vertex_project: Optional[str] = None, + vertex_location: Optional[str] = None, + vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None, + gemini_api_key: Optional[str] = None, + extra_headers: Optional[dict] = None, ) -> Union[ModelResponse, CustomStreamWrapper]: should_use_v1beta1_features = self.is_using_v1beta1_features(optional_params=optional_params) @@ -2877,18 +2881,18 @@ class VertexLLM(VertexBase): logging_obj, optional_params: dict, acompletion: bool, - timeout: Union[float, httpx.Timeout] | None, - vertex_project: str | None, - vertex_location: str | None, - vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, - gemini_api_key: str | None, + timeout: Optional[Union[float, httpx.Timeout]], + vertex_project: Optional[str], + vertex_location: Optional[str], + vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES], + gemini_api_key: Optional[str], litellm_params: dict, logger_fn=None, - extra_headers: dict | None = None, - client: Union[AsyncHTTPHandler, HTTPHandler] | None = None, - api_base: str | None = None, + extra_headers: Optional[dict] = None, + client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, + api_base: Optional[str] = None, ) -> Union[ModelResponse, CustomStreamWrapper]: - stream: bool | None = optional_params.pop("stream", None) # type: ignore + stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore transform_request_params = { "gemini_api_key": gemini_api_key, @@ -3076,7 +3080,7 @@ class ModelResponseIterator: streaming_response, sync_stream: bool, logging_obj: LoggingClass, - response_headers: dict[str, str] | None = None, + response_headers: Optional[Dict[str, str]] = None, response: httpx.Response | None = None, ): from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -3121,9 +3125,9 @@ class ModelResponseIterator: def _apply_stream_candidates( self, - _candidates: list[Candidates], + _candidates: List[Candidates], model_response: Any, - ) -> tuple[list[dict], list[dict], list[dict], list[dict]]: + ) -> Tuple[List[dict], List[dict], List[dict], List[dict]]: ( grounding_metadata, url_context_metadata, @@ -3202,8 +3206,8 @@ class ModelResponseIterator: self, processed_chunk: Any, model_response: Any, - grounding_metadata: list[dict], - ) -> Usage | None: + grounding_metadata: List[dict], + ) -> Optional[Usage]: if "usageMetadata" not in processed_chunk: return None @@ -3250,12 +3254,12 @@ class ModelResponseIterator: if blocked_response is not None: model_response = blocked_response - grounding_metadata: list[dict] = [] - url_context_metadata: list[dict] = [] - safety_ratings: list[dict] = [] - citation_metadata: list[dict] = [] + grounding_metadata: List[dict] = [] + url_context_metadata: List[dict] = [] + safety_ratings: List[dict] = [] + citation_metadata: List[dict] = [] - _candidates: list[Candidates] | None = processed_chunk.get("candidates") + _candidates: Optional[List[Candidates]] = processed_chunk.get("candidates") if _candidates: ( grounding_metadata, diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index dc4f0649331..8962073fe7a 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -3,7 +3,7 @@ import importlib import os from datetime import datetime, timezone -from typing import Any, Literal, Optional, cast +from typing import Any, Dict, List, Literal, Optional, Set, Type, cast from pydantic import ValidationError @@ -54,7 +54,7 @@ guardrail_initializer_registry = { SupportedGuardrailIntegrations.LLM_AS_A_JUDGE.value: initialize_llm_as_a_judge, } -guardrail_class_registry: dict[str, type[CustomGuardrail]] = { +guardrail_class_registry: Dict[str, Type[CustomGuardrail]] = { SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail } @@ -220,7 +220,7 @@ class GuardrailRegistry: ########################################################### ########### In memory management helpers for guardrails ########### ############################################################ - def get_initialized_guardrail_callback(self, guardrail_name: str) -> CustomGuardrail | None: + def get_initialized_guardrail_callback(self, guardrail_name: str) -> Optional[CustomGuardrail]: """ Returns the initialized guardrail callback for a given guardrail name """ @@ -314,7 +314,7 @@ class GuardrailRegistry: @staticmethod async def get_all_guardrails_from_db( prisma_client: PrismaClient, - ) -> list[Guardrail]: + ) -> List[Guardrail]: """ Get all active guardrails from the database. Only rows with status == "active" are returned (pending_review and rejected are excluded). @@ -325,7 +325,7 @@ class GuardrailRegistry: order={"created_at": "desc"}, ) - guardrails: list[Guardrail] = [] + guardrails: List[Guardrail] = [] for guardrail in guardrails_from_db: guardrails.append(Guardrail(**(dict(guardrail)))) # type: ignore @@ -333,7 +333,7 @@ class GuardrailRegistry: except Exception as e: raise Exception(f"Error getting guardrails from DB: {str(e)}") - async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Guardrail | None: + async def get_guardrail_by_id_from_db(self, guardrail_id: str, prisma_client: PrismaClient) -> Optional[Guardrail]: """ Get a guardrail by its ID from the database """ @@ -349,7 +349,9 @@ class GuardrailRegistry: except Exception as e: raise Exception(f"Error getting guardrail from DB: {str(e)}") - async def get_guardrail_by_name_from_db(self, guardrail_name: str, prisma_client: PrismaClient) -> Guardrail | None: + async def get_guardrail_by_name_from_db( + self, guardrail_name: str, prisma_client: PrismaClient + ) -> Optional[Guardrail]: """ Get a guardrail by its name from the database """ @@ -372,17 +374,17 @@ class InMemoryGuardrailHandler: """ def __init__(self): - self.IN_MEMORY_GUARDRAILS: dict[str, Guardrail] = {} + self.IN_MEMORY_GUARDRAILS: Dict[str, Guardrail] = {} """ Guardrail id to Guardrail object mapping """ - self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {} + self.guardrail_id_to_custom_guardrail: Dict[str, Optional[CustomGuardrail]] = {} """ Guardrail id to CustomGuardrail object mapping """ - self._sources: dict[str, Literal["db", "config"]] = {} + self._sources: Dict[str, Literal["db", "config"]] = {} """ Guardrail id to provenance marker. "db" entries are reconciled against the DB on each polling tick; "config" entries are owned by proxy_config.yaml @@ -392,10 +394,10 @@ class InMemoryGuardrailHandler: def initialize_guardrail( self, guardrail: Guardrail, - config_file_path: str | None = None, + config_file_path: Optional[str] = None, llm_router: Optional["Router"] = None, source: Literal["db", "config"] = "config", - ) -> Guardrail | None: + ) -> Optional[Guardrail]: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -411,7 +413,7 @@ class InMemoryGuardrailHandler: self._sources[guardrail_id] = source return self.IN_MEMORY_GUARDRAILS[guardrail_id] - custom_guardrail_callback: CustomGuardrail | None = None + custom_guardrail_callback: Optional[CustomGuardrail] = None litellm_params_data = guardrail["litellm_params"] verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) @@ -487,11 +489,11 @@ class InMemoryGuardrailHandler: def initialize_custom_guardrail( self, - guardrail: dict, + guardrail: Dict, guardrail_type: str, litellm_params: LitellmParams, - config_file_path: str | None = None, - ) -> CustomGuardrail | None: + config_file_path: Optional[str] = None, + ) -> Optional[CustomGuardrail]: """ Initialize a Custom Guardrail from a python file or module path @@ -576,25 +578,25 @@ class InMemoryGuardrailHandler: litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback) - def list_in_memory_guardrails(self) -> list[Guardrail]: + def list_in_memory_guardrails(self) -> List[Guardrail]: """ List all guardrails in memory """ return list(self.IN_MEMORY_GUARDRAILS.values()) - def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None: + def get_guardrail_by_id(self, guardrail_id: str) -> Optional[Guardrail]: """ Get a guardrail by its ID from memory """ return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) - def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None: + def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]: """ Return the provenance of an in-memory guardrail. """ return self._sources.get(guardrail_id) - def reconcile_db_guardrails(self, db_guardrail_ids: set[str]) -> list[str]: + def reconcile_db_guardrails(self, db_guardrail_ids: Set[str]) -> List[str]: """ Drop in-memory entries that originated from the DB but are no longer present in db_guardrail_ids. Config-loaded guardrails are never touched. @@ -617,8 +619,8 @@ class InMemoryGuardrailHandler: @staticmethod def _normalize_litellm_params_for_comparison( - params: Any | None, - ) -> dict[str, Any] | None: + params: Optional[Any], + ) -> Optional[Dict[str, Any]]: """ Render litellm_params to a canonical dict so an in-memory LitellmParams and the raw dict loaded from the DB compare equal when they describe the same @@ -682,9 +684,9 @@ class InMemoryGuardrailHandler: def reinitialize_guardrail( self, guardrail: Guardrail, - config_file_path: str | None = None, + config_file_path: Optional[str] = None, source: Literal["db", "config"] = "config", - ) -> Guardrail | None: + ) -> Optional[Guardrail]: """ Force re-initialization of a guardrail even if it exists in memory. Removes old callback from litellm.callbacks and creates fresh instance. @@ -701,7 +703,9 @@ class InMemoryGuardrailHandler: # Initialize fresh (will add new callback to litellm.callbacks) return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source) - def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None: + def sync_guardrail_from_db( + self, guardrail: Guardrail, config_file_path: Optional[str] = None + ) -> Optional[Guardrail]: """ Sync a guardrail from DB - initializes if new, re-initializes if changed. This is the method to call during DB polling.