From 2c9b0e00acaf2d7cebc962bc9cf6a861796b5c68 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:26:41 -0700 Subject: [PATCH] refactor(types): replace Any with proven types in 8 files (#43551) * refactor(types): replace Any with proven types in 8 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): keep email logger untyped where its alert types differ Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/litellm_logging.py | 24 +++--- litellm/llms/custom_httpx/llm_http_handler.py | 30 ++++++-- litellm/proxy/common_request_processing.py | 6 +- litellm/proxy/proxy_server.py | 12 +-- litellm/proxy/utils.py | 33 ++++++-- litellm/responses/main.py | 4 +- litellm/router.py | 5 +- litellm/utils.py | 77 ++++++++----------- 8 files changed, 112 insertions(+), 79 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 4a27aac9769..5af669591f7 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -2287,7 +2287,7 @@ class Logging(LiteLLMLoggingBaseClass): self.completion_start_time = completion_start_time self.model_call_details["completion_start_time"] = self.completion_start_time - def normalize_logging_result(self, result: Any) -> object: + def normalize_logging_result(self, result: object) -> object: """ Some endpoints return a different type of result than what is expected by the logging system. This function is used to normalize the result to the expected type. @@ -2736,7 +2736,7 @@ class Logging(LiteLLMLoggingBaseClass): start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, - **kwargs: Any, # kwargs-ok: forwarded to _success_handler_body + **kwargs: object, # kwargs-ok: forwarded to _success_handler_body ) -> None: """Restores trace_id/session_id contextvars once this attempt's own success logging (including any nested calls its callbacks trigger) is fully done.""" @@ -3175,7 +3175,7 @@ class Logging(LiteLLMLoggingBaseClass): start_time: datetime.datetime | None = None, end_time: datetime.datetime | None = None, cache_hit: bool | None = None, - **kwargs: Any, # kwargs-ok: forwarded to _async_success_handler_body + **kwargs: object, # kwargs-ok: forwarded to _async_success_handler_body ) -> None: """Restores trace_id/session_id contextvars once this attempt's own success logging (including any nested calls its callbacks trigger) is fully done.""" @@ -3554,7 +3554,7 @@ class Logging(LiteLLMLoggingBaseClass): ) self._handle_callback_failure(callback=callback) - def _handle_callback_failure(self, callback: Any): + def _handle_callback_failure(self, callback: object): """ Handle callback logging failures by incrementing Prometheus metrics. @@ -3959,10 +3959,10 @@ class Logging(LiteLLMLoggingBaseClass): def handle_sync_success_callbacks_for_async_calls( self, - result: Any, + result: object, start_time: datetime.datetime, end_time: datetime.datetime, - cache_hit: Any | None = None, + cache_hit: object | None = None, ) -> None: """ Handles calling success callbacks for Async calls. @@ -4132,10 +4132,10 @@ class Logging(LiteLLMLoggingBaseClass): response so cost calculation and spend tracking see one shape. """ if result.event_type == "interaction.completed" and result.interaction is not None: - return InteractionsAPIResponse(**result.interaction) + return InteractionsAPIResponse.model_validate(result.interaction) if result.status == "completed": - return InteractionsAPIResponse( - **result.model_dump( + return InteractionsAPIResponse.model_validate( + result.model_dump( exclude={ # mutable-ok: pydantic types exclude as set[str], which a frozenset does not satisfy "event_type", "delta", @@ -4165,7 +4165,7 @@ class Logging(LiteLLMLoggingBaseClass): return logged return logged.model_copy(update={"id": streamed_message_id}) - def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse: + def _handle_anthropic_messages_response_logging(self, result: object) -> ModelResponse: """ Handles logging for Anthropic messages responses. @@ -4274,7 +4274,7 @@ class Logging(LiteLLMLoggingBaseClass): ) return model_response - def _handle_non_streaming_google_genai_generate_content_response_logging(self, result: Any) -> ModelResponse: + def _handle_non_streaming_google_genai_generate_content_response_logging(self, result: object) -> ModelResponse: """ Handles logging for Google GenAI generate content responses. """ @@ -4349,7 +4349,7 @@ def _get_masked_values( "passwd", ] - def _mask_value(v: Any) -> Any: + def _mask_value(v: object) -> object: if isinstance(v, dict): if _depth >= _max_depth: return v diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 668e8f1c573..60d12337447 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1,10 +1,10 @@ import asyncio import json import ssl -from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Coroutine, Iterator, Mapping, Sequence from contextlib import asynccontextmanager from functools import lru_cache -from types import MappingProxyType, ModuleType +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, @@ -239,6 +239,26 @@ class _ResponsesClientWebSocket(Protocol): async def close(self, code: int = ..., reason: str | None = ...) -> None: ... +class _WebsocketsExceptions(Protocol): + @property + def WebSocketException(self) -> type[Exception]: ... + + +class _WebsocketsModule(Protocol): + @property + def exceptions(self) -> _WebsocketsExceptions: ... + + def connect( + self, + uri: str, + *, + additional_headers: Mapping[str, str], + max_size: int | None, + ssl: bool | str | ssl.SSLContext, + open_timeout: float, + ) -> Awaitable["ClientConnection"]: ... + + _ResponseT = TypeVar("_ResponseT") @@ -5468,7 +5488,7 @@ class BaseLLMHTTPHandler: max_loops: int, fingerprints: list[str], fingerprint: str, - ) -> Any: + ) -> ModelResponse | CustomStreamWrapper: patch: Final = plan.request_patch or AgenticLoopRequestPatch() if patch.messages is None: raise ValueError("Agentic loop plan missing patched messages") @@ -5978,7 +5998,7 @@ class BaseLLMHTTPHandler: @staticmethod async def _open_realtime_backend_ws( - websockets_module: ModuleType, + websockets_module: _WebsocketsModule, url: str, headers: dict, ssl_context: bool | str | ssl.SSLContext, @@ -6037,7 +6057,7 @@ class BaseLLMHTTPHandler: headers: dict, api_base: str | None = None, api_key: str | None = None, - client: Any | None = None, + client: object | None = None, timeout: float | None = None, user_api_key_dict: object | None = None, litellm_metadata: dict[str, object] | None = None, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 15610da9aec..97de2488b8d 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3,7 +3,7 @@ import contextlib import json import logging import math -from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType @@ -974,7 +974,7 @@ _NO_GENERAL_SETTINGS: Final[Mapping[str, object]] = MappingProxyType({}) async def create_response( - generator: AsyncGenerator[str, None], + generator: AsyncGenerator[str, None] | Coroutine[object, object, AsyncGenerator[str, None]], media_type: str, headers: Mapping[str, str], default_status_code: int = status.HTTP_200_OK, @@ -4166,7 +4166,7 @@ class ProxyBaseLLMRequestProcessing: ) @staticmethod - def _stream_usage_for_event(obj: Mapping[str, object], usage: Mapping[str, Any]) -> Usage | None: + def _stream_usage_for_event(obj: Mapping[str, object], usage: Mapping[str, object]) -> Usage | None: # Anthropic reports input_tokens excluding cache tokens, so reuse the non-streaming # transformation to total the prompt and keep the 5m/1h cache creation split if obj.get("type") == "message_delta": diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f842f2e1e4a..4609d57ff15 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2550,7 +2550,7 @@ def general_settings_view() -> Mapping[str, object]: return _GENERAL_SETTINGS_VIEW.validate_python(general_settings) -config_passthrough_endpoints: list[dict[str, Any]] | None = None +config_passthrough_endpoints: list[dict[str, object]] | None = None log_file: Final = "api_log.json" worker_config: Final = None master_key: str | None = None @@ -2584,7 +2584,7 @@ use_queue = False health_check_interval = None health_check_concurrency = None health_check_details = None -health_check_results: dict[str, int | list[dict[str, Any]]] = {} +health_check_results: dict[str, int | list[dict[str, object]]] = {} background_health_check_loop_active = False background_health_check_cycle_seq = 0 queue: Final[list] = [] @@ -13640,7 +13640,7 @@ async def _try_provider_token_count( model_to_use: str, messages: list | None, contents: list | None, - deployment: dict[str, Any] | None, + deployment: dict[str, object] | None, request_model: str, tools: list | None = None, system: str | None = None, @@ -14439,7 +14439,7 @@ async def _fetch_db_models_for_search( sort_by: str | None, is_byok_outside_caller_teams: Callable[[dict[str, JsonValue]], bool], model_name: str | None = None, -) -> tuple[list[dict[str, Any]], int]: +) -> tuple[list[dict[str, object]], int]: """ Run the bounded DB query that backs `/v2/model/info?search=`. Returns `(decrypted_models, total_count)` where `total_count` is the cheap @@ -14586,7 +14586,7 @@ async def _apply_search_filter_to_models( router_models_count: Final = config_models_count + db_models_in_router_count # Query database for additional models with search term - db_models: list[dict[str, Any]] = [] + db_models: list[dict[str, object]] = [] exact_name_can_match: Final = model_name is None or search_lower in model_name.lower() if prisma_client is not None and exact_name_can_match: try: @@ -14781,7 +14781,7 @@ def _matches_model_info_filters( def _paginate_models_response( - all_models: list[dict[str, Any]], + all_models: Sequence[Mapping[str, object]], page: int, size: int, total_count: int | None, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b336ce1fa27..1fca50e24c9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -303,6 +303,29 @@ class _EndUserSpendBatch(Protocol): def litellm_endusertable(self) -> _EndUserBatchTable: ... +class _BatchUpdateTable(Protocol): + def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + + +class _UpdateManyBatch(Protocol): + @property + def litellm_verificationtoken(self) -> _BatchUpdateTable: ... + + @property + def litellm_usertable(self) -> _EndUserBatchTable: ... + + @property + def litellm_endusertable(self) -> _EndUserBatchTable: ... + + @property + def litellm_budgettable(self) -> _EndUserBatchTable: ... + + @property + def litellm_teamtable(self) -> _EndUserBatchTable: ... + + async def commit(self) -> None: ... + + unified_guardrail: Final = UnifiedLLMGuardrails() NON_OPENAI_STREAM_GUARDRAIL_TRANSLATION_CALL_TYPES: "frozenset[CallTypes]" = frozenset({CallTypes.anthropic_messages}) @@ -1111,7 +1134,7 @@ class _CallbackCapabilities: # Tuple[(resolved_callback, "override" | "apply_guardrail"), ...] # Ordered the same as ``litellm.callbacks``; used to build the streaming # iterator chain without re-scanning per request. - iterator_overrides: tuple[tuple[Any, str], ...] = field(default_factory=tuple) + iterator_overrides: tuple[tuple[CustomLogger, str], ...] = field(default_factory=tuple) # Resolved CustomLogger callbacks in original order. Pre-resolving once # avoids the per-request ``get_custom_logger_compatible_class`` walk for # every string entry in ``litellm.callbacks``. @@ -2761,7 +2784,7 @@ class ProxyLogging: has_pre_call_override = False has_content_enforcer = False has_moderation_override = False - iterator_overrides: Final[list[tuple[Any, str]]] = [] # (callback, kind) + iterator_overrides: Final[list[tuple[CustomLogger, str]]] = [] # (callback, kind) resolved_callbacks: Final[list[CustomLogger]] = [] for callback in callbacks: @@ -5481,7 +5504,7 @@ class PrismaClient: """ Batch write update queries """ - batcher = self.db.batch_() + batcher: _UpdateManyBatch = self.db.batch_() for idx, t in enumerate(data_list): # check if plain text or hash if t.token.startswith("sk-"): @@ -6638,7 +6661,7 @@ class PrismaClient: read-only as a whole (replica, failover in progress) does not get its engine killed on every watchdog cycle or failed write.""" backoff_seconds: Final = min( - self._db_reconnect_cooldown_seconds * 2 ** min(self._db_read_only_recreate_streak, 10), + self._db_reconnect_cooldown_seconds * (1 << min(self._db_read_only_recreate_streak, 10)), _READ_ONLY_RECREATE_BACKOFF_CAP_SECONDS, ) if time.time() - self._db_read_only_recreate_ts < backoff_seconds: @@ -7344,7 +7367,7 @@ class ProxyUpdateSpend: if i >= n_retry_times: await requeue_spend_logs(prisma_client, proxy_logging_obj, logs_to_process) raise - await asyncio.sleep(2**i) + await asyncio.sleep(1 << i) except Exception as e: _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) finally: diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 93c72bc2d3b..f1ed8d3e9b3 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -981,7 +981,7 @@ def _responses_try_dispatch_mcp_gateway( background: bool | None, stream: bool | None, temperature: float | None, - text: Any, + text: Optional["ResponseText"], tool_choice: ToolChoice | None, top_p: float | None, truncation: Literal["auto", "disabled"] | None, @@ -1054,7 +1054,7 @@ def _responses_try_dispatch_emulated_file_search( background: bool | None, stream: bool | None, temperature: float | None, - text: Any, + text: Optional["ResponseText"], tool_choice: ToolChoice | None, top_p: float | None, truncation: Literal["auto", "disabled"] | None, diff --git a/litellm/router.py b/litellm/router.py index 1ef68e60440..cfed5dc81c0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -433,6 +433,7 @@ def model_info_is_active_for_environment(model_info: Mapping[str, object] | None _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT") +_CallbackT = TypeVar("_CallbackT") _ALIAS_PARAMS_NEVER_FORWARDED: Final = frozenset({"model", "api_base", "api_key", "api_version"}) _ALIAS_MARKER_FORWARDED_PARAMS_KWARG: Final = "_alias_marker_forwarded_params" @@ -2217,8 +2218,8 @@ class Router: def _add_encrypted_content_affinity_check(self, enable_global_affinity: bool) -> None: def _move_before_deployment_affinity( - callback_list: list[Any], - callback_to_move: EncryptedContentAffinityCheck, + callback_list: list[_CallbackT], + callback_to_move: _CallbackT, ) -> None: if callback_to_move not in callback_list: return diff --git a/litellm/utils.py b/litellm/utils.py index d45b29c0f16..13a46840431 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -681,7 +681,7 @@ def load_credentials_from_list(kwargs: dict): Updates kwargs with the credentials if credential_name in kwarg """ # Access CredentialAccessor via module to trigger lazy loading if needed - credential_accessor: Final[type[CredentialAccessor]] = getattr(sys.modules[__name__], "CredentialAccessor") + credential_accessor: Final[type[CredentialAccessor]] = litellm_utils.CredentialAccessor credential_name: Final = kwargs.get("litellm_credential_name") if not credential_name: @@ -1015,9 +1015,7 @@ def function_setup( len(litellm.input_callback) > 0 or len(litellm.success_callback) > 0 or len(litellm.failure_callback) > 0 ) and len(callback_list) == 0: callback_list = list(set(litellm.input_callback + litellm.success_callback + litellm.failure_callback)) - get_set_callbacks: Final[Callable[[], Callable[..., None]]] = getattr( - sys.modules[__name__], "get_set_callbacks" - ) + get_set_callbacks: Final[Callable[[], Callable[..., None]]] = litellm_utils.get_set_callbacks get_set_callbacks()(callback_list=callback_list, function_id=function_id) ## ASYNC CALLBACKS - safety net for callbacks added via direct append if len(litellm.input_callback) > 0: @@ -1262,9 +1260,7 @@ def function_setup( call_type=call_type, ): stream = True - get_litellm_logging_class: Final[_LoggingClassGetter] = getattr( - sys.modules[__name__], "get_litellm_logging_class" - ) + get_litellm_logging_class: Final[_LoggingClassGetter] = litellm_utils.get_litellm_logging_class # Victim for object pool logging_obj = get_litellm_logging_class()( # rebind-ok: 2nd assignment to logging_obj (see initial None above) model=model, @@ -1413,8 +1409,8 @@ def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tu if num_retries is None: num_retries = litellm.num_retries if kwargs.get("retry_policy", None): - get_num_retries_from_retry_policy: Final[Callable[..., int | None]] = getattr( - sys.modules[__name__], "get_num_retries_from_retry_policy" + get_num_retries_from_retry_policy: Final[Callable[..., int | None]] = ( + litellm_utils.get_num_retries_from_retry_policy ) reset_retry_policy: Final = litellm_utils.reset_retry_policy retry_policy_num_retries: Final[int | None] = get_num_retries_from_retry_policy( @@ -1802,9 +1798,7 @@ def client(original_function): return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None)) else: # RETURN RESULT - update_response_metadata: _ResponseMetadataUpdater = getattr( - sys.modules[__name__], "update_response_metadata" - ) + update_response_metadata: _ResponseMetadataUpdater = litellm_utils.update_response_metadata update_response_metadata( result=result, logging_obj=logging_obj, @@ -1845,9 +1839,7 @@ def client(original_function): kwargs=kwargs, ) - _update_response_metadata: Final[_ResponseMetadataUpdater] = getattr( - sys.modules[__name__], "update_response_metadata" - ) + _update_response_metadata: Final[_ResponseMetadataUpdater] = litellm_utils.update_response_metadata _update_response_metadata( result=result, logging_obj=logging_obj, @@ -1877,8 +1869,8 @@ def client(original_function): if call_type == CallTypes.completion.value: num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None if kwargs.get("retry_policy", None): - get_num_retries_from_retry_policy: Callable[..., int | None] = getattr( - sys.modules[__name__], "get_num_retries_from_retry_policy" + get_num_retries_from_retry_policy: Callable[..., int | None] = ( + litellm_utils.get_num_retries_from_retry_policy ) reset_retry_policy = litellm_utils.reset_retry_policy num_retries = get_num_retries_from_retry_policy( @@ -1916,9 +1908,7 @@ def client(original_function): elif call_type == CallTypes.responses.value: num_retries = kwargs.get("num_retries", None) or litellm.num_retries or None if kwargs.get("retry_policy", None): - get_num_retries_from_retry_policy = getattr( - sys.modules[__name__], "get_num_retries_from_retry_policy" - ) + get_num_retries_from_retry_policy = litellm_utils.get_num_retries_from_retry_policy reset_retry_policy = litellm_utils.reset_retry_policy num_retries = get_num_retries_from_retry_policy( exception=e, @@ -1955,9 +1945,7 @@ def client(original_function): print_args_passed_to_litellm(original_function, args, kwargs) start_time: Final = datetime.datetime.now() result = None - _update_response_metadata: Final[_ResponseMetadataUpdater] = getattr( - sys.modules[__name__], "update_response_metadata" - ) + _update_response_metadata: Final[_ResponseMetadataUpdater] = litellm_utils.update_response_metadata logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) LLMCachingHandler: Final = _get_cached_llm_caching_handler() _llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler( @@ -3496,7 +3484,8 @@ def get_optional_params_transcription( passed_params.pop("OPENAI_TRANSCRIPTION_PARAMS") custom_llm_provider = passed_params.pop("custom_llm_provider") - drop_params = normalize_drop_params(passed_params.pop("drop_params")) + passed_params.pop("drop_params") + drop_params = normalize_drop_params(drop_params) special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs") for k, v in special_params.items(): passed_params[k] = v @@ -3597,16 +3586,18 @@ def get_optional_params_image_gen( additional_drop_params: list | None = None, provider_config: BaseImageGenerationConfig | None = None, drop_params: bool | None = None, - **kwargs, + **kwargs: object, ): # retrieve all parameters passed to the function passed_params: Final = locals() - model = passed_params.pop("model", None) - custom_llm_provider = passed_params.pop("custom_llm_provider") - provider_config = passed_params.pop("provider_config", None) - drop_params = normalize_drop_params(passed_params.pop("drop_params", None)) + passed_params.pop("model", None) + passed_params.pop("custom_llm_provider") + passed_params.pop("provider_config", None) + passed_params.pop("drop_params", None) + drop_params = normalize_drop_params(drop_params) additional_drop_params = passed_params.pop("additional_drop_params", None) - special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs") + passed_params.pop("kwargs") + special_params: Final[Mapping[str, object]] = kwargs for k, v in special_params.items(): if ( k.startswith("aws_") @@ -3725,18 +3716,18 @@ def get_optional_params_embeddings( **kwargs, ): # Lazy load get_supported_openai_params - get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = getattr( - sys.modules[__name__], "get_supported_openai_params" - ) + get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = litellm_utils.get_supported_openai_params # retrieve all parameters passed to the function passed_params: Final = locals() custom_llm_provider = passed_params.pop("custom_llm_provider", None) special_params: Final = passed_params.pop("kwargs") - drop_params = normalize_drop_params(passed_params.pop("drop_params", None)) - additional_drop_params = passed_params.pop("additional_drop_params", None) - allowed_openai_params = passed_params.pop("allowed_openai_params", None) or [] + passed_params.pop("drop_params", None) + drop_params = normalize_drop_params(drop_params) + passed_params.pop("additional_drop_params", None) + passed_params.pop("allowed_openai_params", None) + allowed_openai_params = allowed_openai_params or [] # Remove function objects from passed_params to avoid JSON serialization errors passed_params.pop("get_supported_openai_params", None) @@ -4502,9 +4493,7 @@ def get_optional_params( message=f"{custom_llm_provider} does not support parameters: {list(unsupported_params.keys())}, for model={model}. To drop these, set `litellm.drop_params=True` or for proxy:\n\n`litellm_settings:\n drop_params: true`\n. \n If you want to use these params dynamically send allowed_openai_params={list(unsupported_params.keys())} in your request.", ) - get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = getattr( - sys.modules[__name__], "get_supported_openai_params" - ) + get_supported_openai_params: Final[_SupportedOpenAIParamsGetter] = litellm_utils.get_supported_openai_params supported_params = get_supported_openai_params( model=model, custom_llm_provider=custom_llm_provider, base_model=base_model ) @@ -4675,7 +4664,7 @@ def get_optional_params( drop_params=bool(drop_params), ) elif custom_llm_provider == "bedrock": - bedrock_model_info: Final[type[BedrockModelInfo]] = getattr(sys.modules[__name__], "BedrockModelInfo") + bedrock_model_info: Final[type[BedrockModelInfo]] = litellm_utils.BedrockModelInfo bedrock_route: Final = bedrock_model_info.get_bedrock_route(model) bedrock_base_model: Final = bedrock_model_info.get_base_model(model) if bedrock_route == "converse" or bedrock_route == "converse_like": @@ -4996,8 +4985,8 @@ def get_optional_params( # Apply nested drops from additional_drop_params if additional_drop_params: - is_nested_path: Final[_NestedPathChecker] = getattr(sys.modules[__name__], "is_nested_path") - delete_nested_value: Final[_NestedValueDeleter] = getattr(sys.modules[__name__], "delete_nested_value") + is_nested_path: Final[_NestedPathChecker] = litellm_utils.is_nested_path + delete_nested_value: Final[_NestedValueDeleter] = litellm_utils.delete_nested_value nested_paths: Final = [p for p in additional_drop_params if is_nested_path(p)] for path in nested_paths: optional_params = delete_nested_value(optional_params, path) @@ -7328,7 +7317,7 @@ class TextCompletionStreamWrapper: except StopIteration: raise StopIteration except Exception as e: - exception_type: Final = getattr(sys.modules[__name__], "exception_type") + exception_type: Final = litellm_utils.exception_type raise exception_type( model=self.model, custom_llm_provider=self.custom_llm_provider or "", @@ -9974,7 +9963,7 @@ def get_end_user_id_for_cost_tracking( service_type: "litellm_logging" or "prometheus" - used to allow prometheus only disable cost tracking. """ - get_litellm_metadata_from_kwargs: Final = getattr(sys.modules[__name__], "get_litellm_metadata_from_kwargs") + get_litellm_metadata_from_kwargs: Final = litellm_utils.get_litellm_metadata_from_kwargs _metadata: Final = cast(dict, get_litellm_metadata_from_kwargs(dict(litellm_params=litellm_params))) end_user_id: Final = cast(