diff --git a/litellm/build-and-push.sh b/litellm/build-and-push.sh new file mode 100644 index 00000000000..4c2bcc00155 --- /dev/null +++ b/litellm/build-and-push.sh @@ -0,0 +1,37 @@ +#!/bin/bash +set -e + +# Configuration +ECR_REPO="420049223852.dkr.ecr.eu-central-1.amazonaws.com/litellm-deepkeep" +TAG="${1:-dev-$(date +%Y%m%d-%H%M%S)}" +FULL_IMAGE="${ECR_REPO}:${TAG}" + +echo "================================================" +echo "Building LiteLLM Docker image with UI" +echo "Image: ${FULL_IMAGE}" +echo "================================================" + +# Login to ECR +echo "Logging in to ECR..." +aws ecr get-login-password --region eu-central-1 | docker login --username AWS --password-stdin 420049223852.dkr.ecr.eu-central-1.amazonaws.com + +# Build the image +echo "Building Docker image..." +docker build -t "${FULL_IMAGE}" -f Dockerfile . + +# Push the image +echo "Pushing image to ECR..." +docker push "${FULL_IMAGE}" + +echo "================================================" +echo "✅ Image pushed successfully!" +echo "" +echo "To use in Tilt, update your tilt_config.yaniv-v2.yaml:" +echo "" +echo "litellm:" +echo " image:" +echo " repository: ${ECR_REPO}" +echo " tag: ${TAG}" +echo "" +echo "Or run: tilt trigger litellm" +echo "================================================" diff --git a/litellm/deepkeep_tilt_config.yaml b/litellm/deepkeep_tilt_config.yaml new file mode 100644 index 00000000000..2cd037a523b --- /dev/null +++ b/litellm/deepkeep_tilt_config.yaml @@ -0,0 +1,22 @@ +model_list: + - model_name: fake-openai-endpoint + litellm_params: + model: openai/fake-model + api_key: fake-key + api_base: https://exampleopenaiendpoint-production.up.railway.app/ + +general_settings: + master_key: sk-1234 + +litellm_settings: + drop_params: True + telemetry: False + +guardrails: + - guardrail_name: deepkeep-firewall + litellm_params: + guardrail: deepkeep + mode: [pre_call, post_call] + api_key: "1D-VOZhxhsgG2o8pnbDcXaDmzLOWSwZpgkGZPDyVQBk" + api_base: "http://localhost:8081/api" + deepkeep_firewall_id: "063e4e5ba5c2be55" diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 296bfb6fc85..b001234f733 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -7,7 +7,7 @@ Users can define """ import copy -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Union, cast from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger @@ -36,17 +36,17 @@ class AnthropicCacheControlHook(CustomPromptManagement): def get_chat_completion_prompt( self, model: str, - messages: List[AllMessageValues], + messages: list[AllMessageValues], non_default_params: dict, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - 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]: + 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]: """ Apply cache control directives based on specified injection points. @@ -56,7 +56,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): - 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: @@ -66,8 +66,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)) @@ -98,10 +98,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 @@ -159,11 +159,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: Optional[Union[int, str]] = point.get("index", None) - targetted_index: Optional[int] = None + _targetted_index: Union[int, str] | None = point.get("index", None) + targetted_index: int | None = None if isinstance(_targetted_index, str): try: targetted_index = int(_targetted_index) @@ -249,8 +249,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): def should_run_prompt_management( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], + prompt_id: str | None, + prompt_spec: PromptSpec | None, dynamic_callback_params: StandardCallbackDynamicParams, ) -> bool: """Always return False since this is not a true prompt management system.""" @@ -258,12 +258,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): def _compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_spec: Optional[PromptSpec], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_spec: PromptSpec | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """Not used - this hook only modifies messages, doesn't fetch prompts.""" return PromptManagementClient( @@ -276,12 +276,12 @@ class AnthropicCacheControlHook(CustomPromptManagement): async def async_compile_prompt_helper( self, - prompt_id: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, - prompt_spec: Optional[PromptSpec] = None, - prompt_label: Optional[str] = None, - prompt_version: Optional[int] = None, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, ) -> PromptManagementClient: """Not used - this hook only modifies messages, doesn't fetch prompts.""" return self._compile_prompt_helper( @@ -296,19 +296,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: Optional[str], - prompt_variables: Optional[dict], + prompt_id: str | None, + prompt_variables: dict | None, dynamic_callback_params: StandardCallbackDynamicParams, litellm_logging_obj: LiteLLMLoggingObj, - 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]: + 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]: """Async version - delegates to sync since no async operations needed.""" return self.get_chat_completion_prompt( model=model, @@ -325,15 +325,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, - ) -> Optional[CustomLogger]: + non_default_params: dict, + ) -> CustomLogger | None: 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 38245a2e5ba..bc55cf2ed63 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -4,11 +4,8 @@ from typing import ( TYPE_CHECKING, Any, ClassVar, - Dict, - List, Literal, Optional, - Type, Union, get_args, ) @@ -59,7 +56,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]) -> Optional[str]: +def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None: """Extract session_id from request data (litellm_session_id or metadata).""" session_id = request_data.get("litellm_session_id") if session_id: @@ -84,20 +81,18 @@ class CustomGuardrail(CustomLogger): def __init__( self, - guardrail_name: Optional[str] = None, - supported_event_hooks: Optional[List[GuardrailEventHooks]] = None, - event_hook: Optional[ - Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode] - ] = None, + guardrail_name: str | None = None, + supported_event_hooks: list[GuardrailEventHooks] | None = None, + event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = None, default_on: bool = False, mask_request_content: bool = False, mask_response_content: bool = False, - 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, + 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, sticky_session_routing: bool = True, **kwargs, ): @@ -120,18 +115,16 @@ class CustomGuardrail(CustomLogger): """ self.guardrail_name = guardrail_name self.supported_event_hooks = supported_event_hooks - self.event_hook: Optional[ - Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode] - ] = event_hook + self.event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None = 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: 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] = ( + 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.sticky_session_routing: bool = sticky_session_routing @@ -142,14 +135,14 @@ class CustomGuardrail(CustomLogger): super().__init__(**kwargs) def render_violation_message( - self, default: str, context: Optional[Dict[str, Any]] = None + self, default: str, context: dict[str, Any] | None = 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: @@ -165,8 +158,8 @@ class CustomGuardrail(CustomLogger): def raise_passthrough_exception( self, violation_message: str, - request_data: Dict[str, Any], - detection_info: Optional[Dict[str, Any]] = None, + request_data: dict[str, Any], + detection_info: dict[str, Any] | None = None, ) -> None: """ Raise a passthrough exception for guardrail violations. @@ -209,8 +202,8 @@ class CustomGuardrail(CustomLogger): def raise_sensitive_data_route_exception( self, route_to_model: str, - request_data: Dict[str, Any], - detection_info: Optional[Dict[str, Any]] = None, + request_data: dict[str, Any], + detection_info: dict[str, Any] | None = None, ) -> None: """ Raise an exception to reroute the request to a different model. @@ -248,8 +241,8 @@ class CustomGuardrail(CustomLogger): ) def _get_session_id_from_request_data( - self, request_data: Dict[str, Any] - ) -> Optional[str]: + self, request_data: dict[str, Any] + ) -> str | None: """Extract session_id from request data.""" return get_session_id_from_request_data(request_data) @@ -265,8 +258,8 @@ class CustomGuardrail(CustomLogger): def handle_sensitive_data_detection( self, - request_data: Dict[str, Any], - detection_info: Optional[Dict[str, Any]] = None, + request_data: dict[str, Any], + detection_info: dict[str, Any] | None = None, ) -> None: """ Handle sensitive data detection based on guardrail configuration. @@ -309,7 +302,7 @@ class CustomGuardrail(CustomLogger): ) @staticmethod - def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + def get_config_model() -> type["GuardrailConfigModel"] | None: """ Returns the config model for the guardrail @@ -319,14 +312,12 @@ class CustomGuardrail(CustomLogger): def _validate_event_hook( self, - event_hook: Optional[ - Union[GuardrailEventHooks, List[GuardrailEventHooks], Mode] - ], - supported_event_hooks: List[GuardrailEventHooks], + event_hook: Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] | None, + 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): @@ -391,7 +382,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) -> Optional[bool]: + def get_disable_global_guardrail(self, data: dict) -> bool | None: """ Returns True if the global guardrail should be disabled. @@ -400,7 +391,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. @@ -433,7 +424,7 @@ class CustomGuardrail(CustomLogger): def get_guardrail_from_metadata( self, data: dict - ) -> Union[List[str], List[Dict[str, DynamicGuardrailParams]]]: + ) -> Union[list[str], list[dict[str, DynamicGuardrailParams]]]: """ Returns the guardrail(s) to be run from the metadata or root """ @@ -454,7 +445,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): @@ -466,13 +457,13 @@ class CustomGuardrail(CustomLogger): return False - def _pre_call_marker(self) -> Optional[str]: + def _pre_call_marker(self) -> str | None: 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. @@ -497,7 +488,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 @@ -510,8 +501,8 @@ class CustomGuardrail(CustomLogger): return False async def async_pre_call_deployment_hook( - self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] - ) -> Optional[dict]: + self, kwargs: dict[str, Any], call_type: CallTypes | None + ) -> dict | None: from litellm.proxy._types import UserAPIKeyAuth # should run guardrail @@ -556,8 +547,8 @@ class CustomGuardrail(CustomLogger): self, request_data: dict, response: LLMResponseTypes, - call_type: Optional[CallTypes], - ) -> Optional[LLMResponseTypes]: + call_type: CallTypes | None, + ) -> LLMResponseTypes | None: """ Allow modifying / reviewing the response just after it's received from the deployment. """ @@ -755,16 +746,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: 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, + 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, ) -> None: """ Builds `StandardLoggingGuardrailInformation` and adds it to the request metadata so it can be used for logging to DataDog, Langfuse, etc. @@ -780,7 +771,7 @@ class CustomGuardrail(CustomLogger): # Use event_type if provided, otherwise fall back to self.event_hook guardrail_mode: Union[ - GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks] + GuardrailEventHooks, GuardrailMode, list[GuardrailEventHooks] ] if event_type is not None: guardrail_mode = event_type @@ -894,13 +885,13 @@ class CustomGuardrail(CustomLogger): def _process_response( self, - response: Optional[Dict], + response: dict | None, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, - original_inputs: Optional[Dict] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, + original_inputs: dict | None = None, ): """ Add StandardLoggingGuardrailInformation to the request data @@ -908,7 +899,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] = ( + guardrail_response: Union[dict[str, Any], str] = ( {} if response is None else response ) @@ -970,10 +961,10 @@ class CustomGuardrail(CustomLogger): self, e: Exception, request_data: dict, - start_time: Optional[float] = None, - end_time: Optional[float] = None, - duration: Optional[float] = None, - event_type: Optional[GuardrailEventHooks] = None, + start_time: float | None = None, + end_time: float | None = None, + duration: float | None = None, + event_type: GuardrailEventHooks | None = None, ): """ Add StandardLoggingGuardrailInformation to the request data @@ -1002,7 +993,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. @@ -1047,8 +1038,8 @@ class CustomGuardrail(CustomLogger): setattr(self, key, value) def get_guardrails_messages_for_call_type( - self, call_type: CallTypes, data: Optional[dict] = None - ) -> Optional[List[AllMessageValues]]: + self, call_type: CallTypes, data: dict | None = None + ) -> list[AllMessageValues] | None: """ Returns the messages for the given call type and data """ @@ -1089,7 +1080,7 @@ class CustomGuardrail(CustomLogger): input=input_data, responses_api_request=data, ) - return cast(List[AllMessageValues], messages) + return cast(list[AllMessageValues], messages) return None @@ -1118,7 +1109,7 @@ def log_guardrail_information(func): def _infer_event_type_from_function_name( func_name: str, - ) -> Optional[GuardrailEventHooks]: + ) -> GuardrailEventHooks | None: """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 95ac939ff7f..cec17013737 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, FrozenSet, Mapping, Optional +from typing import Any, Mapping from opentelemetry.metrics import Histogram, Meter @@ -82,12 +82,12 @@ class GenAIMetricRecorder: """ def __init__( - self, metrics: GenAIMetrics, callback_name: Optional[str] = None + self, metrics: GenAIMetrics, callback_name: str | None = None ) -> None: self._metrics = metrics self._callback_name = callback_name - self._include: Optional[FrozenSet[str]] = None - self._exclude: Optional[FrozenSet[str]] = None + self._include: frozenset[str] | None = None + self._exclude: frozenset[str] | None = None self._filter_resolved = False def record( diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index cf97c946f1c..f8397c38fa9 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -4,11 +4,7 @@ import time from typing import ( TYPE_CHECKING, Any, - Dict, - List, NoReturn, - Optional, - Tuple, Union, cast, ) @@ -152,8 +148,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 @@ -173,7 +169,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 @@ -214,7 +210,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", @@ -237,23 +233,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): to pass metadata to anthropic, it's {"user_id": "any-relevant-information"} """ - 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 + 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 def __init__( self, - 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, + 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, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -261,11 +257,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): setattr(self.__class__, key, value) @property - def custom_llm_provider(self) -> Optional[str]: + def custom_llm_provider(self) -> str | None: return "anthropic" @classmethod - def get_config(cls, *, model: Optional[str] = None): + def get_config(cls, *, model: str | None = None): config = super().get_config() # anthropic requires a default value for max_tokens @@ -275,7 +271,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return config @staticmethod - def get_max_tokens_for_model(model: Optional[str] = None) -> int: + def get_max_tokens_for_model(model: str | None = 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. @@ -292,7 +288,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: """ @@ -317,7 +313,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 @@ -344,7 +340,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) @staticmethod - def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]: + def _validate_effort_for_model(model: str, effort: str | None) -> str | None: """Return ``None`` if ``effort`` is allowed on ``model``, else an error message.""" if effort == "max" and not ( AnthropicConfig._is_claude_4_6_model(model) @@ -434,7 +430,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. @@ -492,7 +488,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): constraint_labels[field].format(schema[field]) ) - result: Dict[str, Any] = {} + result: dict[str, Any] = {} # Update description with removed constraint info if constraint_descriptions: @@ -548,8 +544,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return result def get_json_schema_from_pydantic_object( - self, response_format: Union[Any, Dict, None] - ) -> Optional[dict]: + self, response_format: Union[Any, dict, None] + ) -> dict | None: return type_to_response_format_param( response_format, ref_template="/$defs/{model}" ) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755 @@ -564,10 +560,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tool_choice( self, - tool_choice: Optional[str], - parallel_tool_use: Optional[bool], - ) -> Optional[AnthropicMessagesToolChoice]: - _tool_choice: Optional[AnthropicMessagesToolChoice] = None + tool_choice: str | None, + parallel_tool_use: bool | None, + ) -> AnthropicMessagesToolChoice | None: + _tool_choice: AnthropicMessagesToolChoice | None = None if tool_choice == "auto": _tool_choice = AnthropicMessagesToolChoice( type="auto", @@ -608,9 +604,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _map_tool_helper( self, tool: ChatCompletionToolParam, - ) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]: - returned_tool: Optional[AllAnthropicToolsValues] = None - mcp_server: Optional[AnthropicMcpServerTool] = None + ) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]: + returned_tool: AllAnthropicToolsValues | None = None + mcp_server: AnthropicMcpServerTool | None = None if tool["type"] == "function" or tool["type"] == "custom": _input_schema: dict = tool["function"].get( @@ -665,10 +661,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if "parameters" not in tool["function"]: raise ValueError("Missing required parameter: parameters") - _display_width_px: Optional[int] = tool["function"]["parameters"].get( + _display_width_px: int | None = tool["function"]["parameters"].get( "display_width_px" ) - _display_height_px: Optional[int] = tool["function"]["parameters"].get( + _display_height_px: int | None = tool["function"]["parameters"].get( "display_height_px" ) if _display_width_px is None or _display_height_px is None: @@ -841,14 +837,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration allowed_tools = tool.get("allowed_tools", None) - tool_configuration: Optional[AnthropicMcpServerToolConfiguration] = None + tool_configuration: AnthropicMcpServerToolConfiguration | None = None if allowed_tools is not None: tool_configuration = AnthropicMcpServerToolConfiguration( allowed_tools=tool.get("allowed_tools", None), ) headers = tool.get("headers", {}) - authorization_token: Optional[str] = None + authorization_token: str | None = None if headers is not None: bearer_token = headers.get("Authorization", None) if bearer_token is not None: @@ -868,8 +864,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: @@ -918,9 +914,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. @@ -931,7 +927,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) @@ -979,8 +975,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 @@ -994,7 +990,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 @@ -1011,8 +1007,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}$``. @@ -1036,7 +1032,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 @@ -1058,7 +1054,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) @@ -1084,7 +1080,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return forward, reverse - def _detect_tool_search_tools(self, tools: Optional[List]) -> bool: + def _detect_tool_search_tools(self, tools: list | None) -> bool: """Check if tool search tools are present in the tools list.""" if not tools: return False @@ -1098,7 +1094,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. @@ -1118,9 +1114,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. @@ -1162,9 +1158,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return expanded_content def _map_stop_sequences( - self, stop: Optional[Union[str, List[str]]] - ) -> Optional[List[str]]: - new_stop: Optional[List[str]] = None + self, stop: Union[str, list[str]] | None + ) -> list[str] | None: + new_stop: list[str] | None = None if isinstance(stop, str): if ( stop.isspace() and litellm.drop_params is True @@ -1185,10 +1181,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _map_reasoning_effort( - reasoning_effort: Optional[Union[REASONING_EFFORT, str]], + reasoning_effort: Union[REASONING_EFFORT, str] | None, model: str, llm_provider: str = "anthropic", - ) -> Optional[AnthropicThinkingParam]: + ) -> AnthropicThinkingParam | None: if reasoning_effort is None or reasoning_effort == "none": return None if AnthropicConfig._is_adaptive_thinking_model(model): @@ -1240,11 +1236,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) def _extract_json_schema_from_response_format( - self, value: Optional[dict] - ) -> Optional[dict]: + self, value: dict | None + ) -> dict | None: if value is None: return None - json_schema: Optional[dict] = None + json_schema: dict | None = None if "response_schema" in value: json_schema = value["response_schema"] elif "json_schema" in value: @@ -1253,9 +1249,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return json_schema 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( + self, value: dict | None + ) -> AnthropicOutputSchema | None: + json_schema: dict | None = self._extract_json_schema_from_response_format( value ) if json_schema is None: @@ -1283,15 +1279,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) def map_response_format_to_anthropic_tool( - self, value: Optional[dict], optional_params: dict, is_thinking_enabled: bool - ) -> Optional[AnthropicMessagesTool]: + self, value: dict | None, optional_params: dict, is_thinking_enabled: bool + ) -> AnthropicMessagesTool | None: 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: Optional[dict] = self._extract_json_schema_from_response_format( + json_schema: dict | None = self._extract_json_schema_from_response_format( value ) if json_schema is None: @@ -1342,8 +1338,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def map_openai_context_management_to_anthropic( - context_management: Union[List[Dict[str, Any]], Dict[str, Any]], - ) -> Optional[Dict[str, Any]]: + context_management: Union[list[dict[str, Any]], dict[str, Any]], + ) -> dict[str, Any] | None: """ OpenAI format: [{"type": "compaction", "compact_threshold": 200000}] Anthropic format: { @@ -1374,7 +1370,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( @@ -1438,7 +1434,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: Optional[AnthropicMessagesToolChoice] = ( + _tool_choice: AnthropicMessagesToolChoice | None = ( self._map_tool_choice( tool_choice=non_default_params.get("tool_choice"), parallel_tool_use=non_default_params.get("parallel_tool_calls"), @@ -1584,7 +1580,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _create_json_tool_call_for_response_format( self, - json_schema: Optional[dict] = None, + json_schema: dict | None = None, ) -> AnthropicMessagesTool: """ Handles creating a tool call for getting responses in JSON format. @@ -1622,8 +1618,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return False def translate_system_message( - self, messages: List[AllMessageValues] - ) -> List[AnthropicSystemMessageContent]: + self, messages: list[AllMessageValues] + ) -> list[AnthropicSystemMessageContent]: """ Translate system message to anthropic format. @@ -1631,7 +1627,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) @@ -1691,9 +1687,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: @@ -1832,7 +1828,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, @@ -1938,7 +1934,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ## Add code_execution tool if container_upload is in messages _tools = ( cast( - Optional[List[Union[AllAnthropicToolsValues, Dict]]], + list[Union[AllAnthropicToolsValues, dict]] | None, optional_params.get("tools"), ) or [] @@ -2045,12 +2041,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _resolve_json_mode_non_streaming( self, - json_mode: Optional[bool], - tool_calls: List[ChatCompletionToolCallChunk], - ) -> Tuple[ - Optional[LitellmMessage], - List[ChatCompletionToolCallChunk], - Optional[str], + json_mode: bool | None, + tool_calls: list[ChatCompletionToolCallChunk], + ) -> tuple[ + LitellmMessage | None, + list[ChatCompletionToolCallChunk], + str | None, ]: """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: @@ -2075,38 +2071,30 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): first_json = tool_calls[json_indices[0]] json_msg = AnthropicConfig._convert_tool_response_to_message([first_json]) - extra_content: Optional[str] = ( + extra_content: str | None = ( 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[ + def extract_response_content(self, completion_response: dict) -> tuple[ str, - Optional[List[Any]], - Optional[ - List[ - Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock] - ] - ], - Optional[str], - List[ChatCompletionToolCallChunk], - Optional[List[Any]], - Optional[List[Any]], - Optional[List[Any]], + list[Any] | None, + list[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] | None, + str | None, + list[ChatCompletionToolCallChunk], + list[Any] | None, + list[Any] | None, + list[Any] | None, ]: text_content = "" - 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 + 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 for idx, content in enumerate(completion_response["content"]): if content["type"] == "text": text_content += content["text"] @@ -2171,7 +2159,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if thinking_blocks is not None: reasoning_content = "" for block in thinking_blocks: - thinking_content = cast(Optional[str], block.get("thinking")) + thinking_content = cast(str | None, block.get("thinking")) if thinking_content is not None: reasoning_content += thinking_content @@ -2189,9 +2177,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def calculate_usage( self, usage_object: dict, - reasoning_content: Optional[str], - completion_response: Optional[dict] = None, - speed: Optional[str] = None, + reasoning_content: str | None, + completion_response: dict | None = None, + speed: str | None = 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 @@ -2207,14 +2195,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _usage = usage_object cache_creation_input_tokens: int = 0 cache_read_input_tokens: int = 0 - cache_creation_token_details: Optional[CacheCreationTokenDetails] = None - web_search_requests: Optional[int] = None - tool_search_requests: Optional[int] = None - inference_geo: Optional[str] = None + cache_creation_token_details: CacheCreationTokenDetails | None = None + web_search_requests: int | None = None + tool_search_requests: int | None = None + inference_geo: str | None = None if "inference_geo" in _usage and _usage["inference_geo"] is not None: inference_geo = _usage["inference_geo"] - iterations: Optional[List[Any]] = _usage.get("iterations") + iterations: list[Any] | None = _usage.get("iterations") if iterations: prompt_tokens = sum(it.get("input_tokens", 0) or 0 for it in iterations) completion_tokens = sum( @@ -2328,9 +2316,9 @@ 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] = {} + 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", "{}")) @@ -2344,10 +2332,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _build_code_interpreter_results( self, - tool_results: List[Any], - code_by_id: Dict[str, str], - container_id: Optional[str], - ) -> List[OutputCodeInterpreterCall]: + tool_results: list[Any], + code_by_id: dict[str, str], + container_id: str | None, + ) -> list[OutputCodeInterpreterCall]: code_interpreter_results = [] for tr in tool_results: if tr.get("type") != "bash_code_execution_tool_result": @@ -2370,18 +2358,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): def _build_provider_specific_fields( self, completion_response: dict, - 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: 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": citations, "thinking_blocks": thinking_blocks, } @@ -2422,12 +2406,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): completion_response: dict, raw_response: httpx.Response, model_response: ModelResponse, - json_mode: Optional[bool] = None, - prefix_prompt: Optional[str] = None, - speed: Optional[str] = None, - tool_name_reverse_map: Optional[Dict[str, str]] = None, + json_mode: bool | None = None, + prefix_prompt: str | None = None, + speed: str | None = None, + tool_name_reverse_map: dict[str, str] | None = None, ): - _hidden_params: Dict = {} + _hidden_params: dict = {} _hidden_params["additional_headers"] = process_anthropic_headers( dict(raw_response.headers) ) @@ -2531,7 +2515,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): model_response._hidden_params = _hidden_params return model_response - def get_prefix_prompt(self, messages: List[AllMessageValues]) -> Optional[str]: + def get_prefix_prompt(self, messages: list[AllMessageValues]) -> str | None: """ Get the prefix prompt from the messages. @@ -2560,13 +2544,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: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -2591,7 +2575,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): prefix_prompt = self.get_prefix_prompt(messages=messages) speed = optional_params.get("speed") - tool_name_reverse_map: Optional[Dict[str, str]] = None + tool_name_reverse_map: dict[str, str] | None = None if isinstance(litellm_params, dict): _candidate = litellm_params.get(ANTHROPIC_TOOL_NAME_REVERSE_MAP_KEY) if isinstance(_candidate, dict): @@ -2610,14 +2594,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _convert_tool_response_to_message( - tool_calls: List[ChatCompletionToolCallChunk], - ) -> Optional[LitellmMessage]: + tool_calls: list[ChatCompletionToolCallChunk], + ) -> LitellmMessage | None: """ 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: Optional[str] = tool_calls[0]["function"].get( + json_mode_content_str: str | None = tool_calls[0]["function"].get( "arguments" ) try: @@ -2640,7 +2624,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 162ef1a236c..056a0c989fc 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, Optional, Tuple, Union, cast, overload +from typing import Any, Coroutine, Literal, Union, cast, overload from typing_extensions import override @@ -63,8 +63,8 @@ class ModelScopeChatConfig(OpenAIGPTConfig): ) def _get_openai_compatible_provider_info( - self, api_base: Optional[str], api_key: Optional[str] - ) -> Tuple[Optional[str], Optional[str]]: + self, api_base: str | None, api_key: str | None + ) -> tuple[str | None, str | None]: api_base = ( api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL ) # type: ignore @@ -74,12 +74,12 @@ class ModelScopeChatConfig(OpenAIGPTConfig): @override def get_complete_url( self, - api_base: Optional[str], - api_key: Optional[str], + api_base: str | None, + api_key: str | None, model: str, optional_params: dict, litellm_params: dict, - stream: Optional[bool] = None, + stream: bool | None = 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 dab21e2ce8e..c9131dfb137 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,13 +9,9 @@ from typing import ( TYPE_CHECKING, Any, Callable, - Dict, - List, Literal, Mapping, Optional, - Tuple, - Type, Union, cast, ) @@ -135,7 +131,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 """ @@ -152,7 +148,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 """ @@ -198,29 +194,29 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): Note: Please make sure to modify the default parameters as required for your use case. """ - 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 + 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 def __init__( self, - 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, + 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, ) -> None: locals_ = locals().copy() for key, value in locals_.items(): @@ -232,8 +228,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return super().get_config() def get_json_schema_from_pydantic_object( - self, response_format: Optional[Union[Type["BaseModel"], dict]] - ) -> Optional[dict]: + self, response_format: Union[type["BaseModel"], dict] | None + ) -> dict | None: """ Override to use Pydantic's model_json_schema() instead of OpenAI's to_strict_json_schema(). @@ -292,7 +288,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _forward_gemini_function_call_id( - model: str, custom_llm_provider: Optional[str] = None + model: str, custom_llm_provider: str | None = None ) -> bool: """ Whether to include `id` on function_call / function_response parts. @@ -313,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", @@ -350,7 +346,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def map_tool_choice_values( self, model: str, tool_choice: Union[str, dict] - ) -> Optional[ToolConfig]: + ) -> ToolConfig | None: if tool_choice == "none": return ToolConfig(functionCallingConfig=FunctionCallingConfig(mode="NONE")) elif tool_choice == "required": @@ -494,7 +490,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _extract_google_maps_retrieval_config( self, google_maps_config: dict - ) -> Tuple[dict, Optional[dict]]: + ) -> tuple[dict, dict | None]: """ Extract location configuration from googleMaps tool for Vertex AI toolConfig. @@ -534,7 +530,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return cleaned_config, retrieval_config - def get_tool_value(self, tool: dict, tool_name: str) -> Optional[dict]: + def get_tool_value(self, tool: dict, tool_name: str) -> dict | None: """ Helper function to get tool value handling both camelCase and underscore_case variants @@ -561,10 +557,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _resolve_search_tool_conflict( gtool_func_declarations: list, - googleSearch: Optional[dict], - googleSearchRetrieval: Optional[dict], - enterpriseWebSearch: Optional[dict], - urlContext: Optional[dict], + googleSearch: dict | None, + googleSearchRetrieval: dict | None, + enterpriseWebSearch: dict | None, + urlContext: dict | None, optional_params: dict, ) -> tuple: """ @@ -614,7 +610,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. @@ -630,21 +626,21 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): googleMaps tools contain location data """ gtool_func_declarations = [] - 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 + 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 # 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: Optional[ChatCompletionToolParamFunctionChunk] = ( + openai_function_object: ChatCompletionToolParamFunctionChunk | None = ( None ) if "function" in tool: # tools list @@ -763,7 +759,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, @@ -897,7 +893,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_budget( reasoning_effort: str, - model: Optional[str] = None, + model: str | None = None, ) -> GeminiThinkingConfig: if reasoning_effort == "minimal": # Use model-specific minimum thinking budget or fallback @@ -948,7 +944,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_reasoning_effort_to_thinking_level( reasoning_effort: str, - model: Optional[str] = None, + model: str | None = None, ) -> GeminiThinkingConfig: """ Map reasoning_effort to thinking_level for Gemini 3+ models. @@ -996,12 +992,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): raise ValueError(f"Invalid reasoning effort: {reasoning_effort}") @staticmethod - def _is_thinking_budget_zero(thinking_budget: Optional[int]) -> bool: + def _is_thinking_budget_zero(thinking_budget: int | None) -> 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: @@ -1022,7 +1018,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. @@ -1042,7 +1038,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _map_thinking_param( thinking_param: AnthropicThinkingParam, - model: Optional[str] = None, + model: str | None = None, ) -> GeminiThinkingConfig: thinking_enabled = thinking_param.get("type") == "enabled" thinking_budget = thinking_param.get("budget_tokens") @@ -1153,8 +1149,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. @@ -1173,11 +1169,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 ) @@ -1290,7 +1286,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: Optional[str] = None + effort_value: str | None = None if isinstance(value, str): effort_value = value elif isinstance(value, dict): @@ -1374,7 +1370,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 """ @@ -1411,7 +1407,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return model @staticmethod - def _is_model_gemini_spec_model(model: Optional[str]) -> bool: + def _is_model_gemini_spec_model(model: str | None) -> bool: """ Returns true if user is trying to call custom model in `/gemini` request/response format """ @@ -1434,7 +1430,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 @@ -1471,7 +1467,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. @@ -1495,10 +1491,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return exception_string 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 + self, parts: list[HttpxPartType] + ) -> tuple[str | None, str | None]: + content_str: str | None = None + reasoning_content_str: str | None = None for part in parts: _content_str = "" @@ -1540,8 +1536,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return content_str, reasoning_content_str def _extract_thinking_blocks_from_parts( - self, parts: List[HttpxPartType] - ) -> List[ChatCompletionThinkingBlock]: + 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): @@ -1550,7 +1546,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", "") @@ -1565,8 +1561,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return thinking_blocks def _extract_thought_signatures_from_parts( - self, parts: List[HttpxPartType] - ) -> Optional[List[str]]: + self, parts: list[HttpxPartType] + ) -> list[str] | None: """Extract thoughtSignature values from parts. Per Google's docs, thoughtSignature is returned for multi-turn context preservation @@ -1576,7 +1572,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: @@ -1585,8 +1581,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _extract_server_side_tool_invocations( - parts: List[HttpxPartType], - ) -> Optional[List[Dict[str, Any]]]: + parts: list[HttpxPartType], + ) -> list[dict[str, Any]] | None: """Extract server-side tool invocations (toolCall/toolResponse) from parts. These are returned by Gemini when context circulation is enabled @@ -1597,15 +1593,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"), @@ -1645,10 +1641,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return invocations if invocations else None def _extract_image_response_from_parts( - self, parts: List[HttpxPartType] - ) -> Optional[List[ImageURLListItem]]: + self, parts: list[HttpxPartType] + ) -> list[ImageURLListItem] | None: """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", {}) @@ -1667,8 +1663,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return images def _extract_audio_response_from_parts( - self, parts: List[HttpxPartType] - ) -> Optional[ChatCompletionAudioResponse]: + self, parts: list[HttpxPartType] + ) -> ChatCompletionAudioResponse | None: """Extract audio response from parts if present""" for part in parts: if "text" in part: @@ -1710,16 +1706,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _transform_parts( - parts: List[HttpxPartType], + parts: list[HttpxPartType], cumulative_tool_call_idx: int, - is_function_call: Optional[bool], - ) -> Tuple[ - Optional[ChatCompletionToolCallFunctionChunk], - Optional[List[ChatCompletionToolCallChunk]], + is_function_call: bool | None, + ) -> tuple[ + ChatCompletionToolCallFunctionChunk | None, + list[ChatCompletionToolCallChunk] | None, int, ]: - function: Optional[ChatCompletionToolCallFunctionChunk] = None - _tools: List[ChatCompletionToolCallChunk] = [] + function: ChatCompletionToolCallFunctionChunk | None = None + _tools: list[ChatCompletionToolCallChunk] = [] for part in parts: if "functionCall" in part: _function_chunk: ChatCompletionToolCallFunctionChunk = { @@ -1736,7 +1732,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"] = {} @@ -1769,22 +1765,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): _tools.append(_tool_response_chunk) cumulative_tool_call_idx += 1 if len(_tools) == 0: - tools: Optional[List[ChatCompletionToolCallChunk]] = None + tools: list[ChatCompletionToolCallChunk] | None = None else: tools = _tools return function, tools, cumulative_tool_call_idx @staticmethod def _transform_logprobs( - logprobs_result: Optional[LogprobsResult], - ) -> Optional[ChoiceLogprobs]: + logprobs_result: LogprobsResult | None, + ) -> ChoiceLogprobs | None: 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"] ): @@ -1914,16 +1910,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): raise ValueError( f"usageMetadata not found in completion_response. Got={completion_response}" ) - cached_tokens: Optional[int] = None + cached_tokens: int | None = None # Separate variables for prompt tokens by modality - 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 + 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 usage_metadata = completion_response["usageMetadata"] def _get_token_count(detail: Mapping[str, Any]) -> int: @@ -2021,10 +2017,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: Optional[int] = None - cached_audio_tokens: Optional[int] = None - cached_image_tokens: Optional[int] = None - cached_video_tokens: Optional[int] = None + cached_text_tokens: int | None = None + cached_audio_tokens: int | None = None + cached_image_tokens: int | None = None + cached_video_tokens: int | None = None if "cacheTokensDetails" in usage_metadata: for detail in usage_metadata["cacheTokensDetails"]: @@ -2102,8 +2098,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _check_finish_reason( - chat_completion_message: Optional[ChatCompletionResponseMessage], - finish_reason: Optional[str], + chat_completion_message: ChatCompletionResponseMessage | None, + finish_reason: str | None, ) -> OpenAIChatCompletionFinishReason: from litellm.litellm_core_utils.core_helpers import map_finish_reason @@ -2119,7 +2115,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _check_prompt_level_content_filter( processed_chunk: GenerateContentResponseBody, - response_id: Optional[str], + response_id: str | None, ) -> Optional["ModelResponseStream"]: """ Check if prompt is blocked due to content filtering at the prompt level. @@ -2163,8 +2159,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return None @staticmethod - def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]: - web_search_requests: Optional[int] = None + def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None: + web_search_requests: int | None = None if ( grounding_metadata @@ -2184,10 +2180,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): chat_completion_message: ChatCompletionResponseMessage, candidate: Candidates, idx: int, - tools: Optional[List[ChatCompletionToolCallChunk]], - functions: Optional[ChatCompletionToolCallFunctionChunk], - chat_completion_logprobs: Optional[ChoiceLogprobs], - image_response: Optional[List[ImageURLListItem]], + tools: list[ChatCompletionToolCallChunk] | None, + functions: ChatCompletionToolCallFunctionChunk | None, + chat_completion_logprobs: ChoiceLogprobs | None, + image_response: list[ImageURLListItem] | None, ) -> StreamingChoices: """ Helper method to create a streaming choice object for Vertex AI @@ -2219,7 +2215,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. @@ -2229,10 +2225,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): @@ -2277,10 +2273,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: @@ -2306,10 +2302,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: @@ -2324,14 +2320,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _convert_grounding_metadata_to_annotations( - grounding_metadata: List[dict], - content_text: Optional[str], - ) -> List[ChatCompletionAnnotation]: + grounding_metadata: list[dict], + content_text: str | None, + ) -> 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 @@ -2339,7 +2335,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"] @@ -2379,11 +2375,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 @@ -2399,19 +2395,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) from litellm.types.utils import ModelResponseStream - grounding_metadata: List[dict] = [] - url_context_metadata: List[dict] = [] - image_response: Optional[List[ImageURLListItem]] = None - safety_ratings: List = [] - citation_metadata: List = [] + grounding_metadata: list[dict] = [] + url_context_metadata: list[dict] = [] + image_response: list[ImageURLListItem] | None = None + safety_ratings: list = [] + citation_metadata: list = [] chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} - 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 + 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 for idx, candidate in enumerate(_candidates): if "content" not in candidate: @@ -2470,13 +2466,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) if audio_response is not None: - cast(Dict[str, Any], chat_completion_message)[ + 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)[ + cast(dict[str, Any], chat_completion_message)[ "images" ] = image_response if content is not None: @@ -2583,13 +2579,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: Optional[str] = None, - json_mode: Optional[bool] = None, + api_key: str | None = None, + json_mode: bool | None = None, ) -> ModelResponse: ## LOGGING logging_obj.post_call( @@ -2664,11 +2660,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, @@ -2750,10 +2746,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): def _transform_messages( self, - messages: List[AllMessageValues], - model: Optional[str] = None, - litellm_params: Optional[dict] = None, - ) -> List[ContentType]: + messages: list[AllMessageValues], + model: str | None = None, + litellm_params: dict | None = None, + ) -> list[ContentType]: return _gemini_convert_messages_with_history( messages=messages, model=model, @@ -2762,7 +2758,7 @@ 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 @@ -2771,25 +2767,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): 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: Optional[Dict], + headers: dict | None, model: str, - messages: List[AllMessageValues], - optional_params: Dict, - litellm_params: Dict, - api_key: Optional[Union[str, Dict]] = None, - api_base: Optional[str] = None, - ) -> Dict: + messages: list[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Union[str, dict] | None = None, + api_base: str | None = None, + ) -> dict: default_headers = { "Content-Type": "application/json", } @@ -2804,8 +2800,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): async def make_call( - client: Optional[AsyncHTTPHandler], # module-level client - gemini_client: Optional[AsyncHTTPHandler], # if passed by user + client: AsyncHTTPHandler | None, # module-level client + gemini_client: AsyncHTTPHandler | None, # if passed by user api_base: str, headers: dict, data: str, @@ -2857,8 +2853,8 @@ async def make_call( def make_sync_call( - client: Optional[HTTPHandler], # module-level client - gemini_client: Optional[HTTPHandler], # if passed by user + client: HTTPHandler | None, # module-level client + gemini_client: HTTPHandler | None, # if passed by user api_base: str, headers: dict, data: str, @@ -2914,20 +2910,20 @@ class VertexLLM(VertexBase): model_response: ModelResponse, print_verbose: Callable, data: dict, - timeout: Optional[Union[float, httpx.Timeout]], + timeout: Union[float, httpx.Timeout] | None, encoding, logging_obj, stream, optional_params: dict, litellm_params: dict, logger_fn=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, + 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, ) -> CustomStreamWrapper: should_use_v1beta1_features = self.is_using_v1beta1_features( optional_params=optional_params @@ -3016,20 +3012,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: Optional[Union[float, httpx.Timeout]], + timeout: Union[float, httpx.Timeout] | None, encoding, logging_obj, stream, optional_params: dict, litellm_params: dict, logger_fn=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, + 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, ) -> Union[ModelResponse, CustomStreamWrapper]: should_use_v1beta1_features = self.is_using_v1beta1_features( optional_params=optional_params @@ -3142,18 +3138,18 @@ class VertexLLM(VertexBase): logging_obj, optional_params: dict, acompletion: bool, - 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], + 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, litellm_params: dict, logger_fn=None, - extra_headers: Optional[dict] = None, - client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None, - api_base: Optional[str] = None, + extra_headers: dict | None = None, + client: Union[AsyncHTTPHandler, HTTPHandler] | None = None, + api_base: str | None = None, ) -> Union[ModelResponse, CustomStreamWrapper]: - stream: Optional[bool] = optional_params.pop("stream", None) # type: ignore + stream: bool | None = optional_params.pop("stream", None) # type: ignore transform_request_params = { "gemini_api_key": gemini_api_key, @@ -3347,7 +3343,7 @@ class ModelResponseIterator: streaming_response, sync_stream: bool, logging_obj: LoggingClass, - response_headers: Optional[Dict[str, str]] = None, + response_headers: dict[str, str] | None = None, ): from litellm.litellm_core_utils.prompt_templates.common_utils import ( check_is_function_call, @@ -3390,9 +3386,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, @@ -3477,8 +3473,8 @@ class ModelResponseIterator: self, processed_chunk: Any, model_response: Any, - grounding_metadata: List[dict], - ) -> Optional[Usage]: + grounding_metadata: list[dict], + ) -> Usage | None: if "usageMetadata" not in processed_chunk: return None @@ -3531,12 +3527,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: Optional[List[Candidates]] = processed_chunk.get("candidates") + _candidates: list[Candidates] | None = processed_chunk.get("candidates") if _candidates: ( grounding_metadata, diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 7db7b5cde85..d40124b9921 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -6,7 +6,7 @@ # +-------------------------------------------------------------+ import os -from typing import TYPE_CHECKING, Any, Dict, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal, Optional import httpx @@ -68,11 +68,11 @@ class DeepKeepGuardrail(CustomGuardrail): def __init__( self, - api_key: Optional[str] = None, - api_base: Optional[str] = None, - firewall_id: Optional[str] = None, + api_key: str | None = None, + api_base: str | None = None, + firewall_id: str | None = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", - extra_headers: Optional[Dict[str, str]] = None, + extra_headers: dict[str, str] | None = None, **kwargs, ): self.async_handler = get_async_httpx_client( @@ -114,7 +114,7 @@ class DeepKeepGuardrail(CustomGuardrail): self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( unreachable_fallback ) - self.extra_headers: Dict[str, str] = extra_headers or {} + self.extra_headers: dict[str, str] = extra_headers or {} # Set supported event hooks if "supported_event_hooks" not in kwargs: @@ -133,7 +133,7 @@ class DeepKeepGuardrail(CustomGuardrail): self.firewall_id, ) - def _extract_user_api_key_metadata(self, request_data: dict) -> Dict[str, Any]: + def _extract_user_api_key_metadata(self, request_data: dict) -> dict[str, Any]: """ Extract user API key metadata from request_data for the DeepKeep API. @@ -143,7 +143,7 @@ class DeepKeepGuardrail(CustomGuardrail): Returns: Dictionary with user API key metadata fields. """ - result_metadata: Dict[str, Any] = {} + result_metadata: dict[str, Any] = {} litellm_metadata = request_data.get("litellm_metadata", {}) top_level_metadata = request_data.get("metadata", {}) @@ -177,9 +177,9 @@ class DeepKeepGuardrail(CustomGuardrail): return result_metadata - def _build_request_headers(self) -> Dict[str, str]: + def _build_request_headers(self) -> dict[str, str]: """Build HTTP headers for the DeepKeep API request.""" - headers: Dict[str, str] = { + headers: dict[str, str] = { "Content-Type": "application/json", "X-API-Key": self.deepkeep_api_key, } @@ -194,7 +194,7 @@ class DeepKeepGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"], error: Exception, - http_status_code: Optional[int] = None, + http_status_code: int | None = None, ) -> GenericGuardrailAPIInputs: """Allow the request to proceed when the guardrail is unreachable (fail-open mode).""" status_suffix = ( @@ -241,12 +241,12 @@ class DeepKeepGuardrail(CustomGuardrail): @staticmethod def _build_return_inputs( *, - response_json: Dict[str, Any], + response_json: dict[str, Any], texts: list, - images: Optional[Any], - tools: Optional[Any], - tool_calls: Optional[Any], - structured_messages: Optional[Any], + images: Any | None, + tools: Any | None, + tool_calls: Any | None, + structured_messages: Any | None, ) -> GenericGuardrailAPIInputs: """Merge original inputs with any guardrail-modified values from the API response.""" return_inputs = GenericGuardrailAPIInputs(texts=texts) @@ -311,7 +311,7 @@ class DeepKeepGuardrail(CustomGuardrail): request_body = request_data.get("body") or {} # Merge additional provider-specific params from config and dynamic params - additional_params: Dict[str, Any] = {"firewall_id": self.firewall_id} + additional_params: dict[str, Any] = {"firewall_id": self.firewall_id} dynamic_params = self.get_guardrail_dynamic_request_body_params(request_body) if dynamic_params: additional_params.update( @@ -322,7 +322,7 @@ class DeepKeepGuardrail(CustomGuardrail): user_metadata = self._extract_user_api_key_metadata(request_data) # Build request payload - guardrail_request: Dict[str, Any] = { + guardrail_request: dict[str, Any] = { "litellm_call_id": (logging_obj.litellm_call_id if logging_obj else None), "litellm_trace_id": (logging_obj.litellm_trace_id if logging_obj else None), "texts": texts, @@ -397,7 +397,7 @@ class DeepKeepGuardrail(CustomGuardrail): ) @staticmethod - def get_config_model() -> Optional[type]: + def get_config_model() -> type | None: from litellm.types.proxy.guardrails.guardrail_hooks.deepkeep import ( DeepKeepGuardrailConfigModel, ) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index b99ea8f14a0..1e534590098 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, Dict, List, Literal, Optional, Set, Type, cast +from typing import Any, Literal, Optional, 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 } @@ -224,7 +224,7 @@ class GuardrailRegistry: ############################################################ def get_initialized_guardrail_callback( self, guardrail_name: str - ) -> Optional[CustomGuardrail]: + ) -> CustomGuardrail | None: """ Returns the initialized guardrail callback for a given guardrail name """ @@ -334,7 +334,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). @@ -347,7 +347,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 @@ -357,7 +357,7 @@ class GuardrailRegistry: async def get_guardrail_by_id_from_db( self, guardrail_id: str, prisma_client: PrismaClient - ) -> Optional[Guardrail]: + ) -> Guardrail | None: """ Get a guardrail by its ID from the database """ @@ -375,7 +375,7 @@ class GuardrailRegistry: async def get_guardrail_by_name_from_db( self, guardrail_name: str, prisma_client: PrismaClient - ) -> Optional[Guardrail]: + ) -> Guardrail | None: """ Get a guardrail by its name from the database """ @@ -398,17 +398,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, Optional[CustomGuardrail]] = {} + self.guardrail_id_to_custom_guardrail: dict[str, CustomGuardrail | None] = {} """ 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 @@ -418,10 +418,10 @@ class InMemoryGuardrailHandler: def initialize_guardrail( self, guardrail: Guardrail, - config_file_path: Optional[str] = None, + config_file_path: str | None = None, llm_router: Optional["Router"] = None, source: Literal["db", "config"] = "config", - ) -> Optional[Guardrail]: + ) -> Guardrail | None: """ Initialize a guardrail from a dictionary and add it to the litellm callback manager @@ -439,7 +439,7 @@ class InMemoryGuardrailHandler: self._sources[guardrail_id] = source return self.IN_MEMORY_GUARDRAILS[guardrail_id] - custom_guardrail_callback: Optional[CustomGuardrail] = None + custom_guardrail_callback: CustomGuardrail | None = None litellm_params_data = guardrail["litellm_params"] verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) @@ -520,11 +520,11 @@ class InMemoryGuardrailHandler: def initialize_custom_guardrail( self, - guardrail: Dict, + guardrail: dict, guardrail_type: str, litellm_params: LitellmParams, - config_file_path: Optional[str] = None, - ) -> Optional[CustomGuardrail]: + config_file_path: str | None = None, + ) -> CustomGuardrail | None: """ Initialize a Custom Guardrail from a python file or module path @@ -623,25 +623,25 @@ class InMemoryGuardrailHandler: 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) -> Optional[Guardrail]: + def get_guardrail_by_id(self, guardrail_id: str) -> Guardrail | None: """ Get a guardrail by its ID from memory """ return self.IN_MEMORY_GUARDRAILS.get(guardrail_id) - def get_source(self, guardrail_id: str) -> Optional[Literal["db", "config"]]: + def get_source(self, guardrail_id: str) -> Literal["db", "config"] | None: """ 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. @@ -665,8 +665,8 @@ class InMemoryGuardrailHandler: @staticmethod def _normalize_litellm_params_for_comparison( - params: Optional[Any], - ) -> Optional[Dict[str, Any]]: + params: Any | None, + ) -> dict[str, Any] | None: """ 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 @@ -738,9 +738,9 @@ class InMemoryGuardrailHandler: def reinitialize_guardrail( self, guardrail: Guardrail, - config_file_path: Optional[str] = None, + config_file_path: str | None = None, source: Literal["db", "config"] = "config", - ) -> Optional[Guardrail]: + ) -> Guardrail | None: """ Force re-initialization of a guardrail even if it exists in memory. Removes old callback from litellm.callbacks and creates fresh instance. @@ -762,8 +762,8 @@ class InMemoryGuardrailHandler: ) def sync_guardrail_from_db( - self, guardrail: Guardrail, config_file_path: Optional[str] = None - ) -> Optional[Guardrail]: + self, guardrail: Guardrail, config_file_path: str | None = None + ) -> Guardrail | None: """ Sync a guardrail from DB - initializes if new, re-initializes if changed. This is the method to call during DB polling. diff --git a/scripts/check_any_discipline.py b/scripts/check_any_discipline.py index 3185953d473..de2976c7c7c 100644 --- a/scripts/check_any_discipline.py +++ b/scripts/check_any_discipline.py @@ -151,8 +151,16 @@ class Violation(NamedTuple): # --------------------------------------------------------------------------- # -def contains_any(t: Type, _seen: set[int] | None = None) -> bool: +_MAX_CONTAINS_ANY_DEPTH = 64 + + +def contains_any(t: Type, _seen: set[int] | None = None, _depth: int = 0) -> bool: """True if a *value* of type ``t`` carries `Any` anywhere meaningful.""" + if _depth > _MAX_CONTAINS_ANY_DEPTH: + # Bail out on deeply-nested / potentially-circular types rather than + # overflowing the Python call stack. A type this deep is unlikely to + # carry a *meaningful* Any that the developer could actually fix. + return False seen = _seen if _seen is not None else set() p = get_proper_type(t) if id(p) in seen: @@ -166,11 +174,11 @@ def contains_any(t: Type, _seen: set[int] | None = None) -> bool: if isinstance(p, AnyType): return p.type_of_any not in _HARMLESS_ANY if isinstance(p, UnionType): - return any(contains_any(item, seen) for item in p.items) + return any(contains_any(item, seen, _depth + 1) for item in p.items) if isinstance(p, Instance): - return any(contains_any(arg, seen) for arg in p.args) + return any(contains_any(arg, seen, _depth + 1) for arg in p.args) if isinstance(p, TupleType): - return any(contains_any(item, seen) for item in p.items) + return any(contains_any(item, seen, _depth + 1) for item in p.items) return False @@ -249,17 +257,11 @@ def _reason_ok(reason: str | None) -> bool: return reason is not None and len(reason.strip()) >= MIN_REASON_LEN -def scan_any_ok( - path: Path, source: str -) -> tuple[frozenset[int], tuple[Violation, ...]]: +def scan_any_ok(path: Path, source: str) -> tuple[frozenset[int], tuple[Violation, ...]]: """Return (lines with a valid any-ok suppression, LIT005 violations).""" try: - tokens = tokenize.generate_tokens( - iter(source.splitlines(keepends=True)).__next__ - ) - comments = tuple( - (t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT - ) + tokens = tokenize.generate_tokens(iter(source.splitlines(keepends=True)).__next__) + comments = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT) except tokenize.TokenError: return frozenset(), () @@ -370,9 +372,7 @@ def check_files(rel_paths: Sequence[str]) -> tuple[Violation, ...]: try: source = abs_path.read_text(encoding="utf-8") except (OSError, UnicodeDecodeError) as exc: - out.append( - Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}") - ) + out.append(Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}")) continue ok_lines, ok_violations = scan_any_ok(report_path, source) @@ -495,9 +495,7 @@ def _in_scope(v: Violation, line_map: dict[str, LineScope] | None) -> bool: def main(argv: Sequence[str]) -> int: - parser = argparse.ArgumentParser( - description="Any-discipline gate (changed-only, changed-lines)." - ) + parser = argparse.ArgumentParser(description="Any-discipline gate (changed-only, changed-lines).") parser.add_argument( "paths", nargs="*", @@ -520,9 +518,7 @@ def main(argv: Sequence[str]) -> int: file=sys.stderr, ) return 0 - rel_paths = _to_litellm_relative( - (REPO_ROOT / name).resolve() for name in line_map - ) + rel_paths = _to_litellm_relative((REPO_ROOT / name).resolve() for name in line_map) elif args.paths: rel_paths = _to_litellm_relative((REPO_ROOT / p).resolve() for p in args.paths) else: @@ -546,9 +542,7 @@ def main(argv: Sequence[str]) -> int: file=sys.stderr, ) return 1 - print( - f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines" - ) + print(f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines") return 0