refactor(proxy): expose public names for private proxy helpers (#45170)

* refactor(proxy): expose public names for private proxy helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep original class names behind public aliases

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): preserve internal callback filtering

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep _PROXY_ class names for managed files hooks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep old private names in package exports

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): preserve recursive auth helper name

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): align MCP limiter tests with server enforcement

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): match main's MCP limiter tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep old private names bound in importing modules

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): preserve compatibility imports through strict lint

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): use exact pyright suppression in password helper test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style(proxy): add reasons to compatibility import noqa comments

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): update IN-list baseline for renamed helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Mateo Wang 2026-10-07 18:59:27 -07:00 • committed by GitHub
parent b970e412d9
commit 58258409c9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
377 changed files with 8886 additions and 4979 deletions

View file

@ -255,7 +255,7 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s
_MANAGED_FILES_TARGET: Final = "managed_files"
class PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
# Class variables or attributes
def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient):
self.internal_usage_cache = internal_usage_cache
@ -2110,4 +2110,4 @@ class PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
verbose_logger.debug(
f"Converted file {file_id} from storage backend to base64 with format {content_type}"
)
_PROXY_LiteLLMManagedFiles = PROXY_LiteLLMManagedFiles
PROXY_LiteLLMManagedFiles = _PROXY_LiteLLMManagedFiles

View file

@ -38,7 +38,7 @@ else:
PrismaClient = Any
class PROXY_LiteLLMManagedVectorStores(
class _PROXY_LiteLLMManagedVectorStores(
CustomLogger, BaseManagedResource[VectorStoreCreateResponse]
):
"""
@ -462,4 +462,4 @@ class PROXY_LiteLLMManagedVectorStores(
parent_otel_span=parent_otel_span,
resource_id_key="vector_store_id",
)
_PROXY_LiteLLMManagedVectorStores = PROXY_LiteLLMManagedVectorStores
PROXY_LiteLLMManagedVectorStores = _PROXY_LiteLLMManagedVectorStores

View file

@ -416,11 +416,11 @@ def _provider_output_file_id(output_file_id: str) -> str:
llm_output_file_id, model-encoded ids decode to the raw provider id, raw ids pass through.
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
get_original_file_id,
is_base64_encoded_unified_file_id,
)
unified_file_id: Final = _is_base64_encoded_unified_file_id(output_file_id)
unified_file_id: Final = is_base64_encoded_unified_file_id(output_file_id)
if not unified_file_id:
return get_original_file_id(output_file_id)
try:

View file

@ -90,7 +90,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
end_time: Final = datetime.now()
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
PassThroughStreamingHandler.route_streaming_logging_to_handler(
litellm_logging_obj=self.litellm_logging_obj,
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
url_route="/v1/generateContent",

View file

@ -1845,7 +1845,7 @@ Model Info:
try:
from litellm.proxy.spend_tracking.spend_management_endpoints import (
_get_spend_report_for_time_range,
get_spend_report_for_time_range,
)
# Parse the time range
@ -1862,7 +1862,7 @@ Model Info:
if await self.internal_usage_cache.async_get_cache(key=_event_cache_key):
return
_resp: Final = await _get_spend_report_for_time_range(
_resp: Final = await get_spend_report_for_time_range(
start_date=start_date.strftime("%Y-%m-%d"),
end_date=todays_date.strftime("%Y-%m-%d"),
)
@ -1909,7 +1909,7 @@ Model Info:
from calendar import monthrange
from litellm.proxy.spend_tracking.spend_management_endpoints import (
_get_spend_report_for_time_range,
get_spend_report_for_time_range,
)
todays_date: Final = datetime.datetime.now().date()
@ -1921,7 +1921,7 @@ Model Info:
if await self.internal_usage_cache.async_get_cache(key=_event_cache_key):
return
_resp: Final = await _get_spend_report_for_time_range(
_resp: Final = await get_spend_report_for_time_range(
start_date=first_day_of_month.strftime("%Y-%m-%d"),
end_date=last_day_of_month.strftime("%Y-%m-%d"),
)

View file

@ -44,9 +44,9 @@ class GcsPubSubLogger(CustomBatchLogger):
topic_id (str): Pub/Sub topic ID
credentials_path (str, optional): Path to Google Cloud credentials JSON file
"""
from litellm.proxy.utils import _premium_user_check
from litellm.proxy.utils import premium_user_check
_premium_user_check()
premium_user_check()
self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback)
@ -107,9 +107,9 @@ class GcsPubSubLogger(CustomBatchLogger):
from litellm.proxy.spend_tracking.spend_tracking_utils import (
get_logging_payload,
)
from litellm.proxy.utils import _premium_user_check
from litellm.proxy.utils import premium_user_check
_premium_user_check()
premium_user_check()
try:
verbose_logger.debug("PubSub: Logging - Enters logging function for model %s", kwargs)

View file

