mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
e01d722803
commit
0fcc36c301
52 changed files with 669 additions and 420 deletions
|
|
@ -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 *
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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": (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue