mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
b970e412d9
commit
58258409c9
377 changed files with 8886 additions and 4979 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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="",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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``
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -221,3 +221,6 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
local_only=True,
|
||||
)
|
||||
return new_value
|
||||
|
||||
|
||||
PROXY_MaxIterationsHandler: Final = _PROXY_MaxIterationsHandler
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue