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.
This commit is contained in:
Chesars 2026-03-12 14:23:50 -03:00
parent e01d722803
commit 0fcc36c301
52 changed files with 669 additions and 420 deletions

View file

@ -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 *

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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"

View file

@ -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
):

View file

@ -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)

View file

@ -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")

View file

@ -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}"

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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}",

View file

@ -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:

View file

@ -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}",

View file

@ -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)

View file

@ -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)

View file

@ -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": (

View file

@ -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

View file

@ -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]

View file

@ -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,

View file

@ -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

View file

@ -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"

View file

@ -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)

View file

@ -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:

View file

@ -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]:

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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}

View file

@ -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
)

View file

@ -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:

View file

@ -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:

View file

@ -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"})

View file

@ -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:

View file

@ -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/<deployment-name>
- azure_ai/model-router
- model_router/<deployment-name>
- 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)

View file

@ -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),
):
"""

View file

@ -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

View file

@ -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)
)

View file

@ -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)

View file

@ -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

View file

@ -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},
)
)

View file

@ -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

View file

@ -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"

View file

@ -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 = (

View file

@ -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,

View file

@ -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.

View file

@ -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,

View file

@ -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,

View file

@ -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

View file

@ -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"

View file

@ -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)()