@ -3833,7 +3833,7 @@ class PrometheusLogger(CustomLogger):
"""
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.proxy.management_endpoints.key_management_endpoints import (
_list_key_helper,
list_key_helper,
)
from litellm.proxy.proxy_server import prisma_client
@ -3847,7 +3847,7 @@ class PrometheusLogger(CustomLogger):
list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken],
int | None,
]:
key_list_response: Final = await _list_key_helper(
key_list_response: Final = await list_key_helper(
prisma_client=prisma_client,
page=page,
size=page_size,

View file

@ -633,9 +633,9 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool:
from litellm.exceptions import BudgetExceededError
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
_team_max_budget_check,
_virtual_key_max_budget_check,
get_team_object,
team_max_budget_check,
virtual_key_max_budget_check,
)
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
except ImportError:
@ -645,7 +645,7 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool:
if not isinstance(auth, UserAPIKeyAuth):
return False
try:
await _virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj)
await virtual_key_max_budget_check(valid_token=auth, proxy_logging_obj=proxy_logging_obj)
if auth.team_id:
team: Final = await get_team_object(
team_id=auth.team_id,
@ -653,7 +653,7 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool:
user_api_key_cache=user_api_key_cache,
check_cache_only=True,
)
await _team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj)
await team_max_budget_check(team_object=team, valid_token=auth, proxy_logging_obj=proxy_logging_obj)
except BudgetExceededError:
return True
except Exception as e: # noqa: BLE001 # advisory gate: a failed read must not block sampling

View file

@ -53,8 +53,14 @@ from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook i
VectorStorePreCallHook,
)
from litellm.integrations.zerobus import ZerobusLogger
from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import _PROXY_DynamicRateLimitHandlerV3
from litellm.proxy.hooks.dynamic_rate_limiter import ( # noqa: F401 # legacy module exports
PROXY_DynamicRateLimitHandler,
_PROXY_DynamicRateLimitHandler, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
)
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( # noqa: F401 # legacy module exports
PROXY_DynamicRateLimitHandlerV3,
_PROXY_DynamicRateLimitHandlerV3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
)
class CustomLoggerRegistry:
@ -101,8 +107,8 @@ class CustomLoggerRegistry:
"pointfive": PointFiveLogger,
"zerobus": ZerobusLogger,
"aws_sqs": SQSLogger,
"dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler,
"dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3,
"dynamic_rate_limiter": PROXY_DynamicRateLimitHandler,
"dynamic_rate_limiter_v3": PROXY_DynamicRateLimitHandlerV3,
"vector_store_pre_call_hook": VectorStorePreCallHook,
"dotprompt": DotpromptManager,
"bitbucket": BitBucketPromptManager,

View file

@ -4969,17 +4969,17 @@ def _init_custom_logger_compatible_class(
return _otel_logger
elif logging_integration == "dynamic_rate_limiter":
from litellm.proxy.hooks.dynamic_rate_limiter import (
_PROXY_DynamicRateLimitHandler,
PROXY_DynamicRateLimitHandler,
)
for callback in _in_memory_loggers:
if isinstance(callback, _PROXY_DynamicRateLimitHandler):
if isinstance(callback, PROXY_DynamicRateLimitHandler):
return callback
if internal_usage_cache is None:
raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}")
dynamic_rate_limiter_obj: Final = _PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache)
dynamic_rate_limiter_obj: Final = PROXY_DynamicRateLimitHandler(internal_usage_cache=internal_usage_cache)
if llm_router is not None and isinstance(llm_router, litellm.Router):
dynamic_rate_limiter_obj.update_variables(llm_router=llm_router)
@ -4987,17 +4987,19 @@ def _init_custom_logger_compatible_class(
return dynamic_rate_limiter_obj
elif logging_integration == "dynamic_rate_limiter_v3":
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
_PROXY_DynamicRateLimitHandlerV3,
PROXY_DynamicRateLimitHandlerV3,
)
for callback in _in_memory_loggers:
if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3):
if isinstance(callback, PROXY_DynamicRateLimitHandlerV3):
return callback
if internal_usage_cache is None:
raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}")
dynamic_rate_limiter_obj_v3 = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=internal_usage_cache)
dynamic_rate_limiter_obj_v3: Final = PROXY_DynamicRateLimitHandlerV3(
internal_usage_cache=internal_usage_cache
)
if llm_router is not None and isinstance(llm_router, litellm.Router):
dynamic_rate_limiter_obj_v3.update_variables(llm_router=llm_router)
@ -5546,19 +5548,19 @@ def get_custom_logger_compatible_class(
elif logging_integration == "dynamic_rate_limiter":
from litellm.proxy.hooks.dynamic_rate_limiter import (
_PROXY_DynamicRateLimitHandler,
PROXY_DynamicRateLimitHandler,
)
for callback in _in_memory_loggers:
if isinstance(callback, _PROXY_DynamicRateLimitHandler):
if isinstance(callback, PROXY_DynamicRateLimitHandler):
return callback
elif logging_integration == "dynamic_rate_limiter_v3":
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
_PROXY_DynamicRateLimitHandlerV3,
PROXY_DynamicRateLimitHandlerV3,
)
for callback in _in_memory_loggers:
if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3):
if isinstance(callback, PROXY_DynamicRateLimitHandlerV3):
return callback
elif logging_integration == "langtrace":

View file

@ -565,9 +565,9 @@ def update_messages_with_model_file_ids(
}
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
convert_b64_uid_to_unified_uid,
get_original_file_id,
is_base64_encoded_unified_file_id,
is_model_embedded_id,
)
@ -603,7 +603,7 @@ def update_messages_with_model_file_ids(
if model_file_id_mapping and model_id is not None
else None
)
if not provider_file_id and _is_base64_encoded_unified_file_id(file_id):
if not provider_file_id and is_base64_encoded_unified_file_id(file_id):
unified_file_id = convert_b64_uid_to_unified_uid(file_id)
if "llm_output_file_id," in unified_file_id:
provider_file_id = unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
@ -634,9 +634,9 @@ def update_responses_input_with_model_file_ids(
Format: {"litellm_file_id": {"model_id": "provider_file_id"}}
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
convert_b64_uid_to_unified_uid,
get_original_file_id,
is_base64_encoded_unified_file_id,
is_model_embedded_id,
)
@ -671,7 +671,7 @@ def update_responses_input_with_model_file_ids(
updated_content.append(updated_content_item)
else:
# Check if this is a base64-encoded unified file ID without mapping
is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
is_unified_file_id = is_base64_encoded_unified_file_id(file_id)
if is_unified_file_id:
# Fallback: decode unified file ID
unified_file_id = convert_b64_uid_to_unified_uid(file_id)

View file

@ -322,7 +322,7 @@ class AnthropicMessagesHandler(BaseTranslation):
if not chunks:
return None
try:
return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
return AnthropicPassthroughLoggingHandler.build_usage_only_response_from_chunks(
all_chunks=chunks,
model=str((request_data or {}).get("model") or ""),
)
@ -1226,7 +1226,7 @@ class AnthropicMessagesHandler(BaseTranslation):
has_ended: Final = self._check_streaming_has_ended(responses_so_far)
if has_ended:
# build the model response from the responses_so_far
built_response: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
built_response: Final = AnthropicPassthroughLoggingHandler.build_complete_streaming_response(
all_chunks=responses_so_far,
litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj),
model="",

View file

@ -228,7 +228,7 @@ async def _check_summary_model_access(
try:
from litellm.proxy._types import ProxyException
from litellm.proxy.auth.auth_checks import (
_can_object_call_model,
can_object_call_model,
can_project_access_model,
can_user_call_model,
get_project_object,
@ -258,7 +258,7 @@ async def _check_summary_model_access(
if not models:
continue
try:
_can_object_call_model(
can_object_call_model(
model=summary_model,
llm_router=llm_router,
models=models,
@ -370,7 +370,7 @@ async def _check_summary_model_access(
)
if member_allowed_models:
try:
_can_object_call_model(
can_object_call_model(
model=summary_model,
llm_router=llm_router,
models=list(member_allowed_models),
@ -558,7 +558,7 @@ async def _check_summary_model_rate_limit(
limiter: Final[object] = get_proxy_hook("parallel_request_limiter") if get_proxy_hook is not None else None
should_rate_limit_check: Final[_ShouldRateLimit | None] = getattr(limiter, "should_rate_limit", None)
create_descriptors: Final[_CreateRateLimitDescriptors | None] = getattr(
limiter, "_create_rate_limit_descriptors", None
limiter, "create_rate_limit_descriptors", None
)
add_team_descriptor: Final[_AddModelRateLimitDescriptor | None] = getattr(
limiter, "_add_team_model_rate_limit_descriptor_from_metadata", None

View file

@ -421,7 +421,7 @@ class BaseAnthropicMessagesStreamingIterator:
if self.completion_start_time is not None:
self.litellm_logging_obj.completion_start_time = self.completion_start_time
self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time
logging_coroutine: Final = PassThroughStreamingHandler._route_streaming_logging_to_handler(
logging_coroutine: Final = PassThroughStreamingHandler.route_streaming_logging_to_handler(
litellm_logging_obj=self.litellm_logging_obj,
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
url_route="/v1/messages",

View file

@ -56,12 +56,16 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import (
from litellm.proxy.auth.user_api_key_auth import ( # noqa: F401 # legacy module exports
_get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth
_run_centralized_common_checks,
_run_centralized_common_checks, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
run_centralized_common_checks,
user_api_key_auth,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
from litellm.proxy.common_utils.user_api_key_cache import (
AUTH_OBJECTS_TARGET,
USER_NO_MCP_PERMISSION_SENTINEL,
@ -225,7 +229,7 @@ def _agent_capped_servers(
)
def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool:
def is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool:
"""True when this auth is a keyless subject admitted by the gateway session / bridge user
path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``.
@ -235,6 +239,9 @@ def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> b
return user_api_key_auth is not None and user_api_key_auth.mcp_admitted_user_subject is True
_is_mcp_admitted_user_subject: Final = is_mcp_admitted_user_subject
def _gateway_dcr_challenge_target(
route: str,
mcp_servers: list[str] | None,
@ -452,7 +459,7 @@ class MCPRequestHandler:
HTTPException: If headers are invalid or missing required headers
"""
async with global_manager().catalog.operation():
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope)
# Check if there is an explicit LiteLLM API key (primary header)
has_explicit_litellm_key: Final = (
@ -462,13 +469,19 @@ class MCPRequestHandler:
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
# Get the old mcp_auth_header for backward compatibility
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
mcp_auth_header = MCPRequestHandler.get_mcp_auth_header_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line
headers
)
# Get the new server-specific auth headers
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
mcp_server_auth_headers = MCPRequestHandler.get_mcp_server_auth_headers_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line
headers
)
# Get the oauth2 headers
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
oauth2_headers = MCPRequestHandler.get_oauth2_headers_from_headers( # rebind-ok: pre-existing rebinding on a rename-only line
headers
)
# Parse MCP servers from header
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
@ -621,7 +634,7 @@ class MCPRequestHandler:
mcp_auth_header,
mcp_server_auth_headers,
) = MCPRequestHandler._scrub_gateway_admission_credentials(
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
admitted=is_mcp_admitted_user_subject(validated_user_api_key_auth),
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
@ -1060,7 +1073,7 @@ class MCPRequestHandler:
await pre_db_read_auth_checks(
request=request,
request_data=await _read_request_body(request=request),
request_data=await read_request_body(request=request),
route=route,
)
@ -1351,10 +1364,10 @@ class MCPRequestHandler:
admitted.budget_reservation = None
try:
RouteChecks.should_call_route(route=route, valid_token=admitted, request=request)
await _run_centralized_common_checks(
await run_centralized_common_checks(
user_api_key_auth_obj=admitted,
request=request,
request_data=await _read_request_body(request=request),
request_data=await read_request_body(request=request),
route=route,
)
except (HTTPException, ProxyException):
@ -1401,7 +1414,7 @@ class MCPRequestHandler:
return mcp_servers_header if mcp_servers_header is not None else []
@staticmethod
def _get_mcp_auth_header_from_headers(headers: Headers) -> str | None:
def get_mcp_auth_header_from_headers(headers: Headers) -> str | None:
"""
Get the header passed to LiteLLM to pass to downstream MCP servers
@ -1424,8 +1437,10 @@ class MCPRequestHandler:
)
return auth_header
_get_mcp_auth_header_from_headers = get_mcp_auth_header_from_headers
@staticmethod
def _get_mcp_server_auth_headers_from_headers(
def get_mcp_server_auth_headers_from_headers(
headers: Headers,
) -> dict[str, dict[str, str]]:
"""
@ -1478,8 +1493,10 @@ class MCPRequestHandler:
return server_auth_headers
_get_mcp_server_auth_headers_from_headers = get_mcp_server_auth_headers_from_headers
@staticmethod
def _get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]:
def get_oauth2_headers_from_headers(headers: Headers) -> dict[str, str]:
"""
Get the oauth2 headers from the request headers.
"""
@ -1489,6 +1506,8 @@ class MCPRequestHandler:
oauth2_headers["Authorization"] = header_value
return oauth2_headers
_get_oauth2_headers_from_headers = get_oauth2_headers_from_headers
@staticmethod
def get_mcp_client_side_auth_header_name() -> str:
"""
@ -1535,7 +1554,7 @@ class MCPRequestHandler:
return None
@staticmethod
def _safe_get_headers_from_scope(scope: Scope) -> Headers:
def safe_get_headers_from_scope(scope: Scope) -> Headers:
"""
Safely extract headers from ASGI scope using Starlette's Headers class
which handles case insensitivity and proper header parsing.
@ -1563,6 +1582,8 @@ class MCPRequestHandler:
# Return empty Headers object with empty dict
return Headers({})
_safe_get_headers_from_scope = safe_get_headers_from_scope
@staticmethod
def _reject_duplicate_authorization(raw_headers: object) -> None:
"""Raise 400 when the raw ASGI headers carry more than one ``Authorization`` header."""
@ -1638,7 +1659,7 @@ class MCPRequestHandler:
# matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above
# this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is
# independent; an opt-out silences only its own source, inside the recursive call).
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
if is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
return MCPServerAccess(
server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
)
@ -1950,18 +1971,18 @@ class MCPRequestHandler:
# already-exceeded state; ATTRIBUTION of new spend stays with the user (documented deferral).
from litellm.exceptions import BudgetExceededError
from litellm.proxy.auth.auth_checks import (
_organization_max_budget_check,
_team_max_budget_check,
organization_max_budget_check,
team_max_budget_check,
)
source_view: Final = MCPRequestHandler._scoped_source_auth(
auth, team_id=team_id, org_id=team_obj.organization_id or auth.org_id, carry_user_grants=False
)
try:
await _team_max_budget_check(
await team_max_budget_check(
team_object=team_obj, valid_token=source_view, proxy_logging_obj=proxy_logging_obj
)
await _organization_max_budget_check(
await organization_max_budget_check(
valid_token=source_view,
team_object=team_obj,
prisma_client=prisma_client,
@ -2042,14 +2063,14 @@ class MCPRequestHandler:
owning ``org_id`` so the team's budget accumulates and the right org is charged. Falls back to
user-level attribution (rather than guessing a team) when the tool name does not resolve to a
server, reusing the manager's own tool-name lookup."""
if not _is_mcp_admitted_user_subject(auth):
if not is_mcp_admitted_user_subject(auth):
return auth
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name)
server: Final = global_mcp_server_manager.get_mcp_server_from_tool_name(tool_name)
if server is None:
return auth
source: Final = await MCPRequestHandler.attributing_source_for_server(auth, server.server_id)
@ -2330,7 +2351,7 @@ class MCPRequestHandler:
# FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per
# source and shares nothing with the single-credential prelude below. Ordering is the invariant:
# sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant.
if _is_mcp_admitted_user_subject(user_api_key_auth):
if is_mcp_admitted_user_subject(user_api_key_auth):
return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth)
# Get key and team object permissions (already loaded in main auth flow)
@ -2422,7 +2443,9 @@ class MCPRequestHandler:
# not collapse to allow-all (None); key/JWT auth keeps its prior allow-all-on-error. Both
# keyless_source AND the marker are needed: each source resolves through an UNMARKED auth, so
# without keyless_source a fault under a source returns None and wins the union as allow-all.
deny_all = unreadable_entitlement or keyless_source or _is_mcp_admitted_user_subject(user_api_key_auth)
deny_all: Final = (
unreadable_entitlement or keyless_source or is_mcp_admitted_user_subject(user_api_key_auth)
)
return [] if deny_all else None
@staticmethod
@ -2557,7 +2580,7 @@ class MCPRequestHandler:
global_mcp_server_manager,
)
from litellm.proxy.auth.auth_checks import (
_get_mcp_server_ids_from_access_groups,
get_mcp_server_ids_from_access_groups,
)
from litellm.proxy.proxy_server import (
prisma_client,
@ -2565,7 +2588,7 @@ class MCPRequestHandler:
user_api_key_cache,
)
raw_server_ids: Final = await _get_mcp_server_ids_from_access_groups(
raw_server_ids: Final = await get_mcp_server_ids_from_access_groups(
access_group_ids=user_api_key_auth.access_group_ids or [],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
@ -2647,7 +2670,7 @@ class MCPRequestHandler:
)
# Get MCP servers from access groups
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(
key_object_permission.mcp_access_groups or [],
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
@ -2763,7 +2786,7 @@ class MCPRequestHandler:
return set(team_access_group_servers)
if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []):
return set(global_mcp_server_manager.get_registry().keys())
legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
legacy_access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or [],
requires_fresh_policy=requires_fresh_policy,
)
@ -2797,7 +2820,7 @@ class MCPRequestHandler:
"""
try:
from litellm.proxy.auth.auth_checks import (
_get_mcp_server_ids_from_access_groups,
get_mcp_server_ids_from_access_groups,
get_team_object,
)
from litellm.proxy.proxy_server import (
@ -2825,7 +2848,7 @@ class MCPRequestHandler:
# pinned to a single team_id, but a keyless admitted identity (no team_id) unions
# across all of its teams and would otherwise inherit a blocked team's MCP grants.
return []
team_access_group_servers: Final = await _get_mcp_server_ids_from_access_groups(
team_access_group_servers: Final = await get_mcp_server_ids_from_access_groups(
access_group_ids=team_obj.access_group_ids or [],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
@ -2968,7 +2991,7 @@ class MCPRequestHandler:
# Expand names/aliases to canonical server IDs (consistent with key/team/end-user path)
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or [],
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
@ -3073,7 +3096,7 @@ class MCPRequestHandler:
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permission.mcp_servers or [])
# Get MCP servers from access groups
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(
object_permission.mcp_access_groups or [],
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
)
@ -3212,7 +3235,7 @@ class MCPRequestHandler:
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or [],
requires_fresh_policy=fresh,
)
@ -3318,7 +3341,7 @@ class MCPRequestHandler:
return False
object_permission: Final = user_api_key_auth.object_permission
credential_scoped: Final = (
not _is_mcp_admitted_user_subject(user_api_key_auth)
not is_mcp_admitted_user_subject(user_api_key_auth)
and object_permission is not None
and object_permission.mcp_servers is not None
)
@ -3566,7 +3589,7 @@ class MCPRequestHandler:
expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list(
obj_perm.mcp_servers or []
)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(
obj_perm.mcp_access_groups or [],
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
)
@ -3696,7 +3719,7 @@ class MCPRequestHandler:
return server_ids
@staticmethod
async def _get_mcp_servers_from_access_groups(
async def get_mcp_servers_from_access_groups(
access_groups: list[str],
*,
requires_fresh_policy: bool = False,
@ -3735,6 +3758,8 @@ class MCPRequestHandler:
verbose_logger.warning("Failed to get MCP servers from access groups: %s", e)
return []
_get_mcp_servers_from_access_groups = get_mcp_servers_from_access_groups
@staticmethod
async def get_mcp_access_groups(
user_api_key_auth: UserAPIKeyAuth | None = None,
@ -3863,5 +3888,5 @@ class MCPRequestHandler:
"""
Extract and parse the x-mcp-access-groups header from an ASGI scope.
"""
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope)
return MCPRequestHandler.get_mcp_access_groups_from_headers(headers)

View file

@ -14,8 +14,9 @@ from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
_V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports
_V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
V2_GCM_PREFIX, # pyright: ignore[reportPrivateUsage] # reuse the encrypted credential's format discriminator
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -99,7 +100,7 @@ async def _opaque_bearer_is_gateway_credential(token: str) -> bool:
user_api_key_cache,
)
if is_envelope(token) or is_refresh_envelope(token) or token.startswith(_V2_GCM_PREFIX):
if is_envelope(token) or is_refresh_envelope(token) or token.startswith(V2_GCM_PREFIX):
return True
try:
if ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token) is not None:
@ -283,12 +284,15 @@ async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResol
return _ResolvedKey(key_hash=key_hash, key=key_obj)
async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None":
async def reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None":
"""``None`` when the user is live, else the precise failure ``load_active_user_by_id`` found."""
loaded: Final = await load_active_user_by_id(user_id)
return loaded if isinstance(loaded, str) else None
_reload_active_user_by_id: Final = reload_active_user_by_id
UserRowSource = Literal["cache", "database"]
@ -400,12 +404,12 @@ async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResol
return "no_active_key"
return None
case "user_id":
return await _reload_active_user_by_id(identity.subject)
return await reload_active_user_by_id(identity.subject)
case _:
assert_never(identity.subject_type)
async def _extract_user_id_from_request(request: Request) -> str | None:
async def extract_user_id_from_request(request: Request) -> str | None:
"""Resolve the caller for identity binding without granting credential-write permission."""
from litellm.proxy.auth.handle_jwt import JWTIdentity # noqa: PLC0415 # proxy import cycle
@ -415,6 +419,9 @@ async def _extract_user_id_from_request(request: Request) -> str | None:
return _active_key_user_id(resolved) if resolved is not None else None
_extract_user_id_from_request: Final = extract_user_id_from_request
async def authorize_oauth_credential_request(request: Request, server_id: str) -> str | None:
from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle
@ -448,13 +455,13 @@ async def can_store_oauth_credential(request: Request, auth: "UserAPIKeyAuth", s
)
from litellm.proxy.auth.route_checks import RouteChecks # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.user_api_key_auth import ( # noqa: PLC0415 # proxy import cycle
_run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action
run_centralized_common_checks, # pyright: ignore[reportPrivateUsage] # reuse admission policy for the credential-write action
)
write_route: Final = f"/v1/mcp/server/{server_id}/oauth-user-credential"
try:
RouteChecks.is_virtual_key_allowed_to_call_route(route=write_route, valid_token=auth, request=request)
await _run_centralized_common_checks(
await run_centralized_common_checks(
user_api_key_auth_obj=auth,
request=request,
request_data={},
@ -646,7 +653,7 @@ class _BridgeMintReady:
keys: "EnvelopeKeys"
def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
def bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
"""Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape
(top-level ``error``, no-store headers) for every case, with a status truthful about where the
failure is. The caller's request is 400, a transient gateway outage is 503, a gateway
@ -731,6 +738,9 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
)
_bridge_mint_error_response: Final = bridge_mint_error_response
def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
"""Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays
truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that
@ -759,7 +769,7 @@ def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _Br
assert_never(rejection)
async def _prepare_bridge_mint(
async def prepare_bridge_mint(
request: Request,
mcp_server: MCPServer,
bridge_identity: "_BridgeAuthorizationCode | None" = None,
@ -820,6 +830,9 @@ async def _prepare_bridge_mint(
return _BridgeMintReady(identity=identity, keys=keys)
_prepare_bridge_mint: Final = prepare_bridge_mint
@dataclass(frozen=True, slots=True)
class _BridgeRefreshReady:
"""A validated refresh request: the identity+keys to mint the renewed pair under, the upstream refresh
@ -854,7 +867,7 @@ def _refresh_key_failure_to_mint_error(failure: _KeyResolutionFailure) -> _Bridg
assert_never(failure)
async def _prepare_bridge_refresh(
async def prepare_bridge_refresh(
mcp_server: MCPServer, refresh_value: str | None
) -> "_BridgeRefreshReady | _BridgeMintError":
"""Phase 1 for the refresh_token grant, BEFORE the upstream exchange: open the client's refresh
@ -891,7 +904,10 @@ async def _prepare_bridge_refresh(
)
def _finish_bridge_mint(
_prepare_bridge_refresh: Final = prepare_bridge_refresh
def finish_bridge_mint(
ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime
) -> "JSONResponse | _BridgeMintError":
"""Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held access envelope
@ -933,6 +949,9 @@ def _finish_bridge_mint(
return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
_finish_bridge_mint: Final = finish_bridge_mint
def _upstream_refresh_credential(token_response: object) -> "RefreshCredential | None":
"""Extract the upstream refresh grant from a token response, or ``None`` when there is none to seal.
Each field is isinstance-checked so nothing untyped reaches the refresh envelope; ``refresh_expires_in``

View file

@ -83,12 +83,15 @@ def _oauth_token_error(code: str, status: int = 400) -> JSONResponse:
return JSONResponse(status_code=status, content={"error": code}, headers=TOKEN_NO_CACHE_HEADERS)
def _user_id_from_session_cookie(request: Request) -> str | None:
def user_id_from_session_cookie(request: Request) -> str | None:
"""Return user_id from the UI ``token`` cookie, or None if missing/invalid."""
user_id, _ = _session_identity_from_cookie(request)
return user_id
_user_id_from_session_cookie: Final = user_id_from_session_cookie
def _session_identity_from_cookie(request: Request) -> tuple[str | None, str | None]:
"""Return ``(user_id, session_key)`` from the UI ``token`` cookie
(HS256-signed with ``master_key``), or ``(None, None)`` if missing/invalid.

View file

@ -844,7 +844,7 @@ async def get_filtered_server_tools(
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
if params is None:
page = ListToolsResult(
tools=await global_mcp_server_manager._get_tools_from_server(
tools=await global_mcp_server_manager.get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,

View file

@ -33,13 +33,14 @@ from litellm.proxy._types import (
NewMCPServerRequest,
UpdateMCPServerRequest,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: F401 # legacy module exports
SecretMapDecodeError,
_get_salt_key,
_get_salt_key, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
decode_secret_map,
decrypt_value_helper,
encrypt_secret_map,
encrypt_value_helper,
get_salt_key,
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.config_repository import ConfigRepository
@ -464,7 +465,7 @@ def _prepare_mcp_server_data(
blob_value = credentials.pop(te_field, None)
if blob_value is not None and te_field not in data_dict:
data_dict[te_field] = blob_value
data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key())
data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=get_salt_key())
data_dict["credentials"] = safe_dumps(
_bind_submitted_oauth_client(data_dict["credentials"], data.issuer, data.url)
if not exclude_unset and data.auth_type == "oauth2"
@ -1474,7 +1475,7 @@ async def upsert_mcp_server_oauth_client_credentials(
same way regardless of which store a server's client came from."""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=_get_salt_key())
encrypted: Final = encrypt_credentials(credentials=MCPCredentials(**credentials), encryption_key=get_salt_key())
blob: Final = safe_dumps(encrypted)
await _oauth_client_table_actions(prisma_client).upsert(
where={"server_id": server_id},
@ -1613,11 +1614,14 @@ def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None:
return None
def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
def decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
"""Return the OAuth2 payload dict held in ``stored``, else ``None``."""
return _parse_oauth_payload(_decode_user_credential(stored))
_decode_oauth_payload: Final = decode_oauth_payload
async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str):
"""Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``.
@ -1865,7 +1869,7 @@ async def get_user_oauth_credential(
def _server_user_credential_item(
row: "prisma_db_models.LiteLLM_MCPUserCredentials",
) -> MCPServerUserCredentialListItem:
oauth_payload: Final = _decode_oauth_payload(row.credential_b64)
oauth_payload: Final = decode_oauth_payload(row.credential_b64)
if oauth_payload is None:
return MCPServerUserCredentialListItem(
user_id=row.user_id,
@ -1986,7 +1990,7 @@ async def purge_user_oauth_credentials_for_server(
invalidate_token_cache is injectable for tests; it defaults to the manager's shared
invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
rows: Final = await _db_find_user_credential_rows(prisma_client, {"server_id": server_id})
oauth_rows: Final = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
oauth_rows: Final = [row for row in rows if decode_oauth_payload(row.credential_b64) is not None]
if not oauth_rows:
return 0
deleted_count: Final = await _user_credential_actions(prisma_client).delete_many(
@ -2201,7 +2205,7 @@ async def resolve_user_oauth_access_token(
return None
try:
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
_compute_per_user_token_ttl,
compute_per_user_token_ttl,
mcp_per_user_token_cache,
)
@ -2248,7 +2252,7 @@ async def resolve_user_oauth_access_token(
access_token: Final[str] = cred["access_token"]
if prefetched_creds is None:
ttl: Final = _compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at")))
ttl: Final = compute_per_user_token_ttl(server, _remaining_token_seconds(cred.get("expires_at")))
await mcp_per_user_token_cache.set(
user_id, server_id, access_token, ttl, identity_binding_proof=cred.get("identity_binding_proof")
)

View file

@ -24,18 +24,24 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
TokenEndpointAuthConfigError,
normalize_token_endpoint_auth_method,
)
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_bridge_mint_error_response,
from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( # noqa: F401 # legacy module exports
_bridge_mint_error_response, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_BridgeMintReady,
_BridgeRefreshReady,
_extract_user_id_from_request,
_finish_bridge_mint,
_prepare_bridge_mint,
_prepare_bridge_refresh,
_reload_active_user_by_id,
_extract_user_id_from_request, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_finish_bridge_mint, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_prepare_bridge_mint, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_prepare_bridge_refresh, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_reload_active_user_by_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
authorize_oauth_credential_request,
bridge_mint_error_response,
can_store_oauth_credential,
extract_user_id_from_request,
finish_bridge_mint,
oauth_authorization_uses_gateway_credential,
prepare_bridge_mint,
prepare_bridge_refresh,
reload_active_user_by_id,
)
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._experimental.mcp_server.faults import (
@ -91,7 +97,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
from litellm.types.llms.base import LiteLLMBaseModel
from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer, MCPTokenEndpointAuthMethod
@ -431,10 +440,10 @@ def _session_cookie_user_id(request: Request) -> str | None:
aggregate DCR flow's verbs receive the identity as a plain value instead of parsing
cookies themselves."""
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load
_user_id_from_session_cookie,
user_id_from_session_cookie,
)
return _user_id_from_session_cookie(request)
return user_id_from_session_cookie(request)
def _redirect_to_litellm_login(request: Request) -> RedirectResponse:
@ -637,7 +646,7 @@ async def _store_per_user_token_server_side(
client even when server-side storage fails.
"""
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( # noqa: PLC0415
_compute_per_user_token_ttl,
compute_per_user_token_ttl,
mcp_per_user_token_cache,
)
from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415
@ -693,7 +702,7 @@ async def _store_per_user_token_server_side(
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id)
# Warm the Redis cache so the first subsequent MCP call is a cache hit
ttl: Final = _compute_per_user_token_ttl(server, expires_in)
ttl: Final = compute_per_user_token_ttl(server, expires_in)
await mcp_per_user_token_cache.set(
user_id=user_id,
server_id=server.server_id,
@ -703,7 +712,7 @@ async def _store_per_user_token_server_side(
)
def _raise_if_not_oauth2(mcp_server: MCPServer) -> None:
def raise_if_not_oauth2(mcp_server: MCPServer) -> None:
"""Reject a server without upstream OAuth from the gateway's authorize/token/register flow.
The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are allowed
@ -714,10 +723,10 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None:
Authorize path with ``persist_credentials`` enabled writes nothing to the server row).
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # circular import with mcp_server_manager at module load
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
)
if mcp_server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
if mcp_server.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return
raise HTTPException(
status_code=400,
@ -733,6 +742,9 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None:
)
_raise_if_not_oauth2: Final = raise_if_not_oauth2
def _endpoint_not_configured_detail(
mcp_server: MCPServer,
endpoint_label: str,
@ -919,7 +931,7 @@ async def _resolve_oauth_authorization_user(
) -> str | RedirectResponse:
"""Resolve the authorization subject without replacing denied credentials with cookie grants."""
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # proxy import cycle
_user_id_from_session_cookie,
user_id_from_session_cookie,
)
use_gateway_credential: Final = enforce_binding and await oauth_authorization_uses_gateway_credential(request)
@ -928,7 +940,7 @@ async def _resolve_oauth_authorization_user(
)
if use_gateway_credential and request_user_id is None:
return _bridge_access_denied_redirect(redirect_uri, state, mcp_server)
user_id: Final = request_user_id or _user_id_from_session_cookie(request)
user_id: Final = request_user_id or user_id_from_session_cookie(request)
if user_id is None:
return _redirect_to_litellm_login(request)
if not await _user_can_reach_mcp_server(user_id, mcp_server.server_id):
@ -948,7 +960,7 @@ async def authorize_with_server(
scope: str | None = None,
ephemeral_dcr_client: "EphemeralDcrClient | None" = None,
):
_raise_if_not_oauth2(mcp_server)
raise_if_not_oauth2(mcp_server)
resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint)
if not oauth_client_registration_matches(
resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url
@ -1084,7 +1096,7 @@ async def exchange_token_with_server(
scope: str | None = None,
client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
):
_raise_if_not_oauth2(mcp_server)
raise_if_not_oauth2(mcp_server)
if grant_type not in ("authorization_code", "refresh_token"):
raise HTTPException(status_code=400, detail="Unsupported grant_type")
@ -1131,7 +1143,7 @@ async def exchange_token_with_server(
raise HTTPException(status_code=400, detail=str(exc)) from exc
request_user_id: Final = (
await _extract_user_id_from_request(request)
await extract_user_id_from_request(request)
if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None
else None
)
@ -1148,9 +1160,9 @@ async def exchange_token_with_server(
# identity, and unwrap the real upstream refresh token BEFORE building token_data, so the exchange
# sends the upstream token and never the envelope. A failure returns without touching the upstream.
if is_bridge:
prepared_refresh: Final = await _prepare_bridge_refresh(resolved_server, refresh_token)
prepared_refresh: Final = await prepare_bridge_refresh(resolved_server, refresh_token)
if not isinstance(prepared_refresh, _BridgeRefreshReady):
return _bridge_mint_error_response(prepared_refresh)
return bridge_mint_error_response(prepared_refresh)
bridge_mint_ready = prepared_refresh.ready
bridge_upstream_refresh = prepared_refresh.upstream_refresh_token
bridge_upstream_scope = prepared_refresh.upstream_scope
@ -1227,9 +1239,9 @@ async def exchange_token_with_server(
# Phase 1 for a bridge authorization_code mint: resolve identity (the SSO user recovered above, or
# the presented litellm key) and the envelope keys BEFORE the exchange consumes the single-use code.
if is_bridge:
prepared: Final = await _prepare_bridge_mint(request, resolved_server, bridge_identity)
prepared: Final = await prepare_bridge_mint(request, resolved_server, bridge_identity)
if not isinstance(prepared, _BridgeMintReady):
return _bridge_mint_error_response(prepared)
return bridge_mint_error_response(prepared)
bridge_mint_ready = prepared
refresh_binding: Final = resolved_server.oauth_identity_binding
@ -1269,7 +1281,7 @@ async def exchange_token_with_server(
"re-runs authorization_code rather than an opaque upstream error",
resolved_server.server_id,
)
return _bridge_mint_error_response("invalid_refresh")
return bridge_mint_error_response("invalid_refresh")
return render_token_fault(fault)
token_response = response.json()
@ -1357,10 +1369,10 @@ async def exchange_token_with_server(
token_response = {**token_response, "scope": refresh_request_scope}
# Phase 3: seal the upstream grant into the client-held envelope; failures map through the same
# OAuth-shaped response as the phase-1 preconditions.
minted: Final = _finish_bridge_mint(
minted: Final = finish_bridge_mint(
bridge_mint_ready, resolved_server, token_response, datetime.now(timezone.utc)
)
return minted if isinstance(minted, JSONResponse) else _bridge_mint_error_response(minted)
return minted if isinstance(minted, JSONResponse) else bridge_mint_error_response(minted)
raw_access_token: Final = token_response.get("access_token") if isinstance(token_response, dict) else None
if not isinstance(raw_access_token, str) or not raw_access_token:
@ -1937,7 +1949,7 @@ async def register_client_with_server(
client_redirect_uris: list[str] | None = None,
client_application_type: Literal["native", "web"] | None = None,
):
_raise_if_not_oauth2(mcp_server)
raise_if_not_oauth2(mcp_server)
request_base_url: Final = get_request_base_url(request)
current_redirect_uri: Final = f"{request_base_url}/callback"
client_facing_redirect_uris: Final = client_redirect_uris or [current_redirect_uri]
@ -2111,7 +2123,7 @@ async def authorize(
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
raise HTTPException(status_code=404, detail="MCP server not found")
_raise_if_not_oauth2(mcp_server)
raise_if_not_oauth2(mcp_server)
# Use server's stored client_id when caller doesn't supply one.
# Raise a clear error instead of passing an empty string — an empty
# client_id would silently produce a broken authorization URL.
@ -2181,7 +2193,7 @@ async def token_endpoint(
code_verifier=code_verifier,
refresh_token=refresh_token,
master_key=master_key,
reload_user=_reload_active_user_by_id,
reload_user=reload_active_user_by_id,
cache=user_api_key_cache,
resource=resource,
mint_proxy_credential=mint_proxy_credential,
@ -2303,7 +2315,7 @@ async def introspect_endpoint(token: str = Form(...)) -> Response:
return await introspect_gateway_token(
token=token,
master_key=master_key,
reload_user=_reload_active_user_by_id,
reload_user=reload_active_user_by_id,
cache=user_api_key_cache,
)
@ -3102,7 +3114,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
# Get the correct base URL considering X-Forwarded-* headers
request_base_url: Final = get_request_base_url(request)
request_data: Final = await _read_request_body(request=request)
request_data: Final = await read_request_body(request=request)
data: Final[dict] = {**request_data}
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))

View file

@ -27,16 +27,22 @@ def get_active_mcp_request_ctx() -> "ServerRequestContext | None":
# Set server-side in proxy_server.py route handlers when a request arrives via
# /toolset/{name}/mcp or the toolset fallback in dynamic_mcp_route.
# Never populated from client-supplied headers.
_mcp_active_toolset_id: Final[ContextVar[str | None]] = ContextVar("_mcp_active_toolset_id", default=None)
mcp_active_toolset_id: Final[ContextVar[str | None]] = ContextVar("_mcp_active_toolset_id", default=None)
_mcp_active_toolset_id: Final = mcp_active_toolset_id
# Per-request merged InitializeResult.instructions; set in MCP HTTP/SSE handlers.
_mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar(
mcp_gateway_initialize_instructions: Final[ContextVar[str | None]] = ContextVar(
"_mcp_gateway_initialize_instructions", default=None
)
_mcp_gateway_initialize_instructions: Final = mcp_gateway_initialize_instructions
# Per-request scoped server name; set in MCP HTTP/SSE handlers when the path
# identifies exactly one upstream server. Never populated from client-supplied headers.
_mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None)
mcp_gateway_server_name: Final[ContextVar[str | None]] = ContextVar("_mcp_gateway_server_name", default=None)
_mcp_gateway_server_name: Final = mcp_gateway_server_name
# Set server-side by the /mcp/proxy route. Never populated from client-supplied headers.
_mcp_proxy_mode: Final[ContextVar[bool]] = ContextVar("_mcp_proxy_mode", default=False)

View file

@ -390,7 +390,7 @@ class MCPDebug:
server_auth_type = server.auth_type
break
scope_headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
scope_headers: Final = MCPRequestHandler.safe_get_headers_from_scope(scope)
litellm_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(scope_headers)
return MCPDebug.build_debug_headers(

View file

@ -76,10 +76,11 @@ from litellm.integrations.custom_guardrail import (
)
from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: F401 # legacy module exports
MCPRequestHandler,
MCPServerAccess,
_is_mcp_admitted_user_subject,
_is_mcp_admitted_user_subject, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.catalog import _configuration_identity, _DiscoveryCache, _DiscoveryKey
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
@ -100,10 +101,11 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports
_redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
canonicalize_url_identity,
get_byok_www_authenticate,
redact_mcp_resource_url,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Error,
@ -285,12 +287,14 @@ class ListedToolsCaller:
# gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes.
# OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the
# config-YAML and DB server loaders so the two paths cannot drift on which modes trigger discovery.
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = (
UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final[tuple[MCPAuth, ...]] = (
MCPAuth.oauth2,
MCPAuth.true_passthrough,
MCPAuth.oauth_delegate,
)
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: Final = UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
_MCP_OAUTH_DISCOVERY_ON_STARTUP_ENV: Final = "LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP"
_TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on"))
@ -756,7 +760,7 @@ def _flow_endpoints_missing(
# A configured exchange endpoint replaces discovery entirely; only a server that must
# discover its token endpoint and still has none is unresolved.
return token_exchange_endpoint is None and token_url is None
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
if auth_type not in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return False
if oauth2_flow == "client_credentials":
return token_url is None
@ -926,7 +930,7 @@ def _restrict_discovery_to_corroborated_authorization_server(
def _redacted_origin_list(urls: Sequence[str]) -> str:
return ", ".join(_redact_mcp_resource_url(url) or "<unparseable url>" for url in urls)
return ", ".join(redact_mcp_resource_url(url) or "<unparseable url>" for url in urls)
def _sanitized_error_text(exc: Exception) -> str:
@ -1078,7 +1082,7 @@ def _warn_oauth_endpoints_unresolved(
"(RFC 8414)",
server_ref,
", ".join(unresolved),
_redact_mcp_resource_url(server_url) or "<no url>",
redact_mcp_resource_url(server_url) or "<no url>",
)
return
verbose_logger.warning(
@ -1107,7 +1111,7 @@ def _write_user_env_vars_cache(user_id: str, server_id: str, values: dict[str, s
_user_env_vars_cache[cache_key] = (values, time.monotonic())
def _should_strip_caller_authorization(
def should_strip_caller_authorization(
mcp_server: MCPServer,
raw_headers: dict[str, str] | None,
user_api_key_auth: UserAPIKeyAuth | None,
@ -1171,6 +1175,9 @@ def _should_strip_caller_authorization(
)
_should_strip_caller_authorization: Final = should_strip_caller_authorization
LITELLM_VIRTUAL_KEY_PREFIX: Final = "sk-"
@ -1272,7 +1279,7 @@ def _openapi_forwarded_extra_headers(
if not mcp_server.extra_headers or not raw_headers:
return None
normalized_raw: Final = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)}
skip_caller_authorization: Final = _should_strip_caller_authorization(
skip_caller_authorization: Final = should_strip_caller_authorization(
mcp_server=mcp_server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
@ -1289,7 +1296,7 @@ def _openapi_forwarded_extra_headers(
return forwarded or None
def _resolve_openapi_tool_auth(
def resolve_openapi_tool_auth(
mcp_server: MCPServer,
mcp_auth_header: str | None,
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None, # mutable-ok: sink shape
@ -1339,6 +1346,9 @@ def _resolve_openapi_tool_auth(
return None, forwarded, None
_resolve_openapi_tool_auth: Final = resolve_openapi_tool_auth
async def _resolve_byok_mcp_auth_header(
mcp_server: MCPServer,
user_api_key_auth: UserAPIKeyAuth | None,
@ -1385,7 +1395,7 @@ def _catalog_auth_header(
return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header
def _client_forwarded_authorization_headers(
def client_forwarded_authorization_headers(
mcp_server: MCPServer,
oauth2_headers: dict[str, str] | None,
raw_headers: dict[str, str] | None,
@ -1399,7 +1409,7 @@ def _client_forwarded_authorization_headers(
paths cannot drift, mirroring the ``_should_strip_caller_authorization`` split.
"""
extra_headers: Final = oauth2_headers.copy() if oauth2_headers else None
if extra_headers and _should_strip_caller_authorization(
if extra_headers and should_strip_caller_authorization(
mcp_server=mcp_server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
@ -1408,6 +1418,9 @@ def _client_forwarded_authorization_headers(
return extra_headers
_client_forwarded_authorization_headers: Final = client_forwarded_authorization_headers
async def _materialize_auth_headers(auth: httpx2.Auth | None) -> dict[str, str] | None:
"""Extract the header a resolved ``httpx2.Auth`` would set, as a plain dict, or None.
@ -1472,7 +1485,7 @@ def _redacted_registry_dump(servers: dict[str, MCPServer]) -> dict[str, dict[str
}
def _caller_authorization_fans_out(
def caller_authorization_fans_out(
server: MCPServer,
scope_servers: list[MCPServer] | None,
) -> bool:
@ -1489,6 +1502,9 @@ def _caller_authorization_fans_out(
)
_caller_authorization_fans_out: Final = caller_authorization_fans_out
def _extract_upstream_auth_failure(
exc: BaseException,
) -> tuple[int, str | None] | None:
@ -1986,7 +2002,7 @@ class MCPServerManager:
manual_issuer: Final = _blank_to_none(server.issuer)
manual_authorization_url: Final = _blank_to_none(server.authorization_url)
manual_token_url: Final = _blank_to_none(server.token_url)
is_discovery_auth_type: Final = server.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
is_discovery_auth_type: Final = server.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
use_issuer_anchor: Final = server.issuer_is_anchored
obo_needs_discovery: Final = self._obo_needs_endpoint_discovery(
server.auth_type,
@ -2276,7 +2292,7 @@ class MCPServerManager:
if raw and str(raw).strip():
self._upstream_initialize_instructions_by_server_id[server.server_id] = str(raw).strip()
async def _ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None:
async def ensure_upstream_initialize_instructions_cached(self, server: MCPServer) -> None:
"""
Open one upstream session and cache InitializeResult.instructions if missing.
@ -2326,7 +2342,7 @@ class MCPServerManager:
raise_on_missing=False,
)
extra_headers: dict[str, str] | None = dict(resolved_static_headers) if resolved_static_headers else None
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=None,
extra_headers=extra_headers,
@ -2345,6 +2361,8 @@ class MCPServerManager:
e,
)
_ensure_upstream_initialize_instructions_cached = ensure_upstream_initialize_instructions_cached
def get_registry(self) -> Mapping[str, MCPServer]:
"""
Get the registered MCP Servers from the registry and union with the config MCP Servers
@ -2458,7 +2476,7 @@ class MCPServerManager:
manual_authorization_url = _blank_to_none(server_config.get("authorization_url"))
manual_token_url = _blank_to_none(server_config.get("token_url"))
manual_registration_url = _blank_to_none(server_config.get("registration_url"))
is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
is_discovery_auth_type = auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
obo_needs_discovery = self._obo_needs_endpoint_discovery(
auth_type,
server_config.get("token_exchange_endpoint"),
@ -3110,7 +3128,7 @@ class MCPServerManager:
manual_authorization_url = _blank_to_none(mcp_server.authorization_url)
manual_token_url = _blank_to_none(mcp_server.token_url)
manual_registration_url = _blank_to_none(mcp_server.registration_url)
is_discovery_auth_type: Final = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
is_discovery_auth_type: Final = auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
token_exchange_endpoint: Final = mcp_server.token_exchange_endpoint or (
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None
)
@ -3452,7 +3470,7 @@ class MCPServerManager:
# applying this rule would hide almost every admitted user's OWN submitted servers. Their
# submissions are theirs by authorship, and their scope comes from the per-source union.
has_explicit_object_permission: Final = (
not _is_mcp_admitted_user_subject(user_api_key_auth)
not is_mcp_admitted_user_subject(user_api_key_auth)
and key_object_permission is not None
and (key_object_permission.mcp_servers is not None)
)
@ -3470,7 +3488,7 @@ class MCPServerManager:
the exception fallback, and applied AFTER every union (grants, operator-open,
submitted) because the scope is a ceiling over the whole reachable set; a resolver
fault therefore never widens a scoped bearer to the allow-all set."""
if user_api_key_auth is None or not _is_mcp_admitted_user_subject(user_api_key_auth):
if user_api_key_auth is None or not is_mcp_admitted_user_subject(user_api_key_auth):
return None
return user_api_key_auth.mcp_session_resource_server_id
@ -3505,7 +3523,7 @@ class MCPServerManager:
# rides the HUMAN, not the credential: an admin's session resolves the same registry their
# dashboard shows (connect-page parity), bounded like an admin key by explicit
# object_permission scope, the entitlement ceiling, and the session resource scope below.
is_admitted_subject: Final = _is_mcp_admitted_user_subject(user_api_key_auth)
is_admitted_subject: Final = is_mcp_admitted_user_subject(user_api_key_auth)
# The key explicitly opted out of every MCP server. Return zero before
# layering on allow_all_keys or submitted servers so the opt-out is absolute.
@ -3759,7 +3777,7 @@ class MCPServerManager:
blocked = 0
for sid in server_ids:
s = self.get_mcp_server_by_id(sid)
if s is not None and self._is_server_accessible_from_ip(s, client_ip):
if s is not None and self.is_server_accessible_from_ip(s, client_ip):
allowed.append(sid)
elif s is not None:
blocked += 1
@ -3774,7 +3792,7 @@ class MCPServerManager:
if server is None:
verbose_logger.warning("MCP Server %s not found", server_id)
return []
return list(await self._get_tools_from_server(server))
return list(await self.get_tools_from_server(server))
except Exception as e:
verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e)
return []
@ -3812,7 +3830,7 @@ class MCPServerManager:
try:
tools: Final = list(
await self._get_tools_from_server(
await self.get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
user_api_key_auth=user_api_key_auth,
@ -3861,7 +3879,7 @@ class MCPServerManager:
return None
@staticmethod
def _extract_subject_token(
def extract_subject_token(
oauth2_headers: Mapping[str, str] | None,
raw_headers: Mapping[str, str] | None,
user_api_key_auth: UserAPIKeyAuth | None,
@ -3878,6 +3896,8 @@ class MCPServerManager:
return None
return bearer
_extract_subject_token = extract_subject_token
def _obo_subject_token(
self,
server: MCPServer,
@ -3892,9 +3912,9 @@ class MCPServerManager:
"""
if server.auth_type != MCPAuth.oauth2_token_exchange:
return None
return self._extract_subject_token(None, raw_headers, user_api_key_auth)
return self.extract_subject_token(None, raw_headers, user_api_key_auth)
def _build_stdio_env(
def build_stdio_env(
self,
server: MCPServer,
raw_headers: Mapping[str, str] | None = None,
@ -3921,6 +3941,8 @@ class MCPServerManager:
return resolved_env
_build_stdio_env = build_stdio_env
def _references_per_user_env_var(self, server: MCPServer) -> bool:
"""True when ``server.static_headers`` reference a per-user ``${NAME}`` env var.
@ -4112,7 +4134,7 @@ class MCPServerManager:
Only OBO has a discovery challenge to raise; ID-JAG's failures are plain statuses whose body
already names what the user has to do, so they map through ``raise_public`` as at egress.
"""
subject_token: Final = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
subject_token: Final = self.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
match server.auth_type:
case MCPAuth.oauth2_token_exchange:
if not self._extract_bearer_token(oauth2_headers, None):
@ -4140,7 +4162,7 @@ class MCPServerManager:
)
raise_public(err)
async def _create_mcp_client(
async def create_mcp_client(
self,
server: MCPServer,
mcp_auth_header: str | dict[str, str] | None = None,
@ -4182,7 +4204,9 @@ class MCPServerManager:
elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None),
)
async def _get_tools_from_server(
_create_mcp_client = create_mcp_client
async def get_tools_from_server(
self,
server: MCPServer,
mcp_auth_header: str | dict[str, str] | None = None,
@ -4212,6 +4236,8 @@ class MCPServerManager:
)
return result.tools
_get_tools_from_server = get_tools_from_server
async def get_tools_page(
self,
server: MCPServer,
@ -4302,18 +4328,18 @@ class MCPServerManager:
for_list_tools=True,
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
# token_exchange (OBO) discovery needs the caller's token too: list it with the user's own
# token (mirrors the call path), not v1's deleted client_credentials fallback. Other modes
# never read the inbound bearer, so leave subject_token None to avoid forwarding it.
subject_token: Final = (
self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
self.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
if server.auth_type == MCPAuth.oauth2_token_exchange
else None
)
client = await self._create_mcp_client(
client = await self.create_mcp_client( # rebind-ok: pre-existing rebinding on a rename-only line
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
@ -4458,10 +4484,10 @@ class MCPServerManager:
return None
auth: Final = caller.user_api_key_auth
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
header_env: Final = self.build_stdio_env(server, caller.raw_headers)
stdio_env: Final = None if header_env == self.build_stdio_env(server) else header_env
caller_bearer: Final = (
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
self.extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
else None
)
@ -4608,9 +4634,9 @@ class MCPServerManager:
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
@ -4657,9 +4683,9 @@ class MCPServerManager:
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
@ -4712,9 +4738,9 @@ class MCPServerManager:
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
@ -4767,9 +4793,9 @@ class MCPServerManager:
)
or None
)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=headers,
@ -4820,10 +4846,10 @@ class MCPServerManager:
extra_headers = {}
extra_headers.update(server.static_headers)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
@ -4857,10 +4883,10 @@ class MCPServerManager:
extra_headers = {}
extra_headers.update(server.static_headers)
stdio_env: Final = self._build_stdio_env(server, raw_headers)
stdio_env: Final = self.build_stdio_env(server, raw_headers)
subject_token: Final = self._obo_subject_token(server, raw_headers, user_api_key_auth)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
@ -4946,7 +4972,7 @@ class MCPServerManager:
"MCP OAuth endpoint discovery against %s found no authorization server metadata. Attempts: %s. "
"The MCP server url may be misconfigured, or the upstream may not support OAuth discovery "
"(RFC 9728 / RFC 8414)",
_redact_mcp_resource_url(server_url) or "<unparseable url>",
redact_mcp_resource_url(server_url) or "<unparseable url>",
"; ".join(attempts) if attempts else "none recorded",
)
return metadata
@ -4957,7 +4983,7 @@ class MCPServerManager:
*,
allow_origin_fallback: bool,
) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]:
origin: Final = _redact_mcp_resource_url(server_url) or "<unparseable url>"
origin: Final = redact_mcp_resource_url(server_url) or "<unparseable url>"
try:
client: Final = get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
@ -5006,7 +5032,7 @@ class MCPServerManager:
*,
allow_origin_fallback: bool,
) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]:
origin: Final = _redact_mcp_resource_url(server_url) or "<unparseable url>"
origin: Final = redact_mcp_resource_url(server_url) or "<unparseable url>"
verbose_logger.debug(
"MCP OAuth discovery for %s received status error: %s",
server_url,
@ -5910,14 +5936,14 @@ class MCPServerManager:
}
# Create MCP request object for processing
mcp_request_obj: Final = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
mcp_request_obj: Final = proxy_logging_obj.create_mcp_request_object_from_kwargs(pre_hook_kwargs)
# Convert to LLM format for existing guardrail compatibility.
# Unified guardrails read the seeded logger off the request dict and pass it
# into ``apply_guardrail``, so ``@log_guardrail_information`` bridges their
# evaluations itself; the ``finally`` below covers native guardrails, which
# never receive it. Same seeding the pass-through routes do.
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs)
synthetic_llm_data: Final = proxy_logging_obj.convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs)
synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj
try:
@ -5930,7 +5956,9 @@ class MCPServerManager:
await proxy_logging_obj.enforce_mcp_server_rate_limits(user_api_key_auth, server)
if modified_data:
# Convert response back to MCP format and apply modifications
modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs)
modified_kwargs: Final = proxy_logging_obj.convert_mcp_hook_response_to_kwargs(
modified_data, pre_hook_kwargs
)
if modified_kwargs.get("arguments") != arguments:
hook_result["arguments"] = modified_kwargs["arguments"]
if modified_kwargs.get("extra_headers"):
@ -5990,7 +6018,7 @@ class MCPServerManager:
}
# Seeded for the same reason as in ``pre_call_tool_check``.
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs)
synthetic_llm_data: Final = proxy_logging_obj.convert_mcp_to_llm_format(request_obj, during_hook_kwargs)
synthetic_llm_data["litellm_logging_obj"] = litellm_logging_obj
# Wrapped so the bridge runs inside the task: the caller only holds the task and
@ -6063,7 +6091,7 @@ class MCPServerManager:
spec: Final = to_server_spec(mcp_server)
if spec is not None:
await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec)
retry_client: Final = await self._create_mcp_client(
retry_client: Final = await self.create_mcp_client(
server=mcp_server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
@ -6132,7 +6160,9 @@ class MCPServerManager:
MCPAuth.oauth2_token_exchange,
MCPAuth.oauth2_id_jag,
):
subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
subject_token = self.extract_subject_token( # rebind-ok: pre-existing rebinding on a rename-only line
oauth2_headers, raw_headers, user_api_key_auth
)
elif mcp_server.auth_type == MCPAuth.oauth2:
if mcp_server.has_client_credentials:
# For M2M OAuth servers, Authorization must come from token fetch.
@ -6143,18 +6173,20 @@ class MCPServerManager:
# token, so drop the caller-forwarded Authorization (apply-if-absent would
# otherwise let it shadow the resolved token). Delegate keeps it. Centralized
# via _should_strip_caller_authorization to match _prepare_mcp_server_headers.
if extra_headers and _should_strip_caller_authorization(
if extra_headers and should_strip_caller_authorization(
mcp_server=mcp_server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
):
extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
elif mcp_server.is_client_forwarded_token:
extra_headers = _client_forwarded_authorization_headers(
mcp_server=mcp_server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
extra_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line
client_forwarded_authorization_headers(
mcp_server=mcp_server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
)
if mcp_server.extra_headers and raw_headers:
@ -6162,7 +6194,7 @@ class MCPServerManager:
extra_headers = {}
normalized_raw_headers: Final = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)}
strip_caller_authorization: Final = _should_strip_caller_authorization(
strip_caller_authorization: Final = should_strip_caller_authorization(
mcp_server=mcp_server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
@ -6216,9 +6248,9 @@ class MCPServerManager:
if extra_headers is not None and len(extra_headers) == 0:
extra_headers = None
stdio_env: Final = self._build_stdio_env(mcp_server, raw_headers)
stdio_env: Final = self.build_stdio_env(mcp_server, raw_headers)
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=mcp_server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
@ -6347,7 +6379,9 @@ class MCPServerManager:
) -> MCPServer:
"""Resolve MCP server for call_tool (prefixed name, registry, fallback)."""
prefixed_tool_name: Final = add_server_prefix_to_name(name, server_name)
mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name)
mcp_server = self.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line
prefixed_tool_name
)
resolved_by_server_name_only = False
normalized_server_name: Final = normalize_server_name(server_name)
@ -6368,7 +6402,7 @@ class MCPServerManager:
resolved_by_server_name_only = True
break
if mcp_server is None:
fallback: Final = self._get_mcp_server_from_tool_name(name)
fallback: Final = self.get_mcp_server_from_tool_name(name)
if fallback is not None and (not server_name or _candidate_matches_server_name(fallback)):
mcp_server = fallback
if mcp_server is None:
@ -6495,7 +6529,9 @@ class MCPServerManager:
subject_token: str | None = None
if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
subject_token = self._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth)
subject_token = self.extract_subject_token( # rebind-ok: pre-existing rebinding on a rename-only line
oauth2_headers, raw_headers, user_api_key_auth
)
elif isinstance(spec.config, PassthroughConfig):
inbound_token, forwarded_headers = take_forwarded_authorization(forwarded_headers)
per_server_token: Final = passthrough_token_from_mcp_auth_header(mcp_auth_header)
@ -6637,7 +6673,7 @@ class MCPServerManager:
server_name,
)
auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth(
auth_header_value, openapi_forwarded_headers, upstream_credential = resolve_openapi_tool_auth(
mcp_server=mcp_server,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
@ -6655,21 +6691,21 @@ class MCPServerManager:
async def _call_openapi_via_handler():
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
_request_resolved_auth_headers,
request_auth_header,
request_extra_headers,
request_resolved_auth_headers,
)
auth_token: Final = _request_auth_header.set(auth_header_value)
extra_token: Final = _request_extra_headers.set(forwarded_headers)
resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers)
auth_token: Final = request_auth_header.set(auth_header_value)
extra_token: Final = request_extra_headers.set(forwarded_headers)
resolved_token: Final = request_resolved_auth_headers.set(resolved_auth_headers)
try:
async with self._limit_outbound_concurrency(mcp_server):
return await self._call_openapi_tool_handler(mcp_server, name, arguments, wire_compat)
finally:
_request_auth_header.reset(auth_token)
_request_extra_headers.reset(extra_token)
_request_resolved_auth_headers.reset(resolved_token)
request_auth_header.reset(auth_token)
request_extra_headers.reset(extra_token)
request_resolved_auth_headers.reset(resolved_token)
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
else:
@ -6720,7 +6756,7 @@ class MCPServerManager:
# Skip OAuth2 servers that rely on user-provided tokens
continue
try:
tools = await self._get_tools_from_server(server)
tools = await self.get_tools_from_server(server)
except MCPUpstreamAuthError as e:
# Pass-through servers expect a user-supplied bearer token;
# at startup we have none, so an upstream 401 is normal.
@ -6741,7 +6777,7 @@ class MCPServerManager:
self.tool_name_to_mcp_server_name_mapping[original_name] = server.name
self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name
def _get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None:
def get_mcp_server_from_tool_name(self, tool_name: str) -> MCPServer | None:
"""
Get the MCP Server from the tool name (handles both prefixed and non-prefixed names)
@ -6786,6 +6822,8 @@ class MCPServerManager:
return None
_get_mcp_server_from_tool_name = get_mcp_server_from_tool_name
async def reload_servers_from_database(self):
await self.catalog.reload()
@ -6809,7 +6847,7 @@ class MCPServerManager:
# Fallback if proxy_server not available
return {}
def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool:
def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool:
"""
Check if a server is accessible from the given client IP.
@ -6830,12 +6868,14 @@ class MCPServerManager:
internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges"))
return IPAddressUtils.is_internal_ip(client_ip, internal_networks)
_is_server_accessible_from_ip = is_server_accessible_from_ip
def get_mcp_server_by_id(self, server_id: str, client_ip: str | None = None) -> MCPServer | None:
"""Get the MCP Server from the server id."""
registry: Final = self.get_registry()
for server in registry.values():
if server.server_id == server_id:
if not self._is_server_accessible_from_ip(server, client_ip):
if not self.is_server_accessible_from_ip(server, client_ip):
return None
return server
return None
@ -6961,19 +7001,19 @@ class MCPServerManager:
# Pass 1: Match by alias (highest priority)
for server in registry.values():
if server.alias == server_name:
if not self._is_server_accessible_from_ip(server, client_ip):
if not self.is_server_accessible_from_ip(server, client_ip):
return None
return server
# Pass 2: Match by server_name
for server in registry.values():
if server.server_name == server_name:
if not self._is_server_accessible_from_ip(server, client_ip):
if not self.is_server_accessible_from_ip(server, client_ip):
return None
return server
# Pass 3: Match by name (lowest priority)
for server in registry.values():
if server.name == server_name:
if not self._is_server_accessible_from_ip(server, client_ip):
if not self.is_server_accessible_from_ip(server, client_ip):
return None
return server
return None
@ -6989,7 +7029,7 @@ class MCPServerManager:
registry: Final = self.get_registry()
if client_ip is None:
return registry
return {k: v for k, v in registry.items() if self._is_server_accessible_from_ip(v, client_ip)}
return {k: v for k, v in registry.items() if self.is_server_accessible_from_ip(v, client_ip)}
def _generate_stable_server_id(
self,
@ -7054,7 +7094,7 @@ class MCPServerManager:
if server.spec_path:
spec_status, spec_error, spec_checked_at = await self._openapi_health_probes(server.spec_path).check()
return self._build_mcp_server_table(server).model_copy(
return self.build_mcp_server_table(server).model_copy(
update=MappingProxyType(
{
"status": spec_status,
@ -7086,7 +7126,7 @@ class MCPServerManager:
raise_on_missing=False,
)
extra_headers: Final = dict(resolved_static_headers) if resolved_static_headers else {}
client: Final = await self._create_mcp_client(
client: Final = await self.create_mcp_client(
server=server,
mcp_auth_header=None,
extra_headers=extra_headers,
@ -7207,7 +7247,7 @@ class MCPServerManager:
verbose_logger.warning("MCP Server %s not found in registry", server_id)
continue
mcp_server_table = self._build_mcp_server_table(server)
mcp_server_table = self.build_mcp_server_table(server)
list_mcp_servers.append(mcp_server_table)
return list_mcp_servers
@ -7220,7 +7260,7 @@ class MCPServerManager:
return None
return [MCPEnvVar.model_validate(env_var) for env_var in env_vars]
def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable:
return LiteLLM_MCPServerTable(
server_id=server.server_id,
is_config=self.is_config_declared_server(server.server_id) and server.server_id not in self.registry,
@ -7275,6 +7315,8 @@ class MCPServerManager:
rpm=server.rpm,
)
_build_mcp_server_table = build_mcp_server_table
async def get_all_mcp_servers_unfiltered(self) -> list[LiteLLM_MCPServerTable]:
"""Return all MCP servers from registry without applying access controls."""
@ -7284,7 +7326,7 @@ class MCPServerManager:
servers: Final[list[LiteLLM_MCPServerTable]] = []
for server in registry.values():
servers.append(self._build_mcp_server_table(server))
servers.append(self.build_mcp_server_table(server))
return servers
async def get_all_mcp_servers_with_health_unfiltered(

View file

@ -42,7 +42,11 @@ from typing import Final, Literal, Protocol
from pydantic import JsonValue
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials
from litellm.proxy._experimental.mcp_server.db import ( # noqa: F401 # legacy module exports
_decode_oauth_payload, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
decode_oauth_payload,
decrypt_credentials,
)
from litellm.proxy.utils import PrismaClient
from litellm.types.mcp import MCPCredentials
@ -156,7 +160,7 @@ async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[Backfi
where={"server_id": {"in": server_ids}},
)
server_ids_with_oauth_tokens: Final[set[str]] = {
token_row.server_id for token_row in token_rows if _decode_oauth_payload(token_row.credential_b64) is not None
token_row.server_id for token_row in token_rows if decode_oauth_payload(token_row.credential_b64) is not None
}
classified: Final = tuple(

View file

@ -205,7 +205,7 @@ class MCPOAuth2TokenCache(InMemoryCache):
mcp_oauth2_token_cache: Final = MCPOAuth2TokenCache()
def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int:
def compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) -> int:
"""Compute Redis TTL for a per-user token.
Uses server.token_storage_ttl_seconds when configured, capped at the token's
@ -223,6 +223,9 @@ def _compute_per_user_token_ttl(server: "MCPServer", expires_in: int | None) ->
return MCP_PER_USER_TOKEN_DEFAULT_TTL
_compute_per_user_token_ttl: Final = compute_per_user_token_ttl
class MCPPerUserTokenCache:
"""Redis-backed cache for per-user OAuth2 access tokens.

View file

@ -84,7 +84,7 @@ def _origin_label(scheme: str, netloc: str) -> str:
return f"{scheme}://{netloc}" if netloc else f"{scheme}://"
def _redact_mcp_resource_url(url: str | None) -> str | None:
def redact_mcp_resource_url(url: str | None) -> str | None:
"""Reduce an MCP server URL to its origin (scheme + host + port) for logging.
Everything else is dropped: userinfo (``user:pass@``), the query string, the
@ -107,6 +107,9 @@ def _redact_mcp_resource_url(url: str | None) -> str | None:
return urlunsplit((parts.scheme, netloc, "", "", "")) or None
_redact_mcp_resource_url: Final = redact_mcp_resource_url
def _resolve_proxy_base_url_env() -> str | None:
global _warned_invalid_proxy_base_url
configured: Final = os.environ.get("PROXY_BASE_URL", "").strip()

View file

@ -27,7 +27,9 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
# tag-namespaced operationIds like "actions/download-job-logs-for-workflow-run"
# which include '/'. Sanitize here so the same regex passes everywhere downstream.
_OPENAPI_TOOL_NAME_INVALID_CHARS: Final = re.compile(r"[^a-zA-Z0-9_-]")
_OPENAPI_TOOL_NAME_MAX_LEN: Final = 128
OPENAPI_TOOL_NAME_MAX_LEN: Final = 128
_OPENAPI_TOOL_NAME_MAX_LEN: Final = OPENAPI_TOOL_NAME_MAX_LEN
def sanitize_openapi_tool_name(raw_name: str) -> str:
@ -41,7 +43,7 @@ def sanitize_openapi_tool_name(raw_name: str) -> str:
if not raw_name:
return raw_name
sanitized: Final = _OPENAPI_TOOL_NAME_INVALID_CHARS.sub("_", raw_name).lower()
return sanitized[:_OPENAPI_TOOL_NAME_MAX_LEN]
return sanitized[:OPENAPI_TOOL_NAME_MAX_LEN]
from litellm._logging import verbose_logger
@ -106,23 +108,31 @@ HEADERS: Final[dict[str, str]] = {}
# Per-request auth header override for BYOK servers.
# Set this ContextVar before calling a local tool handler to inject the user's
# stored credential into the HTTP request made by the tool function closure.
_request_auth_header: contextvars.ContextVar[str | None] = contextvars.ContextVar("_request_auth_header", default=None)
request_auth_header: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar(
"_request_auth_header", default=None
)
_request_auth_header: Final = request_auth_header
# Per-request extra headers forwarded from the client request.
# Populated from MCPServer.extra_headers names matched against raw request
# headers in server.py before dispatching to a local/OpenAPI tool handler.
_request_extra_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar(
request_extra_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar(
"_request_extra_headers", default=None
)
_request_extra_headers: Final = request_extra_headers
# Per-request headers carrying the gateway-resolved upstream credential
# (stored per-user OAuth token, minted M2M token, exchanged OBO token).
# Set from MCPServerManager.resolve_openapi_upstream_auth; authoritative
# over every other Authorization source in _merge_openapi_tool_request_headers.
_request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar(
request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | None]] = contextvars.ContextVar(
"_request_resolved_auth_headers", default=None
)
_request_resolved_auth_headers: Final = request_resolved_auth_headers
_request_upstream_url: Final[contextvars.ContextVar[str | None]] = contextvars.ContextVar(
"_request_upstream_url", default=None
)
@ -369,7 +379,7 @@ async def _drop_credential_across_origin(request: httpx.Request) -> None:
built would never be closed.
"""
guard: Final = credential_redirect_hook(
_request_upstream_url.get() or "", custom_credential_slot(_request_resolved_auth_headers.get())
_request_upstream_url.get() or "", custom_credential_slot(request_resolved_auth_headers.get())
)
if guard is not None:
await guard(request)
@ -382,7 +392,7 @@ def _upstream_client() -> AsyncHTTPHandler:
itself, so this arm installs the same hook the MCP client uses. Both variants come from the
shared cache, so a guarded call reuses its connection pool like any other.
"""
if custom_credential_slot(_request_resolved_auth_headers.get()) is None:
if custom_credential_slot(request_resolved_auth_headers.get()) is None:
return get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
return get_async_httpx_client(
llm_provider=httpxSpecialProvider.MCP,
@ -417,20 +427,20 @@ def _merge_openapi_tool_request_headers(
Header names are compared case-insensitively so different casing cannot
bypass the precedence rules.
"""
request_extra: Final = _request_extra_headers.get() or {}
request_extra: Final = request_extra_headers.get() or {}
static: Final = static_headers or {}
static_lower_names: Final = {k.lower() for k in static}
effective_headers: dict[str, str] = {k: v for k, v in request_extra.items() if k.lower() not in static_lower_names}
effective_headers.update(static)
override_auth: Final = _request_auth_header.get()
override_auth: Final = request_auth_header.get()
if override_auth:
for existing in [k for k in effective_headers if k.lower() == "authorization"]:
del effective_headers[existing]
effective_headers["Authorization"] = override_auth
resolved_auth_headers: Final = _request_resolved_auth_headers.get() or {}
resolved_auth_headers: Final = request_resolved_auth_headers.get() or {}
for name, value in resolved_auth_headers.items():
for existing in [k for k in effective_headers if k.lower() == name.lower()]:
del effective_headers[existing]
@ -622,7 +632,7 @@ def register_tools_from_openapi(spec: Mapping[str, Any], base_url: str) -> None:
while unique in used_names:
n += 1
suffix = f"_{n}"
unique = tool_name[: _OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix
unique = tool_name[: OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix
tool_name = unique
used_names.add(tool_name)

View file

@ -83,23 +83,31 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
classify_list_exception,
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports
MCPServerManager,
_caller_authorization_fans_out,
_client_forwarded_authorization_headers,
_resolve_openapi_tool_auth,
_should_strip_caller_authorization,
_caller_authorization_fans_out, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_client_forwarded_authorization_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_resolve_openapi_tool_auth, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_should_strip_caller_authorization, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
caller_authorization_fans_out,
client_forwarded_authorization_headers,
global_mcp_server_manager,
listed_tools_caller_for,
resolve_openapi_tool_auth,
should_strip_caller_authorization,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports
_redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
get_byok_www_authenticate,
redact_mcp_resource_url,
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
_request_resolved_auth_headers,
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( # noqa: F401 # legacy module exports
_request_auth_header, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_request_extra_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_request_resolved_auth_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
request_auth_header,
request_extra_headers,
request_resolved_auth_headers,
)
from litellm.proxy._experimental.mcp_server.result_conversion import (
WireCompat,
@ -523,7 +531,7 @@ async def _get_allowed_mcp_servers_from_mcp_server_names(
if not server_name_matched:
try:
access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups(
access_group_server_ids = await MCPRequestHandler.get_mcp_servers_from_access_groups(
[server_or_group]
)
# Only include servers that the user has access to
@ -882,7 +890,7 @@ def _prepare_mcp_server_headers(
# x-mcp-{alias}-authorization. The decision is computed once so BOTH the forwarding branch and
# the extra_headers copy loop below honor it — otherwise a server that lists Authorization in
# extra_headers would re-copy the withheld bearer from raw_headers and replay it anyway.
withhold_forwarded_authorization: Final = is_client_forwarded_mode and _caller_authorization_fans_out(
withhold_forwarded_authorization: Final = is_client_forwarded_mode and caller_authorization_fans_out(
server, scope_servers
)
if server.auth_type == MCPAuth.oauth2:
@ -897,7 +905,7 @@ def _prepare_mcp_server_headers(
# token, so drop the caller-forwarded Authorization (apply-if-absent would
# otherwise let it shadow the resolved token). Delegate keeps it. Centralized
# via _should_strip_caller_authorization to match _call_regular_mcp_tool.
if extra_headers and _should_strip_caller_authorization(
if extra_headers and should_strip_caller_authorization(
mcp_server=server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
@ -905,11 +913,13 @@ def _prepare_mcp_server_headers(
extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER)
elif is_client_forwarded_mode:
if not withhold_forwarded_authorization:
extra_headers = _client_forwarded_authorization_headers(
mcp_server=server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
extra_headers = ( # rebind-ok: pre-existing rebinding on a rename-only line
client_forwarded_authorization_headers(
mcp_server=server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
)
if server.extra_headers and raw_headers:
@ -922,7 +932,7 @@ def _prepare_mcp_server_headers(
# ``MCPServerManager._call_regular_mcp_tool`` so the two
# code paths cannot drift on this security-sensitive choice.
# See ``_should_strip_caller_authorization`` for the rules.
strip_caller_authorization: Final = _should_strip_caller_authorization(
strip_caller_authorization: Final = should_strip_caller_authorization(
mcp_server=server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
@ -1787,7 +1797,7 @@ def _challenge_missing_token_exchange_subject(
return
if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers):
return
if global_mcp_server_manager._extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) is not None:
if global_mcp_server_manager.extract_subject_token(oauth2_headers, raw_headers, user_api_key_auth) is not None:
return
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 # lazy: adapter pulls MCP subgraph
raise_token_exchange_challenge,
@ -1978,10 +1988,12 @@ async def _execute_mcp_tool(
original_tool_name = name
else:
# Resolve from tool name (MCP JSON-RPC or prefixed REST tool names).
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
mcp_server = global_mcp_server_manager.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line
name
)
if mcp_server is None and requested_server is not None:
for known_prefix in iter_known_server_prefixes(requested_server):
candidate = global_mcp_server_manager._get_mcp_server_from_tool_name(
candidate = global_mcp_server_manager.get_mcp_server_from_tool_name(
add_server_prefix_to_name(name, known_prefix)
)
if candidate is not None:
@ -2034,7 +2046,9 @@ async def _execute_mcp_tool(
# Resolve the MCP server early so BYOK checks and credential injection
# apply to ALL dispatch paths (local tool registry AND managed MCP server).
if mcp_server is None:
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
mcp_server = global_mcp_server_manager.get_mcp_server_from_tool_name( # rebind-ok: pre-existing rebinding on a rename-only line
name
)
client_auth_header: Final = mcp_auth_header
if mcp_server:
@ -2131,7 +2145,7 @@ async def _execute_mcp_tool(
verbose_logger.debug("Executing local registry tool: %s", name)
# The credential rides ContextVars because the tool function has its
# headers baked into the closure at registration time.
auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth(
auth_header_value, openapi_forwarded_headers, upstream_credential = resolve_openapi_tool_auth(
mcp_server=mcp_server,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
@ -2150,15 +2164,15 @@ async def _execute_mcp_tool(
forwarded_headers=openapi_forwarded_headers,
)
_auth_token: Final = _request_auth_header.set(auth_header_value)
_extra_token: Final = _request_extra_headers.set(forwarded_headers)
_resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers)
_auth_token: Final = request_auth_header.set(auth_header_value)
_extra_token: Final = request_extra_headers.set(forwarded_headers)
_resolved_token: Final = request_resolved_auth_headers.set(resolved_auth_headers)
try:
response = await _handle_local_mcp_tool(name, arguments, wire_compat)
finally:
_request_auth_header.reset(_auth_token)
_request_extra_headers.reset(_extra_token)
_request_resolved_auth_headers.reset(_resolved_token)
request_auth_header.reset(_auth_token)
request_extra_headers.reset(_extra_token)
request_resolved_auth_headers.reset(_resolved_token)
# Try managed MCP server tool (the name is bare; the prefix boundary was
# already resolved above against this server's registered prefixes)
@ -2618,7 +2632,7 @@ def _get_standard_logging_mcp_tool_call(
server_name: str | None,
session_id: str | None = None,
) -> StandardLoggingMCPToolCall:
mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name(
mcp_server: Final = global_mcp_server_manager.get_mcp_server_from_tool_name(
add_server_prefix_to_name(name, server_name) if server_name else name
)
namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name
@ -2632,7 +2646,7 @@ def _get_standard_logging_mcp_tool_call(
namespaced_tool_name=namespaced_tool_name,
mcp_session_id=session_id,
mcp_auth_mode=mcp_server.auth_type,
mcp_server_resource=_redact_mcp_resource_url(mcp_server.url),
mcp_server_resource=redact_mcp_resource_url(mcp_server.url),
)
else:
return StandardLoggingMCPToolCall(

View file

@ -37,7 +37,10 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url
from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 # legacy module exports
_redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
redact_mcp_resource_url,
)
from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
acting_user_auth,
@ -64,7 +67,10 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
from litellm.proxy.utils import ProxyLogging
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
safe_get_request_headers,
)
from litellm.types.mcp import MCPAuth
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
@ -144,7 +150,7 @@ def _known_connection_error_message(exc: BaseException, url: str | None, timeout
if isinstance(exc, TimeoutError):
return (
"Failed to connect to MCP server: no valid MCP response received from "
f"{_redact_mcp_resource_url(url) or 'the server'} "
f"{redact_mcp_resource_url(url) or 'the server'} "
f"within {timeout_seconds:.0f}s. Check that the LiteLLM proxy can reach this URL "
"from its network (DNS, egress rules, firewalls) and that the server answers MCP requests."
)
@ -223,8 +229,9 @@ if MCP_AVAILABLE:
from litellm.llms.litellm_proxy.skills.skill_search import (
DEFAULT_SKILL_SEARCH_TOP_K,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: F401 # legacy module exports
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
ListedToolsCaller,
global_mcp_server_manager,
)
@ -242,8 +249,9 @@ if MCP_AVAILABLE:
filter_tools_by_key_team_permissions,
fire_mcp_tool_call_failure_logging,
)
from litellm.proxy._experimental.mcp_server.server import (
_apply_toolset_scope,
from litellm.proxy._experimental.mcp_server.server import ( # noqa: F401 # legacy module exports
_apply_toolset_scope, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
apply_toolset_scope,
reject_disallowed_mcp_client,
)
from litellm.proxy._experimental.mcp_server.tool_catalog_guard import (
@ -369,7 +377,7 @@ if MCP_AVAILABLE:
virtual_mcp_server_auth_headers,
virtual_raw_headers,
) = _extract_mcp_headers_from_request(request, MCPRequestHandler)
virtual_oauth2_headers: Final = MCPRequestHandler._get_oauth2_headers_from_headers(request.headers)
virtual_oauth2_headers: Final = MCPRequestHandler.get_oauth2_headers_from_headers(request.headers)
if tool_name == MCP_TOOL_SEARCH_TOOL_NAME:
return await handle_mcp_tool_search(
query=tool_arguments.get("query", ""),
@ -656,7 +664,7 @@ if MCP_AVAILABLE:
if (
_server is not None
and _rest_client_ip is not None
and not global_mcp_server_manager._is_server_accessible_from_ip(_server, _rest_client_ip)
and not global_mcp_server_manager.is_server_accessible_from_ip(_server, _rest_client_ip)
):
raise HTTPException(
status_code=403,
@ -707,7 +715,7 @@ if MCP_AVAILABLE:
record_listing: bool,
) -> list[MCPTool]:
return list(
await global_mcp_server_manager._get_tools_from_server(
await global_mcp_server_manager.get_tools_from_server(
server=server,
mcp_auth_header=server_auth_header,
extra_headers=extra_headers,
@ -856,7 +864,7 @@ if MCP_AVAILABLE:
if (
_server is not None
and rest_client_ip is not None
and not global_mcp_server_manager._is_server_accessible_from_ip(_server, rest_client_ip)
and not global_mcp_server_manager.is_server_accessible_from_ip(_server, rest_client_ip)
):
raise HTTPException(
status_code=403,
@ -953,7 +961,7 @@ if MCP_AVAILABLE:
status_code=404,
detail=f"Toolset '{toolset_name}' not found",
)
return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id)
return await apply_toolset_scope(user_api_key_dict, toolset.toolset_id)
@router.get("/tools/list", dependencies=[Depends(user_api_key_auth)])
@catalog_operation(global_manager)
@ -1034,8 +1042,8 @@ if MCP_AVAILABLE:
# Extract auth headers from request
headers: Final = request.headers
raw_headers_from_request: Final = dict(headers)
mcp_auth_header: Final = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
mcp_server_auth_headers: Final = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
mcp_auth_header: Final = MCPRequestHandler.get_mcp_auth_header_from_headers(headers)
mcp_server_auth_headers: Final = MCPRequestHandler.get_mcp_server_auth_headers_from_headers(headers)
auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict)
@ -1293,7 +1301,7 @@ if MCP_AVAILABLE:
if target_server is not None:
user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict)
caller_oauth2_headers: Final = (
MCPRequestHandler._get_oauth2_headers_from_headers(request.headers)
MCPRequestHandler.get_oauth2_headers_from_headers(request.headers)
if target_server is not None and target_server.auth_type in _CLIENT_FORWARDED_TOKEN_AUTH_TYPES
else None
)
@ -1396,9 +1404,10 @@ if MCP_AVAILABLE:
# /health/tools/list -> List tools from MCP server
# For these routes users will dynamically pass the MCP connection params, they don't need to be on the MCP registry
########################################################
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
from litellm.proxy.management_endpoints.mcp_management_endpoints import ( # noqa: F401 # legacy module exports
NewMCPServerRequest,
_inherit_credentials_from_existing_server,
_inherit_credentials_from_existing_server, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
inherit_credentials_from_existing_server,
)
def _extract_credentials(
@ -1461,7 +1470,7 @@ if MCP_AVAILABLE:
saved_origin is not None and saved_origin == preview_origin
)
request: Final = (
_inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request
inherit_credentials_from_existing_server(new_mcp_server_request) if may_inherit else new_mcp_server_request
)
mcp_auth_header: Final = (
request.credentials.get("auth_value")
@ -1472,8 +1481,8 @@ if MCP_AVAILABLE:
# when the primary x-litellm-api-key header is absent, the Authorization value is the
# caller's LiteLLM key, not an upstream token, and must never be forwarded upstream.
oauth2_headers: Final = (
MCPRequestHandler._get_oauth2_headers_from_headers(headers)
if request.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
MCPRequestHandler.get_oauth2_headers_from_headers(headers)
if request.auth_type in UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
else None
)
@ -1564,7 +1573,7 @@ if MCP_AVAILABLE:
instructions=request.instructions,
)
stdio_env: Final = global_mcp_server_manager._build_stdio_env(server_model, raw_headers)
stdio_env: Final = global_mcp_server_manager.build_stdio_env(server_model, raw_headers)
# For M2M OAuth servers, drop the incoming Authorization header so that
# resolve_mcp_auth can auto-fetch a token via client_credentials.
@ -1617,7 +1626,7 @@ if MCP_AVAILABLE:
)
with anyio.fail_after(timeout_seconds):
client: Final = await global_mcp_server_manager._create_mcp_client(
client: Final = await global_mcp_server_manager.create_mcp_client(
server=server_model,
mcp_auth_header=mcp_auth_header,
extra_headers=merged_headers,
@ -1648,7 +1657,7 @@ if MCP_AVAILABLE:
async def _preview_openapi_tools(spec_path: str) -> dict:
"""Generate tool previews from an OpenAPI spec without creating a server."""
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_OPENAPI_TOOL_NAME_MAX_LEN,
OPENAPI_TOOL_NAME_MAX_LEN,
build_input_schema,
load_openapi_spec_async,
resolve_operation_params,
@ -1681,7 +1690,7 @@ if MCP_AVAILABLE:
while unique in used_names:
n += 1
suffix = f"_{n}"
unique = op_id[: _OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix
unique = op_id[: OPENAPI_TOOL_NAME_MAX_LEN - len(suffix)] + suffix
op_id = unique
used_names.add(op_id)
summary = operation.get("summary", "")
@ -1738,7 +1747,7 @@ if MCP_AVAILABLE:
_test_connection_operation,
mcp_auth_header=staged.mcp_auth_header,
oauth2_headers=staged.oauth2_headers,
raw_headers=_safe_get_request_headers(request),
raw_headers=safe_get_request_headers(request),
)
@router.post("/test/tools/list", dependencies=[Depends(user_api_key_auth)])
@ -1797,5 +1806,5 @@ if MCP_AVAILABLE:
_list_tools_operation,
mcp_auth_header=staged.mcp_auth_header,
oauth2_headers=staged.oauth2_headers,
raw_headers=_safe_get_request_headers(request),
raw_headers=safe_get_request_headers(request),
)

View file

@ -772,11 +772,11 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N
import litellm
from litellm.proxy._types import ModelAccessDeniedProxyException
from litellm.proxy.auth.auth_checks import (
_check_team_member_model_access,
can_key_call_model,
can_project_access_model,
can_team_access_model,
can_user_call_model,
check_team_member_model_access,
get_project_object,
get_team_object,
get_user_object,
@ -832,7 +832,7 @@ async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | N
team_model_aliases=getattr(user_api_key_auth, "team_model_aliases", None),
)
if _user_id and _proxy_logging_obj:
await _check_team_member_model_access(
await check_team_member_model_access(
model=model,
team_object=team_obj,
valid_token=user_api_key_auth,

View file

@ -117,7 +117,7 @@ class SemanticMCPToolFilter:
self.tool_router = None
raise
def _extract_tool_info(self, tool) -> tuple[str, str]:
def extract_tool_info(self, tool) -> tuple[str, str]:
"""Extract name and description from MCP tool or OpenAI function dict."""
name: str
description: str
@ -133,6 +133,8 @@ class SemanticMCPToolFilter:
return name, description
_extract_tool_info = extract_tool_info
def _build_router(self, tools: list) -> None:
"""Build semantic router with tools (MCPTool objects or OpenAI function dicts)."""
from semantic_router.routers import SemanticRouter
@ -153,7 +155,7 @@ class SemanticMCPToolFilter:
self._tool_map = {}
for tool in tools:
name, description = self._extract_tool_info(tool)
name, description = self.extract_tool_info(tool)
self._tool_map[name] = tool
routes.append(
@ -187,13 +189,13 @@ class SemanticMCPToolFilter:
def _has_tools_missing_from_index(self, tools: Sequence[object]) -> bool:
"""Allocation-free check for any named tool not yet in the semantic index."""
return any(name and name not in self._tool_map for name in (self._extract_tool_info(t)[0] for t in tools))
return any(name and name not in self._tool_map for name in (self.extract_tool_info(t)[0] for t in tools))
def _tools_missing_from_index(self, tools: Sequence[object]) -> Mapping[str, object]:
"""Map name -> tool for every named tool not yet in the semantic index."""
return {
name: tool
for name, tool in ((self._extract_tool_info(t)[0], t) for t in tools)
for name, tool in ((self.extract_tool_info(t)[0], t) for t in tools)
if name and name not in self._tool_map
}
@ -228,7 +230,7 @@ class SemanticMCPToolFilter:
if not missing:
return
descriptions: Final = {name: self._extract_tool_info(tool)[1] for name, tool in missing.items()}
descriptions: Final = {name: self.extract_tool_info(tool)[1] for name, tool in missing.items()}
routes: Final = [
Route(
name=name,
@ -302,7 +304,7 @@ class SemanticMCPToolFilter:
verbose_logger.warning("Semantic router could not be built from the request's tools")
return available_tools
available_names: Final = [name for name in (self._extract_tool_info(t)[0] for t in available_tools) if name]
available_names: Final = [name for name in (self.extract_tool_info(t)[0] for t in available_tools) if name]
if not available_names:
return available_tools
@ -406,7 +408,7 @@ class SemanticMCPToolFilter:
# names happen to be tail-compatible with the same incoming name.
available_by_name: Final[dict[str, object]] = {}
for tool in available_tools:
client_name, _ = self._extract_tool_info(tool)
client_name, _ = self.extract_tool_info(tool)
if client_name and client_name not in available_by_name:
available_by_name[client_name] = tool

View file

@ -32,9 +32,10 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: F401 # legacy module exports
MCPRequestHandler,
_is_mcp_admitted_user_subject,
_is_mcp_admitted_user_subject, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.client_allowlist import (
MCPClientAllowlist,
@ -47,13 +48,16 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
_mcp_gateway_server_name,
from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: F401 # legacy module exports
_mcp_active_toolset_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_mcp_gateway_initialize_instructions, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_mcp_gateway_server_name, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode
active_mcp_request_ctx_var,
get_active_mcp_request_ctx,
mcp_active_toolset_id,
mcp_gateway_initialize_instructions,
mcp_gateway_server_name,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCP_AUTH_DIAGNOSTICS_SCOPE_KEY,
@ -61,9 +65,9 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCPDebug,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
get_passthrough_www_authenticate,
get_route_relative_request_path,
redact_mcp_resource_url,
well_known_root_suffix,
)
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
@ -98,6 +102,8 @@ if TYPE_CHECKING:
from mcp.server.session import ServerSession as _McpServerSession
_redact_mcp_resource_url: Final = redact_mcp_resource_url
_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60
# Upper bound on concurrent stateful sessions a single caller may hold. Each
# `initialize` creates a session that survives until the idle timeout, so
@ -486,6 +492,7 @@ if MCP_AVAILABLE:
"mcp_get_prompt",
"mcp_read_resource",
"raise_denied_scoped_mcp_access",
"redact_mcp_resource_url",
)
from mcp.server import Server
@ -579,10 +586,10 @@ if MCP_AVAILABLE:
else base_options
)
updates: Final[dict[str, str]] = {}
merged: Final = _mcp_gateway_initialize_instructions.get()
merged: Final = mcp_gateway_initialize_instructions.get()
if merged is not None:
updates["instructions"] = merged
scoped_server_name: Final = _mcp_gateway_server_name.get()
scoped_server_name: Final = mcp_gateway_server_name.get()
if scoped_server_name is not None:
updates["server_name"] = scoped_server_name
return opts.model_copy(update=updates) if updates else opts
@ -1028,7 +1035,7 @@ if MCP_AVAILABLE:
# cancel sibling probes or 500 the gateway initialize request.
await asyncio.gather(
*[
operations.global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(s)
operations.global_mcp_server_manager.ensure_upstream_initialize_instructions_cached(s)
for s in allowed
if s is not None
],
@ -1041,13 +1048,13 @@ if MCP_AVAILABLE:
scoped_server_name = (
scoped_server.alias or scoped_server.server_name or scoped_server.name or scoped_server.server_id
)
instructions_token: Final = _mcp_gateway_initialize_instructions.set(merged)
server_name_token: Final = _mcp_gateway_server_name.set(scoped_server_name)
instructions_token: Final = mcp_gateway_initialize_instructions.set(merged)
server_name_token: Final = mcp_gateway_server_name.set(scoped_server_name)
try:
yield
finally:
_mcp_gateway_initialize_instructions.reset(instructions_token)
_mcp_gateway_server_name.reset(server_name_token)
mcp_gateway_initialize_instructions.reset(instructions_token)
mcp_gateway_server_name.reset(server_name_token)
from litellm.proxy._experimental.mcp_server.operations import (
_MCP_CREDENTIAL_REQUEST_FIELDS,
@ -1524,7 +1531,7 @@ if MCP_AVAILABLE:
scope["headers"] = [(k, v) for k, v in _headers if _normalize_header_name(k) != _mcp_session_header]
return False
async def _apply_toolset_scope(
async def apply_toolset_scope(
user_api_key_auth: UserAPIKeyAuth,
toolset_id: str,
acting_user: ActingUser = acting_user_auth,
@ -1544,7 +1551,7 @@ if MCP_AVAILABLE:
of its grant sources. Admins always pass.
"""
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view
# A key scoped to no MCP servers opts out of every MCP path. Enforce it
# here too, since toolset scoping replaces mcp_servers and would otherwise
@ -1558,13 +1565,13 @@ if MCP_AVAILABLE:
)
acting: Final = await acting_user(user_api_key_auth)
is_admin: Final = _user_has_admin_view(acting)
is_admin: Final = user_api_key_has_admin_view(acting)
if not is_admin and toolset_id not in await granted(acting):
raise HTTPException(
status_code=403,
detail=f"API key does not have access to toolset '{toolset_id}'.",
)
if _is_mcp_admitted_user_subject(acting):
if is_mcp_admitted_user_subject(acting):
resource_server_id: Final = acting.mcp_session_resource_server_id
if resource_server_id is not None and resource_server_id not in (
await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(
@ -1600,6 +1607,8 @@ if MCP_AVAILABLE:
)
return acting.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
_apply_toolset_scope: Final = apply_toolset_scope
async def _toolset_server_ids(toolset_id: str) -> set[str]:
return set(
await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id])
@ -1674,7 +1683,7 @@ if MCP_AVAILABLE:
if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth):
continue
if _is_mcp_admitted_user_subject(user_api_key_auth):
if is_mcp_admitted_user_subject(user_api_key_auth):
raise HTTPException(
status_code=401,
detail="Unauthorized",
@ -2033,10 +2042,12 @@ if MCP_AVAILABLE:
# Apply toolset scope if set server-side via ContextVar (set by
# /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py).
active_toolset_id: Final = _mcp_active_toolset_id.get()
active_toolset_id: Final = mcp_active_toolset_id.get()
toolset_allowed_server_ids: set[str] | None = None
if active_toolset_id and user_api_key_auth is not None:
user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
user_api_key_auth = ( # rebind-ok: pre-existing rebinding on a rename-only line
await apply_toolset_scope(user_api_key_auth, active_toolset_id)
)
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
@ -2380,10 +2391,12 @@ if MCP_AVAILABLE:
# Apply toolset scope if set server-side via ContextVar so the
# downstream probe list matches the fully-authorized server set
# (mirrors the streamable HTTP handler).
active_toolset_id: Final = _mcp_active_toolset_id.get()
active_toolset_id: Final = mcp_active_toolset_id.get()
toolset_allowed_server_ids: set[str] | None = None
if active_toolset_id and user_api_key_auth is not None:
user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
user_api_key_auth = ( # rebind-ok: pre-existing rebinding on a rename-only line
await apply_toolset_scope(user_api_key_auth, active_toolset_id)
)
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response

View file

@ -19,9 +19,13 @@ class MCPServerRegistry(Protocol):
def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: ...
def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ...
def is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ...
def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ...
_is_server_accessible_from_ip = is_server_accessible_from_ip
def build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ...
_build_mcp_server_table = build_mcp_server_table
async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth) -> list[str]: ...
@ -50,7 +54,7 @@ async def resolve_mcp_server(
temporary_server: Final[MCPServer | None] = await temp_lookup(server_id)
if temporary_server is not None:
return ResolvedMCPServer(
table=manager._build_mcp_server_table(temporary_server),
table=manager.build_mcp_server_table(temporary_server),
runtime=temporary_server,
source="temp",
)
@ -64,12 +68,12 @@ async def resolve_mcp_server(
registry_server: Final[MCPServer | None] = (
registry_candidate
if registry_candidate is not None
and (id_client_ip is None or manager._is_server_accessible_from_ip(registry_candidate, id_client_ip))
and (id_client_ip is None or manager.is_server_accessible_from_ip(registry_candidate, id_client_ip))
else None
)
if registry_server is not None:
return ResolvedMCPServer(
table=manager._build_mcp_server_table(registry_server),
table=manager.build_mcp_server_table(registry_server),
runtime=registry_server,
source="registry",
)
@ -78,7 +82,7 @@ async def resolve_mcp_server(
named_server: Final[MCPServer | None] = manager.get_mcp_server_by_name(server_id, client_ip=name_client_ip)
if named_server is not None:
return ResolvedMCPServer(
table=manager._build_mcp_server_table(named_server),
table=manager.build_mcp_server_table(named_server),
runtime=named_server,
source="registry",
)

View file

@ -125,9 +125,9 @@ async def acting_user_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth:
if not is_ui_session_credential(user_api_key_auth):
return user_api_key_auth
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view
if _user_has_admin_view(user_api_key_auth):
if user_api_key_has_admin_view(user_api_key_auth):
return user_api_key_auth
admitted: Final = await admitted_user_context(user_api_key_auth)
return admitted if admitted is not None else user_api_key_auth

View file

@ -459,9 +459,9 @@ class AgentRequestHandler:
"""
Resolve unified access group ids to agent IDs.
"""
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
from litellm.proxy.auth.auth_checks import get_agent_ids_from_access_groups
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
return await get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
@staticmethod
async def _get_agents_from_access_groups(

View file

@ -335,9 +335,9 @@ async def _resolve_daily_activity_agent_ids(
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> tuple[str, ...] | None:
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_endpoints.common_utils import user_api_key_has_admin_view
if _user_has_admin_view(user_api_key_dict):
if user_api_key_has_admin_view(user_api_key_dict):
return agent_ids
permitted_agent_ids: Final = await _permitted_daily_activity_agent_ids(
user_api_key_dict=user_api_key_dict,

View file

@ -72,7 +72,7 @@ class _MarketplaceEntry(TypedDict, total=False):
category: object
async def _get_prisma_client() -> object:
async def get_prisma_client() -> object:
"""Get the prisma client from proxy_server."""
from litellm.proxy.proxy_server import prisma_client
@ -84,6 +84,9 @@ async def _get_prisma_client() -> object:
return prisma_client
_get_prisma_client: Final = get_prisma_client
@router.get(
"/claude-code/marketplace.json",
tags=["Claude Code Marketplace"],
@ -111,7 +114,7 @@ async def get_marketplace(request: Request, key: str | None = None):
```
"""
try:
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
caller: Final[UserAPIKeyAuth | None] = (
await user_api_key_auth(request=request, api_key=f"Bearer {key}") if key else None
@ -328,7 +331,7 @@ async def register_plugin(
try:
_require_proxy_admin(user_api_key_dict)
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
if not re.match(r"^[a-z0-9-]+$", request.name):
raise HTTPException(
@ -408,7 +411,7 @@ async def list_plugins(
List of plugins with their metadata.
"""
try:
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
visibility: Final[SkillVisibility] = skill_visibility(user_api_key_dict)
plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many(
@ -478,7 +481,7 @@ async def get_plugin(
Plugin details including source and metadata.
"""
try:
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
where={"name": plugin_name}
@ -579,7 +582,7 @@ async def update_plugin(
try:
_require_proxy_admin(user_api_key_dict)
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
_validate_plugin_source(request.source)
@ -646,7 +649,7 @@ async def enable_plugin(
try:
_require_proxy_admin(user_api_key_dict)
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
where={"name": plugin_name}
@ -695,7 +698,7 @@ async def disable_plugin(
try:
_require_proxy_admin(user_api_key_dict)
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
where={"name": plugin_name}
@ -744,7 +747,7 @@ async def delete_plugin(
try:
_require_proxy_admin(user_api_key_dict)
prisma_client: Final = await _get_prisma_client()
prisma_client: Final = await get_prisma_client()
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
where={"name": plugin_name}

View file

@ -30,7 +30,10 @@ from litellm.proxy.common_request_processing import (
resolve_litellm_call_id,
)
from litellm.proxy.common_utils.error_body_call_id import error_body_call_id
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
from litellm.proxy.common_utils.openai_error_payload import (
LITELLM_CALL_ID_HEADER,
error_status_code,
@ -145,7 +148,7 @@ async def anthropic_response(
version,
)
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
result: Final = await base_llm_response_processor.base_process_llm_request(
@ -319,7 +322,7 @@ async def count_tokens(
litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
try:
request_data: Final = await _read_request_body(request=request)
request_data: Final = await read_request_body(request=request)
data: Final[dict] = {**request_data}
# Extract required fields

View file

@ -52,7 +52,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token i
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
from litellm.proxy.anthropic_endpoints.endpoints import anthropic_response, count_tokens
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import _safe_set_request_parsed_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
safe_set_request_parsed_body,
)
from litellm.proxy.management_endpoints.sso_helper_utils import CLI_SSO_SESSIONS_TARGET
from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail
from litellm.types.llms.base import LiteLLMBaseModel
@ -449,7 +452,7 @@ async def managed_settings(request: Request) -> Response:
async def _skip_otlp_body_parsing(request: Request) -> None:
_safe_set_request_parsed_body(request=request, parsed_body={})
safe_set_request_parsed_body(request=request, parsed_body={})
_OTLP_AUTHENTICATED: Final = (Depends(_skip_otlp_body_parsing), *_AUTHENTICATED)

View file

@ -170,7 +170,7 @@ async def create_skill(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -300,7 +300,7 @@ async def list_skills(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -397,7 +397,7 @@ async def get_skill(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -496,7 +496,7 @@ async def delete_skill(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

@ -96,9 +96,11 @@ from litellm.proxy.auth.model_access_denied import (
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
_safe_get_request_query_params,
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
safe_get_request_headers,
safe_get_request_query_params,
)
from litellm.proxy.common_utils.model_listing_utils import alias_map
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
@ -463,7 +465,7 @@ def _get_router_zero_cost_cache(llm_router: Router) -> dict[str, bool] | None:
return cache if isinstance(cache, dict) else None
def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool:
def is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool:
"""
Check if a model has zero cost (no configured pricing).
@ -582,6 +584,9 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
return True
_is_model_cost_zero: Final = is_model_cost_zero
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
_TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True})
@ -974,7 +979,7 @@ def route_skips_budget_checks(route: str) -> bool:
def request_skips_budget_checks(route: str, model: str | list[str] | None, llm_router: Router | None) -> bool:
return route_skips_budget_checks(route=route) or _is_model_cost_zero(model=model, llm_router=llm_router)
return route_skips_budget_checks(route=route) or is_model_cost_zero(model=model, llm_router=llm_router)
async def common_checks(
@ -1017,8 +1022,8 @@ async def common_checks(
_model: Final[str | list[str] | None] = get_model_from_request(
request_data=request_body,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
request_headers=safe_get_request_headers(request=request),
request_query_params=safe_get_request_query_params(request=request),
llm_router=llm_router,
request=request,
team_id=valid_token.team_id if valid_token is not None else None,
@ -1073,7 +1078,7 @@ async def common_checks(
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
if not await key_access_group_grants_model(
model=_model,
valid_token=valid_token,
team_object=team_object,
@ -1085,7 +1090,7 @@ async def common_checks(
# 2.2. If team member has per-member model scope, enforce it
if _model and team_object and valid_token and valid_token.user_id:
with tracer.trace("litellm.proxy.auth.common_checks.check_team_member_model_access"):
await _check_team_member_model_access(
await check_team_member_model_access(
model=_model,
team_object=team_object,
valid_token=valid_token,
@ -1122,7 +1127,7 @@ async def common_checks(
managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ())
if not isinstance(managed_models, (list, tuple)) or not managed_models:
raise HTTPException(403, "This agent has no model grants")
_can_object_call_model(
can_object_call_model(
model=_resolve_team_alias(
_model, team_model_aliases_for_auth_check(valid_token), valid_token.team_id, llm_router
),
@ -1261,7 +1266,7 @@ async def common_checks(
budget_check_coros: Final = tuple(
coro
for coro in (
_team_max_budget_check(
team_max_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
@ -1273,7 +1278,7 @@ async def common_checks(
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
),
_organization_max_budget_check(
organization_max_budget_check(
valid_token=valid_token,
team_object=team_object,
prisma_client=prisma_client,
@ -1305,7 +1310,7 @@ async def common_checks(
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
),
_check_end_user_budget(end_user_obj=end_user_object, route=route)
check_end_user_budget(end_user_obj=end_user_object, route=route)
if end_user_object is not None and end_user_object.litellm_budget_table is not None
else None,
)
@ -1381,7 +1386,7 @@ def effective_user_role(user_role: str | None) -> LitellmUserRoles:
return LitellmUserRoles.INTERNAL_USER
def _get_user_role(
def get_user_role(
user_obj: LiteLLM_UserTable | None,
) -> LitellmUserRoles | None:
if user_obj is None:
@ -1389,6 +1394,9 @@ def _get_user_role(
return effective_user_role(user_obj.user_role)
_get_user_role: Final = get_user_role
def _is_api_route_allowed(
route: str,
request: Request,
@ -1399,12 +1407,12 @@ def _is_api_route_allowed(
"""
- Route b/w api token check and normal token check
"""
_user_role: Final = _get_user_role(user_obj=user_obj)
_user_role: Final = get_user_role(user_obj=user_obj)
if valid_token is None:
raise Exception("Invalid proxy server token passed. valid_token=None.")
if not _is_user_proxy_admin(user_obj=user_obj): # if non-admin
if not is_user_proxy_admin(user_obj=user_obj): # if non-admin
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=_user_role,
@ -1416,7 +1424,7 @@ def _is_api_route_allowed(
return True
def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None):
def is_user_proxy_admin(user_obj: LiteLLM_UserTable | None) -> bool:
if user_obj is None:
return False
@ -1426,6 +1434,9 @@ def _is_user_proxy_admin(user_obj: LiteLLM_UserTable | None):
return False
_is_user_proxy_admin: Final = is_user_proxy_admin
def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool:
"""
Return if a user is allowed to access route. Helper function for `allowed_routes_check`.
@ -1723,7 +1734,7 @@ async def _apply_default_budget_to_end_user(
return end_user_obj.model_copy(update=MappingProxyType({"litellm_budget_table": default_budget}))
async def _check_end_user_budget(
async def check_end_user_budget(
end_user_obj: LiteLLM_EndUserTable,
route: str,
) -> None:
@ -1765,6 +1776,9 @@ async def _check_end_user_budget(
)
_check_end_user_budget: Final = check_end_user_budget
#: Columns whose non-null value makes an end-user row restrict something auth enforces. ``blocked``
#: is separate: it restricts when true rather than when merely set.
_RESTRICTED_COLUMNS: Final = ("budget_id", "allowed_model_region", "default_model", "object_permission_id")
@ -2900,12 +2914,12 @@ async def _cache_management_object(
@with_service_target(AUTH_OBJECTS_TARGET)
async def _cache_team_object(
async def cache_team_object(
team_id: str,
team_table: LiteLLM_TeamTableCachedObj,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None,
):
) -> None:
## CACHE REFRESH TIME!
team_table.last_refreshed_at = time.time()
@ -2957,6 +2971,9 @@ async def _cache_team_object(
await _invalidate_usage_cache_entry(usage_cache, alias_key, redis_shared=redis_shared, stale="team alias")
_cache_team_object: Final = cache_team_object
@with_service_target(SPEND_COUNTERS_TARGET)
async def _invalidate_usage_cache_entry(
usage_cache: DualCache | None,
@ -3137,18 +3154,18 @@ async def delete_cache_team_object(
await publish_auth_cache_invalidation(cache_key=key)
async def _cache_key_object(
async def cache_key_object(
hashed_token: str,
user_api_key_obj: UserAPIKeyAuth,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None,
):
) -> None:
key: Final = hashed_token
## CACHE REFRESH TIME
user_api_key_obj.last_refreshed_at = time.time()
cached_key_obj: Final = _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_obj)
cached_key_obj: Final = copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_obj)
await _cache_management_object(
key=key,
value=cached_key_obj,
@ -3158,12 +3175,15 @@ async def _cache_key_object(
)
_cache_key_object: Final = cache_key_object
@with_service_target(AUTH_OBJECTS_TARGET)
async def _delete_cache_key_object(
async def delete_cache_key_object(
hashed_token: str,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None,
):
) -> None:
"""
Evict one key object, best-effort, matching `delete_cache_team_object` and
`delete_cache_key_objects`.
@ -3196,6 +3216,9 @@ async def _delete_cache_key_object(
await publish_auth_cache_invalidation(cache_key=key)
_delete_cache_key_object: Final = delete_cache_key_object
async def delete_cache_key_objects(
hashed_tokens: Sequence[str],
user_api_key_cache: UserApiKeyCache,
@ -3215,7 +3238,7 @@ async def delete_cache_key_objects(
"""
results: Final = await asyncio.gather(
*(
_delete_cache_key_object(
delete_cache_key_object(
hashed_token=hashed_token,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
@ -3339,7 +3362,7 @@ async def _get_team_object_from_user_api_key_cache(
)
# save the team object to cache
await _cache_team_object(
await cache_team_object(
team_id=team_id,
team_table=_response,
user_api_key_cache=user_api_key_cache,
@ -3357,7 +3380,7 @@ async def _get_team_object_from_user_api_key_cache(
@with_service_target(AUTH_OBJECTS_TARGET)
async def _get_team_object_from_cache(
async def get_team_object_from_cache(
key: str,
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span | None,
@ -3370,6 +3393,9 @@ async def _get_team_object_from_cache(
return decoded
_get_team_object_from_cache: Final = get_team_object_from_cache
async def get_team_object(
team_id: str,
prisma_client: PrismaClient | None,
@ -3395,7 +3421,7 @@ async def get_team_object(
key: Final = f"team_id:{team_id}"
if not check_db_only:
cached_team_obj: Final = await _get_team_object_from_cache(
cached_team_obj: Final = await get_team_object_from_cache(
key=key,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
@ -3433,12 +3459,12 @@ async def get_team_object(
@with_service_target(AUTH_OBJECTS_TARGET)
async def _cache_access_object(
async def cache_access_object(
access_group_id: str,
access_group_table: LiteLLM_AccessGroupTable,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None = None,
):
) -> None:
key: Final = f"access_group_id:{access_group_id}"
await user_api_key_cache.async_set_cache(
key=key,
@ -3448,12 +3474,15 @@ async def _cache_access_object(
)
_cache_access_object: Final = cache_access_object
@with_service_target(AUTH_OBJECTS_TARGET)
async def _delete_cache_access_object(
async def delete_cache_access_object(
access_group_id: str,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging | None = None,
):
) -> None:
key: Final = f"access_group_id:{access_group_id}"
user_api_key_cache.delete_cache(key=key)
@ -3463,6 +3492,9 @@ async def _delete_cache_access_object(
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key)
_delete_cache_access_object: Final = delete_cache_access_object
@log_db_metrics
@with_service_target(AUTH_OBJECTS_TARGET)
async def get_access_object(
@ -3510,7 +3542,7 @@ async def get_access_object(
_response: Final = LiteLLM_AccessGroupTable.model_validate(response.dict())
# Save to cache
await _cache_access_object(
await cache_access_object(
access_group_id=access_group_id,
access_group_table=_response,
user_api_key_cache=user_api_key_cache,
@ -3566,7 +3598,7 @@ async def get_team_object_by_alias(
# Check cache first (keyed by alias)
cache_key: Final = f"team_alias:{team_alias}"
cached_team_obj: Final = await _get_team_object_from_cache(
cached_team_obj: Final = await get_team_object_from_cache(
key=cache_key,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
@ -3861,7 +3893,7 @@ class ExperimentalUIJWTToken:
raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}")
async def _fetch_key_object_from_db_with_reconnect(
async def fetch_key_object_from_db_with_reconnect(
hashed_token: str,
prisma_client: PrismaClient,
parent_otel_span: Span | None,
@ -3889,6 +3921,9 @@ async def _fetch_key_object_from_db_with_reconnect(
)
_fetch_key_object_from_db_with_reconnect: Final = fetch_key_object_from_db_with_reconnect
async def _fetch_key_object_from_db_unbounded(
hashed_token: str,
prisma_client: PrismaClient,
@ -4031,13 +4066,13 @@ async def get_key_object(
None if check_db_only else await user_api_key_cache.async_get_cache(key=key, model_type=UserAPIKeyAuth)
)
if user_api_key_auth is not None:
return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth)
return copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth)
if check_cache_only:
raise Exception(f"Key doesn't exist in cache + check_cache_only=True. key={key}.")
# else, check db
_valid_token: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect(
_valid_token: Final[BaseModel | None] = await fetch_key_object_from_db_with_reconnect(
hashed_token=hashed_token,
prisma_client=prisma_client,
parent_otel_span=parent_otel_span,
@ -4079,7 +4114,7 @@ async def get_key_object(
return _response
# save the key object to cache
await _cache_key_object(
await cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=_response,
user_api_key_cache=user_api_key_cache,
@ -4089,7 +4124,7 @@ async def get_key_object(
return _response
def _copy_user_api_key_auth_for_cache(
def copy_user_api_key_auth_for_cache(
user_api_key_obj: UserAPIKeyAuth,
) -> UserAPIKeyAuth:
copied_key_obj: Final = user_api_key_obj.model_copy()
@ -4100,6 +4135,9 @@ def _copy_user_api_key_auth_for_cache(
return copied_key_obj
_copy_user_api_key_auth_for_cache: Final = copy_user_api_key_auth_for_cache
@log_db_metrics
@with_service_target(AUTH_OBJECTS_TARGET)
async def get_object_permission(
@ -4424,7 +4462,7 @@ async def _get_resources_from_access_groups(
return list(set(resources))
async def _get_models_from_access_groups(
async def get_models_from_access_groups(
access_group_ids: Sequence[str],
prisma_client: DatabaseClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
@ -4443,7 +4481,10 @@ async def _get_models_from_access_groups(
)
async def _get_mcp_server_ids_from_access_groups(
_get_models_from_access_groups: Final = get_models_from_access_groups
async def get_mcp_server_ids_from_access_groups(
access_group_ids: list[str],
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
@ -4464,7 +4505,10 @@ async def _get_mcp_server_ids_from_access_groups(
)
async def _get_agent_ids_from_access_groups(
_get_mcp_server_ids_from_access_groups: Final = get_mcp_server_ids_from_access_groups
async def get_agent_ids_from_access_groups(
access_group_ids: list[str],
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
@ -4485,6 +4529,9 @@ async def _get_agent_ids_from_access_groups(
)
_get_agent_ids_from_access_groups: Final = get_agent_ids_from_access_groups
def _resolve_all_team_model_sentinel_for_auth_check(
models: list[str],
llm_router: Router | None,
@ -4499,7 +4546,7 @@ def _resolve_all_team_model_sentinel_for_auth_check(
return list(dict.fromkeys(non_sentinel_models + proxy_models))
def _check_model_access_helper(
def check_model_access_helper(
model: str,
llm_router: Router | None,
models: list[str],
@ -4547,6 +4594,9 @@ def _check_model_access_helper(
return True
_check_model_access_helper: Final = check_model_access_helper
def _can_object_call_model(
model: str | list[str],
llm_router: Router | None,
@ -4621,7 +4671,7 @@ def _can_object_call_model(
## check model access for alias + underlying model - allow if either is in allowed models
for m in potential_models:
if _check_model_access_helper(
if check_model_access_helper(
model=m,
llm_router=llm_router,
models=models,
@ -4647,6 +4697,9 @@ def _can_object_call_model(
)
can_object_call_model: Final = _can_object_call_model
def _resolve_team_alias(
model: str | list[str],
team_model_aliases: Mapping[str, str] | None,
@ -4707,7 +4760,7 @@ async def _check_agent_access_group_model_access(
param="model",
code=status.HTTP_403_FORBIDDEN,
)
_can_object_call_model(
can_object_call_model(
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
@ -4751,7 +4804,7 @@ async def _check_agent_caller_model_access(
prisma_client=prisma_client,
key_model_aliases=caller_key_model_aliases,
)
await _check_team_member_model_access(
await check_team_member_model_access(
model=model,
team_object=caller_team,
valid_token=caller_auth,
@ -5112,7 +5165,7 @@ async def can_key_call_model(
"""
key_models: Final = _resolve_key_models_for_auth_check(valid_token=valid_token)
try:
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=key_models,
@ -5125,12 +5178,12 @@ async def can_key_call_model(
# Fallback: check key's access_group_ids
key_access_group_ids: Final = valid_token.access_group_ids or []
if key_access_group_ids:
models_from_groups: Final = await _get_models_from_access_groups(
models_from_groups: Final = await get_models_from_access_groups(
access_group_ids=key_access_group_ids,
prisma_client=prisma_client,
)
if models_from_groups:
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=models_from_groups,
@ -5200,7 +5253,7 @@ async def can_key_call_resolved_model(
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
if not await key_access_group_grants_model(
model=model,
valid_token=valid_token,
team_object=team_object,
@ -5210,7 +5263,7 @@ async def can_key_call_resolved_model(
raise
if valid_token.user_id is not None and team_object_from_lookup:
await _check_team_member_model_access(
await check_team_member_model_access(
model=model,
team_object=team_object,
valid_token=valid_token,
@ -5265,7 +5318,7 @@ def can_org_access_model(
Returns True if the team can access a specific model.
"""
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=org_object.models if org_object else [],
@ -5289,7 +5342,7 @@ async def can_team_access_model(
2. If not allowed natively, falls back to access_group_ids on the team
"""
try:
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=team_object.models if team_object else [],
@ -5302,12 +5355,12 @@ async def can_team_access_model(
# Fallback: check team's access_group_ids
team_access_group_ids: Final = (team_object.access_group_ids or []) if team_object else []
if team_access_group_ids:
models_from_groups: Final = await _get_models_from_access_groups(
models_from_groups: Final = await get_models_from_access_groups(
access_group_ids=team_access_group_ids,
prisma_client=prisma_client,
)
if models_from_groups:
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])),
@ -5367,7 +5420,7 @@ async def get_authorized_resources_from_key_access_groups(
return list(set(authorized_resources))
async def _key_access_group_grants_model(
async def key_access_group_grants_model(
model: str | list[str],
valid_token: UserAPIKeyAuth | None,
team_object: LiteLLM_TeamTable | None,
@ -5387,7 +5440,7 @@ async def _key_access_group_grants_model(
if not authorized_models:
return False
try:
_can_object_call_model(
can_object_call_model(
model=model,
llm_router=llm_router,
models=authorized_models,
@ -5401,6 +5454,9 @@ async def _key_access_group_grants_model(
return False
_key_access_group_grants_model: Final = key_access_group_grants_model
def can_project_access_model(
model: str | list[str],
project_object: LiteLLM_ProjectTable,
@ -5412,7 +5468,7 @@ def can_project_access_model(
Raises ProxyException if access is denied.
"""
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=project_object.models if project_object else [],
@ -5437,7 +5493,7 @@ def can_customer_access_model(
)
if team_target != name and name in (end_user_object.models or ()):
return
_can_object_call_model(
can_object_call_model(
model=team_target,
llm_router=llm_router,
models=end_user_object.models,
@ -5472,7 +5528,7 @@ async def can_user_call_model(
code=status.HTTP_403_FORBIDDEN,
)
return _can_object_call_model(
return can_object_call_model(
model=model,
llm_router=llm_router,
models=user_object.models,
@ -5764,11 +5820,11 @@ def _apply_budget_exceeded_throttle(valid_token: UserAPIKeyAuth) -> bool:
return True
async def _virtual_key_max_budget_check(
async def virtual_key_max_budget_check(
valid_token: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
user_obj: LiteLLM_UserTable | None = None,
):
) -> None:
"""
Raises:
BudgetExceededError if the token is over it's max budget.
@ -5844,6 +5900,9 @@ async def _virtual_key_max_budget_check(
)
_virtual_key_max_budget_check: Final = virtual_key_max_budget_check
async def _virtual_key_multi_budget_check(
valid_token: UserAPIKeyAuth,
):
@ -5888,11 +5947,11 @@ async def _virtual_key_multi_budget_check(
)
async def _virtual_key_soft_budget_check(
async def virtual_key_soft_budget_check(
valid_token: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
user_obj: LiteLLM_UserTable | None = None,
):
) -> None:
"""
Triggers a budget alert if the token is over it's soft budget.
@ -5927,6 +5986,9 @@ async def _virtual_key_soft_budget_check(
)
_virtual_key_soft_budget_check: Final = virtual_key_soft_budget_check
def _parse_email_list(raw: str | Sequence[object] | None) -> list[str]:
"""Parse emails from a list or comma-separated string."""
if isinstance(raw, list):
@ -5968,11 +6030,11 @@ def _merge_budget_alert_email_configs(
}
async def _virtual_key_max_budget_alert_check(
async def virtual_key_max_budget_alert_check(
valid_token: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
user_obj: LiteLLM_UserTable | None = None,
):
) -> None:
"""
Triggers a budget alert if the token has reached EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE
(default 80%) of its max budget.
@ -6050,6 +6112,9 @@ async def _virtual_key_max_budget_alert_check(
)
_virtual_key_max_budget_alert_check: Final = virtual_key_max_budget_alert_check
TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails"
_TEAM_MEMBER_ALERT_CONFIG_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
@ -6074,7 +6139,7 @@ def _valid_alert_threshold_config(raw_config: object) -> Mapping[str, str | Sequ
)
def _team_member_max_budget_alert_check(
def team_member_max_budget_alert_check(
team_id: str,
team_alias: str | None,
team_metadata: Mapping[str, object] | None,
@ -6108,6 +6173,9 @@ def _team_member_max_budget_alert_check(
asyncio.create_task(proxy_logging_obj.budget_alerts(type="max_budget_alert", user_info=call_info))
_team_member_max_budget_alert_check: Final = team_member_max_budget_alert_check
async def _check_team_member_budget(
team_object: LiteLLM_TeamTable | None,
user_object: LiteLLM_UserTable | None,
@ -6176,7 +6244,7 @@ async def _check_team_member_budget(
if not math.isfinite(team_member_budget):
return
_team_member_max_budget_alert_check(
team_member_max_budget_alert_check(
team_id=team_object.team_id,
team_alias=team_object.team_alias,
team_metadata=team_object.metadata,
@ -6198,7 +6266,7 @@ async def _check_team_member_budget(
)
async def _check_team_member_model_access(
async def check_team_member_model_access(
model: str | list[str],
team_object: LiteLLM_TeamTable,
valid_token: UserAPIKeyAuth,
@ -6238,7 +6306,7 @@ async def _check_team_member_model_access(
member_allowed_models: Final[list[str]] = loaded_membership.litellm_budget_table.allowed_models
try:
_can_object_call_model(
can_object_call_model(
model=model,
llm_router=llm_router,
models=member_allowed_models,
@ -6260,11 +6328,14 @@ async def _check_team_member_model_access(
)
async def _team_max_budget_check(
_check_team_member_model_access: Final = check_team_member_model_access
async def team_max_budget_check(
team_object: LiteLLM_TeamTable | None,
valid_token: UserAPIKeyAuth | None,
proxy_logging_obj: ProxyLogging,
):
) -> None:
"""
Check if the team is over it's max budget.
@ -6310,6 +6381,9 @@ async def _team_max_budget_check(
)
_team_max_budget_check: Final = team_max_budget_check
async def _team_multi_budget_check(
team_object: LiteLLM_TeamTable | None,
):
@ -6588,13 +6662,13 @@ async def delete_cached_project_object(
)
async def _organization_max_budget_check(
async def organization_max_budget_check(
valid_token: UserAPIKeyAuth | None,
team_object: LiteLLM_TeamTable | None,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
):
) -> None:
"""
Check if the organization is over its max budget.
@ -6687,6 +6761,9 @@ async def _organization_max_budget_check(
)
_organization_max_budget_check: Final = organization_max_budget_check
async def _tag_max_budget_check(
request_body: dict,
prisma_client: PrismaClient | None,

View file

@ -133,7 +133,7 @@ def get_user_organization_info(
return _user_organizations, _user_organization_role_mapping
def _user_is_org_admin(
def user_is_org_admin(
request_data: dict,
user_object: LiteLLM_UserTable | None = None,
) -> bool:
@ -173,6 +173,9 @@ def _user_is_org_admin(
return all(org_id in admin_org_ids for org_id in candidate_org_ids)
_user_is_org_admin: Final = user_is_org_admin
TEAM_ORG_CONTEXT_ROUTES: Final = frozenset({"/team/update"})
# The RESTful update route carries the team id in the path. Match on the route
# template so the sibling /team/<verb> routes (which share the single-segment

View file

@ -20,8 +20,9 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import (
_get_request_ip_address,
from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports
_get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
get_request_ip_address,
is_invalid_virtual_key_error,
mark_invalid_virtual_key_error,
normalize_request_route,
@ -125,7 +126,7 @@ def _identity_log_suffix(resolved_identity: UserAPIKeyAuth | None) -> str:
class UserAPIKeyAuthExceptionHandler:
@staticmethod
async def _handle_authentication_error(
async def handle_authentication_error(
e: Exception,
request: Request,
request_data: dict[str, object],
@ -180,7 +181,7 @@ class UserAPIKeyAuthExceptionHandler:
)
else:
# raise the exception to the caller
requester_ip: Final = _get_request_ip_address(
requester_ip: Final = get_request_ip_address(
request=request,
use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True,
)
@ -258,3 +259,5 @@ class UserAPIKeyAuthExceptionHandler:
extra=log_extra,
)
raise final_exception
_handle_authentication_error = handle_authentication_error

View file

@ -77,7 +77,7 @@ def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual
return marked_exception
def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None:
def get_request_ip_address(request: Request, use_x_forwarded_for: bool | None = False) -> str | None:
client_ip = None
if use_x_forwarded_for is True and "x-forwarded-for" in request.headers:
client_ip = request.headers["x-forwarded-for"]
@ -89,6 +89,9 @@ def _get_request_ip_address(request: Request, use_x_forwarded_for: bool | None =
return client_ip
_get_request_ip_address: Final = get_request_ip_address
def _check_valid_ip(
allowed_ips: list[str] | None,
request: Request,
@ -101,7 +104,7 @@ def _check_valid_ip(
return True, None
# if general_settings.get("use_x_forwarded_for") is True then use x-forwarded-for
client_ip: Final = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for)
client_ip: Final = get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for)
# Check if IP address is allowed
if client_ip not in allowed_ips:
@ -1884,14 +1887,14 @@ def _extract_models_from_managed_resource_id(
try:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
is_base64_encoded_unified_file_id,
)
_append_model_candidates(candidates=candidates, value=decode_model_from_file_id(resource_id))
unified_file_id: Final = _is_base64_encoded_unified_file_id(resource_id)
unified_file_id: Final = is_base64_encoded_unified_file_id(resource_id)
if unified_file_id:
_append_model_candidates(
candidates=candidates,
@ -2180,7 +2183,7 @@ def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str
def _model_from_bedrock_route(route: str) -> str | None:
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_extract_model_from_bedrock_endpoint,
extract_model_from_bedrock_endpoint,
is_bedrock_count_tokens_endpoint,
)
@ -2188,7 +2191,7 @@ def _model_from_bedrock_route(route: str) -> str | None:
if is_bedrock_count_tokens_endpoint(bedrock_endpoint):
return None
try:
return _extract_model_from_bedrock_endpoint(bedrock_endpoint)
return extract_model_from_bedrock_endpoint(bedrock_endpoint)
except ValueError:
return None

View file

@ -38,10 +38,10 @@ async def resolve_owned_read_scope(
def can_read_team_logs(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable) -> bool:
from litellm.proxy.management.teams.authz import is_team_admin
from litellm.proxy.management_endpoints.common_utils import (
_team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy
team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy
)
return is_team_admin(user_api_key_dict=auth, team_obj=team) or _team_member_has_permission(
return is_team_admin(user_api_key_dict=auth, team_obj=team) or team_member_has_permission(
user_api_key_dict=auth,
team_obj=team,
permission=KeyManagementRoutes.SPEND_LOGS.value,

View file

@ -39,8 +39,9 @@ from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
_is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent
from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports
_is_model_cost_zero, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent
)
from litellm.router import Router
from litellm.types.llms.base import LiteLLMBaseModel
@ -109,7 +110,7 @@ async def is_token_within_budget_for_model(*, model: str, valid_token: UserAPIKe
A zero-cost fallback target is always allowed: refusing it would deny a request on spend some
other model accrued, which is the same reasoning behind the auth-time bypass.
"""
if _is_model_cost_zero(model=model, llm_router=llm_router):
if is_model_cost_zero(model=model, llm_router=llm_router):
return True
key_budget: Final = valid_token.max_budget

View file

@ -16,7 +16,10 @@ from fastapi import Request
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy.auth.auth_utils import _get_request_ip_address
from litellm.proxy.auth.auth_utils import ( # noqa: F401 # legacy module exports
_get_request_ip_address, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
get_request_ip_address,
)
# One-shot warning so operators upgrading from the prior "always trust X-Forwarded-*"
# behaviour see an actionable message in their logs the first time it triggers.
@ -368,4 +371,4 @@ class IPAddressUtils:
return client_ip
case _HopCountUnset():
pass
return _get_request_ip_address(request, use_x_forwarded_for=use_xff)
return get_request_ip_address(request, use_x_forwarded_for=use_xff)

View file

@ -8,10 +8,13 @@ from pydantic import BaseModel
from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import (
_cache_key_object,
_copy_user_api_key_auth_for_cache,
_fetch_key_object_from_db_with_reconnect,
from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports
_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_copy_user_api_key_auth_for_cache, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_fetch_key_object_from_db_with_reconnect, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
cache_key_object,
copy_user_api_key_auth_for_cache,
fetch_key_object_from_db_with_reconnect,
get_object_permission,
)
from litellm.proxy.auth.auth_method import AuthMethod
@ -81,7 +84,7 @@ class IdentityStore:
network: NetworkContext | None = None,
) -> Principal:
key: Final = await self._resolve_key(hashed_token)
return self._principal_from_key(
return self.principal_from_key(
key,
auth_method=auth_method,
network=network,
@ -108,12 +111,12 @@ class IdentityStore:
cached: Final = await self._cache.async_get_cache(key=hashed_token, model_type=UserAPIKeyAuth)
if cached is not None:
return _copy_user_api_key_auth_for_cache(user_api_key_obj=cached)
return copy_user_api_key_auth_for_cache(user_api_key_obj=cached)
if self._check_cache_only:
raise KeyNotInCacheError(hashed_token)
from_db: Final[BaseModel | None] = await _fetch_key_object_from_db_with_reconnect(
from_db: Final[BaseModel | None] = await fetch_key_object_from_db_with_reconnect(
hashed_token=hashed_token,
prisma_client=self._prisma,
parent_otel_span=self._parent_otel_span,
@ -140,7 +143,7 @@ class IdentityStore:
e,
)
await _cache_key_object(
await cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=key,
user_api_key_cache=self._cache,
@ -149,7 +152,7 @@ class IdentityStore:
return key
@staticmethod
def _principal_from_key(
def principal_from_key(
key: UserAPIKeyAuth,
*,
auth_method: AuthMethod,
@ -192,3 +195,5 @@ class IdentityStore:
network=network or NetworkContext(),
source_key=key,
)
_principal_from_key = principal_from_key

View file

@ -15,7 +15,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from .auth_checks_organization import _user_is_org_admin
from .auth_checks_organization import ( # noqa: F401 # legacy module exports
_user_is_org_admin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
user_is_org_admin,
)
# Management write routes denied to PROXY_ADMIN_VIEW_ONLY. Adding a new write
# endpoint to a management router REQUIRES adding it here too — the surrounding
@ -128,7 +131,7 @@ class RouteChecks:
allowed_route in _AUTH_ENFORCED_PASS_THROUGH_ROUTE_GROUPS
and RouteChecks.is_auth_enforced_pass_through_route(
route=route,
method=RouteChecks._get_request_method(request=request),
method=RouteChecks.get_request_method(request=request),
)
):
if RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token):
@ -141,7 +144,7 @@ class RouteChecks:
# For llm_api_routes, also check registered pass-through endpoints
################################################
if allowed_route == "llm_api_routes":
if route == "/auto_router/session" and RouteChecks._get_request_method(request) == "GET":
if route == "/auto_router/session" and RouteChecks.get_request_method(request) == "GET":
return True
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
@ -151,7 +154,7 @@ class RouteChecks:
if InitPassThroughEndpointHelpers.is_registered_pass_through_route(route=route):
if RouteChecks.is_auth_enforced_pass_through_route(
route=route,
method=RouteChecks._get_request_method(request=request),
method=RouteChecks.get_request_method(request=request),
):
if RouteChecks.check_passthrough_route_access(
route=route, user_api_key_dict=valid_token
@ -277,7 +280,7 @@ class RouteChecks:
if RouteChecks.is_auth_enforced_pass_through_route(
route=route,
method=RouteChecks._get_request_method(request=request),
method=RouteChecks.get_request_method(request=request),
):
RouteChecks._require_auth_pass_through_access(
route=route,
@ -329,7 +332,7 @@ class RouteChecks:
elif (
_user_role == LitellmUserRoles.INTERNAL_USER.value
and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.internal_user_routes.value)
or _user_is_org_admin(request_data=request_data, user_object=user_obj)
or user_is_org_admin(request_data=request_data, user_object=user_obj)
and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.org_admin_allowed_routes.value)
or _user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
and RouteChecks.check_route_access(
@ -423,7 +426,7 @@ class RouteChecks:
if RouteChecks._route_matches_pattern(route=route, pattern=openai_route):
return True
# Check for wildcard patterns like "/containers/*"
if RouteChecks._is_wildcard_pattern(pattern=openai_route):
if RouteChecks.is_wildcard_pattern(pattern=openai_route):
if RouteChecks.route_matches_wildcard_pattern(route=route, pattern=openai_route):
return True
@ -549,12 +552,14 @@ class RouteChecks:
return False
@staticmethod
def _is_wildcard_pattern(pattern: str) -> bool:
def is_wildcard_pattern(pattern: str) -> bool:
"""
Check if pattern is a wildcard pattern
"""
return pattern.endswith("*")
_is_wildcard_pattern = is_wildcard_pattern
@staticmethod
def route_matches_wildcard_pattern(route: str, pattern: str) -> bool:
"""
@ -635,7 +640,7 @@ class RouteChecks:
if any(
RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route)
for allowed_route in allowed_routes
if RouteChecks._is_wildcard_pattern(pattern=allowed_route)
if RouteChecks.is_wildcard_pattern(pattern=allowed_route)
):
return True
@ -653,7 +658,7 @@ class RouteChecks:
return False
@staticmethod
def _get_request_method(request: Request | None) -> str | None:
def get_request_method(request: Request | None) -> str | None:
if request is None:
return None
@ -666,6 +671,8 @@ class RouteChecks:
return method.upper()
_get_request_method = get_request_method
@staticmethod
def is_auth_enforced_pass_through_route(route: str, method: str | None = None) -> bool:
"""
@ -829,7 +836,7 @@ class RouteChecks:
return False
@staticmethod
def _is_assistants_api_request(request: Request) -> bool:
def is_assistants_api_request(request: Request) -> bool:
"""
Returns True if `thread` or `assistant` is in the request path
@ -847,6 +854,8 @@ class RouteChecks:
return True
return False
_is_assistants_api_request = is_assistants_api_request
@staticmethod
def is_generate_content_route(route: str) -> bool:
"""

View file

@ -41,22 +41,26 @@ from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.proxy._types import *
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_from_headers
from litellm.proxy.auth.auth_checks import (
from litellm.proxy.auth.auth_checks import ( # noqa: F401 # legacy module exports
ExperimentalUIJWTToken,
TeamNotFoundError,
_cache_key_object,
_can_object_call_model,
_check_end_user_budget,
_delete_cache_key_object,
_get_user_role,
_is_model_cost_zero,
_is_user_proxy_admin,
_team_member_max_budget_alert_check,
_virtual_key_max_budget_alert_check,
_virtual_key_max_budget_check,
_virtual_key_soft_budget_check,
_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_can_object_call_model, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_check_end_user_budget, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_delete_cache_key_object, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_get_user_role, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_is_model_cost_zero, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_is_user_proxy_admin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_team_member_max_budget_alert_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_virtual_key_max_budget_alert_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_virtual_key_max_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_virtual_key_soft_budget_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
cache_key_object,
can_key_call_model,
can_object_call_model,
check_end_user_budget,
common_checks,
delete_cache_key_object,
get_end_user_object,
get_jwt_key_mapping_object,
get_key_end_user_budget_id,
@ -66,11 +70,18 @@ from litellm.proxy.auth.auth_checks import (
get_team_membership,
get_team_object,
get_user_object,
get_user_role,
is_model_cost_zero,
is_user_proxy_admin,
is_valid_fallback_model,
jwt_key_mapping_cache_key,
key_model_aliases_for_auth_check,
resolve_and_validate_end_user_id,
resolve_default_end_user_budget,
team_member_max_budget_alert_check,
virtual_key_max_budget_alert_check,
virtual_key_max_budget_check,
virtual_key_soft_budget_check,
)
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
from litellm.proxy.auth.auth_method import AuthMethod
@ -113,18 +124,25 @@ from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.team_grants import team_grants
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
_safe_get_request_query_params,
_safe_set_request_parsed_body,
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_safe_set_request_parsed_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
is_opaque_audio_pass_through_request,
populate_request_with_path_params,
read_raw_json_body,
read_request_body,
rewrite_request_model,
safe_get_request_headers,
safe_get_request_query_params,
safe_set_request_parsed_body,
)
from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.common_utils.realtime_utils import ( # noqa: F401 # legacy module exports
_realtime_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
realtime_request_body,
)
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
end_user_cache_key,
@ -222,8 +240,8 @@ def _get_model_from_request_context(
return get_model_from_request(
request_data=request_data,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
request_headers=safe_get_request_headers(request=request),
request_query_params=safe_get_request_query_params(request=request),
llm_router=llm_router,
request=request,
team_id=team_id,
@ -477,7 +495,7 @@ async def _check_key_model_budget_with_fallback(
llm_router=llm_router,
)
if valid_token.team_models:
_can_object_call_model(
can_object_call_model(
model=fallback_model,
llm_router=llm_router,
models=valid_token.team_models,
@ -489,9 +507,9 @@ async def _check_key_model_budget_with_fallback(
except ProxyException:
raise e
request_data["model"] = fallback_model
_safe_set_request_parsed_body(request=request, parsed_body=request_data)
request._json = request_data
request._body = orjson.dumps(request_data)
safe_set_request_parsed_body(request=request, parsed_body=request_data)
request._json = request_data # pyright: ignore[reportPrivateUsage] # Starlette JSON cache
request._body = orjson.dumps(request_data) # pyright: ignore[reportPrivateUsage] # Starlette body cache
path_params: Final = request.scope.get("path_params")
if isinstance(path_params, dict) and "model" in path_params:
path_params["model"] = fallback_model
@ -591,9 +609,9 @@ def _should_route_jwt_to_oauth2_override(token: str, jwt_handler: JWTHandler) ->
return False
def _get_bearer_token(
def get_bearer_token(
api_key: str,
):
) -> str:
if api_key.startswith("Bearer "): # ensure Bearer token passed in
api_key = api_key.replace("Bearer ", "") # extract the token
elif api_key.startswith("Basic "):
@ -619,6 +637,9 @@ def _get_bearer_token(
return api_key
_get_bearer_token: Final = get_bearer_token
def _apply_budget_limits_to_end_user_params(
end_user_params: dict,
budget_info: LiteLLM_BudgetTable,
@ -675,10 +696,10 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
synthetic_scope[key] = ws_scope[key]
request: Final = Request(scope=synthetic_scope)
request._url = websocket.url
request._url = websocket.url # pyright: ignore[reportPrivateUsage] # Starlette WebSocket URL storage
async def return_body():
return _realtime_request_body(model)
return realtime_request_body(model)
request.body = return_body
@ -739,7 +760,7 @@ def update_valid_token_with_end_user_params(valid_token: UserAPIKeyAuth, end_use
_global_spend_coordinator: Final = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEND]")
async def _fetch_global_spend_with_event_coordination(
async def fetch_global_spend_with_event_coordination(
cache_key: str,
user_api_key_cache: UserApiKeyCache,
prisma_client: PrismaClient,
@ -766,6 +787,9 @@ async def _fetch_global_spend_with_event_coordination(
)
_fetch_global_spend_with_event_coordination: Final = fetch_global_spend_with_event_coordination
async def get_global_proxy_spend(
litellm_proxy_admin_name: str,
user_api_key_cache: UserApiKeyCache,
@ -777,10 +801,12 @@ async def get_global_proxy_spend(
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
# Use event-driven coordination to prevent cache stampede
cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
global_proxy_spend = ( # rebind-ok: pre-existing rebinding on a rename-only line
await fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
)
if global_proxy_spend is not None:
user_info: Final = CallInfo(
@ -824,7 +850,7 @@ def get_api_key(
"""
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_query_params,
safe_get_request_query_params,
)
api_key = api_key
@ -834,7 +860,7 @@ def get_api_key(
api_key = _get_bearer_token_or_received_api_key(custom_litellm_key_header)
elif isinstance(api_key, str) and len(api_key) > 0:
passed_in_key = api_key
api_key = _get_bearer_token(api_key=api_key)
api_key = get_bearer_token(api_key=api_key)
elif isinstance(azure_api_key_header, str):
passed_in_key = azure_api_key_header
api_key = azure_api_key_header
@ -850,9 +876,9 @@ def get_api_key(
elif (
RouteChecks.is_generate_content_route(route=route)
and request is not None
and _safe_get_request_query_params(request).get("key")
and safe_get_request_query_params(request).get("key")
):
google_auth_key: Final[str] = _safe_get_request_query_params(request).get("key") or ""
google_auth_key: Final[str] = safe_get_request_query_params(request).get("key") or ""
passed_in_key = google_auth_key
api_key = google_auth_key
elif pass_through_endpoints is not None:
@ -1388,7 +1414,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
return
parent_otel_span: Final = open_telemetry_logger.create_litellm_proxy_request_started_span(
start_time=start_time,
headers=_safe_get_request_headers(request),
headers=safe_get_request_headers(request),
)
# Under V2 the FastAPI instrumentor stamps http.route / url.path on the server
# span; only the legacy logger needs these set explicitly.
@ -1413,12 +1439,12 @@ async def _read_request_body_deferring_parse_failure(
"""
if is_opaque_audio_pass_through_request(
route=get_request_route(request=request),
content_type=_safe_get_request_headers(request=request).get("content-type", ""),
content_type=safe_get_request_headers(request=request).get("content-type", ""),
):
_safe_set_request_parsed_body(request=request, parsed_body={})
safe_set_request_parsed_body(request=request, parsed_body={})
return {}, None
try:
parsed_body: Final = await _read_request_body(request=request)
parsed_body: Final = await read_request_body(request=request)
except ProxyException as parse_exception:
return {}, parse_exception
return populate_request_with_path_params(request_data=parsed_body, request=request), None
@ -1481,7 +1507,7 @@ async def _refresh_session_token_grants(
{
**valid_token.model_dump(exclude_none=True),
**team_grants(team_object, team_membership, user_object.user_id),
"user_role": _get_user_role(user_object),
"user_role": get_user_role(user_object),
"models": () if team_object is not None else user_models(user_object),
}
)
@ -1515,7 +1541,7 @@ async def _resolve_object_permission_for_unresolvable_team(
)
async def _user_api_key_auth_builder(
async def user_api_key_auth_builder(
request: Request,
api_key: str,
azure_api_key_header: str,
@ -1765,8 +1791,8 @@ async def _user_api_key_auth_builder(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
parent_otel_span=parent_otel_span,
request_headers=_safe_get_request_headers(request),
request_method=RouteChecks._get_request_method(request=request),
request_headers=safe_get_request_headers(request),
request_method=RouteChecks.get_request_method(request=request),
)
is_proxy_admin: Final = result["is_proxy_admin"]
@ -1862,9 +1888,11 @@ async def _user_api_key_auth_builder(
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
from litellm.proxy.auth.auth_checks import is_model_cost_zero
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
skip_budget_checks = ( # rebind-ok: pre-existing rebinding on a rename-only line
is_model_cost_zero(model=model, llm_router=llm_router)
)
if skip_budget_checks:
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
@ -1950,7 +1978,7 @@ async def _user_api_key_auth_builder(
_end_user_object = None
end_user_params: Final = {}
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request))
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request))
end_user_id = await resolve_and_validate_end_user_id(
raw_end_user_id=raw_end_user_id,
prisma_client=prisma_client,
@ -2070,7 +2098,7 @@ async def _user_api_key_auth_builder(
if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None:
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
if expiry_time < current_time:
await _delete_cache_key_object(
await delete_cache_key_object(
hashed_token=hash_token(api_key),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
@ -2151,7 +2179,7 @@ async def _user_api_key_auth_builder(
start_time=start_time,
)
asyncio.create_task(
_cache_key_object(
cache_key_object(
hashed_token=hash_token(master_key),
user_api_key_obj=_user_api_key_obj,
user_api_key_cache=user_api_key_cache,
@ -2248,7 +2276,7 @@ async def _user_api_key_auth_builder(
_end_user_object=_end_user_object,
)
except Exception as e:
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
return await UserAPIKeyAuthExceptionHandler.handle_authentication_error(
e=e,
request=request,
request_data=request_data,
@ -2259,6 +2287,9 @@ async def _user_api_key_auth_builder(
)
_user_api_key_auth_builder: Final = user_api_key_auth_builder
async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of existing shared authorization checks
request: Request,
request_data: dict[str, object],
@ -2303,7 +2334,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
## base case ## key is disabled
if valid_token.blocked is True:
raise Exception("Key is blocked. Update via `/key/unblock` if you're an admin.")
await _enforce_key_and_fallback_model_access(
await enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route=route,
@ -2359,9 +2390,11 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
)
skip_budget_checks = False
if model is not None and llm_router is not None:
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
from litellm.proxy.auth.auth_checks import is_model_cost_zero
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
skip_budget_checks = is_model_cost_zero( # rebind-ok: pre-existing rebinding on a rename-only line
model=model, llm_router=llm_router
)
if skip_budget_checks:
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
@ -2412,7 +2445,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
if team_member_spend >= team_member_budget:
# common_checks sends this alert on requests that get past here, so only the
# request rejected here sends it from the builder.
_team_member_max_budget_alert_check(
team_member_max_budget_alert_check(
team_id=_team_id,
team_alias=valid_token.team_alias,
team_metadata=valid_token.team_metadata,
@ -2461,7 +2494,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
# Check 4. Max Budget Alert Check (runs before budget enforcement
# so multi-threshold 100% alerts fire on the request that crosses
# max_budget, before BudgetExceededError is raised below)
await _virtual_key_max_budget_alert_check(
await virtual_key_max_budget_alert_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
@ -2469,14 +2502,14 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
# Check 5. Token Spend is under budget
if RouteChecks.is_llm_api_route(route=route):
await _virtual_key_max_budget_check(
await virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 6. Soft Budget Check
await _virtual_key_soft_budget_check(
await virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
@ -2613,7 +2646,7 @@ async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of e
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
global_proxy_spend = await fetch_global_spend_with_event_coordination( # rebind-ok: pre-existing rebinding on a rename-only line
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
@ -2803,7 +2836,7 @@ def is_no_auth_dev_mode(master_key: str | None, general_settings: Mapping[str, o
@tracer.wrap()
async def _run_centralized_common_checks(
async def run_centralized_common_checks(
user_api_key_auth_obj: UserAPIKeyAuth,
request: Request,
request_data: dict[str, object],
@ -2884,7 +2917,7 @@ async def _run_centralized_common_checks(
key_end_user_budget_id: Final = get_key_end_user_budget_id(user_api_key_auth_obj.metadata)
end_user_id = user_api_key_auth_obj.end_user_id
if end_user_id is None:
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request))
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request))
end_user_id = await resolve_and_validate_end_user_id(
raw_end_user_id=raw_end_user_id,
prisma_client=prisma_client,
@ -3164,6 +3197,9 @@ async def _run_centralized_common_checks(
release_spend_counter_batch()
_run_centralized_common_checks: Final = run_centralized_common_checks
async def _noop_none() -> None:
"""Sentinel coroutine for asyncio.gather when a fetch is unnecessary
(e.g. token has no team_id). Keeps the result tuple positional."""
@ -3271,7 +3307,7 @@ def _should_skip_budget_checks(
team_id=team_id,
)
if model is not None and llm_router is not None:
return _is_model_cost_zero(model=model, llm_router=llm_router)
return is_model_cost_zero(model=model, llm_router=llm_router)
return False
@ -3290,7 +3326,7 @@ def _resolve_request_principal(request: Request, valid_token: UserAPIKeyAuth) ->
TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs),
)
auth_method: Final = AuthMethod.BEARER_JWT if valid_token.jwt_claims else AuthMethod.API_KEY
return IdentityStore._principal_from_key(
return IdentityStore.principal_from_key(
valid_token,
auth_method=auth_method,
network=network,
@ -3362,7 +3398,7 @@ async def _authorize_authenticated_request(
billable=request_data.get("method")
in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
)
await _run_centralized_common_checks(
await run_centralized_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request=request,
request_data=authorized_data,
@ -3370,7 +3406,7 @@ async def _authorize_authenticated_request(
force_virtual_key_checks=force_virtual_key_checks,
)
except Exception as e:
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
return await UserAPIKeyAuthExceptionHandler.handle_authentication_error(
e=e,
request=request,
request_data=request_data,
@ -3393,7 +3429,7 @@ async def _authorize_authenticated_request(
user_api_key_cache,
)
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, _safe_get_request_headers(request))
raw_end_user_id: Final = get_end_user_id_from_request_body(request_data, safe_get_request_headers(request))
if raw_end_user_id is not None:
resolved_end_user_id: Final = await resolve_and_validate_end_user_id(
raw_end_user_id=raw_end_user_id,
@ -3475,7 +3511,7 @@ def _seed_request_destinations(user_api_key_dict: UserAPIKeyAuth, request: Reque
set_request_destinations(
deliverable_destinations(
resolve_tenant_otel_destinations(user_api_key_dict, _safe_get_request_headers(request)),
resolve_tenant_otel_destinations(user_api_key_dict, safe_get_request_headers(request)),
fan_out_provider(),
)
)
@ -3518,7 +3554,7 @@ async def user_api_key_auth(
spend_counter_batch_scope(_spend_counter_redis_cache()),
):
try:
user_api_key_auth_obj: Final = await _user_api_key_auth_builder(
user_api_key_auth_obj: Final = await user_api_key_auth_builder(
request=request,
api_key=api_key,
azure_api_key_header=azure_api_key_header,
@ -3537,7 +3573,7 @@ async def user_api_key_auth(
raise
user_api_key_auth_obj.budget_reservation = None
user_api_key_auth_obj.agent_caller = agent_caller_from_headers(
_safe_get_request_headers(request), user_api_key_auth_obj
safe_get_request_headers(request), user_api_key_auth_obj
)
_seed_request_destinations(user_api_key_auth_obj, request)
@ -3611,7 +3647,7 @@ async def _return_user_api_key_auth_obj(
)
)
retrieved_user_role: Final = user_role or _get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER
retrieved_user_role: Final = user_role or get_user_role(user_obj=user_obj) or LitellmUserRoles.INTERNAL_USER
user_api_key_kwargs: Final = {
"api_key": api_key,
@ -3628,7 +3664,7 @@ async def _return_user_api_key_auth_obj(
user_max_budget=getattr(user_obj, "max_budget", None),
user_model_max_budget=getattr(user_obj, "model_max_budget", None),
)
if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj):
if user_obj is not None and is_user_proxy_admin(user_obj=user_obj):
user_api_key_kwargs.update(
user_role=LitellmUserRoles.PROXY_ADMIN,
)
@ -3659,7 +3695,7 @@ def get_api_key_from_custom_header(request: Request, custom_litellm_key_header_n
)
custom_api_key: Final = _headers.get(custom_litellm_key_header_name)
if custom_api_key:
api_key = _get_bearer_token(api_key=custom_api_key)
api_key = get_bearer_token(api_key=custom_api_key) # rebind-ok: pre-existing rebinding on a rename-only line
verbose_proxy_logger.debug(
"Found custom API key using header: %s, setting api_key=%s",
custom_litellm_key_header_name,
@ -3756,7 +3792,7 @@ async def _lookup_end_user_and_apply_budget(
return valid_token, end_user_object
async def _enforce_key_and_fallback_model_access(
async def enforce_key_and_fallback_model_access(
*,
valid_token: UserAPIKeyAuth,
request_data: dict,
@ -3816,6 +3852,9 @@ async def _enforce_key_and_fallback_model_access(
)
_enforce_key_and_fallback_model_access: Final = enforce_key_and_fallback_model_access
async def _run_post_custom_auth_checks(
valid_token: UserAPIKeyAuth,
request: Request,
@ -3849,7 +3888,7 @@ async def _run_post_custom_auth_checks(
# custom_auth_run_common_checks is set. Enforce it here on that path
# so an over-budget end user can't keep making requests.
if end_user_object is not None and not general_settings.get("custom_auth_run_common_checks", False):
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
await check_end_user_budget(end_user_obj=end_user_object, route=route)
# 2. Check token expiry
if valid_token.expires is not None:
@ -3869,7 +3908,7 @@ async def _run_post_custom_auth_checks(
)
if general_settings.get("custom_auth_run_common_checks", False):
await _enforce_key_and_fallback_model_access(
await enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route=route,
@ -3892,7 +3931,7 @@ async def _run_post_custom_auth_checks(
# every budget check for these; this path did not, so the same request could
# be refused under custom auth and served under the other two.
skip_budget_checks: Final = (
_is_model_cost_zero(model=current_model, llm_router=llm_router)
is_model_cost_zero(model=current_model, llm_router=llm_router)
if current_model is not None and llm_router is not None
else False
)

View file

@ -36,14 +36,17 @@ from litellm.proxy.common_request_processing import (
request_litellm_call_id,
)
from litellm.proxy.common_utils.callback_utils import sanitize_openai_provider_metadata
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports
BATCH_CREATE_HIDDEN_PARAM,
_is_base64_encoded_unified_file_id,
_is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
add_deployment_model_info,
add_internal_model_credentials,
apply_team_provider_credentials,
@ -59,6 +62,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
get_model_id_from_unified_batch_id,
get_models_from_unified_file_id,
get_original_file_id,
is_base64_encoded_unified_file_id,
is_litellm_executed_batch,
prepare_data_with_credentials,
update_batch_in_database,
@ -269,7 +273,7 @@ async def create_batch(
data: dict = {}
try:
data = await _read_request_body(request=request)
data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line
verbose_proxy_logger.debug(
"Request received by LiteLLM:\n%s",
json.dumps(data, indent=4),
@ -341,7 +345,9 @@ async def create_batch(
model_from_file_id = None
if input_file_id:
model_from_file_id = decode_model_from_file_id(input_file_id)
unified_file_id = _is_base64_encoded_unified_file_id(input_file_id)
unified_file_id = ( # rebind-ok: pre-existing rebinding on a rename-only line
is_base64_encoded_unified_file_id(input_file_id)
)
# SCENARIO 1: File ID is encoded with model info
if model_from_file_id is not None and input_file_id:
@ -587,7 +593,7 @@ async def retrieve_batch(
)
data = cast(dict, _retrieve_batch_request)
unified_batch_id: Final = _is_base64_encoded_unified_file_id(batch_id)
unified_batch_id: Final = is_base64_encoded_unified_file_id(batch_id)
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
(
@ -891,7 +897,7 @@ async def list_batches(
)
# Include original request and headers in the data
data = await _read_request_body(request=request)
data = await read_request_body(request=request) # rebind-ok: pre-existing rebinding on a rename-only line
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
(
data,
@ -1084,7 +1090,7 @@ async def cancel_batch(
)
data = cast(dict, _cancel_batch_request)
unified_batch_id: Final = _is_base64_encoded_unified_file_id(batch_id)
unified_batch_id: Final = is_base64_encoded_unified_file_id(batch_id)
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
(

View file

@ -28,6 +28,7 @@ from fastapi import HTTPException, Request, status
from fastapi.responses import JSONResponse, Response, StreamingResponse
from pydantic import BaseModel, TypeAdapter, ValidationError
from starlette.types import Receive, Scope, Send
from typing_extensions import Never
import litellm
from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger
@ -118,7 +119,11 @@ from litellm.proxy.native_compaction import with_proxy_compaction_executor
from litellm.proxy.route_llm_request import (
route_request,
)
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
from litellm.proxy.utils import ( # noqa: F401 # legacy module exports
ProxyLogging,
_check_and_merge_model_level_guardrails, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
check_and_merge_model_level_guardrails,
)
from litellm.router import Router
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
from litellm.router_utils.common_utils import resolve_model_group_alias
@ -290,13 +295,16 @@ def resolve_litellm_call_id(client_call_id: str | None) -> str:
return str(uuid.uuid4())
def _should_return_raw_model_name(request_data: dict[str, object]) -> bool:
def should_return_raw_model_name(request_data: dict[str, object]) -> bool:
return any(
isinstance(metadata, dict) and metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY) is True
for metadata in (request_data.get("metadata"), request_data.get("litellm_metadata"))
)
_should_return_raw_model_name: Final = should_return_raw_model_name
def _apply_client_disconnect_metadata(target_metadata: dict[str, object] | None) -> None:
if target_metadata is None:
return
@ -1302,7 +1310,7 @@ async def open_sse_before_first_byte(
)
def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool:
def is_azure_model_router_request(model: str, hidden_params: Mapping[str, object] | None = None) -> bool:
"""
Check if a request went down the Azure Model Router route.
@ -1323,6 +1331,9 @@ def _is_azure_model_router_request(model: str, hidden_params: Mapping[str, objec
return AzureFoundryModelInfo.is_model_router_call(model=model, hidden_params=hidden_params)
_is_azure_model_router_request: Final = is_azure_model_router_request
def _override_openai_response_model(
*,
response_obj: object,
@ -1388,7 +1399,7 @@ def _override_openai_response_model(
return
# Check if this is an Azure Model Router request - if so, preserve the actual model used
if _is_azure_model_router_request(requested_model, hidden_params):
if is_azure_model_router_request(requested_model, hidden_params):
verbose_proxy_logger.debug(
"%s: Azure Model Router detected - preserving actual model used from response instead of overriding to router model.",
log_context,
@ -2074,9 +2085,9 @@ class ProxyBaseLLMRequestProcessing:
# Store queue time in metadata after add_litellm_data_to_request to ensure it's preserved
if queue_time_seconds is not None:
from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name
from litellm.proxy.litellm_pre_call_utils import get_metadata_variable_name
_metadata_variable_name: Final = _get_metadata_variable_name(request)
_metadata_variable_name: Final = get_metadata_variable_name(request)
if _metadata_variable_name not in self.data:
self.data[_metadata_variable_name] = {}
if not isinstance(self.data[_metadata_variable_name], dict):
@ -2200,11 +2211,11 @@ class ProxyBaseLLMRequestProcessing:
merged_for_requested: Final = (
self.data
if rate_limited_model is None
else _check_and_merge_model_level_guardrails(
else check_and_merge_model_level_guardrails(
data=self.data, llm_router=llm_router, trust_client_model_info=False, model_alias=rate_limited_model
)
)
self.data = _check_and_merge_model_level_guardrails(
self.data = check_and_merge_model_level_guardrails(
data=merged_for_requested,
llm_router=llm_router,
trust_client_model_info=False,
@ -2803,7 +2814,7 @@ class ProxyBaseLLMRequestProcessing:
proxy_logging_obj=proxy_logging_obj,
request=request,
restamp_model=(
None if _should_return_raw_model_name(self.data) else requested_model_from_client
None if should_return_raw_model_name(self.data) else requested_model_from_client
),
)
selected_data_generator = wrap_sse_stream_with_keepalive_pings(
@ -2942,7 +2953,7 @@ class ProxyBaseLLMRequestProcessing:
response_obj=response,
requested_model=requested_model_from_client,
log_context=f"litellm_call_id={logging_obj.litellm_call_id}",
return_raw_model_name=_should_return_raw_model_name(self.data),
return_raw_model_name=should_return_raw_model_name(self.data),
)
fastapi_response.headers.update(
@ -3235,9 +3246,9 @@ class ProxyBaseLLMRequestProcessing:
because should_run_guardrail treats it as matching every hook.
"""
from litellm.proxy.proxy_server import llm_router
from litellm.proxy.utils import _check_and_merge_model_level_guardrails
from litellm.proxy.utils import check_and_merge_model_level_guardrails
guardrail_data: Final = _check_and_merge_model_level_guardrails(data=self.data, llm_router=llm_router)
guardrail_data: Final = check_and_merge_model_level_guardrails(data=self.data, llm_router=llm_router)
for cb in litellm.callbacks:
if not isinstance(cb, CustomGuardrail):
continue
@ -3558,11 +3569,13 @@ class ProxyBaseLLMRequestProcessing:
try:
from litellm.proxy.proxy_server import llm_router as _global_llm_router
from litellm.proxy.utils import (
_check_and_merge_model_level_guardrails,
check_and_merge_model_level_guardrails,
stream_gated_guardrail_names,
)
guardrail_data = _check_and_merge_model_level_guardrails(data=captured_data, llm_router=_global_llm_router)
guardrail_data: Final = check_and_merge_model_level_guardrails(
data=captured_data, llm_router=_global_llm_router
)
stream_gated: Final = stream_gated_guardrail_names(captured_data, captured_user_api_key_dict)
for cb in litellm.callbacks:
if not isinstance(cb, CustomGuardrail):
@ -3636,13 +3649,13 @@ class ProxyBaseLLMRequestProcessing:
if isinstance(e, RouterRateLimitError) and e.cooldown_time > 0:
headers["retry-after"] = str(math.ceil(e.cooldown_time))
async def _handle_llm_api_exception(
async def handle_llm_api_exception(
self,
e: Exception,
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
version: str | None = None,
):
) -> Never:
"""Raises ProxyException (OpenAI API compatible) if an exception is raised"""
log_llm_api_exception(e, self.litellm_call_id)
# Allow callbacks to transform the error response
@ -3787,6 +3800,8 @@ class ProxyBaseLLMRequestProcessing:
headers=safe_headers,
)
_handle_llm_api_exception = handle_llm_api_exception
#########################################################
# Proxy Level Streaming Data Generator
#########################################################
@ -3820,7 +3835,7 @@ class ProxyBaseLLMRequestProcessing:
return serialize
@staticmethod
async def _finalize_streaming_generator_cleanup(
async def finalize_streaming_generator_cleanup(
request: Request | None,
request_data: dict,
response: Any,
@ -3840,7 +3855,7 @@ class ProxyBaseLLMRequestProcessing:
)
if recorded_client_disconnect:
deferred_stream_logging_armed: Final = _deferred_stream_logging_is_armed(request_data)
ProxyLogging._fire_deferred_stream_logging(request_data)
ProxyLogging.fire_deferred_stream_logging(request_data)
# A disconnect-time success event (the deferred-guardrail flush
# above, or the partial-spend billing below) releases the
# request's max_parallel_requests slot through the limiter's
@ -3858,7 +3873,7 @@ class ProxyBaseLLMRequestProcessing:
and proxy_logging_obj is not None
and user_api_key_dict is not None
):
await proxy_logging_obj._arelease_max_parallel_requests_on_disconnect(user_api_key_dict)
await proxy_logging_obj.arelease_max_parallel_requests_on_disconnect(user_api_key_dict)
if hasattr(response, "aclose"):
try:
@ -3877,6 +3892,8 @@ class ProxyBaseLLMRequestProcessing:
):
await logging_obj.invalidate_baseline_cache_estimate("incomplete_response", completed=True)
_finalize_streaming_generator_cleanup = finalize_streaming_generator_cleanup
@staticmethod
async def async_streaming_data_generator(
response: object,
@ -3913,7 +3930,7 @@ class ProxyBaseLLMRequestProcessing:
# consumed, and cost injection is a no-op -- so the per-chunk coroutine
# await, response-string materialization, and cost-injection call are
# pure overhead on the streaming hot path (the default config).
caps: Final = ProxyLogging._callback_capabilities()
caps: Final = ProxyLogging.callback_capabilities()
cost_injection_enabled: Final = bool(getattr(litellm, "include_cost_in_streaming_usage", False))
fast_path = not caps.has_streaming_chunk_override and not caps.has_guardrail and not cost_injection_enabled
debug_enabled: Final = verbose_proxy_logger.isEnabledFor(logging.DEBUG)
@ -3956,7 +3973,7 @@ class ProxyBaseLLMRequestProcessing:
str_so_far += str(chunk.get("content", ""))
model_name = request_data.get("model", "")
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
chunk = ProxyBaseLLMRequestProcessing.process_chunk_with_cost_injection(
chunk, model_name, request_data.get("litellm_logging_obj")
)
@ -4022,7 +4039,7 @@ class ProxyBaseLLMRequestProcessing:
seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail)
yield seal + error_frame if seal else error_frame
finally:
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
await ProxyBaseLLMRequestProcessing.finalize_streaming_generator_cleanup(
request=request,
request_data=request_data,
response=response,
@ -4071,18 +4088,18 @@ class ProxyBaseLLMRequestProcessing:
@overload
@staticmethod
def _process_chunk_with_cost_injection(
def process_chunk_with_cost_injection(
chunk: bytes, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> bytes: ...
@overload
@staticmethod
def _process_chunk_with_cost_injection(
def process_chunk_with_cost_injection(
chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> object: ...
@staticmethod
def _process_chunk_with_cost_injection(
def process_chunk_with_cost_injection(
chunk: object, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None
) -> object:
"""
@ -4131,6 +4148,8 @@ class ProxyBaseLLMRequestProcessing:
return chunk
_process_chunk_with_cost_injection = process_chunk_with_cost_injection
@staticmethod
def _inject_cost_into_sse_frame_str(
frame_str: str, model_name: str, litellm_logging_obj: LiteLLMLoggingObj | None = None

View file

@ -6,10 +6,12 @@ from typing import TYPE_CHECKING, Final
from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.config_sync_pubsub import (
_ConfigSyncPubSub,
_pubsub_capable_client,
from litellm.proxy.common_utils.config_sync_pubsub import ( # noqa: F401 # legacy module exports
ConfigSyncPubSub,
_ConfigSyncPubSub, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_pubsub_capable_client, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
coordination_redis_cache,
pubsub_capable_client,
)
from litellm.proxy.common_utils.user_api_key_cache import AUTH_OBJECTS_TARGET
@ -75,7 +77,7 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None:
async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None:
try:
client: Final = _pubsub_capable_client(redis_cache)
client: Final = pubsub_capable_client(redis_cache)
async with _in_flight_publishes:
await client.publish(auth_cache_invalidation_channel(redis_cache), message)
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
@ -181,7 +183,7 @@ class AuthCacheInvalidationSubscriber:
backoff_seconds = _BACKOFF_INITIAL_SECONDS # rebind-ok: exponential backoff accumulator across reconnects
while True:
try:
client = _pubsub_capable_client(self._redis_cache)
client = pubsub_capable_client(self._redis_cache)
pubsub = client.pubsub()
try:
await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache))
@ -200,7 +202,7 @@ class AuthCacheInvalidationSubscriber:
await asyncio.sleep(backoff_seconds)
backoff_seconds = min(backoff_seconds * 2, _BACKOFF_MAX_SECONDS)
async def _consume(self, pubsub: _ConfigSyncPubSub) -> None:
async def _consume(self, pubsub: ConfigSyncPubSub) -> None:
while True:
message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS)
if message is None:
@ -222,7 +224,7 @@ class AuthCacheInvalidationSubscriber:
additional_cache.delete_cache(parsed.cache_key)
@staticmethod
async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None:
async def _close_pubsub(pubsub: ConfigSyncPubSub) -> None:
try:
await pubsub.aclose()
except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection

View file

@ -243,7 +243,7 @@ async def _choose_cached_model(
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner
PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner
)
from litellm.router_strategy.complexity_router.context_compaction import compaction_pending
@ -289,7 +289,7 @@ async def _choose_cached_model(
if body is None:
return None
limiter: Final = proxy_server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter")
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
if not isinstance(limiter, PROXY_MaxParallelRequestsHandler_v3):
return None
def counter_for_model(model_name: str) -> TokenCounter:

View file

@ -181,7 +181,7 @@ def initialize_callbacks_on_proxy(
imported_list.append(callback)
elif isinstance(callback, str) and callback == "presidio":
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
_OPTIONAL_PresidioPIIMasking,
OPTIONAL_PresidioPIIMasking,
)
presidio_logging_only: bool | None = litellm_settings.get("presidio_logging_only", None)
@ -196,7 +196,7 @@ def initialize_callbacks_on_proxy(
"logging_only": presidio_logging_only,
**_presidio_params,
}
pii_masking_object = _OPTIONAL_PresidioPIIMasking(**params)
pii_masking_object = OPTIONAL_PresidioPIIMasking(**params)
imported_list.append(pii_masking_object)
elif isinstance(callback, str) and callback == "llamaguard_moderations":
try:
@ -324,7 +324,7 @@ def initialize_callbacks_on_proxy(
imported_list.append(banned_keywords_obj)
elif isinstance(callback, str) and callback == "detect_prompt_injection":
from litellm.proxy.hooks.prompt_injection_detection import (
_OPTIONAL_PromptInjectionDetection,
OPTIONAL_PromptInjectionDetection,
)
prompt_injection_params = None
@ -332,20 +332,20 @@ def initialize_callbacks_on_proxy(
prompt_injection_params_in_config = litellm_settings["prompt_injection_params"]
prompt_injection_params = LiteLLMPromptInjectionParams(**prompt_injection_params_in_config)
prompt_injection_detection_obj = _OPTIONAL_PromptInjectionDetection(
prompt_injection_detection_obj = OPTIONAL_PromptInjectionDetection(
prompt_injection_params=prompt_injection_params,
)
imported_list.append(prompt_injection_detection_obj)
elif isinstance(callback, str) and callback == "batch_redis_requests":
from litellm.proxy.hooks.batch_redis_get import (
_PROXY_BatchRedisRequests,
PROXY_BatchRedisRequests,
)
batch_redis_obj = _PROXY_BatchRedisRequests()
batch_redis_obj = PROXY_BatchRedisRequests()
imported_list.append(batch_redis_obj)
elif isinstance(callback, str) and callback == "azure_content_safety":
from litellm.proxy.hooks.azure_content_safety import (
_PROXY_AzureContentSafety,
PROXY_AzureContentSafety,
)
azure_content_safety_params = litellm_settings["azure_content_safety_params"]
@ -353,7 +353,7 @@ def initialize_callbacks_on_proxy(
if v is not None and isinstance(v, str) and v.startswith("os.environ/"):
azure_content_safety_params[k] = get_secret(v)
azure_content_safety_obj = _PROXY_AzureContentSafety(
azure_content_safety_obj = PROXY_AzureContentSafety(
**azure_content_safety_params,
)
imported_list.append(azure_content_safety_obj)

View file

@ -21,10 +21,13 @@ class _ConfigSyncPubSub(Protocol):
def aclose(self) -> Awaitable[object]: ...
ConfigSyncPubSub = _ConfigSyncPubSub
class _ConfigSyncPubSubClient(Protocol):
def publish(self, channel: str, message: str) -> Awaitable[int]: ...
def pubsub(self) -> _ConfigSyncPubSub: ...
def pubsub(self) -> ConfigSyncPubSub: ...
CONFIG_SYNC_CHANNEL: Final = "litellm_proxy.config_change"
@ -82,13 +85,16 @@ def config_sync_channel(redis_cache: "RedisCache") -> str:
return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}"
def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient:
def pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient:
return cast( # cast-ok: protocol view of the pub/sub-capable async redis client
_ConfigSyncPubSubClient,
redis_cache.init_pubsub_client(), # pyright: ignore[reportUnknownMemberType] # redis generics
)
_pubsub_capable_client: Final = pubsub_capable_client
@dataclass(frozen=True, slots=True)
class _ConfigChangeMessage:
object_type: str
@ -102,7 +108,7 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s
if redis_cache is None:
return
try:
client: Final = _pubsub_capable_client(redis_cache)
client: Final = pubsub_capable_client(redis_cache)
await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type))
except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors
verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e)
@ -222,7 +228,7 @@ class ConfigSyncSubscriber:
backoff_seconds = self._backoff_initial_seconds
while True:
try:
client = _pubsub_capable_client(self._redis_cache)
client = pubsub_capable_client(self._redis_cache)
pubsub = client.pubsub()
try:
await pubsub.subscribe(config_sync_channel(self._redis_cache))
@ -241,7 +247,7 @@ class ConfigSyncSubscriber:
await self._sleep(backoff_seconds)
backoff_seconds = min(backoff_seconds * 2, self._backoff_max_seconds)
async def _consume(self, pubsub: _ConfigSyncPubSub) -> None:
async def _consume(self, pubsub: ConfigSyncPubSub) -> None:
while True:
message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=_POLL_TIMEOUT_SECONDS)
if message is None:
@ -265,7 +271,7 @@ class ConfigSyncSubscriber:
await self._sleep(seconds_until_next_resync)
@staticmethod
async def _drain_pending(pubsub: _ConfigSyncPubSub) -> None:
async def _drain_pending(pubsub: ConfigSyncPubSub) -> None:
while await pubsub.get_message(ignore_subscribe_messages=True, timeout=0) is not None:
pass
@ -277,7 +283,7 @@ class ConfigSyncSubscriber:
verbose_proxy_logger.warning("config sync resync callback failed: %s", e)
@staticmethod
async def _close_pubsub(pubsub: _ConfigSyncPubSub) -> None:
async def _close_pubsub(pubsub: ConfigSyncPubSub) -> None:
try:
await pubsub.aclose()
except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection

View file

@ -12,18 +12,24 @@ from litellm._logging import verbose_proxy_logger
# Legacy XSalsa20-Poly1305 (nacl) values carry no marker; the colon in the
# prefix can never appear in base64url(nacl output), so the prefix check is an
# unambiguous discriminator between the two formats on read.
_V2_GCM_PREFIX: Final = "v2:gcm:"
V2_GCM_PREFIX: Final = "v2:gcm:"
_V2_GCM_PREFIX: Final = V2_GCM_PREFIX
# general_settings key selecting the at-rest encryption algorithm for new writes.
# Default preserves the legacy algorithm so existing deployments are byte-for-byte
# unchanged until they explicitly opt in. Decrypt is always format-detecting, so
# flipping this flag forward (or back) never strands previously-written data.
_ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm"
_ALGO_AES_GCM: Final = "aes-256-gcm"
ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm"
_ENCRYPTION_ALGORITHM_SETTING: Final = ENCRYPTION_ALGORITHM_SETTING
ALGO_AES_GCM: Final = "aes-256-gcm"
_ALGO_AES_GCM: Final = ALGO_AES_GCM
_ALGO_XSALSA20: Final = "xsalsa20-poly1305"
def _get_salt_key():
def get_salt_key() -> str | None:
from litellm.proxy.proxy_server import master_key
salt_key = os.getenv("LITELLM_SALT_KEY", None)
@ -34,6 +40,9 @@ def _get_salt_key():
return salt_key
_get_salt_key: Final = get_salt_key
def _get_encryption_algorithm() -> str:
"""
Resolve the configured at-rest encryption algorithm for *new writes*.
@ -45,14 +54,14 @@ def _get_encryption_algorithm() -> str:
try:
from litellm.proxy.proxy_server import general_settings
algo: Final = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20)
algo: Final = general_settings.get(ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20)
except Exception:
# general_settings may not be importable in some contexts (e.g. SDK-only
# use of these helpers). Fall back to the legacy algorithm.
return _ALGO_XSALSA20
if isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM:
return _ALGO_AES_GCM
if isinstance(algo, str) and algo.lower() == ALGO_AES_GCM:
return ALGO_AES_GCM
return _ALGO_XSALSA20
@ -92,18 +101,18 @@ def _open_aes_gcm(sealed: bytes, signing_key: str, aad: bytes | None) -> str:
def _encrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string."""
sealed: Final = _seal_aes_gcm(value=value, signing_key=signing_key, aad=None)
return _V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8")
return V2_GCM_PREFIX + base64.urlsafe_b64encode(sealed).decode("utf-8")
def _decrypt_aes_gcm(value: str, signing_key: str) -> str:
"""Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`."""
sealed: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :])
sealed: Final = base64.urlsafe_b64decode(value[len(V2_GCM_PREFIX) :])
return _open_aes_gcm(sealed=sealed, signing_key=signing_key, aad=None)
def encrypt_bearer_token(value: str, prefix: str) -> str:
"""AES-256-GCM as unpadded base64url behind ``prefix``, which is also the AAD so a token can't change kind."""
salt_key: Final = _get_salt_key()
salt_key: Final = get_salt_key()
if not isinstance(salt_key, str):
raise ValueError("Set LITELLM_SALT_KEY or a master key to mint bearer tokens")
sealed: Final = _seal_aes_gcm(value=value, signing_key=salt_key, aad=prefix.encode("utf-8"))
@ -112,7 +121,7 @@ def encrypt_bearer_token(value: str, prefix: str) -> str:
def decrypt_bearer_token(token: str, prefix: str) -> str | None:
"""None unless ``token`` came from :func:`encrypt_bearer_token` with the same ``prefix``."""
salt_key: Final = _get_salt_key()
salt_key: Final = get_salt_key()
if not isinstance(salt_key, str) or not token.startswith(prefix):
return None
encoded: Final = token.removeprefix(prefix)
@ -124,11 +133,11 @@ def decrypt_bearer_token(token: str, prefix: str) -> str | None:
def encrypt_value_helper(value: str, new_encryption_key: str | None = None):
signing_key: Final = new_encryption_key or _get_salt_key()
signing_key: Final = new_encryption_key or get_salt_key()
try:
if isinstance(value, str):
if _get_encryption_algorithm() == _ALGO_AES_GCM:
if _get_encryption_algorithm() == ALGO_AES_GCM:
# AES path: the v2:gcm: output is already a base64url string, so it
# is returned directly with no extra base64 wrapper.
return _encrypt_aes_gcm(value=value, signing_key=cast(str, signing_key))
@ -160,7 +169,7 @@ def _legacy_ciphertext_bytes(value: str) -> bytes:
def _decrypt_with_signing_key(value: str, signing_key: str) -> str:
# Versioned AES-256-GCM values are detected before any base64 decode.
# The prefix is the algorithm tag the legacy nacl format never carried.
if value.startswith(_V2_GCM_PREFIX):
if value.startswith(V2_GCM_PREFIX):
return _decrypt_aes_gcm(value=value, signing_key=signing_key)
return decrypt_value(value=_legacy_ciphertext_bytes(value), signing_key=signing_key)
@ -171,7 +180,7 @@ def decrypt_if_encrypted_with(value: str, signing_key: str) -> str | None:
try:
# base64 decoding skips characters outside its alphabet, so "" and "*" decode to no bytes,
# which decrypt_value reads as an empty plaintext under any key.
decodes_to_nothing: Final = not value.startswith(_V2_GCM_PREFIX) and not _legacy_ciphertext_bytes(value)
decodes_to_nothing: Final = not value.startswith(V2_GCM_PREFIX) and not _legacy_ciphertext_bytes(value)
return None if decodes_to_nothing else _decrypt_with_signing_key(value=value, signing_key=signing_key)
except Exception: # noqa: BLE001 # base64, nacl and AES-GCM each raise their own "not a ciphertext" type
return None
@ -183,7 +192,7 @@ def decrypt_value_helper(
exception_type: Literal["debug", "error"] = "error",
return_original_value: bool = False,
) -> str | None:
signing_key: Final = _get_salt_key()
signing_key: Final = get_salt_key()
try:
if isinstance(value, str):

View file

@ -186,7 +186,7 @@ def is_otlp_trace_request(request: Request) -> bool:
return request.method == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"}
async def _read_request_body(request: Request | None) -> dict:
async def read_request_body(request: Request | None) -> dict:
"""
Safely read the request body and parse it as JSON.
@ -208,7 +208,7 @@ async def _read_request_body(request: Request | None) -> dict:
if _cached_request_body is not None:
return _cached_request_body
_request_headers: Final[dict] = _safe_get_request_headers(request=request)
_request_headers: Final[dict] = safe_get_request_headers(request=request)
content_type: Final = _request_headers.get("content-type", "")
if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES:
@ -285,7 +285,7 @@ async def _read_request_body(request: Request | None) -> dict:
)
# Cache the parsed result
_safe_set_request_parsed_body(request=request, parsed_body=parsed_body)
safe_set_request_parsed_body(request=request, parsed_body=parsed_body)
return parsed_body
except (json.JSONDecodeError, orjson.JSONDecodeError, ProxyException) as e:
@ -298,6 +298,9 @@ async def _read_request_body(request: Request | None) -> dict:
return {}
_read_request_body: Final = read_request_body
def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool:
"""Azure Speech bodies (raw audio, multipart uploads) are forwarded byte for byte, so auth must not consume them."""
media_type: Final = _normalize_media_type(content_type)
@ -309,7 +312,7 @@ def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool:
async def read_raw_json_body(request: Request | None) -> bytes | None:
if request is None or _safe_get_request_parsed_body(request=request) is None:
return None
content_type: Final = _safe_get_request_headers(request=request).get("content-type", "")
content_type: Final = safe_get_request_headers(request=request).get("content-type", "")
if _is_form_content_type(content_type):
return None
try:
@ -334,7 +337,7 @@ def get_client_requested_model(request: Request | None) -> str | None:
return model if isinstance(model, str) else None
def _safe_get_request_query_params(request: Request | None) -> dict:
def safe_get_request_query_params(request: Request | None) -> dict:
if request is None:
return {}
try:
@ -346,7 +349,10 @@ def _safe_get_request_query_params(request: Request | None) -> dict:
return {}
def _safe_set_request_parsed_body(
_safe_get_request_query_params: Final = safe_get_request_query_params
def safe_set_request_parsed_body(
request: Request | None,
parsed_body: dict,
) -> None:
@ -358,6 +364,9 @@ def _safe_set_request_parsed_body(
verbose_proxy_logger.debug("Unexpected error setting request parsed body - %s", e)
_safe_set_request_parsed_body: Final = safe_set_request_parsed_body
def rewrite_request_model(
request_data: dict[str, object], # mutable-ok: the request body is rewritten in place for every downstream reader
request: Request | None,
@ -371,12 +380,12 @@ def rewrite_request_model(
return
cached_body: Final = _safe_get_request_parsed_body(request=request)
body: Final = {**cached_body, "model": model} if cached_body is not None else request_data
_safe_set_request_parsed_body(request=request, parsed_body=body)
request._json = body
request._body = orjson.dumps(body)
safe_set_request_parsed_body(request=request, parsed_body=body)
request._json = body # pyright: ignore[reportPrivateUsage] # Starlette JSON cache
request._body = orjson.dumps(body) # pyright: ignore[reportPrivateUsage] # Starlette body cache
def _safe_get_request_headers(request: Request | None) -> dict:
def safe_get_request_headers(request: Request | None) -> dict:
"""
[Non-Blocking] Safely get the request headers.
Caches the result on request.state to avoid re-creating dict(request.headers) per call.
@ -405,6 +414,9 @@ def _safe_get_request_headers(request: Request | None) -> dict:
return headers
_safe_get_request_headers: Final = safe_get_request_headers
def check_file_size_under_limit(
request_data: dict,
file: UploadFile,
@ -537,7 +549,7 @@ async def get_request_body(request: Request) -> dict[str, Any]:
if request.method == "POST":
content_type: Final = request.headers.get("content-type", "")
if is_json_content_type(content_type):
return await _read_request_body(request)
return await read_request_body(request)
elif _is_form_content_type(content_type):
return await get_form_data(request)
else:
@ -685,7 +697,7 @@ def populate_request_with_path_params(request_data: dict, request: Request) -> d
dict: Updated request_data with path parameters and query parameters added
"""
# Add query parameters to request_data (for GET requests, etc.)
query_params: Final = _safe_get_request_query_params(request)
query_params: Final = safe_get_request_query_params(request)
if query_params:
for key, value in query_params.items():
# Don't overwrite existing values from request body

View file

@ -20,8 +20,9 @@ from litellm.proxy._types import (
RegenerateKeyRequest,
)
from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks
from litellm.proxy.management_endpoints.key_management_endpoints import (
_calculate_key_rotation_time,
from litellm.proxy.management_endpoints.key_management_endpoints import ( # noqa: F401 # legacy module exports
_calculate_key_rotation_time, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
calculate_key_rotation_time,
regenerate_key_fn,
)
from litellm.proxy.utils import PrismaClient
@ -187,7 +188,7 @@ class KeyRotationManager:
if isinstance(response, GenerateKeyResponse) and response.token_id and key.rotation_interval:
# Calculate next rotation time using helper function
now: Final = datetime.now(timezone.utc)
next_rotation_time: Final = _calculate_key_rotation_time(key.rotation_interval)
next_rotation_time: Final = calculate_key_rotation_time(key.rotation_interval)
await VerificationTokenRepository(self.prisma_client).table.update(
where={"token": response.token_id},
data={

View file

@ -7,7 +7,10 @@ from typing import Final
from fastapi import Request
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
SENSITIVE_DATA_MASKER: Final = SensitiveDataMasker()
@ -55,7 +58,7 @@ async def get_custom_llm_provider_from_request_body(request: Request) -> str | N
Safely reads the request body
"""
request_body: Final[dict] = await _read_request_body(request=request) or {}
request_body: Final[dict] = await read_request_body(request=request) or {}
if "custom_llm_provider" in request_body:
return request_body["custom_llm_provider"]
return None

View file

@ -40,7 +40,9 @@ def get_openapi_schema_with_compat(
from pydantic_core import core_schema
# Store original method
original_unknown_type_schema: Final = GenerateSchema._unknown_type_schema
original_unknown_type_schema: Final = (
GenerateSchema._unknown_type_schema # pyright: ignore[reportPrivateUsage] # Pydantic schema internals
)
def patched_unknown_type_schema(self, obj):
"""Patch to handle openai.Timeout and other non-serializable types"""

View file

@ -51,10 +51,10 @@ async def check_feature_access_for_user(
# Feature is disabled. Check if team/org admins are exempted.
if general_settings.get(allow_team_admins_flag, False):
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_privileges,
user_has_admin_privileges,
)
is_admin: Final = await _user_has_admin_privileges(
is_admin: Final = await user_has_admin_privileges(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,

View file

@ -1,12 +1,16 @@
from functools import lru_cache
from typing import Final
from litellm.constants import _REALTIME_BODY_CACHE_SIZE
@lru_cache(maxsize=_REALTIME_BODY_CACHE_SIZE)
def _realtime_request_body(model: str | None) -> bytes:
def realtime_request_body(model: str | None) -> bytes:
"""
Generate the realtime websocket request body. Cached with LRU semantics to avoid repeated
string formatting work while keeping memory usage bounded.
"""
return f'{{"model": "{model or ""}"}}'.encode()
_realtime_request_body: Final = realtime_request_body

View file

@ -9,7 +9,10 @@ from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_body,
get_custom_llm_provider_from_request_headers,
@ -89,7 +92,7 @@ async def create_container(
)
# Read request body
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
# Extract custom_llm_provider using priority chain
# Priority: headers > query params > request body > default
@ -125,7 +128,7 @@ async def create_container(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -245,7 +248,7 @@ async def list_containers(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -360,7 +363,7 @@ async def retrieve_container(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -466,7 +469,7 @@ async def delete_container(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

@ -260,7 +260,7 @@ async def _process_binary_request(
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -349,7 +349,7 @@ async def _process_multipart_upload_request(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -438,7 +438,7 @@ async def _process_request(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

@ -5,7 +5,10 @@ from fastapi_sso.sso.base import OpenID
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_safe_get_request_headers, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
safe_get_request_headers,
)
class CustomSSOLoginHandler(CustomLogger):
@ -22,7 +25,7 @@ class CustomSSOLoginHandler(CustomLogger):
self,
request: Request,
) -> OpenID:
request_headers_dict: Final = _safe_get_request_headers(request)
request_headers_dict: Final = safe_get_request_headers(request)
verbose_logger.debug("inside custom ui sso sign in hook...")
return OpenID(
id=request_headers_dict.get("x-litellm-user-id") or "123",

View file

@ -22,7 +22,12 @@ from typing import Final, TypeVar
from litellm._logging import verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.proxy.db.log_db_metrics import _is_exception_related_to_db, claim_db_io, db_io_claimed
from litellm.proxy.db.log_db_metrics import ( # noqa: F401 # legacy module exports
_is_exception_related_to_db, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
claim_db_io,
db_io_claimed,
is_exception_related_to_db,
)
_T = TypeVar("_T")
@ -75,7 +80,7 @@ async def db_span(call_type: str, table: str | None, operation: str | None = Non
try:
yield
except Exception as e:
if service_logging is not None and _is_exception_related_to_db(e):
if service_logging is not None and is_exception_related_to_db(e):
await _emit_failure(service_logging, call_type, event_metadata, start_time, e)
raise
if service_logging is None or not witness.touched:

View file

@ -1546,7 +1546,7 @@ class DBSpendUpdateWriter:
else:
- Regular flow of this method
"""
if RedisUpdateBuffer._should_commit_spend_updates_to_redis():
if RedisUpdateBuffer.should_commit_spend_updates_to_redis():
await self._commit_spend_updates_to_db_with_redis(
prisma_client=prisma_client,
n_retry_times=n_retry_times,
@ -1885,12 +1885,12 @@ class DBSpendUpdateWriter:
################## Tool Registry Upserts ##################
await self._flush_tool_discovery_queue(prisma_client=prisma_client)
async def _commit_daily_tag_spend_to_db(
async def commit_daily_tag_spend_to_db(
self,
prisma_client: PrismaClient,
n_retry_times: int,
proxy_logging_obj: ProxyLogging,
):
) -> None:
"""
Commit only tag spend updates to database.
This is called by a separate scheduler job at a longer interval.
@ -1904,12 +1904,14 @@ class DBSpendUpdateWriter:
proxy_logging_obj=proxy_logging_obj,
)
async def _commit_daily_tag_spend_to_db_with_redis(
_commit_daily_tag_spend_to_db = commit_daily_tag_spend_to_db
async def commit_daily_tag_spend_to_db_with_redis(
self,
prisma_client: PrismaClient,
n_retry_times: int,
proxy_logging_obj: ProxyLogging,
):
) -> None:
"""
Commit daily tag spend updates using Redis buffering.
@ -1942,6 +1944,8 @@ class DBSpendUpdateWriter:
cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME,
)
_commit_daily_tag_spend_to_db_with_redis = commit_daily_tag_spend_to_db_with_redis
@staticmethod
async def _commit_window_spend_updates(
prisma_client: PrismaClient,
@ -2027,7 +2031,7 @@ class DBSpendUpdateWriter:
verbose_proxy_logger.debug("_flush_tool_discovery_queue error (non-blocking): %s", e)
@staticmethod
async def _handle_spend_update_failure(
async def handle_spend_update_failure(
e: Exception,
attempt: int,
n_retry_times: int,
@ -2038,7 +2042,7 @@ class DBSpendUpdateWriter:
``lock_timeout`` (55P03), else re-raise. All three roll the transaction back before any
increment applied, so re-sending the same batch cannot double-count."""
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.utils import _raise_failed_update_spend_exception
from litellm.proxy.utils import raise_failed_update_spend_exception
is_retryable = (
isinstance(e, DB_RETRY_SAFE_ERROR_TYPES)
@ -2046,7 +2050,7 @@ class DBSpendUpdateWriter:
or PrismaDBExceptionHandler.is_lock_timeout_error(e)
)
if not is_retryable or attempt >= n_retry_times:
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
verbose_proxy_logger.warning(
"Retrying spend update after retryable DB error (attempt %s/%s): %s",
attempt + 1,
@ -2055,6 +2059,8 @@ class DBSpendUpdateWriter:
)
await asyncio.sleep(random.uniform(2**attempt, 2 ** (attempt + 1)))
_handle_spend_update_failure = handle_spend_update_failure
async def _commit_spend_updates_to_db(
self,
prisma_client: PrismaClient,
@ -2088,7 +2094,7 @@ class DBSpendUpdateWriter:
)
break
except Exception as e:
await self._handle_spend_update_failure(
await self.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2132,7 +2138,7 @@ class DBSpendUpdateWriter:
)
break
except Exception as e:
await self._handle_spend_update_failure(
await self.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2162,7 +2168,7 @@ class DBSpendUpdateWriter:
)
break
except Exception as e:
await self._handle_spend_update_failure(
await self.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2192,7 +2198,7 @@ class DBSpendUpdateWriter:
# Transaction succeeded, break out of retry loop
break
except Exception as e:
await self._handle_spend_update_failure(
await self.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2234,7 +2240,7 @@ class DBSpendUpdateWriter:
)
break
except Exception as e:
await self._handle_spend_update_failure(
await self.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2262,7 +2268,7 @@ class DBSpendUpdateWriter:
)
break
except Exception as e:
await self._handle_spend_update_failure(
await self.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2389,7 +2395,7 @@ class DBSpendUpdateWriter:
)
break
except Exception as e:
await DBSpendUpdateWriter._handle_spend_update_failure(
await DBSpendUpdateWriter.handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
@ -2489,7 +2495,7 @@ class DBSpendUpdateWriter:
"""
Generic function to update daily spend for any entity type (user, team, org, tag, end_user, agent)
"""
from litellm.proxy.utils import _raise_failed_update_spend_exception
from litellm.proxy.utils import raise_failed_update_spend_exception
verbose_proxy_logger.debug(
"Daily %s Spend transactions: %s", entity_type.capitalize(), len(daily_spend_transactions)
@ -2589,7 +2595,7 @@ class DBSpendUpdateWriter:
if not is_retryable:
raise
if i >= n_retry_times:
_raise_failed_update_spend_exception(
raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
@ -2604,7 +2610,7 @@ class DBSpendUpdateWriter:
)
except Exception as e:
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
@staticmethod
async def update_daily_user_spend(

View file

@ -147,7 +147,7 @@ class RedisUpdateBuffer:
self.redis_cache = redis_cache
@staticmethod
def _should_commit_spend_updates_to_redis() -> bool:
def should_commit_spend_updates_to_redis() -> bool:
"""
Checks if the Pod should commit spend updates to Redis
@ -163,6 +163,8 @@ class RedisUpdateBuffer:
return False
return _use_redis_transaction_buffer
_should_commit_spend_updates_to_redis = should_commit_spend_updates_to_redis
@with_service_target(SPEND_QUEUE_TARGET)
async def _store_transactions_in_redis(
self,
@ -556,7 +558,7 @@ class RedisUpdateBuffer:
max_rows: int = REDIS_SPEND_LOGS_BUFFER_MAX_ROWS,
) -> bool:
"""Park spend-log rows in Redis so they outlive this pod, dropping the oldest past ``max_rows``."""
if self.redis_cache is None or len(rows) == 0 or not self._should_commit_spend_updates_to_redis():
if self.redis_cache is None or len(rows) == 0 or not self.should_commit_spend_updates_to_redis():
return False
try:
buffer_size: Final = await self.redis_cache.async_rpush_and_trim(
@ -582,7 +584,7 @@ class RedisUpdateBuffer:
@with_service_target(SPEND_QUEUE_TARGET)
async def get_spend_logs_from_redis_buffer(self, limit: int) -> tuple[dict[str, object], ...]:
"""Atomically take up to ``limit`` parked spend-log rows out of Redis."""
if self.redis_cache is None or not self._should_commit_spend_updates_to_redis():
if self.redis_cache is None or not self.should_commit_spend_updates_to_redis():
return ()
popped: Final[str | list[str] | None] = await self.redis_cache.async_lpop(
key=REDIS_SPEND_LOGS_BUFFER_KEY,

View file

@ -175,7 +175,7 @@ def log_db_metrics(func):
return wrapper
def _is_exception_related_to_db(e: Exception) -> bool:
def is_exception_related_to_db(e: Exception) -> bool:
"""
Returns True if the exception is related to the DB
"""
@ -186,6 +186,9 @@ def _is_exception_related_to_db(e: Exception) -> bool:
return isinstance(e, (PrismaError, httpx.TransportError))
_is_exception_related_to_db: Final = is_exception_related_to_db
async def _handle_logging_db_exception(
e: Exception,
func: Callable,
@ -198,7 +201,7 @@ async def _handle_logging_db_exception(
from litellm.proxy.proxy_server import proxy_logging_obj
# don't log this as a DB Service Failure, if the DB did not raise an exception
if _is_exception_related_to_db(e) is not True:
if is_exception_related_to_db(e) is not True:
return False
try:

View file

@ -87,14 +87,14 @@ async def decisions(
model=str(data.get("model", "")),
llm_provider="",
)
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=bad_request_error,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
except Exception as error:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=error,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

@ -16,8 +16,9 @@ from litellm.litellm_core_utils.hidden_params import set_hidden_param
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
from litellm.proxy.openai_files_endpoints.common_utils import ( # noqa: F401 # legacy module exports
_is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
is_base64_encoded_unified_file_id,
validate_managed_id_requirement,
)
from litellm.proxy.utils import handle_exception_on_proxy
@ -150,7 +151,9 @@ async def create_fine_tuning_job(
)
response: LiteLLMFineTuningJob | None = None
if training_file:
unified_file_id = _is_base64_encoded_unified_file_id(training_file)
unified_file_id = ( # rebind-ok: pre-existing rebinding on a rename-only line
is_base64_encoded_unified_file_id(training_file)
)
## IF SO, Route based on that
if unified_file_id:
""" """
@ -292,7 +295,9 @@ async def retrieve_fine_tuning_job(
unified_finetuning_job_id: str | Literal[False] = False
response: LiteLLMFineTuningJob | None = None
if fine_tuning_job_id:
unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id)
unified_finetuning_job_id = ( # rebind-ok: pre-existing rebinding on a rename-only line
is_base64_encoded_unified_file_id(fine_tuning_job_id)
)
if unified_finetuning_job_id:
if llm_router is None:
raise HTTPException(
@ -565,7 +570,9 @@ async def cancel_fine_tuning_job(
unified_finetuning_job_id: str | Literal[False] = False
response: LiteLLMFineTuningJob | None = None
if fine_tuning_job_id:
unified_finetuning_job_id = _is_base64_encoded_unified_file_id(fine_tuning_job_id)
unified_finetuning_job_id = ( # rebind-ok: pre-existing rebinding on a rename-only line
is_base64_encoded_unified_file_id(fine_tuning_job_id)
)
if unified_finetuning_job_id:
if llm_router is None:
raise HTTPException(

View file

@ -23,9 +23,11 @@ from fastapi.responses import ORJSONResponse
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_query_params,
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_safe_get_request_query_params, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
safe_get_request_query_params,
)
router: Final = APIRouter(tags=["gemini managed agents"])
@ -92,7 +94,7 @@ def _merge_query_params_into_data(data: dict, request: Request) -> dict:
headers. Use the ``litellm_params_template`` JSON body field on POST
requests, or the JSON-encoded query parameter above for GET/DELETE.
"""
query_params: Final = _safe_get_request_query_params(request)
query_params: Final = safe_get_request_query_params(request)
if not query_params:
return data
@ -172,7 +174,7 @@ async def create_gemini_agent(
```
"""
srv: Final = _proxy_server_imports()
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
# Merge litellm_params_template (e.g. custom_llm_provider, api_key) into the request
litellm_params_template: Final = data.pop("litellm_params_template", None) or {}
if isinstance(litellm_params_template, dict):
@ -203,7 +205,7 @@ async def create_gemini_agent(
version=srv["version"],
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=srv["proxy_logging_obj"],
@ -260,7 +262,7 @@ async def list_gemini_agents(
version=srv["version"],
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=srv["proxy_logging_obj"],
@ -318,7 +320,7 @@ async def get_gemini_agent(
version=srv["version"],
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=srv["proxy_logging_obj"],
@ -376,7 +378,7 @@ async def delete_gemini_agent(
version=srv["version"],
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=srv["proxy_logging_obj"],
@ -434,7 +436,7 @@ async def list_gemini_agent_versions(
version=srv["version"],
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=srv["proxy_logging_obj"],

View file

@ -6,7 +6,10 @@ from fastapi.responses import ORJSONResponse
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import ( # noqa: F401 # legacy module exports
_read_request_body, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
read_request_body,
)
from litellm.types.llms.vertex_ai import TokenCountDetailsResponse
router: Final = APIRouter(
@ -42,7 +45,7 @@ async def google_generate_content(
version,
)
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
@ -67,7 +70,7 @@ async def google_generate_content(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -103,7 +106,7 @@ async def google_stream_generate_content(
version,
)
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
data["stream"] = True
@ -132,7 +135,7 @@ async def google_stream_generate_content(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -166,10 +169,10 @@ async def google_count_tokens(request: Request, model_name: str):
```
"""
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import read_request_body
from litellm.proxy.proxy_server import token_counter as internal_token_counter
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
contents: Final = data.get("contents", [])
# Create TokenCountRequest for the internal endpoint
from litellm.proxy._types import TokenCountRequest
@ -268,7 +271,7 @@ async def create_interaction(
version,
)
data: Final = await _read_request_body(request=request)
data: Final = await read_request_body(request=request)
# Default to gemini provider for interactions
if "custom_llm_provider" not in data:
@ -295,7 +298,7 @@ async def create_interaction(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -363,7 +366,7 @@ async def get_interaction(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -431,7 +434,7 @@ async def delete_interaction(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
@ -499,7 +502,7 @@ async def cancel_interaction(
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
raise await processor.handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

@ -43,7 +43,10 @@ from litellm.proxy.guardrails.guardrail_registry import (
parse_tolerant_litellm_params,
)
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.proxy.management_endpoints.common_utils import ( # noqa: F401 # legacy module exports
_user_has_admin_view, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
user_api_key_has_admin_view,
)
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import GuardrailsRepository
from litellm.types.guardrails import (
@ -247,7 +250,7 @@ async def list_guardrails_v2(
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
from litellm.proxy.proxy_server import prisma_client
is_admin: Final = _user_has_admin_view(user_api_key_dict)
is_admin: Final = user_api_key_has_admin_view(user_api_key_dict)
try:
guardrails = (
@ -942,7 +945,7 @@ async def list_guardrail_submissions(
# Admin Viewer follows the read-parity rule: see all submissions like a
# Proxy Admin would (no writes — registration / approval still gated
# elsewhere by their own per-action checks).
is_admin: Final = _user_has_admin_view(user_api_key_dict)
is_admin: Final = user_api_key_has_admin_view(user_api_key_dict)
visible_team_ids: list[str] | None = None
if not is_admin:
visible_team_ids = await _get_user_team_ids(user_api_key_dict)
@ -1021,7 +1024,7 @@ async def get_guardrail_submission(
if prisma_client is None:
raise HTTPException(status_code=500, detail="Prisma client not initialized")
is_admin: Final = _user_has_admin_view(user_api_key_dict)
is_admin: Final = user_api_key_has_admin_view(user_api_key_dict)
try:
row: Final = await _guardrails_table(prisma_client).find_unique(where={"guardrail_id": guardrail_id})

View file

@ -26,7 +26,9 @@ AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000
AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01"
JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1"
_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
_RESPONSES_API_CALL_TYPES: Final = RESPONSES_API_CALL_TYPES
def resolve_content_safety_api_version(configured: str | None) -> str:
@ -155,7 +157,7 @@ class AzureGuardrailBase:
return get_last_user_message(messages)
def get_user_prompt_from_request(self, data: Mapping[str, object], call_type: CallTypesLiteral) -> str | None:
if call_type in _RESPONSES_API_CALL_TYPES:
if call_type in RESPONSES_API_CALL_TYPES:
responses_input: Final = data.get("input")
if not isinstance(responses_input, (str, list)):
return None

View file

@ -27,7 +27,12 @@ from litellm.types.utils import (
GuardrailTracingDetail,
)
from .base import _RESPONSES_API_CALL_TYPES, AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase
from .base import ( # noqa: F401 # legacy module exports
_RESPONSES_API_CALL_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH,
RESPONSES_API_CALL_TYPES,
AzureGuardrailBase,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -249,7 +254,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
"Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s",
call_type,
)
if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None:
if call_type not in RESPONSES_API_CALL_TYPES and data.get("messages") is None:
verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data")
return data
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)

View file

@ -16,7 +16,11 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypesLiteral, GenericGuardrailAPIInputs, LLMResponseTypes
from .base import _RESPONSES_API_CALL_TYPES, AzureGuardrailBase
from .base import ( # noqa: F401 # legacy module exports
_RESPONSES_API_CALL_TYPES, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
RESPONSES_API_CALL_TYPES,
AzureGuardrailBase,
)
if TYPE_CHECKING:
from litellm.caching.caching import DualCache
@ -231,7 +235,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
"Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
call_type,
)
if call_type not in _RESPONSES_API_CALL_TYPES and data.get("messages") is None:
if call_type not in RESPONSES_API_CALL_TYPES and data.get("messages") is None:
verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data")
return data
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)

View file

@ -52,7 +52,10 @@ from litellm.types.utils import (
TextCompletionResponse,
)
from .cisco_ai_defense_mcp import _CiscoAIDefenseMcpMixin
from .cisco_ai_defense_mcp import ( # noqa: F401 # legacy module exports
CiscoAIDefenseMcpMixin,
_CiscoAIDefenseMcpMixin, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
)
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import (
@ -116,7 +119,7 @@ class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
"""Base-class constructor options this guardrail forwards untouched to CustomGuardrail."""
class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail):
class CiscoAIDefenseGuardrail(CiscoAIDefenseMcpMixin, CustomGuardrail):
"""
Cisco AI Defense guardrail integration.

View file

@ -218,7 +218,7 @@ class _CiscoAIDefenseMcpMixin:
inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
if inner is not None:
if _CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj):
if CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj):
return True
try:
setattr(response_obj, "mcp_tool_call_response", replacement)
@ -229,7 +229,7 @@ class _CiscoAIDefenseMcpMixin:
content: Final = getattr(response_obj, "content", None)
if isinstance(content, list):
content[:] = replacement
structured_replacement: Final = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
structured_replacement: Final = CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
if hasattr(response_obj, "structured_content"):
try:
setattr(response_obj, "structured_content", structured_replacement)
@ -250,12 +250,12 @@ class _CiscoAIDefenseMcpMixin:
result: Final = response_obj.get("result")
if isinstance(result, dict):
result["content"] = replacement
result["structuredContent"] = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
result["structuredContent"] = CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
result["isError"] = True
return True
response_obj["result"] = {
"content": replacement,
"structuredContent": _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement),
"structuredContent": CiscoAIDefenseMcpMixin._replacement_structured_content(replacement),
"isError": True,
}
return True
@ -472,7 +472,7 @@ class _CiscoAIDefenseMcpMixin:
return {
"jsonrpc": "2.0",
"id": response.get("id") or "litellm-mcp",
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response),
"result": CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response),
}
if isinstance(response, list):
if response and all(
@ -484,7 +484,7 @@ class _CiscoAIDefenseMcpMixin:
return {
"jsonrpc": "2.0",
"id": "litellm-mcp",
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(
"result": CiscoAIDefenseMcpMixin._build_mcp_result(
content=inner_content, source=response_fields
),
}
@ -493,7 +493,7 @@ class _CiscoAIDefenseMcpMixin:
return {
"jsonrpc": "2.0",
"id": "litellm-mcp",
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=response),
"result": CiscoAIDefenseMcpMixin._build_mcp_result(content=response),
}
model_dump: Final = getattr(response, "model_dump", None)
if callable(model_dump):
@ -502,13 +502,13 @@ class _CiscoAIDefenseMcpMixin:
except TypeError:
dumped = model_dump()
if isinstance(dumped, dict):
return _CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped)
return CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped)
content = getattr(response, "content", None)
if isinstance(content, list):
return {
"jsonrpc": "2.0",
"id": "litellm-mcp",
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response),
"result": CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response),
}
return None
@ -536,9 +536,9 @@ class _CiscoAIDefenseMcpMixin:
inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
if inner is not None:
return _CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text)
return CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text)
content_list: Final = _CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj)
content_list: Final = CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj)
replaced = False
if isinstance(content_list, list):
@ -586,7 +586,7 @@ class _CiscoAIDefenseMcpMixin:
return None
inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
if inner is not None:
return _CiscoAIDefenseMcpMixin._coerce_to_content_list(inner)
return CiscoAIDefenseMcpMixin._coerce_to_content_list(inner)
content: Final = getattr(response_obj, "content", None)
if isinstance(content, list):
return content
@ -643,3 +643,6 @@ class _CiscoAIDefenseMcpMixin:
if isinstance(direct, dict) and direct:
return dict(direct)
return None
CiscoAIDefenseMcpMixin = _CiscoAIDefenseMcpMixin

View file

@ -13,11 +13,14 @@ import json
from pathlib import Path
from typing import Any, Final
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import (
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.competitor_intent.base import ( # noqa: F401 # legacy module exports
BaseCompetitorIntentChecker,
_compile_marker,
_count_signals,
_word_boundary_match,
_compile_marker, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_count_signals, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_word_boundary_match, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
compile_marker,
count_signals,
word_boundary_match,
)
# Location/travel context: prepositions, travel verbs, booking nouns, entry/geo nouns.
@ -170,8 +173,8 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker):
self._other_meaning_signals = list(merged.get("other_meaning_signals") or [])
self._competitor_signals = list(merged.get("competitor_signals") or [])
self._other_meaning_anchors = list(merged.get("other_meaning_anchors") or [])
self._explicit_competitor_marker = _compile_marker(merged.get("explicit_competitor_marker"))
self._explicit_other_meaning_marker = _compile_marker(merged.get("explicit_other_meaning_marker"))
self._explicit_competitor_marker = compile_marker(merged.get("explicit_competitor_marker"))
self._explicit_other_meaning_marker = compile_marker(merged.get("explicit_other_meaning_marker"))
def _classify_ambiguous(self, text: str, token: str) -> tuple[str, float]:
"""Other meaning vs competitor using airline signals and explicit markers."""
@ -179,21 +182,25 @@ class AirlineCompetitorIntentChecker(BaseCompetitorIntentChecker):
if (
self._explicit_competitor_marker
and self._explicit_competitor_marker.search(text_lower)
and _word_boundary_match(text_lower, token.lower())
and word_boundary_match(text_lower, token.lower())
):
return "COMPETITOR", 0.85
if self._explicit_other_meaning_marker and self._explicit_other_meaning_marker.search(text_lower):
return "OTHER_MEANING", 0.85
# Operational-only: baggage/lounge/check-in/refund with no comparison → product query
has_comparison: Final = _count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0
operational_count: Final = _count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS)
has_comparison: Final = count_signals(text_lower, AIRLINE_COMPARISON_SIGNALS) > 0
operational_count: Final = count_signals(text_lower, AIRLINE_OPERATIONAL_SIGNALS)
if not has_comparison and operational_count > 0:
return "OTHER_MEANING", 0.85
# Score: location/travel context vs airline context (no place-name list)
other_count = _count_signals(text_lower, self._other_meaning_signals)
other_count = count_signals( # rebind-ok: pre-existing rebinding on a rename-only line
text_lower, self._other_meaning_signals
)
if self._other_meaning_anchors:
other_count += _count_signals(text_lower, self._other_meaning_anchors)
comp_count: Final = _count_signals(text_lower, self._competitor_signals)
other_count += count_signals( # rebind-ok: pre-existing rebinding on a rename-only line
text_lower, self._other_meaning_anchors
)
comp_count: Final = count_signals(text_lower, self._competitor_signals)
total: Final = other_count + comp_count
if total == 0:
return "OTHER_MEANING", 0.5

View file

@ -31,17 +31,23 @@ def normalize(text: str) -> str:
return re.sub(r"\s+", " ", t)
def _word_boundary_match(text: str, token: str) -> bool:
def word_boundary_match(text: str, token: str) -> bool:
"""True if token appears as a word in text."""
return bool(re.search(r"\b" + re.escape(token) + r"\b", text))
def _count_signals(text: str, patterns: list[str]) -> int:
_word_boundary_match: Final = word_boundary_match
def count_signals(text: str, patterns: list[str]) -> int:
"""Count how many of the patterns appear in text."""
return sum(1 for p in patterns if re.search(p, text, re.IGNORECASE))
def _compile_marker(pattern: str | None) -> Pattern[str] | None:
_count_signals: Final = count_signals
def compile_marker(pattern: str | None) -> Pattern[str] | None:
"""Compile optional regex string to a pattern."""
if not pattern or not pattern.strip():
return None
@ -51,6 +57,9 @@ def _compile_marker(pattern: str | None) -> Pattern[str] | None:
return None
_compile_marker: Final = compile_marker
def text_for_entity_matching(text: str) -> str:
"""Letters-only variant for entity matching (e.g. split punctuation)."""
t: Final = re.sub(r"[^\w\s]", " ", text)
@ -117,7 +126,7 @@ class BaseCompetitorIntentChecker:
found: Final[list[tuple[str, str, bool]]] = []
seen: Final[set[tuple[str, str]]] = set()
for token in self._competitor_tokens:
if not _word_boundary_match(normalized, token):
if not word_boundary_match(normalized, token):
continue
canonical = self.competitor_canonical.get(token, token)
key = (token, canonical)
@ -139,7 +148,7 @@ class BaseCompetitorIntentChecker:
}
for b in self.brand_self:
if _word_boundary_match(normalized, b):
if word_boundary_match(normalized, b):
entities["brand_self"].append(b)
evidence.append({"type": "entity", "key": "brand_self", "value": b, "match": b})

View file

@ -193,7 +193,7 @@ class MCPEndUserPermissionGuardrail(CustomGuardrail):
MCPRequestHandler,
)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups)
access_group_servers: Final = await MCPRequestHandler.get_mcp_servers_from_access_groups(mcp_access_groups)
return list(set(direct_mcp_servers + access_group_servers))

View file

@ -19,7 +19,7 @@ from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
import aiohttp
from pydantic import ConfigDict, TypeAdapter, with_config
from pydantic import ConfigDict, JsonValue, TypeAdapter, with_config
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
@ -279,7 +279,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
self,
presidio_analyzer_api_base: str | None = None,
presidio_anonymizer_api_base: str | None = None,
):
) -> None:
self.presidio_analyzer_api_base: str | None = presidio_analyzer_api_base or get_secret(
"PRESIDIO_ANALYZER_API_BASE", None
)
@ -922,7 +922,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
def raise_exception_if_blocked_entities_detected(
self, analyze_results: list[PresidioAnalyzeResponseItem] | _PresidioAnonymizeResponse
):
) -> None:
"""
Raise an exception if blocked entities are detected
"""
@ -1022,7 +1022,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
cache: DualCache,
data: dict,
call_type: str,
):
) -> dict[str, object]:
"""
- Check if request turned off pii
- Check if user allowed to turn off pii (key permissions -> 'allow_pii_controls')
@ -1212,10 +1212,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
async def async_post_call_success_hook(
self,
data: dict,
data: dict[str, object],
user_api_key_dict: UserAPIKeyAuth,
response: ModelResponse | EmbeddingResponse | ImageResponse,
):
) -> dict[str, JsonValue] | ModelResponse | EmbeddingResponse | ImageResponse:
"""
Output parse the response object to replace the masked tokens with user sent values
"""
@ -1546,7 +1546,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
and delta.get("type") == "text_delta"
and isinstance(delta.get("text"), str)
):
unmasked = _OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens)
unmasked = OPTIONAL_PresidioPIIMasking._unmask_pii_text(delta["text"], pii_tokens)
if unmasked != delta["text"]:
event["delta"]["text"] = unmasked
line = "data: " + json.dumps(event, ensure_ascii=False)
@ -1707,7 +1707,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
return None
def print_verbose(self, print_statement):
def print_verbose(self, print_statement) -> None:
try:
verbose_proxy_logger.debug(print_statement)
if litellm.set_verbose:
@ -1771,3 +1771,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
self.presidio_analyze_chunk_size_bytes = self._coerce_analyze_chunk_size(
litellm_params.presidio_analyze_chunk_size_bytes
)
OPTIONAL_PresidioPIIMasking = _OPTIONAL_PresidioPIIMasking

View file

@ -134,7 +134,7 @@ def _presidio_output_mode(mode: str | list[str] | Mode, *, include_mcp: bool) ->
def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> tuple[CustomGuardrail, ...]:
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
_OPTIONAL_PresidioPIIMasking,
OPTIONAL_PresidioPIIMasking,
)
explicit_filter_scope: Final = litellm_params.presidio_filter_scope
@ -163,7 +163,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) ->
params.update(overrides)
# Passed outside the heterogeneous params dict so the argument keeps
# its precise int | None type.
callback: Final = _OPTIONAL_PresidioPIIMasking(
callback: Final = OPTIONAL_PresidioPIIMasking(
presidio_analyze_chunk_size_bytes=litellm_params.presidio_analyze_chunk_size_bytes,
**params,
)

View file

@ -36,8 +36,9 @@ from litellm.proxy.guardrails.guardrail_hooks.grayswan import (
)
from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import lakeraAI_Moderation
from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import LakeraAIGuardrail
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
_OPTIONAL_PresidioPIIMasking,
from litellm.proxy.guardrails.guardrail_hooks.presidio import ( # noqa: F401 # legacy module exports
OPTIONAL_PresidioPIIMasking,
_OPTIONAL_PresidioPIIMasking, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
)
from litellm.proxy.guardrails.guardrail_hooks.tool_permission import (
ToolPermissionGuardrail,
@ -226,7 +227,7 @@ guardrail_class_registry: Final[dict[str, type[CustomGuardrail]]] = {
SupportedGuardrailIntegrations.GRAYSWAN.value: GraySwanGuardrail,
SupportedGuardrailIntegrations.LAKERA.value: lakeraAI_Moderation,
SupportedGuardrailIntegrations.LAKERA_V2.value: LakeraAIGuardrail,
SupportedGuardrailIntegrations.PRESIDIO.value: _OPTIONAL_PresidioPIIMasking,
SupportedGuardrailIntegrations.PRESIDIO.value: OPTIONAL_PresidioPIIMasking,
SupportedGuardrailIntegrations.TOOL_PERMISSION.value: ToolPermissionGuardrail,
}

View file

@ -166,7 +166,7 @@ def _get_random_llm_message():
return [{"role": "user", "content": random.choice(messages)}]
def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True):
def clean_endpoint_data(endpoint_data: Mapping[str, object], details: bool | None = True) -> dict[str, object]:
"""
Keep only the explicitly approved, JSON-safe diagnostic fields for display to users.
"""
@ -174,6 +174,9 @@ def _clean_endpoint_data(endpoint_data: dict, details: bool | None = True):
return {k: v for k, v in endpoint_data.items() if k in displayed}
_clean_endpoint_data: Final = clean_endpoint_data
def health_check_filter_kwargs_from_general_settings(
general_settings: dict | None,
) -> dict:
@ -543,7 +546,9 @@ async def _run_model_health_check(model: dict):
model_info,
litellm_params, # any-ok: untyped router config dict
)
litellm_params = _update_litellm_params_for_health_check(model_info, litellm_params)
litellm_params = update_litellm_params_for_health_check( # rebind-ok: pre-existing rebinding on a rename-only line
model_info, litellm_params
)
timeout: Final = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS
return await run_with_timeout(
@ -649,12 +654,12 @@ async def _perform_health_check(
_model_id = (model.get("model_info") or {}).get("id")
if isinstance(is_healthy, dict) and "error" not in is_healthy:
cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details)
cleaned = clean_endpoint_data({**litellm_params, **is_healthy}, details)
if _model_id:
cleaned["model_id"] = _model_id
healthy_endpoints.append(cleaned)
elif isinstance(is_healthy, dict):
cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details)
cleaned = clean_endpoint_data({**litellm_params, **is_healthy}, details)
if _model_id:
cleaned["model_id"] = _model_id
if "exception" in is_healthy:
@ -665,7 +670,7 @@ async def _perform_health_check(
cleaned["exception_status"] = getattr(exc, "status_code", 500)
unhealthy_endpoints.append(cleaned)
else:
cleaned = _clean_endpoint_data(litellm_params, details)
cleaned = clean_endpoint_data(litellm_params, details)
if _model_id:
cleaned["model_id"] = _model_id
if isinstance(is_healthy, Exception):
@ -772,7 +777,7 @@ def _resolve_health_check_max_tokens(model_info: dict, litellm_params: dict) ->
return None
def _update_litellm_params_for_health_check(model_info: dict, litellm_params: dict) -> dict:
def update_litellm_params_for_health_check(model_info: dict, litellm_params: dict) -> dict:
"""
Update the litellm params for health check.
@ -865,6 +870,9 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di
return litellm_params
_update_litellm_params_for_health_check: Final = update_litellm_params_for_health_check
async def perform_health_check(
model_list: list,
model: str | None = None,

View file

@ -53,15 +53,17 @@ from litellm.proxy.db.health_check_latest import (
query_latest_health_checks,
)
from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers
from litellm.proxy.health_check import (
from litellm.proxy.health_check import ( # noqa: F401 # legacy module exports
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS,
_clean_endpoint_data,
_update_litellm_params_for_health_check,
_clean_endpoint_data, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
_update_litellm_params_for_health_check, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
clean_endpoint_data,
deployments_targeted_by_name,
health_check_filter_kwargs_from_general_settings,
perform_health_check,
resolve_health_check_mode,
run_with_timeout,
update_litellm_params_for_health_check,
)
from litellm.proxy.middleware.admission_control_middleware import (
get_admission_control_stats,
@ -634,7 +636,7 @@ async def health_services_endpoint(
)
def _convert_health_check_to_dict(check) -> dict:
def convert_health_check_to_dict(check) -> dict:
"""Convert health check database record to dictionary format"""
return {
"health_check_id": check.health_check_id,
@ -652,6 +654,9 @@ def _convert_health_check_to_dict(check) -> dict:
}
_convert_health_check_to_dict: Final = convert_health_check_to_dict
def _check_prisma_client():
"""Helper to check if prisma_client is available and raise appropriate error"""
from litellm.proxy.proxy_server import prisma_client
@ -883,7 +888,7 @@ async def _save_health_check_results_if_changed(
return all(row is not None for row in rows)
async def _save_background_health_checks_to_db(
async def save_background_health_checks_to_db(
prisma_client,
model_list: list,
healthy_endpoints: list,
@ -941,6 +946,9 @@ async def _save_background_health_checks_to_db(
return False
_save_background_health_checks_to_db: Final = save_background_health_checks_to_db
_PROXY_ADMIN_ROLES: Final = frozenset(
{
LitellmUserRoles.PROXY_ADMIN.value,
@ -1353,7 +1361,7 @@ async def health_check_history_endpoint(
)
# Convert to dict format for JSON response using helper function
history_data: Final = [_convert_health_check_to_dict(check) for check in history]
history_data: Final = [convert_health_check_to_dict(check) for check in history]
return {
"health_checks": history_data,
@ -1385,7 +1393,7 @@ async def latest_health_checks_endpoint(
# Convert to dict format for JSON response using helper function
checks_data: Final = {
(check.model_id if check.model_id else check.model_name): _convert_health_check_to_dict(check)
(check.model_id if check.model_id else check.model_name): convert_health_check_to_dict(check)
for check in latest_checks
}
@ -2234,7 +2242,7 @@ async def test_model_connection(
stored_params=_OBJECT_MAPPING.validate_python(config_litellm_params),
request_params=_OBJECT_MAPPING.validate_python(request_litellm_params),
)
litellm_params = _update_litellm_params_for_health_check(
litellm_params = update_litellm_params_for_health_check(
model_info=dict(probe_model_info),
litellm_params=litellm_params,
)
@ -2272,7 +2280,7 @@ async def test_model_connection(
)
# Clean the result for display
cleaned_result: Final = _clean_endpoint_data({**litellm_params, **result}, details=True)
cleaned_result: Final = clean_endpoint_data({**litellm_params, **result}, details=True)
return {
"status": "error" if "error" in result else "success",

View file

@ -3,35 +3,53 @@ from typing import Final, Literal
from . import *
from .autorouter_baseline_cache import AutoRouterBaselineCache
from .cache_control_check import _PROXY_CacheControlCheck
from .cache_control_check import ( # noqa: F401 # backwards-compatible package export
PROXY_CacheControlCheck,
_PROXY_CacheControlCheck, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export
)
from .litellm_skills import SkillsInjectionHook
from .max_budget_per_session_limiter import _PROXY_MaxBudgetPerSessionHandler
from .max_iterations_limiter import _PROXY_MaxIterationsHandler
from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from .max_budget_per_session_limiter import ( # noqa: F401 # backwards-compatible package export
PROXY_MaxBudgetPerSessionHandler,
_PROXY_MaxBudgetPerSessionHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export
)
from .max_iterations_limiter import ( # noqa: F401 # backwards-compatible package export
PROXY_MaxIterationsHandler,
_PROXY_MaxIterationsHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export
)
from .parallel_request_limiter import ( # noqa: F401 # backwards-compatible package export
PROXY_MaxParallelRequestsHandler,
_PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export
)
from .parallel_request_limiter_v3 import ( # noqa: F401 # backwards-compatible package export
PROXY_MaxParallelRequestsHandler_v3,
_PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export
)
from .prompt_cache_prediction import PromptCacheObserver
from .responses_id_security import ResponsesIDSecurity
from .sensitive_data_routing import _PROXY_SensitiveDataRoutingHandler
from .sensitive_data_routing import ( # noqa: F401 # backwards-compatible package export
PROXY_SensitiveDataRoutingHandler,
_PROXY_SensitiveDataRoutingHandler, # pyright: ignore[reportPrivateUsage] # backwards-compatible package export
)
# List of all available hooks that can be enabled.
# Defined before the enterprise import below so that any module re-imported
# transitively through `enterprise.enterprise_hooks` can resolve `PROXY_HOOKS`
# and `get_proxy_hook` from this partially-initialized module without circling.
PROXY_HOOKS: Final = {
"parallel_request_limiter": _PROXY_MaxParallelRequestsHandler_v3,
"cache_control_check": _PROXY_CacheControlCheck,
"parallel_request_limiter": PROXY_MaxParallelRequestsHandler_v3,
"cache_control_check": PROXY_CacheControlCheck,
"responses_id_security": ResponsesIDSecurity,
"litellm_skills": SkillsInjectionHook,
"max_iterations_limiter": _PROXY_MaxIterationsHandler,
"max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler,
"sensitive_data_routing": _PROXY_SensitiveDataRoutingHandler,
"max_iterations_limiter": PROXY_MaxIterationsHandler,
"max_budget_per_session_limiter": PROXY_MaxBudgetPerSessionHandler,
"sensitive_data_routing": PROXY_SensitiveDataRoutingHandler,
"prompt_cache_prediction": PromptCacheObserver,
"autorouter_baseline_cache": AutoRouterBaselineCache,
}
## FEATURE FLAG HOOKS ##
if os.getenv("LEGACY_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true":
PROXY_HOOKS["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler
PROXY_HOOKS["parallel_request_limiter"] = PROXY_MaxParallelRequestsHandler
def get_proxy_hook(

View file

@ -77,7 +77,7 @@ class _PROXY_AzureContentSafety(
return result
async def test_violation(self, content: str, source: str | None = None):
async def test_violation(self, content: str, source: str | None = None) -> None:
verbose_proxy_logger.debug("Testing Azure Content-Safety for: %s", content)
# Construct a request
@ -115,7 +115,7 @@ class _PROXY_AzureContentSafety(
cache: DualCache,
data: dict,
call_type: str, # "completion", "embeddings", "image_generation", "moderation"
):
) -> None:
verbose_proxy_logger.debug("Inside Azure Content-Safety Pre-Call Hook")
try:
if is_text_content_call_type(call_type):
@ -135,7 +135,7 @@ class _PROXY_AzureContentSafety(
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response,
):
) -> None:
verbose_proxy_logger.debug("Inside Azure Content-Safety Post-Call Hook")
if not isinstance(response, litellm.ModelResponse):
return
@ -148,10 +148,13 @@ class _PROXY_AzureContentSafety(
if isinstance(content, str):
await self.test_violation(content=content, source="output")
# async def async_post_call_streaming_hook(
# self,
# user_api_key_dict: UserAPIKeyAuth,
# response: str,
# ):
# verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook")
# await self.test_violation(content=response, source="output")
PROXY_AzureContentSafety: Final = _PROXY_AzureContentSafety
# async def async_post_call_streaming_hook(
# self,
# user_api_key_dict: UserAPIKeyAuth,
# response: str,
# ):
# verbose_proxy_logger.debug("Inside Azure Content-Safety Call-Stream Hook")
# await self.test_violation(content=response, source="output")

View file

@ -148,12 +148,12 @@ def resolve_batch_enqueued_token_scopes(
def canonical_provider_batch_id(batch_id: str) -> str:
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper
get_batch_id_from_unified_batch_id,
get_original_file_id,
is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper
)
decoded: Final = _is_base64_encoded_unified_file_id(batch_id)
decoded: Final = is_base64_encoded_unified_file_id(batch_id)
if isinstance(decoded, str):
if "llm_batch_id" in decoded or "generic_response_id" in decoded:
return get_batch_id_from_unified_batch_id(decoded)

View file

@ -68,15 +68,15 @@ if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.caching.caching import DualCache
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
RateLimitDescriptor as _RateLimitDescriptor,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
RateLimitStatus as _RateLimitStatus,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _ParallelRequestLimiter,
)
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
from litellm.router import Router as _Router
from litellm.types.llms.openai import HttpxBinaryResponseContent
@ -164,16 +164,16 @@ class _PROXY_BatchRateLimiter(CustomLogger):
return None
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_models_from_unified_file_id,
is_base64_encoded_unified_file_id,
)
model_from_file_id: Final = decode_model_from_file_id(input_file_id)
if model_from_file_id:
return model_from_file_id
unified_file_id: Final = _is_base64_encoded_unified_file_id(input_file_id)
unified_file_id: Final = is_base64_encoded_unified_file_id(input_file_id)
if unified_file_id:
target_model_names: Final = get_models_from_unified_file_id(unified_file_id)
if target_model_names:
@ -250,7 +250,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
minute with the submission. The daily descriptor uses its own key so
its 24h window never collides with the online limiter's counters.
"""
descriptors: Final = self.parallel_request_limiter._create_rate_limit_descriptors(
descriptors: Final = self.parallel_request_limiter.create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data=data,
rpm_limit_type=None,
@ -892,13 +892,13 @@ class _PROXY_BatchRateLimiter(CustomLogger):
try:
# Check if this is a managed file (base64 encoded unified file ID)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
get_models_from_unified_file_id,
is_base64_encoded_unified_file_id,
)
# Managed files require bypassing the HTTP endpoint (which runs access-check hooks)
# and calling the managed files hook directly with the user's credentials.
is_managed_file: Final = _is_base64_encoded_unified_file_id(file_id)
is_managed_file: Final = is_base64_encoded_unified_file_id(file_id)
# For managed files the unified file id encodes the proxy model
# alias(es) the file was uploaded for; auth validates against those.
target_model_names: Final = get_models_from_unified_file_id(is_managed_file) if is_managed_file else []
@ -1039,11 +1039,11 @@ class _PROXY_BatchRateLimiter(CustomLogger):
enforces on `/chat/completions` apply here.
"""
from litellm.proxy.auth.auth_checks import (
_check_team_member_model_access,
_key_access_group_grants_model,
can_key_call_model,
can_team_access_model,
check_team_member_model_access,
get_team_object,
key_access_group_grants_model,
)
from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_logging_obj, user_api_key_cache
@ -1092,14 +1092,14 @@ class _PROXY_BatchRateLimiter(CustomLogger):
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
raise
if not await _key_access_group_grants_model(
if not await key_access_group_grants_model(
model=model_to_check,
valid_token=user_api_key_dict,
team_object=team_object,
llm_router=llm_router,
):
raise
await _check_team_member_model_access(
await check_team_member_model_access(
model=model_to_check,
team_object=team_object,
valid_token=user_api_key_dict,
@ -1281,3 +1281,6 @@ class _PROXY_BatchRateLimiter(CustomLogger):
verbose_proxy_logger.error("Error in batch rate limiting: %s", e, exc_info=True)
# Don't block the request if rate limiting fails
return data
PROXY_BatchRateLimiter = _PROXY_BatchRateLimiter

View file

@ -25,7 +25,7 @@ class _PROXY_BatchRedisRequests(CustomLogger):
self.async_get_cache
) # map the litellm 'get_cache' function to our custom function
def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG"):
def print_verbose(self, print_statement, debug_level: Literal["INFO", "DEBUG"] = "DEBUG") -> None:
if debug_level == "DEBUG" or debug_level == "INFO":
verbose_proxy_logger.debug(print_statement)
if litellm.set_verbose is True:
@ -37,7 +37,7 @@ class _PROXY_BatchRedisRequests(CustomLogger):
cache: DualCache,
data: dict,
call_type: str,
):
) -> None:
try:
"""
Get the user key
@ -83,7 +83,7 @@ class _PROXY_BatchRedisRequests(CustomLogger):
)
verbose_proxy_logger.debug(traceback.format_exc())
async def async_get_cache(self, *args, **kwargs):
async def async_get_cache(self, *args, **kwargs) -> object | None:
"""
- Check if the cache key is in-memory
@ -113,3 +113,6 @@ class _PROXY_BatchRedisRequests(CustomLogger):
return litellm.cache.get_cache_logic(cached_result=cached_result, max_age=max_age)
except Exception:
return None
PROXY_BatchRedisRequests: Final = _PROXY_BatchRedisRequests

View file

@ -24,7 +24,7 @@ class _PROXY_CacheControlCheck(CustomLogger):
cache: DualCache,
data: dict,
call_type: str,
):
) -> None:
try:
verbose_proxy_logger.debug("Inside Cache Control Check Pre-Call Hook")
allowed_cache_controls: Final = user_api_key_dict.allowed_cache_controls
@ -56,3 +56,6 @@ class _PROXY_CacheControlCheck(CustomLogger):
verbose_logger.exception(
"litellm.proxy.hooks.cache_control_check.py::async_pre_call_hook(): Exception occured - %s", e
)
PROXY_CacheControlCheck: Final = _PROXY_CacheControlCheck

View file

@ -1,7 +1,6 @@
# What is this?
## Allocates dynamic tpm/rpm quota for a project based on current traffic
## Tracks num active projects per minute
import asyncio
import os
from collections.abc import Callable
@ -23,7 +22,7 @@ from litellm.proxy.hooks.rate_limiter_utils import (
resolve_llm_provider_for_rate_limit,
)
from litellm.types.router import ModelGroupInfo
from litellm.types.utils import CallTypesLiteral
from litellm.types.utils import CallTypesLiteral, LLMResponseTypes
from litellm.utils import get_utc_datetime
@ -83,7 +82,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
def __init__(self, internal_usage_cache: DualCache, time_fn: Callable[[], datetime] = get_utc_datetime):
self.internal_usage_cache = DynamicRateLimiterCache(cache=internal_usage_cache, time_fn=time_fn)
def update_variables(self, llm_router: Router):
def update_variables(self, llm_router: Router) -> None:
self.llm_router = llm_router
@with_service_target("rate_limits")
@ -241,7 +240,9 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
return None
@with_service_target("rate_limits")
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
) -> LLMResponseTypes | None:
try:
if isinstance(response, ModelResponse):
model_id: Final = response.hidden_params["model_id"]
@ -281,3 +282,6 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
"litellm.proxy.hooks.dynamic_rate_limiter.py::async_post_call_success_hook(): Exception occured - %s", e
)
return response
PROXY_DynamicRateLimitHandler: Final = _PROXY_DynamicRateLimitHandler

View file

@ -20,11 +20,12 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import (
ProxyRateLimitError,
map_v3_rate_limit_type,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
from litellm.proxy.hooks.parallel_request_limiter_v3 import ( # noqa: F401 # legacy module exports
PROXY_MaxParallelRequestsHandler_v3,
RateLimitDescriptor,
RateLimitDescriptorRateLimitObject,
RateLimitResponse,
_PROXY_MaxParallelRequestsHandler_v3,
_PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
claim_request_stash_for_data,
get_or_create_request_stash,
)
@ -38,7 +39,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
response_has_hidden_params,
)
from litellm.types.router import ModelGroupInfo
from litellm.types.utils import CallTypesLiteral
from litellm.types.utils import CallTypesLiteral, LLMResponseTypes
if TYPE_CHECKING:
from litellm.types.utils import PriorityReservationSettings
@ -90,9 +91,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
time_provider: Callable[[], datetime] | None = None,
):
self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache)
self.v3_limiter = _PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider)
self.v3_limiter = PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider)
def update_variables(self, llm_router: Router):
def update_variables(self, llm_router: Router) -> None:
self.llm_router = llm_router
def _get_saturation_check_cache_ttl(self) -> int:
@ -659,7 +660,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
return None
@with_service_target("rate_limits")
async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
async def async_post_call_success_hook(
self, data: dict, user_api_key_dict: UserAPIKeyAuth, response
) -> LLMResponseTypes:
"""
Post-call hook to add rate limit headers to response.
Leverages v3 limiter's post-call hook functionality.
@ -689,7 +692,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
return response
@with_service_target("rate_limits")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
"""
Update token usage for priority-based rate limiting after successful API calls.
@ -804,3 +807,6 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
except Exception as e:
verbose_proxy_logger.exception("Error in dynamic rate limiter success event: %s", e)
PROXY_DynamicRateLimitHandlerV3: Final = _PROXY_DynamicRateLimitHandlerV3

View file

@ -21,7 +21,10 @@ from litellm.proxy._types import (
UpdateKeyRequest,
UserAPIKeyAuth,
)
from litellm.proxy.utils import _hash_token_if_needed
from litellm.proxy.utils import ( # noqa: F401 # legacy module exports
_hash_token_if_needed, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
hash_token_if_needed,
)
from litellm.secret_managers.base_secret_manager import BaseSecretManager
if TYPE_CHECKING:
@ -140,7 +143,7 @@ class KeyManagementEventHooks:
),
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=_hash_token_if_needed(data.key),
object_id=hash_token_if_needed(data.key),
action="updated",
updated_values=json.dumps(updated_fields, default=str),
before_value=json.dumps(existing_key_row.json(exclude_none=True), default=str),

View file

@ -130,7 +130,7 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
return None
@with_service_target("session_budgets")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
"""
After a successful LLM call, increment the session spend by the response cost.
"""
@ -271,3 +271,6 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
local_only=True,
)
return new_value
PROXY_MaxBudgetPerSessionHandler: Final = _PROXY_MaxBudgetPerSessionHandler

View file

@ -221,3 +221,6 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
local_only=True,
)
return new_value
PROXY_MaxIterationsHandler: Final = _PROXY_MaxIterationsHandler

View file

@ -169,7 +169,7 @@ class SemanticToolFilterHook(CustomLogger):
def _selected_tool_names(self, filtered_tools: Sequence[object]) -> list[str]:
"""Names of the semantically selected tools, as produced by the MCP expansion."""
names: Final = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools)
names: Final = (self.filter.extract_tool_info(tool)[0] for tool in filtered_tools)
return [name for name in names if name]
@staticmethod
@ -219,7 +219,7 @@ class SemanticToolFilterHook(CustomLogger):
return False
if isinstance(tool, dict) and tool.get("type") == "function" and isinstance(tool.get("name"), str):
return False
name, _ = self.filter._extract_tool_info(tool)
name, _ = self.filter.extract_tool_info(tool)
return bool(name) and name in self.filter._tool_map
def _get_metadata_variable_name(self, data: dict) -> str:
@ -397,14 +397,14 @@ class SemanticToolFilterHook(CustomLogger):
filtered_mcp_names: Final[set[str]] = set()
for t in filtered_mcp_tools:
name, _ = self.filter._extract_tool_info(t)
name, _ = self.filter.extract_tool_info(t)
if name:
filtered_mcp_names.add(name)
filtered_tools: Final[list[object]] = []
for i, t in enumerate(tools):
if i in mcp_indices:
name, _ = self.filter._extract_tool_info(t)
name, _ = self.filter.extract_tool_info(t)
if name in filtered_mcp_names:
filtered_tools.append(t)
else:

Some files were not shown because too many files have changed in this diff Show more