From 0fcc36c3017ede5c49da61ab459776cbdd80cd6e Mon Sep 17 00:00:00 2001 From: Chesars Date: Thu, 12 Mar 2026 14:23:50 -0300 Subject: [PATCH] style: run black formatter on 52 non-enterprise files Formats all files flagged by CI except 12 enterprise/ files which require separate access to format. --- litellm/__init__.py | 6 +- .../transformation.py | 12 +- litellm/files/main.py | 19 +- litellm/images/main.py | 28 +-- .../litellm_core_utils/get_model_cost_map.py | 4 +- litellm/litellm_core_utils/litellm_logging.py | 5 +- .../prompt_templates/factory.py | 24 +- .../llms/azure/chat/gpt_5_transformation.py | 12 +- .../azure/realtime/http_transformation.py | 22 +- .../azure_model_router/transformation.py | 6 +- .../base_llm/realtime/http_transformation.py | 4 +- .../bedrock/chat/converse_transformation.py | 4 +- .../black_forest_labs/image_edit/handler.py | 14 +- .../image_edit/transformation.py | 21 +- .../image_generation/handler.py | 14 +- .../image_generation/transformation.py | 4 +- litellm/llms/custom_httpx/llm_http_handler.py | 30 ++- .../audio_transcription/transformation.py | 8 +- .../llms/openai/chat/gpt_5_transformation.py | 20 +- .../openai/realtime/http_transformation.py | 8 +- .../openai/responses/count_tokens/handler.py | 4 +- .../responses/count_tokens/transformation.py | 30 +-- litellm/llms/openai_like/dynamic_config.py | 9 +- .../perplexity/responses/transformation.py | 9 +- .../llms/vertex_ai/gemini/transformation.py | 4 +- .../vertex_and_google_ai_studio_gemini.py | 25 +- .../batch_embed_content_handler.py | 32 +-- .../batch_embed_content_transformation.py | 78 +++--- litellm/main.py | 19 +- litellm/proxy/_experimental/mcp_server/db.py | 4 +- .../mcp_server/discoverable_endpoints.py | 4 +- .../mcp_server/rest_endpoints.py | 24 +- .../proxy/_experimental/mcp_server/server.py | 12 +- litellm/proxy/_types.py | 78 +++--- litellm/proxy/auth/model_checks.py | 8 +- litellm/proxy/common_request_processing.py | 15 +- .../proxy/credential_endpoints/endpoints.py | 24 +- .../guardrails/guardrail_hooks/presidio.py | 4 +- .../customer_endpoints.py | 5 +- .../mcp_management_endpoints.py | 53 ++-- .../management_endpoints/team_endpoints.py | 228 +++++++++++------- .../object_permission_utils.py | 42 ++-- .../pass_through_endpoints.py | 19 +- litellm/proxy/proxy_cli.py | 4 +- litellm/proxy/proxy_server.py | 4 +- litellm/proxy/realtime_endpoints/endpoints.py | 10 +- litellm/proxy/utils.py | 8 +- .../proxy/vector_store_endpoints/endpoints.py | 24 +- litellm/realtime_api/main.py | 26 +- .../transformation.py | 6 +- litellm/types/utils.py | 4 +- litellm/utils.py | 8 +- 52 files changed, 669 insertions(+), 420 deletions(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index 439f88f205a..8b3723cb2b0 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1261,7 +1261,11 @@ from .containers.main import * from .ocr.main import * from .rag.main import * from .search.main import * -from .realtime_api.main import _arealtime, acreate_realtime_client_secret, arealtime_calls +from .realtime_api.main import ( + _arealtime, + acreate_realtime_client_secret, + arealtime_calls, +) from .responses.main import _aresponses_websocket from .fine_tuning.main import * from .files.main import * diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index f1e9b18450b..42359afef4d 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -398,7 +398,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ResponseOutputMessage, ResponseReasoningItem, ) - from openai.types.responses.response_output_item import ResponseApplyPatchToolCall + from openai.types.responses.response_output_item import ( + ResponseApplyPatchToolCall, + ) from litellm.types.utils import Choices, Message @@ -448,11 +450,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): LiteLLMCompletionResponsesConfig, ) - tool_call_dict = ( - LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call( - tool_call_item=item, - index=tool_call_index, - ) + tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call( + tool_call_item=item, + index=tool_call_index, ) accumulated_tool_calls.append(tool_call_dict) tool_call_index += 1 diff --git a/litellm/files/main.py b/litellm/files/main.py index b55dd021bc8..f7c89e0ba3b 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -15,11 +15,24 @@ from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast import httpx # Type aliases for provider parameters -FileCreateProvider = Literal["openai", "azure", "gemini", "vertex_ai", "bedrock", "hosted_vllm", "manus", "anthropic"] -FileRetrieveProvider = Literal["openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus", "anthropic"] +FileCreateProvider = Literal[ + "openai", + "azure", + "gemini", + "vertex_ai", + "bedrock", + "hosted_vllm", + "manus", + "anthropic", +] +FileRetrieveProvider = Literal[ + "openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "manus", "anthropic" +] FileDeleteProvider = Literal["openai", "azure", "gemini", "manus", "anthropic"] FileListProvider = Literal["openai", "azure", "manus", "anthropic"] -FileContentProvider = Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"] +FileContentProvider = Literal[ + "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus" +] import litellm from litellm import get_secret_str diff --git a/litellm/images/main.py b/litellm/images/main.py index f0c68ef6c70..a3ae97b57dd 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -929,20 +929,20 @@ def image_edit( # noqa: PLR0915 elif custom_llm_provider == "stability": image_edit_request_params.update(non_default_params) return base_llm_http_handler.image_edit_handler( - model=model, - image=images, - prompt=prompt, - image_edit_provider_config=image_edit_provider_config, - image_edit_optional_request_params=image_edit_request_params, - custom_llm_provider=custom_llm_provider, - litellm_params=litellm_params, - logging_obj=litellm_logging_obj, - extra_headers=extra_headers, - extra_body=extra_body, - timeout=timeout or DEFAULT_REQUEST_TIMEOUT, - _is_async=_is_async, - client=kwargs.get("client"), - ) + model=model, + image=images, + prompt=prompt, + image_edit_provider_config=image_edit_provider_config, + image_edit_optional_request_params=image_edit_request_params, + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + logging_obj=litellm_logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout or DEFAULT_REQUEST_TIMEOUT, + _is_async=_is_async, + client=kwargs.get("client"), + ) elif custom_llm_provider == "black_forest_labs": # Route to BFL-specific handler (polling required) if model is None: diff --git a/litellm/litellm_core_utils/get_model_cost_map.py b/litellm/litellm_core_utils/get_model_cost_map.py index da2908c8586..7679358bbc6 100644 --- a/litellm/litellm_core_utils/get_model_cost_map.py +++ b/litellm/litellm_core_utils/get_model_cost_map.py @@ -289,7 +289,9 @@ def get_model_cost_map(url: str) -> dict: url, ) _cost_map_source_info.source = "local" - _cost_map_source_info.fallback_reason = "Remote data failed integrity validation" + _cost_map_source_info.fallback_reason = ( + "Remote data failed integrity validation" + ) return _expand_model_aliases(GetModelCostMap.load_local_model_cost_map()) _cost_map_source_info.source = "remote" diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 566ff3218ad..e22d057bb69 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -5354,7 +5354,10 @@ def get_standard_logging_object_payload( requested_model = kwargs.get("model") if ( isinstance(requested_model, str) - and ("model_router" in requested_model.lower() or "model-router" in requested_model.lower()) + and ( + "model_router" in requested_model.lower() + or "model-router" in requested_model.lower() + ) and isinstance(response_model_name, str) and response_model_name ): diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 31716479b15..0df10ae9f90 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -2514,7 +2514,11 @@ def anthropic_messages_pt( # noqa: PLR0915 if isinstance(_tc, dict) else getattr(_tc, "id", None) ) - if _tc_id and isinstance(_tc_id, str) and _tc_id.startswith("srvtoolu_"): + if ( + _tc_id + and isinstance(_tc_id, str) + and _tc_id.startswith("srvtoolu_") + ): _has_server_tool_calls = True break @@ -2590,9 +2594,9 @@ def anthropic_messages_pt( # noqa: PLR0915 original_content_element=dict(assistant_content_block), ) if "cache_control" in _content_element: - _anthropic_text_content_element["cache_control"] = ( - _content_element["cache_control"] - ) + _anthropic_text_content_element[ + "cache_control" + ] = _content_element["cache_control"] text_element = _anthropic_text_content_element # Interleave: each thinking block precedes its server tool group. @@ -2681,13 +2685,15 @@ def anthropic_messages_pt( # noqa: PLR0915 _list_has_thinking = False if _content_is_list: for _item in assistant_content_block["content"]: - if isinstance(_item, dict) and _item.get("type") in ("thinking", "redacted_thinking"): + if isinstance(_item, dict) and _item.get("type") in ( + "thinking", + "redacted_thinking", + ): _list_has_thinking = True break if ( - thinking_blocks is not None - and not _list_has_thinking + thinking_blocks is not None and not _list_has_thinking ): # IMPORTANT: ADD THIS FIRST, ELSE ANTHROPIC WILL RAISE AN ERROR assistant_content.extend(thinking_blocks) if _content_is_list: @@ -2745,9 +2751,9 @@ def anthropic_messages_pt( # noqa: PLR0915 ) if "cache_control" in _content_element: - _anthropic_text_content_element["cache_control"] = _content_element[ + _anthropic_text_content_element[ "cache_control" - ] + ] = _content_element["cache_control"] assistant_content.append(_anthropic_text_content_element) diff --git a/litellm/llms/azure/chat/gpt_5_transformation.py b/litellm/llms/azure/chat/gpt_5_transformation.py index 81c3dfded71..6310df9cecc 100644 --- a/litellm/llms/azure/chat/gpt_5_transformation.py +++ b/litellm/llms/azure/chat/gpt_5_transformation.py @@ -84,10 +84,9 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): drop_params: bool, api_version: str = "", ) -> dict: - reasoning_effort_value = ( - non_default_params.get("reasoning_effort") - or optional_params.get("reasoning_effort") - ) + reasoning_effort_value = non_default_params.get( + "reasoning_effort" + ) or optional_params.get("reasoning_effort") effective_effort = _get_effort_level(reasoning_effort_value) # gpt-5.1/5.2/5.4 support reasoning_effort='none', but other gpt-5 models don't @@ -100,7 +99,10 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): ): non_default_params = non_default_params.copy() optional_params = optional_params.copy() - if _get_effort_level(non_default_params.get("reasoning_effort")) == "none": + if ( + _get_effort_level(non_default_params.get("reasoning_effort")) + == "none" + ): non_default_params.pop("reasoning_effort") if _get_effort_level(optional_params.get("reasoning_effort")) == "none": optional_params.pop("reasoning_effort") diff --git a/litellm/llms/azure/realtime/http_transformation.py b/litellm/llms/azure/realtime/http_transformation.py index ef9a2d92d48..df1e2707af2 100644 --- a/litellm/llms/azure/realtime/http_transformation.py +++ b/litellm/llms/azure/realtime/http_transformation.py @@ -9,22 +9,14 @@ from litellm.secret_managers.main import get_secret_str class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig): def get_api_base(self, api_base: Optional[str], **kwargs) -> str: - return ( - api_base - or litellm.api_base - or get_secret_str("AZURE_API_BASE") - or "" - ) + return api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") or "" def get_api_key(self, api_key: Optional[str], **kwargs) -> str: - return ( - api_key - or litellm.api_key - or get_secret_str("AZURE_API_KEY") - or "" - ) + return api_key or litellm.api_key or get_secret_str("AZURE_API_KEY") or "" - def get_complete_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_complete_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: base = self.get_api_base(api_base).rstrip("/") version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" return f"{base}/openai/realtime/client_secrets?api-version={version}" @@ -41,7 +33,9 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig): "Content-Type": "application/json", } - def get_realtime_calls_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_realtime_calls_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: base = self.get_api_base(api_base).rstrip("/") version = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17" return f"{base}/openai/realtime/calls?api-version={version}" diff --git a/litellm/llms/azure_ai/azure_model_router/transformation.py b/litellm/llms/azure_ai/azure_model_router/transformation.py index b5eb0edba59..57acb147063 100644 --- a/litellm/llms/azure_ai/azure_model_router/transformation.py +++ b/litellm/llms/azure_ai/azure_model_router/transformation.py @@ -63,7 +63,7 @@ class AzureModelRouterConfig(AzureAIStudioConfig): ) -> ModelResponse: """ Transform response for Model Router. - + Extracts the actual model used from the Azure response (e.g., gpt-5-nano-2025-08-07) and returns it with the azure_ai/ prefix for proper display and cost tracking. """ @@ -71,8 +71,8 @@ class AzureModelRouterConfig(AzureAIStudioConfig): # Get base model for the parent call (strips routing prefixes for API compatibility) base_model: str = AzureFoundryModelInfo.get_base_model(model) - - # Call parent transform_response first - this will extract the actual model + + # Call parent transform_response first - this will extract the actual model # from the raw response (e.g., "gpt-5-nano-2025-08-07") model_response = super().transform_response( model=base_model, diff --git a/litellm/llms/base_llm/realtime/http_transformation.py b/litellm/llms/base_llm/realtime/http_transformation.py index 7aadd49ffd3..712ec42380f 100644 --- a/litellm/llms/base_llm/realtime/http_transformation.py +++ b/litellm/llms/base_llm/realtime/http_transformation.py @@ -54,7 +54,9 @@ class BaseRealtimeHTTPConfig(ABC): # ------------------------------------------------------------------ # @abstractmethod - def get_complete_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_complete_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: """Return the full URL for POST /realtime/client_secrets.""" @abstractmethod diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 4898ada1c2f..229457a73b4 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1209,7 +1209,9 @@ class AmazonConverseConfig(BaseConfig): if request_metadata is not None: self._validate_request_metadata(request_metadata) - output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None) + output_config: Optional[OutputConfigBlock] = inference_params.pop( + "outputConfig", None + ) inference_params.pop( "output_config", None ) # Bedrock Converse doesn't support it diff --git a/litellm/llms/black_forest_labs/image_edit/handler.py b/litellm/llms/black_forest_labs/image_edit/handler.py index 44a102ec48d..dea2683a049 100644 --- a/litellm/llms/black_forest_labs/image_edit/handler.py +++ b/litellm/llms/black_forest_labs/image_edit/handler.py @@ -356,7 +356,12 @@ class BlackForestLabsImageEdit: if status == "Ready": return response - elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + elif status in [ + "Error", + "Failed", + "Content Moderated", + "Request Moderated", + ]: raise BlackForestLabsError( status_code=400, message=f"Image generation failed: {status}", @@ -436,7 +441,12 @@ class BlackForestLabsImageEdit: if status == "Ready": return response - elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + elif status in [ + "Error", + "Failed", + "Content Moderated", + "Request Moderated", + ]: raise BlackForestLabsError( status_code=400, message=f"Image generation failed: {status}", diff --git a/litellm/llms/black_forest_labs/image_edit/transformation.py b/litellm/llms/black_forest_labs/image_edit/transformation.py index 78898345bf6..610fee18899 100644 --- a/litellm/llms/black_forest_labs/image_edit/transformation.py +++ b/litellm/llms/black_forest_labs/image_edit/transformation.py @@ -179,11 +179,7 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): """ Get the complete URL for the Black Forest Labs API request. """ - base_url: str = ( - api_base - or get_secret_str("BFL_API_BASE") - or DEFAULT_API_BASE - ) + base_url: str = api_base or get_secret_str("BFL_API_BASE") or DEFAULT_API_BASE base_url = base_url.rstrip("/") endpoint = self._get_model_endpoint(model) @@ -247,9 +243,18 @@ class BlackForestLabsImageEditConfig(BaseImageEditConfig): # Add optional params (only BFL-recognized parameters) bfl_request_params = [ - "seed", "output_format", "safety_tolerance", "prompt_upsampling", - "aspect_ratio", "steps", "guidance", "grow_mask", - "top", "bottom", "left", "right", + "seed", + "output_format", + "safety_tolerance", + "prompt_upsampling", + "aspect_ratio", + "steps", + "guidance", + "grow_mask", + "top", + "bottom", + "left", + "right", ] for key, value in image_edit_optional_request_params.items(): if key in bfl_request_params and value is not None: diff --git a/litellm/llms/black_forest_labs/image_generation/handler.py b/litellm/llms/black_forest_labs/image_generation/handler.py index 99dc2feca3c..5a1d885e527 100644 --- a/litellm/llms/black_forest_labs/image_generation/handler.py +++ b/litellm/llms/black_forest_labs/image_generation/handler.py @@ -342,7 +342,12 @@ class BlackForestLabsImageGeneration: if status == "Ready": return response - elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + elif status in [ + "Error", + "Failed", + "Content Moderated", + "Request Moderated", + ]: raise BlackForestLabsError( status_code=400, message=f"Image generation failed: {status}", @@ -422,7 +427,12 @@ class BlackForestLabsImageGeneration: if status == "Ready": return response - elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + elif status in [ + "Error", + "Failed", + "Content Moderated", + "Request Moderated", + ]: raise BlackForestLabsError( status_code=400, message=f"Image generation failed: {status}", diff --git a/litellm/llms/black_forest_labs/image_generation/transformation.py b/litellm/llms/black_forest_labs/image_generation/transformation.py index fd664b3ea7e..a6ed77f5359 100644 --- a/litellm/llms/black_forest_labs/image_generation/transformation.py +++ b/litellm/llms/black_forest_labs/image_generation/transformation.py @@ -203,9 +203,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): """ Get the complete URL for the Black Forest Labs API request. """ - base_url: str = ( - api_base or get_secret_str("BFL_API_BASE") or DEFAULT_API_BASE - ) + base_url: str = api_base or get_secret_str("BFL_API_BASE") or DEFAULT_API_BASE base_url = base_url.rstrip("/") endpoint = self._get_model_endpoint(model) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 1442be71c30..4e6c3cba684 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -4835,7 +4835,9 @@ class BaseLLMHTTPHandler: async_httpx_client = client if provider_config is not None: - url = provider_config.get_complete_url(api_base=api_base, model=model or "", api_version=api_version) + url = provider_config.get_complete_url( + api_base=api_base, model=model or "", api_version=api_version + ) headers: Dict[str, Any] = provider_config.validate_environment( headers={}, model=model or "", api_key=api_key ) @@ -4905,7 +4907,9 @@ class BaseLLMHTTPHandler: async_httpx_client = client if provider_config is not None: - url = provider_config.get_realtime_calls_url(api_base=api_base, model=model or "", api_version=api_version) + url = provider_config.get_realtime_calls_url( + api_base=api_base, model=model or "", api_version=api_version + ) headers: Dict[str, Any] = provider_config.get_realtime_calls_headers( ephemeral_key=openai_ephemeral_key ) @@ -7910,9 +7914,7 @@ class BaseLLMHTTPHandler: ) try: - response = await async_httpx_client.get( - url=url, headers=headers - ) + response = await async_httpx_client.get(url=url, headers=headers) except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_provider_config) @@ -8023,7 +8025,7 @@ class BaseLLMHTTPHandler: ) url = api_base - + params = {} if after is not None: params["after"] = after @@ -8105,7 +8107,7 @@ class BaseLLMHTTPHandler: ) url = api_base - + params = {} if after is not None: params["after"] = after @@ -8167,14 +8169,15 @@ class BaseLLMHTTPHandler: ) url = f"{api_base}/{vector_store_id}" - + request_body = dict(vector_store_update_optional_params) - + # Clean metadata to only include string values (OpenAI requirement) if "metadata" in request_body and request_body["metadata"] is not None: from litellm.utils import add_openai_metadata + request_body["metadata"] = add_openai_metadata(request_body["metadata"]) - + if extra_body: request_body.update(extra_body) @@ -8249,14 +8252,15 @@ class BaseLLMHTTPHandler: ) url = f"{api_base}/{vector_store_id}" - + request_body = dict(vector_store_update_optional_params) - + # Clean metadata to only include string values (OpenAI requirement) if "metadata" in request_body and request_body["metadata"] is not None: from litellm.utils import add_openai_metadata + request_body["metadata"] = add_openai_metadata(request_body["metadata"]) - + if extra_body: request_body.update(extra_body) diff --git a/litellm/llms/mistral/audio_transcription/transformation.py b/litellm/llms/mistral/audio_transcription/transformation.py index fd84d63c4fa..4d294063499 100644 --- a/litellm/llms/mistral/audio_transcription/transformation.py +++ b/litellm/llms/mistral/audio_transcription/transformation.py @@ -60,9 +60,7 @@ class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig): stream: Optional[bool] = None, ) -> str: api_base = ( - "https://api.mistral.ai/v1" - if api_base is None - else api_base.rstrip("/") + "https://api.mistral.ai/v1" if api_base is None else api_base.rstrip("/") ) return f"{api_base}/audio/transcriptions" @@ -121,7 +119,9 @@ class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig): openai_params=self.get_supported_openai_params(model), ) for key, value in provider_specific_params.items(): - form_fields[key] = str(value).lower() if isinstance(value, bool) else str(value) + form_fields[key] = ( + str(value).lower() if isinstance(value, bool) else str(value) + ) files = { "file": ( diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index f186bc60859..f19d4891c6e 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -183,16 +183,19 @@ class OpenAIGPT5Config(OpenAIGPTConfig): # Use effective_effort (extracted string) for xhigh validation, "none" checks, and # tool/sampling guards — dict inputs like {"effort": "none", "summary": "detailed"} # must be treated as effort="none" to avoid incorrect tool-drop or sampling errors. - raw_reasoning_effort = ( - non_default_params.get("reasoning_effort") - or optional_params.get("reasoning_effort") - ) + raw_reasoning_effort = non_default_params.get( + "reasoning_effort" + ) or optional_params.get("reasoning_effort") effective_effort = _get_effort_level(raw_reasoning_effort) # Normalize to string for Chat Completions API when dict has only "effort". # Preserve full dict (e.g. {"effort": "high", "summary": "detailed"}) for Responses API. - if isinstance(raw_reasoning_effort, dict) and set(raw_reasoning_effort.keys()) <= {"effort"}: - normalized = _normalize_reasoning_effort_for_chat_completion(raw_reasoning_effort) + if isinstance(raw_reasoning_effort, dict) and set( + raw_reasoning_effort.keys() + ) <= {"effort"}: + normalized = _normalize_reasoning_effort_for_chat_completion( + raw_reasoning_effort + ) if normalized is not None: if "reasoning_effort" in non_default_params: non_default_params["reasoning_effort"] = normalized @@ -237,7 +240,6 @@ class OpenAIGPT5Config(OpenAIGPTConfig): if not self.is_model_gpt_5_4_plus_model(model): non_default_params.pop("reasoning_effort", None) optional_params.pop("reasoning_effort", None) - reasoning_effort = None # gpt-5.1/5.2 support logprobs, top_p, top_logprobs only when reasoning_effort="none" supports_none = self._supports_reasoning_effort_level(model, "none") @@ -262,7 +264,9 @@ class OpenAIGPT5Config(OpenAIGPTConfig): temperature_value: Optional[float] = non_default_params.pop("temperature") if temperature_value is not None: # models supporting reasoning_effort="none" also support flexible temperature - if supports_none and (effective_effort == "none" or effective_effort is None): + if supports_none and ( + effective_effort == "none" or effective_effort is None + ): optional_params["temperature"] = temperature_value elif temperature_value == 1: optional_params["temperature"] = temperature_value diff --git a/litellm/llms/openai/realtime/http_transformation.py b/litellm/llms/openai/realtime/http_transformation.py index ff69ef987db..1663fcd1fcd 100644 --- a/litellm/llms/openai/realtime/http_transformation.py +++ b/litellm/llms/openai/realtime/http_transformation.py @@ -25,13 +25,17 @@ class OpenAIRealtimeHTTPConfig(BaseRealtimeHTTPConfig): or "" ) - def get_complete_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_complete_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: base = self.get_api_base(api_base).rstrip("/") if base.endswith("/v1"): base = base[:-3] return f"{base}/v1/realtime/client_secrets" - def get_realtime_calls_url(self, api_base: Optional[str], model: str, api_version: Optional[str] = None) -> str: + def get_realtime_calls_url( + self, api_base: Optional[str], model: str, api_version: Optional[str] = None + ) -> str: base = self.get_api_base(api_base).rstrip("/") if base.endswith("/v1"): base = base[:-3] diff --git a/litellm/llms/openai/responses/count_tokens/handler.py b/litellm/llms/openai/responses/count_tokens/handler.py index 721d07796ee..7fb5f6dad78 100644 --- a/litellm/llms/openai/responses/count_tokens/handler.py +++ b/litellm/llms/openai/responses/count_tokens/handler.py @@ -66,7 +66,9 @@ class OpenAICountTokensHandler(OpenAICountTokensConfig): llm_provider=litellm.LlmProviders.OPENAI ) - request_timeout = timeout if timeout is not None else litellm.request_timeout + request_timeout = ( + timeout if timeout is not None else litellm.request_timeout + ) response = await async_client.post( endpoint_url, diff --git a/litellm/llms/openai/responses/count_tokens/transformation.py b/litellm/llms/openai/responses/count_tokens/transformation.py index 3893775fc01..41d1a01ec66 100644 --- a/litellm/llms/openai/responses/count_tokens/transformation.py +++ b/litellm/llms/openai/responses/count_tokens/transformation.py @@ -52,9 +52,7 @@ class OpenAICountTokensConfig: "Authorization": f"Bearer {api_key}", } - def validate_request( - self, model: str, input: Union[str, List[Any]] - ) -> None: + def validate_request(self, model: str, input: Union[str, List[Any]]) -> None: if not model: raise ValueError("model parameter is required") @@ -139,20 +137,24 @@ class OpenAICountTokensConfig: if tool_calls: for tc in tool_calls: func = tc.get("function", {}) - input_items.append({ - "type": "function_call", - "call_id": tc.get("id", ""), - "name": func.get("name", ""), - "arguments": func.get("arguments", ""), - }) + input_items.append( + { + "type": "function_call", + "call_id": tc.get("id", ""), + "name": func.get("name", ""), + "arguments": func.get("arguments", ""), + } + ) elif not content: input_items.append({"role": "assistant", "content": content}) elif role == "tool": - input_items.append({ - "type": "function_call_output", - "call_id": msg.get("tool_call_id", ""), - "output": content if isinstance(content, str) else str(content), - }) + input_items.append( + { + "type": "function_call_output", + "call_id": msg.get("tool_call_id", ""), + "output": content if isinstance(content, str) else str(content), + } + ) instructions = "\n".join(instructions_parts) if instructions_parts else None return input_items, instructions diff --git a/litellm/llms/openai_like/dynamic_config.py b/litellm/llms/openai_like/dynamic_config.py index 8f216fe2144..3d66556e522 100644 --- a/litellm/llms/openai_like/dynamic_config.py +++ b/litellm/llms/openai_like/dynamic_config.py @@ -203,10 +203,7 @@ def create_responses_config_class(provider: SimpleProviderConfig): litellm_params: Optional[GenericLiteLLMParams], ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() - api_key = ( - litellm_params.api_key - or get_secret_str(provider.api_key_env) - ) + api_key = litellm_params.api_key or get_secret_str(provider.api_key_env) if api_key: headers["Authorization"] = f"Bearer {api_key}" return headers @@ -223,9 +220,7 @@ def create_responses_config_class(provider: SimpleProviderConfig): api_base = provider.base_url if api_base is None: - raise ValueError( - f"api_base is required for provider {provider.slug}" - ) + raise ValueError(f"api_base is required for provider {provider.slug}") api_base = api_base.rstrip("/") return f"{api_base}/responses" diff --git a/litellm/llms/perplexity/responses/transformation.py b/litellm/llms/perplexity/responses/transformation.py index f365ef07a61..cacdcdb9d7b 100644 --- a/litellm/llms/perplexity/responses/transformation.py +++ b/litellm/llms/perplexity/responses/transformation.py @@ -23,7 +23,6 @@ from litellm.types.utils import LlmProviders class PerplexityResponsesConfig(OpenAIResponsesAPIConfig): - def get_supported_openai_params(self, model: str) -> list: """Ref: https://docs.perplexity.ai/api-reference/responses-post""" return [ @@ -55,7 +54,11 @@ class PerplexityResponsesConfig(OpenAIResponsesAPIConfig): return headers def get_complete_url(self, api_base: Optional[str], litellm_params: dict) -> str: - api_base = api_base or get_secret_str("PERPLEXITY_API_BASE") or "https://api.perplexity.ai" + api_base = ( + api_base + or get_secret_str("PERPLEXITY_API_BASE") + or "https://api.perplexity.ai" + ) return f"{api_base.rstrip('/')}/v1/responses" def _ensure_message_type( @@ -86,7 +89,7 @@ class PerplexityResponsesConfig(OpenAIResponsesAPIConfig): if model.startswith("preset/"): input = self._validate_input_param(input) data: Dict = { - "preset": model[len("preset/"):], + "preset": model[len("preset/") :], "input": input, } data.update(response_api_optional_request_params) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index cbb7c7172f1..d7b96b4db7b 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -704,7 +704,9 @@ def _transform_request_body( # noqa: PLR0915 max_media_resolution ) if media_resolution_value and generation_config is not None: - generation_config["mediaResolution"] = media_resolution_value["level"] + generation_config["mediaResolution"] = media_resolution_value[ + "level" + ] data = RequestBody(contents=content) if system_instructions is not None: diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 7e686ee2799..3f1bccaccfc 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 @@ -1227,12 +1227,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "IMAGE_PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for prohibited image content.", } - _GEMINI_FINISH_REASON_KEYS = frozenset({ - "STOP", "MAX_TOKENS", "SAFETY", "RECITATION", "FINISH_REASON_UNSPECIFIED", - "MALFORMED_FUNCTION_CALL", "LANGUAGE", "OTHER", "BLOCKLIST", - "PROHIBITED_CONTENT", "SPII", "IMAGE_SAFETY", "IMAGE_PROHIBITED_CONTENT", - "TOO_MANY_TOOL_CALLS", "MALFORMED_RESPONSE", - }) + _GEMINI_FINISH_REASON_KEYS = frozenset( + { + "STOP", + "MAX_TOKENS", + "SAFETY", + "RECITATION", + "FINISH_REASON_UNSPECIFIED", + "MALFORMED_FUNCTION_CALL", + "LANGUAGE", + "OTHER", + "BLOCKLIST", + "PROHIBITED_CONTENT", + "SPII", + "IMAGE_SAFETY", + "IMAGE_PROHIBITED_CONTENT", + "TOO_MANY_TOOL_CALLS", + "MALFORMED_RESPONSE", + } + ) @staticmethod def get_finish_reason_mapping() -> Dict[str, OpenAIChatCompletionFinishReason]: diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index 68901340c7c..3ec0bdf22a5 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -40,37 +40,37 @@ class GoogleBatchEmbeddings(VertexLLM): ) -> Dict[str, Dict[str, str]]: """ Resolve Gemini file references (files/...) to get mime_type and uri. - + Args: input: EmbeddingInput that may contain file references api_key: Gemini API key sync_handler: HTTP client - + Returns: Dict mapping file name to {mime_type, uri} """ input_list = [input] if isinstance(input, str) else input resolved_files: Dict[str, Dict[str, str]] = {} - + for element in input_list: if isinstance(element, str) and _is_file_reference(element): url = f"https://generativelanguage.googleapis.com/v1beta/{element}" headers = {"x-goog-api-key": api_key} response = sync_handler.get(url=url, headers=headers) - + if response.status_code != 200: raise Exception( f"Error fetching file {element}: {response.status_code} {response.text}" ) - + file_data = response.json() resolved_files[element] = { "mime_type": file_data.get("mimeType", ""), "uri": file_data.get("uri", element), } - + return resolved_files - + async def _async_resolve_file_references( self, input: EmbeddingInput, @@ -79,37 +79,37 @@ class GoogleBatchEmbeddings(VertexLLM): ) -> Dict[str, Dict[str, str]]: """ Async version of _resolve_file_references. - + Args: input: EmbeddingInput that may contain file references api_key: Gemini API key async_handler: Async HTTP client - + Returns: Dict mapping file name to {mime_type, uri} """ input_list = [input] if isinstance(input, str) else input resolved_files: Dict[str, Dict[str, str]] = {} - + for element in input_list: if isinstance(element, str) and _is_file_reference(element): url = f"https://generativelanguage.googleapis.com/v1beta/{element}" headers = {"x-goog-api-key": api_key} response = await async_handler.get(url=url, headers=headers) - + if response.status_code != 200: raise Exception( f"Error fetching file {element}: {response.status_code} {response.text}" ) - + file_data = response.json() resolved_files[element] = { "mime_type": file_data.get("mimeType", ""), "uri": file_data.get("uri", element), } - + return resolved_files - + def batch_embeddings( self, model: str, @@ -238,7 +238,7 @@ class GoogleBatchEmbeddings(VertexLLM): raise Exception(f"Error: {response.status_code} {response.text}") _json_response = response.json() - + if use_embed_content: return process_embed_content_response( input=input, @@ -327,7 +327,7 @@ class GoogleBatchEmbeddings(VertexLLM): raise Exception(f"Error: {response.status_code} {response.text}") _json_response = response.json() - + if use_embed_content: return process_embed_content_response( input=input, diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index 41f477d9db9..0f6d85525d9 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -43,13 +43,13 @@ def _is_gcs_url(s: str) -> bool: def _infer_mime_type_from_gcs_url(gcs_url: str) -> str: """ Infer MIME type from GCS URL file extension. - + Args: gcs_url: GCS URL like gs://bucket/path/to/file.png - + Returns: str: Inferred MIME type - + Raises: ValueError: If file extension is not supported """ @@ -63,12 +63,12 @@ def _infer_mime_type_from_gcs_url(gcs_url: str) -> str: ".mov": "video/quicktime", ".pdf": "application/pdf", } - + gcs_url_lower = gcs_url.lower() for ext, mime_type in extension_to_mime.items(): if gcs_url_lower.endswith(ext): return mime_type - + raise ValueError( f"Unable to infer MIME type from GCS URL: {gcs_url}. " f"Supported extensions: {', '.join(extension_to_mime.keys())}" @@ -78,49 +78,49 @@ def _infer_mime_type_from_gcs_url(gcs_url: str) -> str: def _parse_data_url(data_url: str) -> Tuple[str, str]: """ Parse a data URL to extract the media type and base64 data. - + Args: data_url: Data URL in format: data:image/jpeg;base64,/9j/4AAQ... - + Returns: tuple: (media_type, base64_data) media_type: e.g., "image/jpeg", "video/mp4", "audio/mpeg" base64_data: The base64-encoded data without the prefix - + Raises: ValueError: If data URL format is invalid or MIME type is unsupported """ if not data_url.startswith("data:"): raise ValueError(f"Invalid data URL format: {data_url[:50]}...") - + if "," not in data_url: raise ValueError(f"Invalid data URL format (missing comma): {data_url[:50]}...") - + metadata, base64_data = data_url.split(",", 1) - + metadata = metadata[5:] - + if ";" in metadata: media_type = metadata.split(";")[0] else: media_type = metadata - + if media_type not in SUPPORTED_EMBEDDING_MIME_TYPES: raise ValueError( f"Unsupported MIME type for embedding: {media_type}. " f"Supported types: {', '.join(sorted(SUPPORTED_EMBEDDING_MIME_TYPES))}" ) - + return media_type, base64_data def _is_multimodal_input(input: EmbeddingInput) -> bool: """ Check if the input contains multimodal data (data URIs, file references, or GCS URLs). - + Args: input: EmbeddingInput (str or List[str]) - + Returns: bool: True if any element is a data URI, file reference, or GCS URL """ @@ -128,7 +128,7 @@ def _is_multimodal_input(input: EmbeddingInput) -> bool: input_list = [input] else: input_list = input - + for element in input_list: if isinstance(element, str): if element.startswith("data:") and ";base64," in element: @@ -137,7 +137,7 @@ def _is_multimodal_input(input: EmbeddingInput) -> bool: return True if _is_gcs_url(element): return True - + return False @@ -148,17 +148,17 @@ def transform_openai_input_gemini_content( The content to embed. Only the parts.text fields will be counted. """ gemini_model_name = "models/{}".format(model) - + gemini_params = optional_params.copy() if "dimensions" in gemini_params: gemini_params["outputDimensionality"] = gemini_params.pop("dimensions") - + requests: List[EmbedContentRequest] = [] if isinstance(input, str): request = EmbedContentRequest( model=gemini_model_name, content=ContentType(parts=[PartType(text=input)]), - **gemini_params + **gemini_params, ) requests.append(request) else: @@ -166,7 +166,7 @@ def transform_openai_input_gemini_content( request = EmbedContentRequest( model=gemini_model_name, content=ContentType(parts=[PartType(text=i)]), - **gemini_params + **gemini_params, ) requests.append(request) @@ -181,29 +181,29 @@ def transform_openai_input_gemini_embed_content( ) -> dict: """ Transform OpenAI embedding input to Gemini embedContent format (multimodal). - + Args: input: EmbeddingInput (str or List[str]) with text, data URIs, or file references model: Model name optional_params: Additional parameters (taskType, outputDimensionality, etc.) resolved_files: Dict mapping file names (files/abc) to {mime_type, uri} - + Returns: dict: Gemini embedContent request body with content.parts """ resolved_files = resolved_files or {} - + gemini_params = optional_params.copy() if "dimensions" in gemini_params: gemini_params["outputDimensionality"] = gemini_params.pop("dimensions") - + input_list = [input] if isinstance(input, str) else input parts: List[PartType] = [] - + for element in input_list: if not isinstance(element, str): raise ValueError(f"Unsupported input type: {type(element)}") - + if element.startswith("data:") and ";base64," in element: mime_type, base64_data = _parse_data_url(element) blob: BlobType = {"mime_type": mime_type, "data": base64_data} @@ -226,12 +226,12 @@ def transform_openai_input_gemini_embed_content( parts.append(PartType(file_data=file_data_ref)) else: parts.append(PartType(text=element)) - + request_body: dict = { "content": ContentType(parts=parts), **gemini_params, } - + return request_body @@ -243,30 +243,32 @@ def process_embed_content_response( ) -> EmbeddingResponse: """ Process Gemini embedContent response (single embedding for multimodal input). - + Args: input: Original input model_response: EmbeddingResponse to populate model: Model name response_json: Raw JSON response from embedContent endpoint - + Returns: EmbeddingResponse with single embedding """ if "embedding" not in response_json: - raise ValueError(f"embedContent response missing 'embedding' field: {response_json}") - + raise ValueError( + f"embedContent response missing 'embedding' field: {response_json}" + ) + embedding_data = response_json["embedding"] - + openai_embedding = Embedding( embedding=embedding_data["values"], index=0, object="embedding", ) - + model_response.data = [openai_embedding] model_response.model = model - + if _is_multimodal_input(input): prompt_tokens = 0 else: @@ -275,7 +277,7 @@ def process_embed_content_response( model_response.usage = Usage( prompt_tokens=prompt_tokens, total_tokens=prompt_tokens ) - + return model_response diff --git a/litellm/main.py b/litellm/main.py index 0448ceceb87..f2ce894ba38 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5197,7 +5197,9 @@ def embedding( # noqa: PLR0915 ) try: - model_info = get_model_info(model=model, custom_llm_provider="vertex_ai") + model_info = get_model_info( + model=model, custom_llm_provider="vertex_ai" + ) uses_embed_content = model_info.get("uses_embed_content", False) except Exception: uses_embed_content = False @@ -7634,12 +7636,15 @@ async def acount_tokens( from litellm.utils import ProviderConfigManager # Determine provider from model string - resolved_model, custom_llm_provider, dynamic_api_key, dynamic_api_base = ( - get_llm_provider( - model=model, - api_base=api_base, - api_key=api_key, - ) + ( + resolved_model, + custom_llm_provider, + dynamic_api_key, + dynamic_api_base, + ) = get_llm_provider( + model=model, + api_base=api_base, + api_key=api_key, ) # Use dynamic key/base if not explicitly provided diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 8a65936d5b1..45ec1bcebdf 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -409,7 +409,9 @@ async def update_mcp_server( # Pre-fetch existing record once if we need it for auth_type or credential logic existing = None - has_credentials = "credentials" in data_dict and data_dict["credentials"] is not None + has_credentials = ( + "credentials" in data_dict and data_dict["credentials"] is not None + ) if data.auth_type or has_credentials: existing = await prisma_client.db.litellm_mcpservertable.find_unique( where={"server_id": data.server_id} diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 4e68e3f28c1..af3a715051b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -329,7 +329,9 @@ async def authorize( lookup_name: Optional[str] = mcp_server_name or client_id client_ip = IPAddressUtils.get_mcp_client_ip(request) mcp_server = ( - global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) + global_mcp_server_manager.get_mcp_server_by_name( + lookup_name, client_ip=client_ip + ) if lookup_name else None ) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 164a48ded62..948d16ceff2 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -119,7 +119,9 @@ if MCP_AVAILABLE: prisma_client = get_prisma_client_or_throw( "Database not connected. Connect a database to use OAuth2 MCP tools." ) - cred = await get_user_oauth_credential(prisma_client, user_id, server_id) + cred = await get_user_oauth_credential( + prisma_client, user_id, server_id + ) if cred and cred.get("access_token"): if is_oauth_credential_expired(cred): verbose_logger.debug( @@ -192,7 +194,9 @@ if MCP_AVAILABLE: if c.get("access_token") and c.get("server_id") } except Exception: - verbose_logger.debug("Failed to bulk-fetch OAuth credentials", exc_info=True) + verbose_logger.debug( + "Failed to bulk-fetch OAuth credentials", exc_info=True + ) return {} def _create_tool_response_objects(tools, server_mcp_info): @@ -429,8 +433,12 @@ if MCP_AVAILABLE: # IP-filter error reporting if the resolved UUID is not in allowed_server_ids. _name_resolved = None if server_id not in allowed_server_ids: - _name_resolved = global_mcp_server_manager.get_mcp_server_by_name(server_id) - if _name_resolved is not None and _name_resolved.server_id in set(allowed_server_ids): + _name_resolved = global_mcp_server_manager.get_mcp_server_by_name( + server_id + ) + if _name_resolved is not None and _name_resolved.server_id in set( + allowed_server_ids + ): server_id = _name_resolved.server_id if server_id not in allowed_server_ids: @@ -478,7 +486,9 @@ if MCP_AVAILABLE: server, mcp_server_auth_headers, mcp_auth_header ) # Single-server request: targeted lookup is more efficient than a bulk fetch. - user_oauth_extra_headers = await _get_user_oauth_extra_headers(server, user_api_key_dict) + user_oauth_extra_headers = await _get_user_oauth_extra_headers( + server, user_api_key_dict + ) try: list_tools_result = await _get_tools_for_single_server( @@ -541,7 +551,9 @@ if MCP_AVAILABLE: server, mcp_server_auth_headers, mcp_auth_header ) user_oauth_extra_headers = await _get_user_oauth_extra_headers( - server, user_api_key_dict, prefetched_creds=prefetched_oauth_creds + server, + user_api_key_dict, + prefetched_creds=prefetched_oauth_creds, ) try: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4ce4881840f..da5f18d82c5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -915,7 +915,9 @@ if MCP_AVAILABLE: prisma_client = get_prisma_client_or_throw( "Database not connected. Connect a database to use OAuth2 MCP tools." ) - cred = await get_user_oauth_credential(prisma_client, user_id, server_id) + cred = await get_user_oauth_credential( + prisma_client, user_id, server_id + ) if cred and cred.get("access_token"): if is_oauth_credential_expired(cred): verbose_logger.debug( @@ -938,7 +940,9 @@ if MCP_AVAILABLE: Returns a dict keyed by server_id to avoid N+1 queries in asyncio.gather loops. """ - user_id = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + user_id = ( + getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + ) if not user_id: return {} try: @@ -1131,7 +1135,9 @@ if MCP_AVAILABLE: # If no OAuth2 token came from request headers, fall back to pre-fetched creds if extra_headers is None and server.auth_type == MCPAuth.oauth2: extra_headers = await _get_user_oauth_extra_headers_from_db( - server, user_api_key_auth, prefetched_creds=_prefetched_oauth_creds + server, + user_api_key_auth, + prefetched_creds=_prefetched_oauth_creds, ) try: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a64ceffda1e..1174740948b 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,40 +1,59 @@ import enum import json from datetime import datetime -from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal, - Optional, Union) +from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union import httpx -from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator, - model_validator) +from pydantic import ( + BaseModel, + ConfigDict, + Field, + Json, + field_validator, + model_validator, +) from typing_extensions import Required, TypedDict from litellm._uuid import uuid from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject, - ResponsesAPIResponse) -from litellm.types.mcp import (MCPAuthType, MCPCredentials, MCPTransport, - MCPTransportType) +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIFileObject, + ResponsesAPIResponse, +) +from litellm.types.mcp import ( + MCPAuthType, + MCPCredentials, + MCPTransport, + MCPTransportType, +) from litellm.types.mcp_server.mcp_server_manager import MCPInfo from litellm.types.router import RouterErrors, UpdateRouterConfig from litellm.types.secret_managers.main import KeyManagementSystem -from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse, - GenericBudgetConfigType, ImageResponse, - LiteLLMBatch, LiteLLMFineTuningJob, - LiteLLMPydanticObjectBase, ModelResponse, - ProviderField, StandardCallbackDynamicParams, - StandardLoggingGuardrailInformation, - StandardLoggingMCPToolCall, - StandardLoggingModelInformation, - StandardLoggingPayloadErrorInformation, - StandardLoggingPayloadStatus, - StandardLoggingVectorStoreRequest, - StandardPassThroughResponseObject, - TextCompletionResponse) +from litellm.types.utils import ( + CallTypes, + CostBreakdown, + EmbeddingResponse, + GenericBudgetConfigType, + ImageResponse, + LiteLLMBatch, + LiteLLMFineTuningJob, + LiteLLMPydanticObjectBase, + ModelResponse, + ProviderField, + StandardCallbackDynamicParams, + StandardLoggingGuardrailInformation, + StandardLoggingMCPToolCall, + StandardLoggingModelInformation, + StandardLoggingPayloadErrorInformation, + StandardLoggingPayloadStatus, + StandardLoggingVectorStoreRequest, + StandardPassThroughResponseObject, + TextCompletionResponse, +) from litellm.types.videos.main import VideoObject -from .types_utils.utils import (get_instance_fn, - validate_custom_validate_return_type) +from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -2445,7 +2464,9 @@ class UserAPIKeyAuth( user_max_budget: Optional[float] = None request_route: Optional[str] = None user: Optional[Any] = None # Expanded user object when expand=user is used - created_by_user: Optional[Any] = None # Expanded created_by user when expand=user is used + created_by_user: Optional[ + Any + ] = None # Expanded created_by user when expand=user is used end_user_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None model_config = ConfigDict(arbitrary_types_allowed=True) @@ -2489,8 +2510,7 @@ class UserAPIKeyAuth( This is used to track number of requests/spend for health check calls. """ - from litellm.constants import \ - LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME return cls( api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, @@ -2522,8 +2542,7 @@ class UserAPIKeyAuth( This is used to track actions performed by automated system jobs. """ - from litellm.constants import \ - LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME return cls( api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, @@ -2929,8 +2948,7 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): @model_validator(mode="after") def mask_api_keys(self): - from litellm.litellm_core_utils.sensitive_data_masker import \ - SensitiveDataMasker + from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker masker = SensitiveDataMasker(sensitive_patterns={"key"}) diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index d988579b919..bf76f99db69 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -107,9 +107,13 @@ def get_key_models( """ all_models: List[str] = [] if len(user_api_key_dict.models) > 0: - all_models = list(user_api_key_dict.models) # copy to avoid mutating cached objects + all_models = list( + user_api_key_dict.models + ) # copy to avoid mutating cached objects if SpecialModelNames.all_team_models.value in all_models: - all_models = list(user_api_key_dict.team_models) # copy to avoid mutating cached objects + all_models = list( + user_api_key_dict.team_models + ) # copy to avoid mutating cached objects if SpecialModelNames.all_proxy_models.value in all_models: all_models = list(proxy_model_list) # copy to avoid mutating caller's list if include_model_access_groups: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 5a3f3a984b3..7b9e0a43731 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -249,24 +249,21 @@ async def create_response( def _is_azure_model_router_request(model: str) -> bool: """ Check if the requested model is an Azure Model Router. - + Azure Model Router models follow the pattern: - azure_ai/model_router/ - azure_ai/model-router - model_router/ - model-router - + Args: model: The requested model name - + Returns: bool: True if this is an Azure Model Router request """ model_lower = model.lower() - return ( - "model-router" in model_lower - or "model_router" in model_lower - ) + return "model-router" in model_lower or "model_router" in model_lower def _override_openai_response_model( @@ -1233,7 +1230,9 @@ class ProxyBaseLLMRequestProcessing: data=self.data, user_api_key_dict=user_api_key_dict, response=None, - request_headers=(self.data.get("proxy_server_request") or {}).get("headers", {}), + request_headers=(self.data.get("proxy_server_request") or {}).get( + "headers", {} + ), ) if callback_headers: headers.update(callback_headers) diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 5fa9546e006..64f860fc4f1 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -21,11 +21,15 @@ router = APIRouter() class CredentialHelperUtils: @staticmethod - def encrypt_credential_values(credential: CredentialItem, new_encryption_key: Optional[str] = None) -> CredentialItem: + def encrypt_credential_values( + credential: CredentialItem, new_encryption_key: Optional[str] = None + ) -> CredentialItem: """Encrypt values in credential.credential_values and add to DB""" encrypted_credential_values = {} for key, value in (credential.credential_values or {}).items(): - encrypted_credential_values[key] = encrypt_value_helper(value, new_encryption_key) + encrypted_credential_values[key] = encrypt_value_helper( + value, new_encryption_key + ) # Return a new object to avoid mutating the caller's credential, which # is kept in memory and should remain unencrypted. @@ -145,7 +149,9 @@ async def get_credentials( async def get_credential_by_name( request: Request, fastapi_response: Response, - credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"), + credential_name: str = Path( + ..., description="The credential name, percent-decoded; may contain slashes" + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -223,7 +229,9 @@ async def get_credential_by_model( async def delete_credential( request: Request, fastapi_response: Response, - credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"), + credential_name: str = Path( + ..., description="The credential name, percent-decoded; may contain slashes" + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -253,7 +261,9 @@ async def delete_credential( def update_db_credential( - db_credential: CredentialItem, updated_patch: CredentialItem, new_encryption_key: Optional[str] = None + db_credential: CredentialItem, + updated_patch: CredentialItem, + new_encryption_key: Optional[str] = None, ) -> CredentialItem: """ Update a credential in the DB. @@ -300,7 +310,9 @@ async def update_credential( request: Request, fastapi_response: Response, credential: CredentialItem, - credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"), + credential_name: str = Path( + ..., description="The credential name, percent-decoded; may contain slashes" + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index b84c74bee4c..8d94a7051a3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -1172,9 +1172,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): yield mock_response_stream except Exception as e: - verbose_proxy_logger.error( - f"Error masking streaming PII output: {str(e)}" - ) + verbose_proxy_logger.error(f"Error masking streaming PII output: {str(e)}") for chunk in all_chunks: yield chunk diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index 148be10da8f..084c2f47d0f 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -167,7 +167,10 @@ def new_budget_request(data: NewCustomerRequest) -> Optional[BudgetNewRequest]: if budget_kv_pairs: budget_request = BudgetNewRequest(**budget_kv_pairs) - if budget_request.budget_reset_at is None and budget_request.budget_duration is not None: + if ( + budget_request.budget_reset_at is None + and budget_request.budget_duration is not None + ): budget_request.budget_reset_at = datetime.utcnow() + timedelta( seconds=duration_in_seconds(duration=budget_request.budget_duration) ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 77dc1c9724d..3e5b729cea6 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -424,11 +424,17 @@ if MCP_AVAILABLE: inherited_credentials["scopes"] = existing_server.scopes # AWS SigV4 fields if existing_server.aws_access_key_id: - inherited_credentials["aws_access_key_id"] = existing_server.aws_access_key_id + inherited_credentials[ + "aws_access_key_id" + ] = existing_server.aws_access_key_id if existing_server.aws_secret_access_key: - inherited_credentials["aws_secret_access_key"] = existing_server.aws_secret_access_key + inherited_credentials[ + "aws_secret_access_key" + ] = existing_server.aws_secret_access_key if existing_server.aws_session_token: - inherited_credentials["aws_session_token"] = existing_server.aws_session_token + inherited_credentials[ + "aws_session_token" + ] = existing_server.aws_session_token if existing_server.aws_region_name: inherited_credentials["aws_region_name"] = existing_server.aws_region_name if existing_server.aws_service_name: @@ -734,8 +740,7 @@ if MCP_AVAILABLE: check_db_only=True, ) user_in_team = any( - m.user_id is not None - and m.user_id == user_api_key_dict.user_id + m.user_id is not None and m.user_id == user_api_key_dict.user_id for m in team_obj.members_with_roles ) if not user_in_team: @@ -744,20 +749,26 @@ if MCP_AVAILABLE: detail="You do not have permission to view MCP servers for this team.", ) - redacted_mcp_servers = await _get_team_scoped_mcp_server_list(sanitized_team_id) + redacted_mcp_servers = await _get_team_scoped_mcp_server_list( + sanitized_team_id + ) else: user_mcp_management_mode = _get_user_mcp_management_mode() if user_mcp_management_mode == "view_all" and not is_restricted_virtual_key: - servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered() + servers = ( + await global_mcp_server_manager.get_all_mcp_servers_unfiltered() + ) redacted_mcp_servers = _redact_mcp_credentials_list(servers) else: auth_contexts = await build_effective_auth_contexts(user_api_key_dict) aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {} for auth_context in auth_contexts: - servers = await global_mcp_server_manager.get_all_allowed_mcp_servers( - user_api_key_auth=auth_context + servers = ( + await global_mcp_server_manager.get_all_allowed_mcp_servers( + user_api_key_auth=auth_context + ) ) for server in servers: if server.server_id not in aggregated_servers: @@ -1084,8 +1095,11 @@ if MCP_AVAILABLE: client_ip = IPAddressUtils.get_mcp_client_ip(request) registry_server = global_mcp_server_manager.get_mcp_server_by_id(server_id) - if registry_server is not None and not global_mcp_server_manager._is_server_accessible_from_ip( - registry_server, client_ip + if ( + registry_server is not None + and not global_mcp_server_manager._is_server_accessible_from_ip( + registry_server, client_ip + ) ): registry_server = None if registry_server is None: @@ -1120,8 +1134,10 @@ if MCP_AVAILABLE: exists = does_mcp_server_exist(mcp_server_records, server_id) else: # Registry/config server: use same access logic as list endpoint - allowed_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers( - user_api_key_dict + allowed_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers( + user_api_key_dict + ) ) exists = mcp_server.server_id in allowed_server_ids @@ -1319,10 +1335,9 @@ if MCP_AVAILABLE: global_mcp_server_manager, ) - server = ( - global_mcp_server_manager.get_mcp_server_by_id(server_id) - or global_mcp_server_manager.get_mcp_server_by_name(server_id) - ) + server = global_mcp_server_manager.get_mcp_server_by_id( + server_id + ) or global_mcp_server_manager.get_mcp_server_by_name(server_id) if server is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -1647,7 +1662,9 @@ if MCP_AVAILABLE: # Only delete if the stored credential is actually an OAuth2 token. # This prevents accidentally deleting a BYOK credential if one exists # for the same (user_id, server_id) pair. - cred_to_delete = await get_user_oauth_credential(prisma_client, user_id, server_id) + cred_to_delete = await get_user_oauth_credential( + prisma_client, user_id, server_id + ) if cred_to_delete is not None: try: await delete_user_credential(prisma_client, user_id, server_id) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index ee1868fc740..9c8e6f7282b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -759,19 +759,25 @@ async def new_team( # noqa: PLR0915 if data.max_budget is not None and data.max_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"} + detail={ + "error": f"max_budget cannot be negative. Received: {data.max_budget}" + }, ) if data.team_member_budget is not None and data.team_member_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"team_member_budget cannot be negative. Received: {data.team_member_budget}"} + detail={ + "error": f"team_member_budget cannot be negative. Received: {data.team_member_budget}" + }, ) if data.soft_budget is not None and data.soft_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"} + detail={ + "error": f"soft_budget cannot be negative. Received: {data.soft_budget}" + }, ) - + if data.soft_budget is not None: if data.max_budget is not None: # If max_budget is set, soft_budget must be strictly lower than max_budget @@ -780,7 +786,7 @@ async def new_team( # noqa: PLR0915 status_code=400, detail={ "error": f"soft_budget ({data.soft_budget}) must be strictly lower than max_budget ({data.max_budget})" - } + }, ) # Check if license is over limit @@ -940,12 +946,16 @@ async def new_team( # noqa: PLR0915 complete_team_data.members_with_roles = [] complete_team_data_dict = complete_team_data.model_dump(exclude_none=True) - + # Serialize router_settings to JSON (matching key creation pattern) router_settings_value = getattr(data, "router_settings", None) - router_settings_json = safe_dumps(router_settings_value) if router_settings_value is not None else safe_dumps({}) + router_settings_json = ( + safe_dumps(router_settings_value) + if router_settings_value is not None + else safe_dumps({}) + ) complete_team_data_dict["router_settings"] = router_settings_json - + complete_team_data_dict = prisma_client.jsonify_team_object( db_data=complete_team_data_dict ) @@ -1121,7 +1131,9 @@ async def fetch_and_validate_organization( validate_team_org_change( team=LiteLLM_TeamTable(**existing_team_row.model_dump()), - organization=LiteLLM_OrganizationTableWithMembers(**organization_row.model_dump()), + organization=LiteLLM_OrganizationTableWithMembers( + **organization_row.model_dump() + ), llm_router=llm_router, ) @@ -1129,7 +1141,9 @@ async def fetch_and_validate_organization( def validate_team_org_change( - team: LiteLLM_TeamTable, organization: LiteLLM_OrganizationTableWithMembers, llm_router: Router + team: LiteLLM_TeamTable, + organization: LiteLLM_OrganizationTableWithMembers, + llm_router: Router, ) -> bool: """ Validate that a team can be moved to an organization. @@ -1180,7 +1194,9 @@ def validate_team_org_change( # Check if the team's user_id is a member of the org team_members = [m.user_id for m in team.members_with_roles] - org_members = [m.user_id for m in organization.members] if organization.members else [] + org_members = ( + [m.user_id for m in organization.members] if organization.members else [] + ) not_in_org = [ m for m in team_members @@ -1226,7 +1242,7 @@ def validate_team_org_change( "/team/update", tags=["team management"], dependencies=[Depends(user_api_key_auth)] ) @management_endpoint_wrapper -async def update_team( # noqa: PLR0915 +async def update_team( # noqa: PLR0915 data: UpdateTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -1314,24 +1330,32 @@ async def update_team( # noqa: PLR0915 ) if data.team_id is None: - raise HTTPException(status_code=400, detail={"error": "No team id passed in"}) + raise HTTPException( + status_code=400, detail={"error": "No team id passed in"} + ) verbose_proxy_logger.debug("/team/update - %s", data) # Validate budget values are not negative if data.max_budget is not None and data.max_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"} + detail={ + "error": f"max_budget cannot be negative. Received: {data.max_budget}" + }, ) if data.team_member_budget is not None and data.team_member_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"team_member_budget cannot be negative. Received: {data.team_member_budget}"} + detail={ + "error": f"team_member_budget cannot be negative. Received: {data.team_member_budget}" + }, ) if data.soft_budget is not None and data.soft_budget < 0: raise HTTPException( status_code=400, - detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"} + detail={ + "error": f"soft_budget cannot be negative. Received: {data.soft_budget}" + }, ) existing_team_row = await prisma_client.db.litellm_teamtable.find_unique( @@ -1343,28 +1367,38 @@ async def update_team( # noqa: PLR0915 status_code=404, detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - + if data.soft_budget is not None: - max_budget_to_check = data.max_budget if data.max_budget is not None else existing_team_row.max_budget + max_budget_to_check = ( + data.max_budget + if data.max_budget is not None + else existing_team_row.max_budget + ) if max_budget_to_check is not None: if data.soft_budget >= max_budget_to_check: raise HTTPException( status_code=400, detail={ "error": f"soft_budget ({data.soft_budget}) must be strictly lower than max_budget ({max_budget_to_check})" - } + }, ) - + if data.max_budget is not None: - existing_soft_budget = getattr(existing_team_row, 'soft_budget', None) - soft_budget_to_check = data.soft_budget if data.soft_budget is not None else existing_soft_budget - if soft_budget_to_check is not None and isinstance(soft_budget_to_check, (int, float)): + existing_soft_budget = getattr(existing_team_row, "soft_budget", None) + soft_budget_to_check = ( + data.soft_budget + if data.soft_budget is not None + else existing_soft_budget + ) + if soft_budget_to_check is not None and isinstance( + soft_budget_to_check, (int, float) + ): if data.max_budget <= soft_budget_to_check: raise HTTPException( status_code=400, detail={ "error": f"max_budget ({data.max_budget}) must be strictly greater than soft_budget ({soft_budget_to_check})" - } + }, ) if ( @@ -1465,16 +1499,19 @@ async def update_team( # noqa: PLR0915 updated_kv["model_id"] = _model_id # Serialize router_settings to JSON if present (matching key update pattern) - if "router_settings" in updated_kv and updated_kv["router_settings"] is not None: + if ( + "router_settings" in updated_kv + and updated_kv["router_settings"] is not None + ): updated_kv["router_settings"] = safe_dumps(updated_kv["router_settings"]) updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv) - team_row: Optional[LiteLLM_TeamTable] = ( - await prisma_client.db.litellm_teamtable.update( - where={"team_id": data.team_id}, - data=updated_kv, - include={"litellm_model_table": True}, # type: ignore - ) + team_row: Optional[ + LiteLLM_TeamTable + ] = await prisma_client.db.litellm_teamtable.update( + where={"team_id": data.team_id}, + data=updated_kv, + include={"litellm_model_table": True}, # type: ignore ) if team_row is None or team_row.team_id is None: @@ -1483,7 +1520,9 @@ async def update_team( # noqa: PLR0915 detail={"error": "Team doesn't exist. Got={}".format(team_row)}, ) - verbose_proxy_logger.info("Successfully updated team - %s, info", team_row.team_id) + verbose_proxy_logger.info( + "Successfully updated team - %s, info", team_row.team_id + ) await _cache_team_object( team_id=team_row.team_id, team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()), @@ -1834,14 +1873,14 @@ async def _validate_and_populate_member_user_info( ) -> Member: """ Validate and populate user_email/user_id for a member. - + Logic: 1. If both user_email and user_id are provided, verify they belong to the same user (use user_email as source of truth) 2. If only user_email is provided, populate user_id from DB 3. If only user_id is provided, populate user_email from DB (if user exists) 4. If only user_id is provided and doesn't exist, allow it to pass with user_email as None (will be upserted later) 5. If user_email and user_id mismatch, throw error - + Returns a Member with user_email and user_id populated (user_email may be None if only user_id provided and user doesn't exist). """ if member.user_email is None and member.user_id is None: @@ -1849,7 +1888,7 @@ async def _validate_and_populate_member_user_info( status_code=400, detail={"error": "Either user_id or user_email must be provided"}, ) - + # Case 1: Both user_email and user_id provided - verify they match if member.user_email is not None and member.user_id is not None: # Use user_email as source of truth @@ -1859,13 +1898,13 @@ async def _validate_and_populate_member_user_info( table_name="user", query_type="find_all", ) - + if users_by_email is None or ( isinstance(users_by_email, list) and len(users_by_email) == 0 ): # User doesn't exist yet - this is fine, will be created later return member - + if isinstance(users_by_email, list) and len(users_by_email) > 1: raise HTTPException( status_code=400, @@ -1873,10 +1912,10 @@ async def _validate_and_populate_member_user_info( "error": f"Multiple users found with email '{member.user_email}'. Please use 'user_id' instead." }, ) - + # Get the single user user_by_email = users_by_email[0] - + # Verify the user_id matches if user_by_email.user_id != member.user_id: raise HTTPException( @@ -1885,56 +1924,61 @@ async def _validate_and_populate_member_user_info( "error": f"user_email '{member.user_email}' and user_id '{member.user_id}' do not belong to the same user." }, ) - + # Both match, return as is return member - + # Case 2: Only user_email provided - populate user_id from DB if member.user_email is not None and member.user_id is None: user_by_email = await prisma_client.db.litellm_usertable.find_first( where={"user_email": {"equals": member.user_email, "mode": "insensitive"}} ) - + if user_by_email is None: # User doesn't exist yet - this is fine, will be created later return member - + # Check for multiple users with same email users_by_email = await prisma_client.get_data( key_val={"user_email": member.user_email}, table_name="user", query_type="find_all", ) - - if users_by_email and isinstance(users_by_email, list) and len(users_by_email) > 1: + + if ( + users_by_email + and isinstance(users_by_email, list) + and len(users_by_email) > 1 + ): raise HTTPException( status_code=400, detail={ "error": f"Multiple users found with email '{member.user_email}'. Please use 'user_id' instead." }, ) - + # Populate user_id member.user_id = user_by_email.user_id return member - + # Case 3: Only user_id provided - populate user_email from DB if user exists if member.user_id is not None and member.user_email is None: user_by_id = await prisma_client.db.litellm_usertable.find_unique( where={"user_id": member.user_id} ) - + if user_by_id is None: # User doesn't exist yet - allow it to pass with user_email as None # Will be upserted later with just user_id and null email return member - + # Populate user_email member.user_email = user_by_id.user_email return member - + return member + @router.post( "/team/member_add", tags=["team management"], @@ -2023,14 +2067,16 @@ async def team_member_add( prisma_client=prisma_client, ) - updated_team, updated_users, updated_team_memberships = ( - await _add_team_members_to_team( - data=data, - complete_team_data=complete_team_data, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - ) + ( + updated_team, + updated_users, + updated_team_memberships, + ) = await _add_team_members_to_team( + data=data, + complete_team_data=complete_team_data, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name=litellm_proxy_admin_name, ) # Check if updated_team is None @@ -2212,15 +2258,15 @@ async def team_member_delete( ) # Fetch keys before deletion to persist them - keys_to_delete: List[LiteLLM_VerificationToken] = ( - await prisma_client.db.litellm_verificationtoken.find_many( - where={ - "user_id": {"in": list(user_ids_to_delete)}, - "team_id": data.team_id, - } - ) + keys_to_delete: List[ + LiteLLM_VerificationToken + ] = await prisma_client.db.litellm_verificationtoken.find_many( + where={ + "user_id": {"in": list(user_ids_to_delete)}, + "team_id": data.team_id, + } ) - + if keys_to_delete: await _persist_deleted_verification_tokens( keys=keys_to_delete, @@ -2602,10 +2648,10 @@ async def delete_team( team_rows: List[LiteLLM_TeamTable] = [] for team_id in data.team_ids: try: - team_row_base: Optional[BaseModel] = ( - await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) + team_row_base: Optional[ + BaseModel + ] = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} ) if team_row_base is None: raise Exception @@ -2664,10 +2710,10 @@ async def delete_team( _persist_deleted_verification_tokens, ) - keys_to_delete: List[LiteLLM_VerificationToken] = ( - await prisma_client.db.litellm_verificationtoken.find_many( - where={"team_id": {"in": data.team_ids}} - ) + keys_to_delete: List[ + LiteLLM_VerificationToken + ] = await prisma_client.db.litellm_verificationtoken.find_many( + where={"team_id": {"in": data.team_ids}} ) if keys_to_delete: @@ -2706,7 +2752,6 @@ async def delete_team( return deleted_teams - def _transform_teams_to_deleted_records( teams: List[LiteLLM_TeamTable], user_api_key_dict: UserAPIKeyAuth, @@ -2729,7 +2774,13 @@ def _transform_teams_to_deleted_records( ) record = deleted_record.model_dump() - for json_field in ["members_with_roles", "metadata", "model_spend", "model_max_budget", "router_settings"]: + for json_field in [ + "members_with_roles", + "metadata", + "model_spend", + "model_max_budget", + "router_settings", + ]: if json_field in record and record[json_field] is not None: record[json_field] = json.dumps(record[json_field]) @@ -2748,9 +2799,7 @@ async def _save_deleted_team_records( """Save deleted team records to the database.""" if not records: return - await prisma_client.db.litellm_deletedteamtable.create_many( - data=records - ) + await prisma_client.db.litellm_deletedteamtable.create_many(data=records) async def _persist_deleted_team_records( @@ -2770,6 +2819,7 @@ async def _persist_deleted_team_records( prisma_client=prisma_client, ) + async def validate_membership( user_api_key_dict: UserAPIKeyAuth, team_table: LiteLLM_TeamTable ): @@ -2806,9 +2856,7 @@ async def validate_membership( ) # Check direct team membership - if user_api_key_dict.user_id in [ - m.user_id for m in team_table.members_with_roles - ]: + if user_api_key_dict.user_id in [m.user_id for m in team_table.members_with_roles]: return # Check if user is an org admin for the team's organization @@ -2827,8 +2875,6 @@ async def validate_membership( ) - - async def _add_team_member_budget_table( team_member_budget_id: str, prisma_client: PrismaClient, @@ -2887,11 +2933,11 @@ async def team_info( ) try: - team_info: Optional[BaseModel] = ( - await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id}, - include={"object_permission": True}, - ) + team_info: Optional[ + BaseModel + ] = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id}, + include={"object_permission": True}, ) if team_info is None: raise Exception @@ -3346,7 +3392,9 @@ async def list_team_v2( order=order_by if order_by else {"created_at": "desc"}, # Default sort ) # Get total count for pagination - total_count = await prisma_client.db.litellm_teamtable.count(where=where_conditions) + total_count = await prisma_client.db.litellm_teamtable.count( + where=where_conditions + ) # Calculate total pages total_pages = -(-total_count // page_size) # Ceiling division diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 0f426bf6045..8aba8307b9d 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -18,7 +18,6 @@ if TYPE_CHECKING: LiteLLM_ObjectPermissionTable, LiteLLM_TeamTableCachedObj, ) - async def attach_object_permission_to_dict( @@ -27,30 +26,32 @@ async def attach_object_permission_to_dict( ) -> Dict: """ Helper method to attach object_permission to a dictionary if object_permission_id is set. - + This function: 1. Checks if the dictionary has an object_permission_id 2. If found, queries the database for the corresponding object permission 3. Converts the object permission to a dictionary format 4. Attaches it to the input dictionary under the 'object_permission' key - + Args: data_dict: The dictionary to attach object_permission to prisma_client: The database client - + Returns: Dict: The input dictionary with object_permission attached if found - + Raises: ValueError: If prisma_client is None """ if prisma_client is None: raise ValueError("Prisma client not found") - + object_permission_id = data_dict.get("object_permission_id") if object_permission_id: - object_permission = await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": object_permission_id}, + object_permission = ( + await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": object_permission_id}, + ) ) if object_permission: # Convert to dict if needed @@ -168,21 +169,24 @@ async def _set_object_permission( if not isinstance(permission_data, dict): data_json.pop("object_permission") return data_json - + # Clean data: exclude None values and object_permission_id clean_data = { - k: v for k, v in permission_data.items() + k: v + for k, v in permission_data.items() if v is not None and k != "object_permission_id" } - + # Serialize mcp_tool_permissions to JSON string for GraphQL compatibility if "mcp_tool_permissions" in clean_data: - clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"]) - + clean_data["mcp_tool_permissions"] = safe_dumps( + clean_data["mcp_tool_permissions"] + ) + created_permission = await prisma_client.db.litellm_objectpermissiontable.create( data=clean_data ) - + data_json["object_permission_id"] = created_permission.object_permission_id data_json.pop("object_permission") return data_json @@ -204,10 +208,10 @@ async def _resolve_team_allowed_mcp_servers( ) direct_servers: List[str] = team_object_permission.mcp_servers or [] - access_group_servers: List[str] = ( - await MCPRequestHandler._get_mcp_servers_from_access_groups( - team_object_permission.mcp_access_groups or [] - ) + access_group_servers: List[ + str + ] = await MCPRequestHandler._get_mcp_servers_from_access_groups( + team_object_permission.mcp_access_groups or [] ) raw_tool_perms = team_object_permission.mcp_tool_permissions or {} if isinstance(raw_tool_perms, str): @@ -359,4 +363,4 @@ async def validate_key_mcp_servers_against_team( raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={"error": detail}, - ) \ No newline at end of file + ) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 8de9099687c..eb9d303eb91 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -404,7 +404,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): headers=headers, params=requested_query_params, ) - elif HttpPassThroughEndpointHelpers.is_multipart(request) is True and not _parsed_body: + elif ( + HttpPassThroughEndpointHelpers.is_multipart(request) is True + and not _parsed_body + ): # Only use multipart handler if we don't have a parsed body # (parsed body means it was JSON despite multipart content-type header) return await HttpPassThroughEndpointHelpers.make_multipart_http_request( @@ -681,8 +684,10 @@ async def pass_through_request( # noqa: PLR0915 # Skip body parsing for multipart requests - make_multipart_http_request will handle it # But if custom_body is provided (e.g., JSON parsed despite multipart content-type), use it - is_multipart = HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body - + is_multipart = ( + HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body + ) + if custom_body: _parsed_body = custom_body elif is_multipart: @@ -1133,7 +1138,9 @@ def create_pass_through_route( fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), subpath: str = "", # captures sub-paths when include_subpath=True - custom_body: Optional[dict] = None, # caller-supplied body takes precedence over request-parsed body + custom_body: Optional[ + dict + ] = None, # caller-supplied body takes precedence over request-parsed body ): from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( InitPassThroughEndpointHelpers, @@ -2062,7 +2069,9 @@ class InitPassThroughEndpointHelpers: """ ## CHECK IF MAPPED PASS THROUGH ENDPOINT for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value: - full_mapped_route = InitPassThroughEndpointHelpers._build_full_path_with_root(mapped_route) + full_mapped_route = ( + InitPassThroughEndpointHelpers._build_full_path_with_root(mapped_route) + ) if route.startswith(full_mapped_route): return True diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index abe65257268..153c3c2dba0 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -854,7 +854,9 @@ def run_server( # noqa: PLR0915 ): check_prisma_schema_diff(db_url=None) else: - if not PrismaManager.setup_database(use_migrate=not use_prisma_db_push): + if not PrismaManager.setup_database( + use_migrate=not use_prisma_db_push + ): print( # noqa "\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. " "The proxy cannot start safely. Please check your database connection and migration status.\033[0m" diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0c24e8b6b5b..e01ee8e9d9e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5437,9 +5437,7 @@ def _restamp_streaming_chunk_model( return chunk, model_mismatch_logged # For Azure Model Router, preserve the actual model used in each chunk - if _is_azure_model_router_request( - requested_model_from_client - ): + if _is_azure_model_router_request(requested_model_from_client): return chunk, model_mismatch_logged downstream_model = ( diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index bb286d1fd0d..d2975ba5fcf 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -114,14 +114,14 @@ async def create_realtime_client_secret( ) data = {"model": model} - + # If session is provided, use it; otherwise create one from model if req.session: data["session"] = req.session.model_dump(exclude_none=True) elif req.model: # User provided model at root level, convert to session format data["session"] = {"type": "realtime", "model": model} - + if req.expires_after: data["expires_after"] = req.expires_after.model_dump(exclude_none=True) @@ -275,7 +275,7 @@ async def proxy_realtime_calls( status_code=http_status.HTTP_401_UNAUTHORIZED, media_type="application/json", ) - + openai_ephemeral_key = decoded_payload.get("ephemeral_key", "") model = ( decoded_payload.get("model_id") @@ -328,9 +328,7 @@ async def proxy_realtime_calls( call_type="arealtime_calls", ) - verbose_proxy_logger.debug( - "WebRTC: /v1/realtime/calls (model=%s)", model - ) + verbose_proxy_logger.debug("WebRTC: /v1/realtime/calls (model=%s)", model) llm_call = await route_request( data=data, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 3019de617f0..369b56c0f58 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1992,7 +1992,9 @@ class ProxyLogging: merged_headers: Dict[str, str] = {} try: # Build litellm_call_info — normalized routing metadata for callbacks - litellm_call_info = self._build_litellm_call_info(data=data, response=response) + litellm_call_info = self._build_litellm_call_info( + data=data, response=response + ) for callback in litellm.callbacks: _callback: Optional[CustomLogger] = None @@ -2029,9 +2031,7 @@ class ProxyLogging: return merged_headers @staticmethod - def _build_litellm_call_info( - data: dict, response: Any - ) -> Dict[str, Any]: + def _build_litellm_call_info(data: dict, response: Any) -> Dict[str, Any]: """ Build a normalized dict of routing metadata from response._hidden_params and data, abstracting away the metadata vs litellm_metadata split. diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 0b688077744..63ce5c104dc 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -288,8 +288,12 @@ async def vector_store_create( ) -@router.get("/v1/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)]) -@router.get("/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)]) +@router.get( + "/v1/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)] +) +@router.get( + "/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)] +) async def vector_store_retrieve( request: Request, vector_store_id: str, @@ -421,8 +425,12 @@ async def vector_store_list( ) -@router.post("/v1/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)]) -@router.post("/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)]) +@router.post( + "/v1/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)] +) +@router.post( + "/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)] +) async def vector_store_update( request: Request, vector_store_id: str, @@ -487,8 +495,12 @@ async def vector_store_update( ) -@router.delete("/v1/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)]) -@router.delete("/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)]) +@router.delete( + "/v1/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)] +) +@router.delete( + "/vector_stores/{vector_store_id}", dependencies=[Depends(user_api_key_auth)] +) async def vector_store_delete( request: Request, vector_store_id: str, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 2e5efe1c338..7dcdc0d8d9a 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -78,11 +78,7 @@ def _get_realtime_http_provider_config( resolved_api_key = provider_config.get_api_key(api_key=raw_api_key) else: # Fallback for providers without a dedicated HTTP config (treated as OpenAI-compatible). - resolved_api_base = ( - raw_api_base - or litellm.api_base - or "https://api.openai.com" - ) + resolved_api_base = raw_api_base or litellm.api_base or "https://api.openai.com" resolved_api_key = ( raw_api_key or litellm.api_key @@ -115,12 +111,21 @@ async def acreate_realtime_client_secret( litellm_logging_obj: LiteLLMLogging = kwargs.get("litellm_logging_obj") # type: ignore litellm_params = GenericLiteLLMParams(**kwargs) - model_name, custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider( + ( + model_name, + custom_llm_provider, + dynamic_api_key, + dynamic_api_base, + ) = get_llm_provider( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, ) - provider_config, resolved_api_base, resolved_api_key = _get_realtime_http_provider_config( + ( + provider_config, + resolved_api_base, + resolved_api_key, + ) = _get_realtime_http_provider_config( custom_llm_provider=custom_llm_provider, dynamic_api_base=dynamic_api_base, dynamic_api_key=dynamic_api_key, @@ -160,7 +165,12 @@ async def arealtime_calls( litellm_logging_obj: LiteLLMLogging = kwargs.get("litellm_logging_obj") # type: ignore litellm_params = GenericLiteLLMParams(**kwargs) - model_name, custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider( + ( + model_name, + custom_llm_provider, + dynamic_api_key, + dynamic_api_base, + ) = get_llm_provider( model=model_name, api_base=litellm_params.api_base, api_key=litellm_params.api_key, diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 8e5dd2bd067..0310d758956 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -859,8 +859,10 @@ class LiteLLMCompletionResponsesConfig: str(tool_call_id_raw) if tool_call_id_raw is not None else "" ) - prev_assistant_idx = LiteLLMCompletionResponsesConfig._find_previous_assistant_idx( - fixed_messages, i + prev_assistant_idx = ( + LiteLLMCompletionResponsesConfig._find_previous_assistant_idx( + fixed_messages, i + ) ) # Try to recover empty tool_call_id from previous assistant message diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 957905935be..de8e7074234 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1339,7 +1339,9 @@ class Choices(SafeAttributeModel, OpenAIObject): mapped = map_finish_reason(finish_reason) params["finish_reason"] = mapped if finish_reason != mapped: - provider_specific_fields = dict(provider_specific_fields) if provider_specific_fields else {} + provider_specific_fields = ( + dict(provider_specific_fields) if provider_specific_fields else {} + ) provider_specific_fields["native_finish_reason"] = finish_reason else: params["finish_reason"] = "stop" diff --git a/litellm/utils.py b/litellm/utils.py index 67c7fb3e82a..860e33f0472 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8350,7 +8350,9 @@ class ProviderConfigManager: from litellm.llms.openai_like.json_loader import JSONProviderRegistry # Resolve provider string for JSON lookup - provider_str = provider.value if isinstance(provider, LlmProviders) else str(provider) + provider_str = ( + provider.value if isinstance(provider, LlmProviders) else str(provider) + ) # Try to convert to enum for Python class lookup first. # Python classes take priority over JSON (they have custom overrides). @@ -8371,7 +8373,9 @@ class ProviderConfigManager: return result # Fall back to JSON providers (generic OpenAI-compatible) - if JSONProviderRegistry.exists(provider_str) and JSONProviderRegistry.supports_responses_api(provider_str): + if JSONProviderRegistry.exists( + provider_str + ) and JSONProviderRegistry.supports_responses_api(provider_str): provider_config = JSONProviderRegistry.get(provider_str) if provider_config is not None: return create_responses_config_class(provider_config)()