diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261005000000_add_mcp_server_rpm/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261005000000_add_mcp_server_rpm/migration.sql new file mode 100644 index 00000000000..9ba76aa6db7 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261005000000_add_mcp_server_rpm/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "rpm" INTEGER; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index dbd44934ca8..dbfa8c9ce92 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -417,6 +417,7 @@ model LiteLLM_MCPServerTable { source_url String? timeout Float? max_concurrent_requests Int? + rpm Int? // BYOM submission lifecycle approval_status String? @default("active") submitted_by String? diff --git a/litellm/models/mcp_server.py b/litellm/models/mcp_server.py index dac79145644..cdbda7b1b70 100644 --- a/litellm/models/mcp_server.py +++ b/litellm/models/mcp_server.py @@ -111,6 +111,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): source_url: str | None = None timeout: float | None = None max_concurrent_requests: int | None = None + rpm: int | None = None approval_status: str | None = Field( default="active", description="Approval status: 'pending_review', 'active', 'rejected'", diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index a957b352705..9331f017582 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -952,24 +952,48 @@ async def aggregate_gateway_tools( prefetched: Mapping[str, OAuthCredentialPayload], *, record_listing: bool = False, + enforce_rate_limits: bool = True, ) -> AggregateToolListing: import time - from mcp.types import PaginatedRequestParams + from mcp.types import ListToolsResult, PaginatedRequestParams from pydantic import TypeAdapter from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( SERVER_OUTCOMES_META_KEY, AggregateToolListing, ServerOutcome, + classify_list_exception, ) - from litellm.proxy._experimental.mcp_server.operations import _aggregate_server_key, global_mcp_server_manager + from litellm.proxy._experimental.mcp_server.operations import ( + _aggregate_server_key, + _mcp_server_rate_limit_rejection, + global_mcp_server_manager, + ) + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError async with global_mcp_server_manager.catalog.operation() as snapshot: servers: Final = {server.server_id: server for server in allowed} listing_updates: Final = ExitStack() + rejections: Final[list[ProxyRateLimitError]] = [] # mutable-ok: concurrent fetches share first-page errors async def fetch(server_id: str, cursor: str | None) -> ListToolsResult: + if enforce_rate_limits: + error: Final = await _mcp_server_rate_limit_rejection(servers[server_id], context.user_api_key_auth) + if error is not None: + if cursor is not None: + raise error + rejections.append(error) + return ListToolsResult( + tools=[], + _meta={ + SERVER_OUTCOMES_META_KEY: { + _aggregate_server_key(servers[server_id]): classify_list_exception(error).model_dump( + mode="json" + ) + } + }, + ) result, outcome = await get_filtered_server_tools( servers[server_id], context=context, @@ -1003,6 +1027,8 @@ async def aggregate_gateway_tools( fetch=fetch, now=int(time.time()), ) + if params.cursor is None and servers and len(rejections) == len(servers): + raise rejections[0] listing_updates.close() return AggregateToolListing( tools=result.tools, @@ -1063,6 +1089,7 @@ async def list_gateway_catalog( global_mcp_server_manager, raise_denied_scoped_mcp_access, ) + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError context = replace(context, _caller=await MCPRequestHandler.refresh_catalog_authority(context.user_api_key_auth)) params: Final = request.params or PaginatedRequestParams() @@ -1080,6 +1107,7 @@ async def list_gateway_catalog( requested_names=list(scope), user_api_key_auth=caller, client_ip=client_ip ) servers: Final = {server.server_id: server for server in allowed} + rejections: Final[list[ProxyRateLimitError]] = [] # mutable-ok: concurrent fetches share first-page errors async def fetch(server_id: str, cursor: str | None) -> CatalogListResult: server: Final = servers[server_id] @@ -1090,7 +1118,26 @@ async def list_gateway_catalog( SERVER_OUTCOMES_META_KEY, classify_list_exception, ) - from litellm.proxy._experimental.mcp_server.operations import _aggregate_server_key + from litellm.proxy._experimental.mcp_server.operations import ( + _aggregate_server_key, + _mcp_server_rate_limit_rejection, + ) + + error: Final = await _mcp_server_rate_limit_rejection(server, caller) + if error is not None: + if cursor is not None: + raise error + rejections.append(error) + return combine_optional_catalog( + request, + (), + None, + { + SERVER_OUTCOMES_META_KEY: { + _aggregate_server_key(server): classify_list_exception(error).model_dump(mode="json") + } + }, + ) try: page: Final = await fetch_optional_catalog_page(context, request, server, allowed, cursor) @@ -1126,6 +1173,8 @@ async def list_gateway_catalog( fetch=fetch, now=int(time.time()), ) + if params.cursor is None and servers and len(rejections) == len(servers): + raise rejections[0] from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( SERVER_OUTCOMES_META_KEY, ServerOutcome, diff --git a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py index 0c1f7599718..6c74ef68118 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py +++ b/litellm/proxy/_experimental/mcp_server/faults/list_outcomes.py @@ -24,11 +24,13 @@ from litellm.proxy._experimental.mcp_server.exceptions import ( MCPUpstreamAuthError, ) from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.types.llms.base import LiteLLMBaseModel ListFaultCategory: TypeAlias = Literal[ "auth_required", "forbidden", + "rate_limited", "timeout", "unreachable", "upstream_error", @@ -125,6 +127,8 @@ def classify_list_exception(exc: BaseException) -> ServerListFault: if isinstance(exc, MCPUpstreamAuthError): tag: Final = "forbidden" if exc.status_code == 403 else "auth_required" return ServerListFault(tag=tag, status_code=exc.status_code) + if isinstance(exc, ProxyRateLimitError): + return ServerListFault(tag="rate_limited", status_code=429) if isinstance(exc, TimeoutError): return ServerListFault(tag="timeout") if isinstance(exc, ConnectionError): @@ -152,7 +156,7 @@ def outcome_wire_value(outcome: ServerOutcome) -> dict[str, object]: match outcome.tag: case "ok": return {"status": "ok", "tool_count": outcome.tool_count} - case "auth_required" | "forbidden" | "timeout" | "unreachable" | "upstream_error" | "internal": + case "auth_required" | "forbidden" | "rate_limited" | "timeout" | "unreachable" | "upstream_error" | "internal": return { "status": outcome.tag, **({"http_status": outcome.status_code} if outcome.status_code is not None else {}), @@ -170,6 +174,8 @@ def list_fault_http_status(fault: ServerListFault) -> int: return fault.status_code or 401 case "forbidden": return 403 + case "rate_limited": + return 429 case "timeout": return 504 case "unreachable" | "upstream_error": diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index cb0eed8823c..f057c471a26 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -426,6 +426,7 @@ class MCPServerConfig(TypedDict, total=False): client_assertion_signing_alg: str timeout: float max_concurrent_requests: int + rpm: ReadOnly[int | None] class _ProtectedResourceMetadataPayload(TypedDict, total=False): @@ -2738,6 +2739,7 @@ class MCPServerManager: allow_elicitation=bool(server_config.get("allow_elicitation", False)), timeout=server_config.get("timeout", None), max_concurrent_requests=server_config.get("max_concurrent_requests", None), + rpm=server_config.get("rpm", None), token_validation=server_config.get("token_validation", None), oauth_identity_binding=server_config.get("oauth_identity_binding", None), ) @@ -3327,6 +3329,7 @@ class MCPServerManager: or "rfc8693", timeout=getattr(mcp_server, "timeout", None), max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None), + rpm=getattr(mcp_server, "rpm", None), ) _warn_legacy_delegate_auth_if_applicable(new_server, source="database") if register_oauth_discovery: @@ -5973,6 +5976,7 @@ class MCPServerManager: data=synthetic_llm_data, call_type=CallTypes.call_mcp_tool.value, ) + 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) @@ -7193,6 +7197,7 @@ class MCPServerManager: instructions=server.instructions, timeout=server.timeout, max_concurrent_requests=server.max_concurrent_requests, + rpm=server.rpm, ) async def get_all_mcp_servers_with_health_and_teams( @@ -7316,6 +7321,7 @@ class MCPServerManager: instructions=server.instructions, timeout=server.timeout, max_concurrent_requests=server.max_concurrent_requests, + rpm=server.rpm, ) async def get_all_mcp_servers_unfiltered(self) -> list[LiteLLM_MCPServerTable]: diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 2164dd332ac..e226b7f3fcb 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -5,6 +5,8 @@ import traceback import types import uuid from collections.abc import Mapping, Sequence +from contextvars import ContextVar +from dataclasses import dataclass from datetime import datetime from functools import partial from typing import Any, Final, NoReturn, TypeAlias, overload @@ -78,6 +80,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( AggregateToolListing, ServerListOk, ServerOutcome, + classify_list_exception, outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -128,6 +131,7 @@ from litellm.proxy._types import ( from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( publish_auth_cache_invalidation, ) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, get_chain_id_from_headers, @@ -224,6 +228,65 @@ class ListMCPToolsRestAPIResponseObject(MCPTool): model_config = ConfigDict(arbitrary_types_allowed=True) +@dataclass(frozen=True, slots=True) +class _MCPServerRateLimitAdmission: + admitted_servers: tuple[MCPServer, ...] + rejected_servers: tuple[tuple[MCPServer, ProxyRateLimitError], ...] + + +_mcp_server_admission_memo: Final[ContextVar[dict[str, asyncio.Task[ProxyRateLimitError | None]] | None]] = ContextVar( + "mcp_server_admission_memo", default=None +) + + +async def _enforce_mcp_server_rate_limit( + user_api_key_auth: UserAPIKeyAuth | None, + server: MCPServer, +) -> None: + from litellm.proxy.proxy_server import proxy_logging_obj + + if proxy_logging_obj is not None: + await proxy_logging_obj.enforce_mcp_server_rate_limits(user_api_key_auth, server) + + +async def _admit_mcp_servers( + servers: Sequence[MCPServer], + user_api_key_auth: UserAPIKeyAuth | None, +) -> _MCPServerRateLimitAdmission: + memo: Final = _mcp_server_admission_memo.get() + + async def _server_rate_limit_error(server: MCPServer) -> ProxyRateLimitError | None: + try: + await _enforce_mcp_server_rate_limit(user_api_key_auth, server) + except ProxyRateLimitError as error: + return error + return None + + async def _admit_server(server: MCPServer) -> tuple[MCPServer, ProxyRateLimitError | None]: + if memo is None: + return server, await _server_rate_limit_error(server) + admission_task: Final = memo.get(server.server_id) + if admission_task is not None: + return server, await admission_task + created_task: Final = asyncio.create_task(_server_rate_limit_error(server)) + memo[server.server_id] = created_task + return server, await created_task + + results: Final = await asyncio.gather(*(_admit_server(server) for server in servers)) + return _MCPServerRateLimitAdmission( + admitted_servers=tuple(server for server, error in results if error is None), + rejected_servers=tuple((server, error) for server, error in results if error is not None), + ) + + +async def _mcp_server_rate_limit_rejection( + server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, +) -> ProxyRateLimitError | None: + admission: Final = await _admit_mcp_servers((server,), user_api_key_auth) + return admission.rejected_servers[0][1] if admission.rejected_servers else None + + async def _build_virtual_call_logging_obj( name: str, arguments: dict[str, object], @@ -961,6 +1024,7 @@ async def _get_tools_from_mcp_servers( protocol_version: str | None = None, *, record_listing: bool = False, + enforce_rate_limits: bool = True, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -1093,20 +1157,37 @@ async def _get_tools_from_mcp_servers( return page.tools, outcome if params is None: + server_admission: Final = ( + await _admit_mcp_servers(allowed_mcp_servers, user_api_key_auth) + if enforce_rate_limits + else _MCPServerRateLimitAdmission(tuple(allowed_mcp_servers), ()) + ) + if not server_admission.admitted_servers and server_admission.rejected_servers: + raise server_admission.rejected_servers[0][1] + admitted_servers: Final = server_admission.admitted_servers results: Final = await asyncio.gather( - *(_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers) + *(_fetch_and_filter_server_tools(server) for server in admitted_servers) ) aggregated = AggregateToolListing( tools=[tool for tools, _ in results for tool in tools], outcomes={ - _aggregate_server_key(server): outcome for server, (_, outcome) in zip(allowed_mcp_servers, results) + _aggregate_server_key(server): outcome for server, (_, outcome) in zip(admitted_servers, results) + } + | { + _aggregate_server_key(server): classify_list_exception(error) + for server, error in server_admission.rejected_servers }, ) else: from litellm.proxy._experimental.mcp_server.catalog import aggregate_gateway_tools aggregated = await aggregate_gateway_tools( - context, params, allowed_mcp_servers, _prefetched_oauth_creds, record_listing=record_listing + context, + params, + allowed_mcp_servers, + _prefetched_oauth_creds, + record_listing=record_listing, + enforce_rate_limits=enforce_rate_limits, ) all_tools: Final = aggregated.tools server_outcomes: Final = aggregated.outcomes @@ -1751,6 +1832,7 @@ async def _list_tools_before_first_call( raw_headers=raw_headers, client_ip=client_ip, record_listing=False, + enforce_rate_limits=False, ) except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e) @@ -2464,6 +2546,7 @@ async def mcp_get_prompt( user_api_key_auth=user_api_key_auth, ) + await _enforce_mcp_server_rate_limit(user_api_key_auth, server) return await global_mcp_server_manager.get_prompt_from_server( server=server, user_api_key_auth=user_api_key_auth, @@ -2517,6 +2600,7 @@ async def mcp_read_resource( user_api_key_auth=user_api_key_auth, ) + await _enforce_mcp_server_rate_limit(user_api_key_auth, server) return await global_mcp_server_manager.read_resource_from_server( server=server, user_api_key_auth=user_api_key_auth, @@ -3125,6 +3209,7 @@ class GatewayOperations: if context.mcp_proxy_mode else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest()) ) + memo_token: Final = _mcp_server_admission_memo.set({}) tasks: Final = ( asyncio.create_task( _execute_handle_list_tools( @@ -3139,9 +3224,12 @@ class GatewayOperations: try: results: Final = await asyncio.gather(*tasks) finally: - for task in tasks: - task.cancel() - await asyncio.gather(*tasks, return_exceptions=True) + try: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + finally: + _mcp_server_admission_memo.reset(memo_token) return build_discovery( configured=configured_versions(), revision=context.protocol_version or "2025-11-25", diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index e2807d06fbb..4a1ae19b2ca 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -55,6 +55,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.responses.mcp.request_context import MCPRequestContext if TYPE_CHECKING: @@ -737,6 +738,8 @@ if MCP_AVAILABLE: """ from litellm.proxy.proxy_server import proxy_logging_obj + if apply_tool_filters and proxy_logging_obj is not None: + await proxy_logging_obj.enforce_mcp_server_rate_limits(user_api_key_auth, server) listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) tools: Final = await _list_server_tools( server, @@ -900,6 +903,8 @@ if MCP_AVAILABLE: # matching status code and WWW-Authenticate challenge; that is what # lets standards-compliant MCP clients run the upstream OAuth flow. raise + except ProxyRateLimitError: + raise except MCPServerListError as e: fault: Final = classify_list_exception(e) verbose_logger.info("Listing tools from %s failed with a %s fault", server.name, fault.tag) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 9e4773390bf..9ea05d65129 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -31962,6 +31962,18 @@ ], "title": "Registration Url" }, + "rpm": { + "anyOf": [ + { + "minimum": 0.0, + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "anyOf": [ { @@ -33814,6 +33826,17 @@ ], "title": "Reviewed At" }, + "rpm": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "title": "Server Id", "type": "string" @@ -34700,6 +34723,18 @@ ], "title": "Registration Url" }, + "rpm": { + "anyOf": [ + { + "minimum": 0.0, + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "anyOf": [ { @@ -37098,6 +37133,17 @@ ], "title": "Reviewed At" }, + "rpm": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "title": "Server Id", "type": "string" @@ -38828,6 +38874,18 @@ ], "title": "Registration Url" }, + "rpm": { + "anyOf": [ + { + "minimum": 0.0, + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "anyOf": [ { @@ -39371,6 +39429,18 @@ ], "title": "Registration Url" }, + "rpm": { + "anyOf": [ + { + "minimum": 0.0, + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "title": "Server Id", "type": "string" @@ -42217,6 +42287,18 @@ ], "title": "Registration Url" }, + "rpm": { + "anyOf": [ + { + "minimum": 0.0, + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Rpm" + }, "server_id": { "anyOf": [ { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 1627c5b63ec..0ed2b38cfa6 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1675,6 +1675,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): source_url: str | None = None timeout: float | None = None max_concurrent_requests: int | None = None + rpm: int | None = Field(default=None, ge=0) # BYOM submission fields — set by the endpoint, not by the caller. # Any caller-provided values are silently overridden before persistence. approval_status: str | None = Field( @@ -1782,6 +1783,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): source_url: str | None = None timeout: float | None = None max_concurrent_requests: int | None = None + rpm: int | None = Field(default=None, ge=0) @model_validator(mode="after") def validate_protocol_transport(self) -> "UpdateMCPServerRequest": diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 3430a4a8863..1a3452d9f87 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -76,6 +76,7 @@ from litellm.router_utils.add_retry_fallback_headers import ( from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType from litellm.types.utils import ( CallTypes, @@ -2984,6 +2985,48 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + async def enforce_mcp_server_rate_limits( + self, + user_api_key_dict: UserAPIKeyAuth | None, + server: MCPServer, + ) -> None: + descriptors: Final[list[RateLimitDescriptor]] = [] # mutable-ok: existing descriptor helpers append in place + mcp_server_name: Final = server.alias or server.server_name or server.name + if user_api_key_dict is not None: + self._add_mcp_per_key_rate_limit_descriptor( + user_api_key_dict=user_api_key_dict, + mcp_server_name=mcp_server_name, + descriptors=descriptors, + ) + self._add_mcp_per_team_rate_limit_descriptor( + user_api_key_dict=user_api_key_dict, + mcp_server_name=mcp_server_name, + descriptors=descriptors, + ) + if server.rpm is not None: + descriptors.append( + RateLimitDescriptor( + key="mcp_server", + value=server.server_id, + rate_limit={ + "requests_per_unit": server.rpm, + "tokens_per_unit": None, + "window_size": self.window_size, + }, + ) + ) + if not descriptors: + return + + parent_otel_span: Final = user_api_key_dict.parent_otel_span if user_api_key_dict is not None else None + response: Final = await self.atomic_check_and_increment_by_n( + descriptors=descriptors, + increments=[{"requests": 1} for _ in descriptors], + parent_otel_span=parent_otel_span, + ) + if response["overall_code"] == "OVER_LIMIT": + self._handle_rate_limit_error(response, descriptors) + def _should_enforce_rate_limit( self, limit_type: str | None, @@ -3270,21 +3313,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptors=descriptors, ) - # REST MCP calls pass the raw body through this hook before server - # resolution; only the later synthetic hook payload may carry this key. - if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data: - mcp_server_name: Final = data.get("mcp_server_name", None) - self._add_mcp_per_key_rate_limit_descriptor( - user_api_key_dict=user_api_key_dict, - mcp_server_name=mcp_server_name, - descriptors=descriptors, - ) - self._add_mcp_per_team_rate_limit_descriptor( - user_api_key_dict=user_api_key_dict, - mcp_server_name=mcp_server_name, - descriptors=descriptors, - ) - self._add_team_model_rate_limit_descriptor_from_metadata( user_api_key_dict=user_api_key_dict, requested_model=requested_model if isinstance(requested_model, str) else None, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 6391cadd96d..82022639095 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1048,6 +1048,7 @@ if MCP_AVAILABLE: available_on_public_internet=payload.available_on_public_internet, timeout=payload.timeout, max_concurrent_requests=payload.max_concurrent_requests, + rpm=payload.rpm, ) def get_prisma_client_or_throw(message: str): diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index dbd44934ca8..dbfa8c9ce92 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -417,6 +417,7 @@ model LiteLLM_MCPServerTable { source_url String? timeout Float? max_concurrent_requests Int? + rpm Int? // BYOM submission lifecycle approval_status String? @default("active") submitted_by String? diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b775a1770b3..3a3c3fa2928 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -274,6 +274,7 @@ if TYPE_CHECKING: from litellm.proxy.db.model_usage_rollup import ModelUsageTransaction from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction from litellm.repositories.prisma_protocols import TableActions + from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline Span = _Span | object @@ -4210,6 +4211,16 @@ class ProxyLogging: return await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict) + async def enforce_mcp_server_rate_limits( + self, + user_api_key_dict: UserAPIKeyAuth | None, + server: "MCPServer", + ) -> None: + limiter: Final = self.get_proxy_hook("parallel_request_limiter") + if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + return + await limiter.enforce_mcp_server_rate_limits(user_api_key_dict, server) + def _init_response_taking_too_long_task(self, data: dict | None = None): """ Initialize the response taking too long task if user is using slack alerting diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 5b1519bbccb..2f5d983cdc0 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -237,6 +237,7 @@ class MCPServer(LiteLLMBaseModel): # Max concurrent outbound tool calls to this server; excess calls queue. # None or a value <= 0 means unlimited. max_concurrent_requests: int | None = None + rpm: int | None = None # Resolved short-ID tool prefix when LITELLM_USE_SHORT_MCP_TOOL_PREFIX is # enabled. Set by ``MCPServerManager.assign_unique_short_prefix`` at # registration time so that natural-hash collisions between two diff --git a/schema.prisma b/schema.prisma index dbd44934ca8..dbfa8c9ce92 100644 --- a/schema.prisma +++ b/schema.prisma @@ -417,6 +417,7 @@ model LiteLLM_MCPServerTable { source_url String? timeout Float? max_concurrent_requests Int? + rpm Int? // BYOM submission lifecycle approval_status String? @default("active") submitted_by String? diff --git a/tests/integration/mcp/test_mcp_rate_limits.py b/tests/integration/mcp/test_mcp_rate_limits.py new file mode 100644 index 00000000000..273419d624f --- /dev/null +++ b/tests/integration/mcp/test_mcp_rate_limits.py @@ -0,0 +1,271 @@ +import asyncio +import os +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import scratch_database +from integration._support.mcp import McpCaller, McpPeer, mcp_peer, paginated_mcp_peer, register_mcp, tool_calls +from integration._support.process import owned_proxy +from integration._support.redis_process import OwnedRedis, owned_redis +from mcp import ClientSession, MCPError +from mcp.client.streamable_http import streamable_http_client +from mcp.types import ListToolsResult, PaginatedRequestParams + +REMOVE_DATABASE: Final = ("DATABASE_URL", "DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH") + + +@asynccontextmanager +async def _catalog_session(gateway: Gateway) -> AsyncIterator[ClientSession]: + async with httpx.AsyncClient( + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=15, + trust_env=False, + ) as client: + async with streamable_http_client( + f"{str(gateway.client.base_url).rstrip('/')}/mcp/", + http_client=client, + ) as streams: + async with ClientSession(streams[0], streams[1]) as session: + await session.initialize() + yield session + + +async def _list_tools(gateway: Gateway, cursor: str | None = None) -> ListToolsResult | MCPError: + async with _catalog_session(gateway) as session: + try: + if cursor is None: + return await session.list_tools() + return await session.list_tools(params=PaginatedRequestParams(cursor=cursor)) + except MCPError as error: + return error + + +def _tool_items(result: ListToolsResult) -> tuple[dict[str, object], ...]: + return tuple(tool.model_dump(mode="json") for tool in result.tools) + + +def _has_method(calls: tuple[dict[str, object], ...], method: str) -> bool: + return any( + isinstance(call.get("body"), dict) and isinstance(call["body"], dict) and call["body"].get("method") == method + for call in calls + ) + + +def _config_file( + directory: Path, + master_key: str, + redis: OwnedRedis, + *, + upstream: str | None = None, + rpm: int | None = None, + allowed_tools: tuple[str, ...] = (), + store_model_in_db: bool, +) -> Path: + config: dict[str, object] = { + "model_list": [], + "general_settings": { + "master_key": master_key, + "store_model_in_db": store_model_in_db, + "coordination_redis": {"host": redis.host, "port": redis.port}, + }, + } + if upstream is not None: + server: dict[str, object] = {"url": upstream, "transport": "http", "rpm": rpm} + if allowed_tools: + server["allowed_tools"] = list(allowed_tools) + config["mcp_servers"] = {"rpm": server} + path: Final = directory / "proxy.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _call_tool(gateway: Gateway, key: str, server_id: str, name: str) -> httpx.Response: + return gateway.client.post( + "/mcp-rest/tools/call", + headers={"x-litellm-api-key": key}, + json={"name": name, "arguments": {"a": 1, "b": 2}, "server_id": server_id}, + ) + + +def _assert_rate_limit(response: httpx.Response, descriptor: str) -> None: + assert response.status_code == 429, response.text + assert descriptor in response.text + + +def test_shared_redis_enforces_paginated_tools_and_rest_listings( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + results_directory: Final = tmp_path / "results" + results_directory.mkdir() + monkeypatch.setenv("INTEGRATION_RESULTS_DIR", str(results_directory)) + + async def exercise( + first_replica: Gateway, second_replica: Gateway, peer: McpPeer + ) -> tuple[str, tuple[dict[str, object], ...]]: + first_page: Final = await _list_tools(first_replica) + assert isinstance(first_page, ListToolsResult) + assert first_page.next_cursor is not None + first_cursor: Final = first_page.next_cursor + + continued_page: Final = await _list_tools(second_replica, first_cursor) + assert isinstance(continued_page, ListToolsResult) + continued_items: Final = _tool_items(continued_page) + assert continued_items + + peer.drain() + rejected: Final = await _list_tools(first_replica, first_cursor) + assert isinstance(rejected, MCPError) + assert "mcp_server" in str(rejected) + assert not _has_method(peer.drain(), "tools/list") + + peer.drain() + rest_rejected: Final = second_replica.client.get( + "/mcp-rest/tools/list", + headers={"x-litellm-api-key": second_replica.key}, + params={"server_id": "rpm"}, + ) + assert rest_rejected.status_code == 429, rest_rejected.text + assert not _has_method(peer.drain(), "tools/list") + + return first_cursor, continued_items + + with paginated_mcp_peer(page_size=1) as peer, owned_redis(tmp_path) as redis, httpx.Client() as client: + seed: Final = Gateway(client, "sk-mcp-pagination-rate-limit", peer.url) + config: Final = _config_file( + tmp_path, + seed.key, + redis, + upstream=peer.url, + rpm=2, + store_model_in_db=False, + ) + environment: Final = { + "STORE_MODEL_IN_DB": "False", + "DISABLE_SCHEMA_UPDATE": "true", + "LITELLM_SALT_KEY": "shared-mcp-pagination-rate-limit", + "LITELLM_RATE_LIMIT_WINDOW_SIZE": "10", + } + options: Final = { + "config": config, + "database_setup": (), + "remove_environment": REMOVE_DATABASE, + } + with ( + owned_proxy(seed, tmp_path / "first", environment, **options) as first_replica, + owned_proxy(seed, tmp_path / "second", environment, **options) as second_replica, + ): + first_cursor, continued_items = asyncio.run(exercise(first_replica, second_replica, peer)) + + def retry() -> tuple[dict[str, object], ...] | None: + result: Final = asyncio.run(_list_tools(second_replica, first_cursor)) + return _tool_items(result) if isinstance(result, ListToolsResult) else None + + retried_items: Final = eventually(retry, lambda items: items is not None, seconds=30) + assert retried_items == continued_items + + +def test_mcp_key_team_and_server_rpm_limits_share_redis(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + results_directory: Final = tmp_path / "results" + results_directory.mkdir() + monkeypatch.setenv("INTEGRATION_RESULTS_DIR", str(results_directory)) + assert os.environ.get("DATABASE_URL"), "This integration case requires disposable-database access" + + with ( + owned_redis(tmp_path) as redis, + scratch_database() as database_url, + mcp_peer() as peer, + httpx.Client() as client, + ): + monkeypatch.setenv("DATABASE_URL", database_url) + seed: Final = Gateway(client, "sk-mcp-key-team-server-rate-limit", peer.url) + config: Final = _config_file( + tmp_path, + seed.key, + redis, + store_model_in_db=True, + ) + environment: Final = { + "DATABASE_URL": database_url, + "LITELLM_SALT_KEY": "shared-mcp-key-team-server-rate-limit", + "LITELLM_RATE_LIMIT_WINDOW_SIZE": "30", + } + second_environment: Final = {**environment, "DISABLE_SCHEMA_UPDATE": "true"} + options: Final = { + "config": config, + "remove_environment": ("DATABASE_URL_READ_REPLICA", "LITELLM_LICENSE", "LITELLM_LICENSE_PATH"), + } + with ( + owned_proxy(seed, tmp_path / "first", environment, **options) as first_replica, + first_replica.scenario() as scenario, + ): + server_id: Final = register_mcp( + scenario, + peer, + "rpm", + rpm=5, + allowed_tools=["add"], + ) + permission: Final = {"mcp_servers": [server_id]} + team_id: Final = scenario.team( + mcp_rpm_limit={"rpm": 3}, + object_permission=permission, + ) + key_one: Final = scenario.key( + team_id=team_id, + mcp_rpm_limit={"rpm": 1}, + object_permission=permission, + ) + key_two: Final = scenario.key(team_id=team_id, object_permission=permission) + key_three: Final = scenario.key(object_permission=permission) + key_four: Final = scenario.key(rpm_limit=1, object_permission=permission) + + with owned_proxy( + seed, tmp_path / "second", second_environment, database_setup=(), **options + ) as second_replica: + first_call: Final = _call_tool(first_replica, key_one, server_id, "rpm-add") + assert first_call.status_code == 200, first_call.text + + peer.drain() + key_one_rejected: Final = _call_tool(second_replica, key_one, server_id, "rpm-add") + _assert_rate_limit(key_one_rejected, "mcp_per_key") + assert tool_calls(peer.drain()) == () + + second_call: Final = _call_tool(first_replica, key_two, server_id, "rpm-add") + assert second_call.status_code == 200, second_call.text + third_call: Final = _call_tool(second_replica, key_two, server_id, "rpm-add") + assert third_call.status_code == 200, third_call.text + + peer.drain() + key_two_rejected: Final = _call_tool(first_replica, key_two, server_id, "rpm-add") + _assert_rate_limit(key_two_rejected, "mcp_per_team") + assert tool_calls(peer.drain()) == () + + key_four_second_replica: Final = McpCaller(second_replica, key_four, "mcp") + key_four_first_replica: Final = McpCaller(first_replica, key_four, "mcp") + key_four_first_call: Final = key_four_second_replica.call("rpm-add", {"a": 1, "b": 2}) + assert key_four_first_call.ok, key_four_first_call.raw + + peer.drain() + key_four_rejected: Final = key_four_first_replica.call("rpm-add", {"a": 1, "b": 2}) + assert not key_four_rejected.ok, key_four_rejected.raw + assert "api_key" in (key_four_rejected.error or "") + assert tool_calls(peer.drain()) == () + + peer.drain() + forbidden: Final = _call_tool(second_replica, key_three, server_id, "rpm-multiply") + assert forbidden.status_code == 403, forbidden.text + assert tool_calls(peer.drain()) == () + + key_three_first_call: Final = _call_tool(first_replica, key_three, server_id, "rpm-add") + assert key_three_first_call.status_code == 200, key_three_first_call.text + + peer.drain() + server_rejected: Final = _call_tool(second_replica, key_three, server_id, "rpm-add") + _assert_rate_limit(server_rejected, "mcp_server") + assert tool_calls(peer.drain()) == () diff --git a/tests/unit/proxy/_experimental/mcp_server/faults/test_list_outcomes.py b/tests/unit/proxy/_experimental/mcp_server/faults/test_list_outcomes.py index f951499e18f..b41e786881d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/faults/test_list_outcomes.py +++ b/tests/unit/proxy/_experimental/mcp_server/faults/test_list_outcomes.py @@ -3,6 +3,7 @@ to exactly one category, wire values never carry upstream prose, and single-upst stay truthful to who failed.""" import sys +from typing import Final if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 from exceptiongroup import BaseExceptionGroup @@ -23,6 +24,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( list_fault_http_status, outcome_wire_value, ) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError def test_carried_fault_passes_through(): @@ -35,6 +37,11 @@ def test_upstream_auth_error_maps_to_auth_required_and_forbidden(): assert classify_list_exception(MCPUpstreamAuthError(403, None, "srv")).tag == "forbidden" +def test_proxy_rate_limit_error_maps_to_rate_limited() -> None: + fault: Final = classify_list_exception(ProxyRateLimitError(detail="server RPM exceeded")) + assert fault == ServerListFault(tag="rate_limited", status_code=429) + + def test_timeout_and_connection_errors_classify_without_status(): assert classify_list_exception(TimeoutError()).tag == "timeout" assert classify_list_exception(ConnectionError()).tag == "unreachable" @@ -100,6 +107,10 @@ def test_wire_value_carries_no_prose(): assert outcome_wire_value(fault) == {"status": "upstream_error", "http_status": 500} assert outcome_wire_value(ServerListOk(tool_count=7)) == {"status": "ok", "tool_count": 7} assert outcome_wire_value(ServerListFault(tag="timeout")) == {"status": "timeout"} + assert outcome_wire_value(ServerListFault(tag="rate_limited", status_code=429)) == { + "status": "rate_limited", + "http_status": 429, + } @pytest.mark.parametrize( @@ -108,6 +119,7 @@ def test_wire_value_carries_no_prose(): ("auth_required", 401, 401), ("auth_required", None, 401), ("forbidden", 403, 403), + ("rate_limited", 429, 429), ("timeout", None, 504), ("unreachable", None, 502), ("upstream_error", 500, 502), diff --git a/tests/unit/proxy/_experimental/mcp_server/test_catalog.py b/tests/unit/proxy/_experimental/mcp_server/test_catalog.py index 5365bda516a..5c029d36c7d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_catalog.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_catalog.py @@ -1,11 +1,37 @@ import asyncio from collections.abc import Sequence +from types import SimpleNamespace +from typing import Final, Literal +from unittest.mock import AsyncMock import pytest from mcp.shared.exceptions import MCPError -from mcp.types import ListToolsResult, Tool +from mcp.types import ( + ListPromptsRequest, + ListPromptsResult, + ListResourcesRequest, + ListResourcesResult, + ListResourceTemplatesRequest, + ListResourceTemplatesResult, + ListToolsResult, + PaginatedRequestParams, + Tool, +) from litellm.proxy._experimental.mcp_server import catalog +from litellm.proxy._experimental.mcp_server.contracts import OperationContext +from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( + SERVER_OUTCOMES_META_KEY, + AggregateToolListing, + ServerListOk, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.types.mcp import MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer + +CatalogKind = Literal["tools", "prompts", "resources", "templates"] +OptionalCatalogResult = ListPromptsResult | ListResourcesResult | ListResourceTemplatesResult def page(name: str, cursor: str | None = None, revision: str = "stable") -> ListToolsResult: @@ -16,6 +42,77 @@ def page(name: str, cursor: str | None = None, revision: str = "stable") -> List ) +def rate_limit_catalog_setup( + monkeypatch: pytest.MonkeyPatch, rejected_server_ids: frozenset[str] +) -> tuple[tuple[MCPServer, MCPServer], OperationContext, AsyncMock]: + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import operations + + monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-rate-limit-test-key") + servers: Final = ( + MCPServer(server_id="catalog-a", name="catalog-a", transport=MCPTransport.http), + MCPServer(server_id="catalog-b", name="catalog-b", transport=MCPTransport.http), + ) + caller: Final = UserAPIKeyAuth(api_key="catalog-rate-limit-key", user_id="catalog-rate-limit-user") + + async def enforce_rate_limit(_user: UserAPIKeyAuth | None, server: MCPServer) -> None: + if server.server_id in rejected_server_ids: + raise ProxyRateLimitError(detail=f"{server.server_id} RPM exceeded") + + limiter: Final = AsyncMock(side_effect=enforce_rate_limit) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", SimpleNamespace(enforce_mcp_server_rate_limits=limiter)) + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(operations.global_mcp_server_manager, "registry", {server.server_id: server for server in servers}) + monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=list(servers))) + context: Final = operations.prepare_context( + user_api_key_auth=caller, + mcp_servers=[server.server_id for server in servers], + ) + return servers, context, limiter + + +def optional_catalog_page( + kind: CatalogKind, server_id: str, next_cursor: str | None = None +) -> OptionalCatalogResult: + from mcp import types + + if kind == "prompts": + return types.ListPromptsResult( + prompts=[types.Prompt(name=f"{server_id}-item")], next_cursor=next_cursor + ) + if kind == "resources": + return types.ListResourcesResult( + resources=[types.Resource(name=f"{server_id}-item", uri=f"https://example.com/{server_id}")], + next_cursor=next_cursor, + ) + return types.ListResourceTemplatesResult( + resource_templates=[ + types.ResourceTemplate(name=f"{server_id}-item", uri_template=f"https://example.com/{server_id}/{{name}}") + ], + next_cursor=next_cursor, + ) + + +async def run_catalog_listing( + kind: CatalogKind, + context: OperationContext, + servers: Sequence[MCPServer], + cursor: str | None = None, +) -> AggregateToolListing | OptionalCatalogResult: + from mcp import types + + if kind == "tools": + return await catalog.aggregate_gateway_tools( + context, PaginatedRequestParams(cursor=cursor), servers, {} + ) + request: Final = { + "prompts": types.ListPromptsRequest, + "resources": types.ListResourcesRequest, + "templates": types.ListResourceTemplatesRequest, + }[kind](params=PaginatedRequestParams(cursor=cursor)) + return await catalog.list_gateway_catalog(context, request) + + async def listing( fetch, cursor: str | None = None, @@ -374,3 +471,119 @@ async def test_optional_gateway_catalog_reports_initial_failure_and_rejects_fail assert "upstream secret" not in str(denied.value) assert fetch.await_count == 3 assert fetch.await_args.args[-1] == "next" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["tools", "prompts", "resources", "templates"]) +async def test_first_catalog_page_keeps_admitted_items_and_reports_rate_limited_servers( + monkeypatch: pytest.MonkeyPatch, kind: CatalogKind +) -> None: + from mcp import types + + servers, context, limiter = rate_limit_catalog_setup(monkeypatch, frozenset({"catalog-a"})) + + async def fetch_tools(server: MCPServer, **_kwargs: object) -> tuple[ListToolsResult, ServerListOk]: + return ListToolsResult( + tools=[types.Tool(name=f"{server.server_id}-item", inputSchema={"type": "object"})] + ), ServerListOk(tool_count=1) + + async def fetch_optional( + _context: OperationContext, + _request: ListPromptsRequest | ListResourcesRequest | ListResourceTemplatesRequest, + server: MCPServer, + _allowed: Sequence[MCPServer], + _cursor: str | None, + ) -> OptionalCatalogResult: + return optional_catalog_page(kind, server.server_id) + + if kind == "tools": + monkeypatch.setattr(catalog, "get_filtered_server_tools", fetch_tools) + else: + monkeypatch.setattr(catalog, "fetch_optional_catalog_page", fetch_optional) + + result: Final = await run_catalog_listing(kind, context, servers) + if isinstance(result, AggregateToolListing): + assert [tool.name for tool in result.tools] == ["catalog-b-item"] + assert result.outcomes["catalog-a"].tag == "rate_limited" + else: + field: Final = { + "prompts": "prompts", + "resources": "resources", + "templates": "resource_templates", + }[kind] + assert [item.name for item in getattr(result, field)] == ["catalog-b-item"] + assert result.meta[SERVER_OUTCOMES_META_KEY]["catalog-a"]["status"] == "rate_limited" + assert {call.args[1].server_id for call in limiter.await_args_list} == {"catalog-a", "catalog-b"} + assert len(limiter.await_args_list) == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["tools", "prompts", "resources", "templates"]) +async def test_first_catalog_page_raises_when_every_server_is_rate_limited( + monkeypatch: pytest.MonkeyPatch, kind: CatalogKind +) -> None: + servers, context, limiter = rate_limit_catalog_setup(monkeypatch, frozenset({"catalog-a", "catalog-b"})) + fetch_tools: Final = AsyncMock(return_value=(ListToolsResult(tools=[]), ServerListOk(tool_count=0))) + fetch_optional: Final = AsyncMock(return_value=optional_catalog_page(kind, "catalog-a")) + if kind == "tools": + monkeypatch.setattr(catalog, "get_filtered_server_tools", fetch_tools) + else: + monkeypatch.setattr(catalog, "fetch_optional_catalog_page", fetch_optional) + + with pytest.raises(ProxyRateLimitError, match="RPM exceeded"): + await run_catalog_listing(kind, context, servers) + + if kind == "tools": + fetch_tools.assert_not_awaited() + else: + fetch_optional.assert_not_awaited() + assert len(limiter.await_args_list) == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["tools", "prompts", "resources", "templates"]) +async def test_continuation_rate_limit_raises_without_refetching_completed_servers( + monkeypatch: pytest.MonkeyPatch, kind: CatalogKind +) -> None: + from mcp import types + + servers, context, limiter = rate_limit_catalog_setup(monkeypatch, frozenset()) + + async def fetch_tools(server: MCPServer, **kwargs: object) -> tuple[ListToolsResult, ServerListOk]: + params: Final = PaginatedRequestParams.model_validate(kwargs["params"]) + next_cursor: Final = "next" if server.server_id == "catalog-a" and params.cursor is None else None + return ( + ListToolsResult( + tools=[types.Tool(name=f"{server.server_id}-item", inputSchema={"type": "object"})], + next_cursor=next_cursor, + ), + ServerListOk(tool_count=1), + ) + + async def fetch_optional( + _context: OperationContext, + _request: ListPromptsRequest | ListResourcesRequest | ListResourceTemplatesRequest, + server: MCPServer, + _allowed: Sequence[MCPServer], + cursor: str | None, + ) -> OptionalCatalogResult: + return optional_catalog_page( + kind, + server.server_id, + "next" if server.server_id == "catalog-a" and cursor is None else None, + ) + + if kind == "tools": + monkeypatch.setattr(catalog, "get_filtered_server_tools", fetch_tools) + else: + monkeypatch.setattr(catalog, "fetch_optional_catalog_page", fetch_optional) + + first_page: Final = await run_catalog_listing(kind, context, servers) + assert first_page.next_cursor is not None + limiter.side_effect = ProxyRateLimitError(detail="catalog-a RPM exceeded") + with pytest.raises(ProxyRateLimitError, match="catalog-a RPM exceeded"): + await run_catalog_listing(kind, context, servers, first_page.next_cursor) + + charged_server_ids: Final = [call.args[1].server_id for call in limiter.await_args_list] + assert charged_server_ids.count("catalog-a") == 2 + assert charged_server_ids.count("catalog-b") == 1 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py index 8e87837611a..1578ea8e601 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_guardrail_usage_monitor.py @@ -67,6 +67,7 @@ def _fake_proxy_logging(capture: dict, *, guardrail_effect=None): way a blocking guardrail does). """ plo = mock.MagicMock() + plo.enforce_mcp_server_rate_limits = mock.AsyncMock() plo._create_mcp_request_object_from_kwargs.return_value = mock.MagicMock() # Mirror the real conversion's metadata bucket so a test can prove it survives. plo._convert_mcp_to_llm_format.side_effect = lambda *_a, **_k: { diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 417c6cad3b1..5a608a77931 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -5853,7 +5853,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) user_api_key_auth: Final = UserAPIKeyAuth() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) @@ -5886,7 +5886,7 @@ class TestMCPServerManager: # Mock dependencies user_api_key_auth = MagicMock() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # This should raise an HTTPException with pytest.raises(HTTPException) as exc_info: @@ -5920,7 +5920,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) user_api_key_auth: Final = UserAPIKeyAuth() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) @@ -5953,7 +5953,7 @@ class TestMCPServerManager: # Mock dependencies user_api_key_auth = MagicMock() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # This should raise an HTTPException with pytest.raises(HTTPException) as exc_info: @@ -5987,7 +5987,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) user_api_key_auth: Final = UserAPIKeyAuth() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) @@ -6022,7 +6022,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) user_api_key_auth: Final = UserAPIKeyAuth() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) @@ -6890,7 +6890,7 @@ class TestMCPServerManager: object_permission=object_permission, ) - proxy_logging = MagicMock() + proxy_logging = _mock_proxy_logging() proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) @@ -6933,7 +6933,7 @@ class TestMCPServerManager: object_permission=object_permission, ) - proxy_logging = MagicMock() + proxy_logging = _mock_proxy_logging() proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging.pre_call_hook = AsyncMock(return_value=None) @@ -7074,7 +7074,7 @@ class TestMCPServerManager: # Mock dependencies - set object_permission and object_permission_id to None # so permission checks return None (no restrictions) user_api_key_auth: Final = UserAPIKeyAuth() - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() # Mock the async methods that pre_call_tool_check calls proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) @@ -7159,7 +7159,7 @@ class TestMCPServerManager: user_api_key_auth: Final = UserAPIKeyAuth(api_key="sk-test") # Mock proxy logging - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -7205,7 +7205,7 @@ class TestMCPServerManager: mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) manager._create_mcp_client = AsyncMock(return_value=mock_client) - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -7690,7 +7690,7 @@ class TestMCPServerManager: manager._fetch_tools_with_timeout = AsyncMock( return_value=[MCPTool(name="turn", description="stored cred catalog", inputSchema={})] ) - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -8024,7 +8024,7 @@ class TestMCPServerManager: record_listing=True, ) - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -9553,19 +9553,28 @@ class TestMCPServerTimestamps: @pytest.mark.asyncio async def test_load_servers_from_config_preserves_timeout(self, config_only_mcp_manager_factory): - """timeout from proxy config is loaded into MCPServer.""" + """MCP server request limits from proxy config are loaded into MCPServer.""" manager = config_only_mcp_manager_factory() config = { "my_server": { "url": "https://example.com/mcp", "transport": MCPTransport.http, "timeout": 90.0, + "max_concurrent_requests": 4, + "rpm": 7, + }, + "unlimited_server": { + "url": "https://example.com/other-mcp", + "transport": MCPTransport.http, } } await manager.load_servers_from_config(config) servers = list(manager.config_mcp_servers.values()) - assert len(servers) == 1 + assert len(servers) == 2 assert servers[0].timeout == 90.0 + assert servers[0].max_concurrent_requests == 4 + assert servers[0].rpm == 7 + assert servers[1].rpm is None @pytest.mark.asyncio async def test_call_regular_mcp_tool_timeout_returns_504(self): @@ -12850,8 +12859,14 @@ def _unrestricted_auth() -> UserAPIKeyAuth: return UserAPIKeyAuth() +def _mock_proxy_logging() -> MagicMock: + proxy_logging_obj: Final = MagicMock() + proxy_logging_obj.enforce_mcp_server_rate_limits = AsyncMock() + return proxy_logging_obj + + def _permissive_proxy_logging() -> MagicMock: - proxy_logging_obj = MagicMock() + proxy_logging_obj: Final = _mock_proxy_logging() proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) @@ -18348,7 +18363,7 @@ class TestToolCatalogGuard: manager = MCPServerManager() server = _notes_server({"list_notes": _pin(LIST_NOTES)}) user_api_key_auth = MagicMock(object_permission=None, object_permission_id=None) - proxy_logging_obj = MagicMock() + proxy_logging_obj = _mock_proxy_logging() proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={}) proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 59e21404301..df3f1d78a7e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -1706,6 +1706,56 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error( assert exc_info.value.error.message == denial_message +@pytest.mark.asyncio +@pytest.mark.parametrize( + "handler_name", + [ + "list_prompts", + "list_resources", + "list_resource_templates", + ], +) +async def test_rate_limited_catalog_lists_return_mcp_errors(handler_name): + from mcp.shared.exceptions import MCPError + + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + + user_api_key_auth: Final = UserAPIKeyAuth(api_key="test_key", user_id="test_user") + server_config: Final = MCPServer( + server_id="rate-limited", + name="rate-limited", + server_name="rate-limited", + transport=MCPTransport.http, + rpm=1, + ) + rate_limit_error: Final = ProxyRateLimitError(detail="server RPM exceeded") + enforce_rate_limit: Final = AsyncMock(side_effect=rate_limit_error) + proxy_logging: Final = MagicMock(enforce_mcp_server_rate_limits=enforce_rate_limit) + execute_list: Final = { + "list_prompts": mcp_operations._execute_list_prompts, + "list_resources": mcp_operations._execute_list_resources, + "list_resource_templates": mcp_operations._execute_list_resource_templates, + }[handler_name] + context: Final = mcp_operations.prepare_context( + user_api_key_auth, + mcp_servers=[server_config.server_id], + ) + + with ( + patch.object(mcp_operations, "_get_allowed_mcp_servers", new=AsyncMock(return_value=[server_config])), + patch("litellm.proxy.proxy_server.proxy_logging_obj", new=proxy_logging), + patch("litellm.proxy.proxy_server.prisma_client", None), + ): + with pytest.raises(MCPError) as exc_info: + await execute_list(context, _paged_params()) + + assert exc_info.value.error.code == INVALID_REQUEST + assert exc_info.value.error.message == "server RPM exceeded" + assert enforce_rate_limit.await_count == 1 + assert enforce_rate_limit.await_args.args[0].api_key == user_api_key_auth.api_key + assert enforce_rate_limit.await_args.args[1] is server_config + + @pytest.mark.asyncio async def test_mcp_server_tool_call_renders_denial_message_not_detail_dict(_mcp_request_ctx): try: @@ -9150,6 +9200,7 @@ def _mock_mcp_logging_obj() -> MagicMock: def _mock_mcp_proxy_logging() -> MagicMock: """ProxyLogging stand-in whose post_mcp_call_hook passes the result through.""" proxy_logging_mock = MagicMock() + proxy_logging_mock.enforce_mcp_server_rate_limits = AsyncMock() proxy_logging_mock.post_call_failure_hook = AsyncMock() proxy_logging_mock.post_mcp_call_hook = AsyncMock(side_effect=lambda response, **_: response) return proxy_logging_mock diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index 8274b01ceeb..089426eb5cb 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -1,4 +1,5 @@ import asyncio +from collections.abc import Sequence from typing import Final from unittest.mock import AsyncMock, patch @@ -11,12 +12,19 @@ from litellm.caching.dual_cache import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._experimental.mcp_server import operations from litellm.proxy._experimental.mcp_server import rest_endpoints -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + HTTPException as MCPServerManagerHTTPException, + ListedToolsCaller, +) from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache -from litellm.proxy.utils import ProxyLogging +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, +) +from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.mcp import MCPAuth, MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -284,6 +292,271 @@ def _catalog_case(method): return cases[method] +def _mcp_rate_limited_proxy_logging() -> ProxyLogging: + proxy_logging: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache()) + proxy_logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + return proxy_logging + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "operation", + ["tools/list", "prompts/list", "resources/list", "resources/templates/list", "prompts/get", "resources/read"], +) +async def test_mcp_server_rpm_limits_every_catalog_operation(operation: str) -> None: + from unittest.mock import MagicMock + + from mcp import types + from mcp.shared.exceptions import MCPError + from mcp.types import INVALID_REQUEST + + from litellm.proxy._experimental.mcp_server import catalog + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListOk + + server: Final = MCPServer( + server_id="catalog-rpm", + name="catalog", + server_name="catalog", + transport=MCPTransport.http, + rpm=1, + ) + caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-catalog-rpm")) + operation_to_manager_method: Final = { + "tools/list": "_get_tools_from_server", + "prompts/list": "get_prompts_from_server", + "resources/list": "get_resources_from_server", + "resources/templates/list": "get_resource_templates_from_server", + "prompts/get": "get_prompt_from_server", + "resources/read": "read_resource_from_server", + } + upstream_results: Final = { + "tools/list": ( + types.ListToolsResult(tools=[types.Tool(name="echo", inputSchema={"type": "object"})]), + ServerListOk(tool_count=1), + ), + "prompts/list": types.ListPromptsResult(prompts=[types.Prompt(name="catalog-prompt")]), + "resources/list": types.ListResourcesResult( + resources=[types.Resource(name="document", uri="https://example.com/document")] + ), + "resources/templates/list": types.ListResourceTemplatesResult( + resource_templates=[ + types.ResourceTemplate(name="document", uri_template="https://example.com/{name}") + ] + ), + "prompts/get": GetPromptResult(messages=[]), + "resources/read": types.ReadResourceResult(contents=[]), + } + upstream: Final = AsyncMock(return_value=upstream_results[operation]) + manager_method: Final = operation_to_manager_method[operation] + manager: Final = operations.global_mcp_server_manager + rate_limit_error: Final = ProxyRateLimitError(detail="server RPM exceeded") + enforce_rate_limit: Final = AsyncMock(side_effect=[None, rate_limit_error, rate_limit_error]) + proxy_logging: Final = MagicMock(enforce_mcp_server_rate_limits=enforce_rate_limit) + is_protocol_listing: Final = operation.endswith("/list") + context: Final = prepare_context(caller, mcp_servers=[server.server_id]) + + async def invoke() -> object: + if operation == "tools/list": + return await GatewayOperations().execute(types.ListToolsRequest(), context) + if operation == "prompts/list": + return await GatewayOperations().execute(types.ListPromptsRequest(), context) + if operation == "resources/list": + return await GatewayOperations().execute(types.ListResourcesRequest(), context) + if operation == "resources/templates/list": + return await GatewayOperations().execute(types.ListResourceTemplatesRequest(), context) + if operation == "prompts/get": + return await operations.mcp_get_prompt( + name=f"{server.name}-catalog-prompt", + user_api_key_auth=caller, + mcp_servers=[server.server_id], + ) + return await operations.mcp_read_resource( + url="https://example.com/document", + user_api_key_auth=caller, + mcp_servers=[server.server_id], + ) + + with ( + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch.dict(manager.registry, {server.server_id: server}), + patch.object(manager, manager_method, upstream), + patch.object(catalog, "get_filtered_server_tools", upstream), + patch.object(catalog, "fetch_optional_catalog_page", upstream), + ): + await invoke() + if is_protocol_listing: + with pytest.raises(MCPError) as rejected: + await invoke() + assert rejected.value.error.code == INVALID_REQUEST + assert rejected.value.error.message == "server RPM exceeded" + if operation == "tools/list": + with pytest.raises(ProxyRateLimitError) as rejected: + await operations._get_tools_from_mcp_servers( + user_api_key_auth=caller, + mcp_auth_header=None, + mcp_servers=[server.server_id], + params=None, + ) + assert rejected.value is rate_limit_error + assert upstream.await_count == 1 + assert enforce_rate_limit.await_count == 3 + else: + with pytest.raises(ProxyRateLimitError): + await invoke() + + assert upstream.await_count == 1 + if operation != "tools/list": + assert enforce_rate_limit.await_count == 2 + + +@pytest.mark.asyncio +async def test_tools_call_warmup_does_not_consume_mcp_server_rpm() -> None: + from mcp import types + + server: Final = MCPServer( + server_id="catalog-warmup", + name="catalog-warmup", + server_name="catalog-warmup", + transport=MCPTransport.http, + rpm=1, + ) + caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-catalog-warmup")) + proxy_logging: Final = _mcp_rate_limited_proxy_logging() + upstream: Final = AsyncMock( + return_value=[types.Tool(name="echo", inputSchema={"type": "object"})] + ) + with ( + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch.object(operations.global_mcp_server_manager, "server_exposes_tool", return_value=False), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + ): + await operations._list_tools_before_first_call( + server=server, + tool_name="echo", + allowed_mcp_servers=[server], + user_api_key_auth=caller, + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + ) + listing: Final = await operations._get_tools_from_mcp_servers( + user_api_key_auth=caller, + mcp_auth_header=None, + mcp_servers=[server.server_id], + ) + + assert [tool.name for tool in listing.tools] == ["echo"] + + +@pytest.mark.asyncio +async def test_tools_call_pre_call_check_enforces_mcp_server_rpm() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + server: Final = MCPServer( + server_id="catalog-call", + name="catalog-call", + server_name="catalog-call", + transport=MCPTransport.http, + rpm=1, + ) + proxy_logging: Final = _mcp_rate_limited_proxy_logging() + manager: Final = MCPServerManager() + + await manager.pre_call_tool_check( + name="echo", + arguments={}, + server_name=server.name, + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + with pytest.raises(ProxyRateLimitError): + await manager.pre_call_tool_check( + name="echo", + arguments={}, + server_name=server.name, + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + + +@pytest.mark.asyncio +async def test_tools_call_pre_call_hook_rejection_does_not_enforce_mcp_server_rpm() -> None: + from unittest.mock import MagicMock + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + server: Final = MCPServer( + server_id="catalog-call-pre-hook-rejected", + name="catalog-call-pre-hook-rejected", + server_name="catalog-call-pre-hook-rejected", + transport=MCPTransport.http, + rpm=1, + ) + rate_limit_error: Final = ProxyRateLimitError(detail="ordinary key rate limit") + proxy_logging: Final = MagicMock() + proxy_logging._create_mcp_request_object_from_kwargs.return_value = {} + proxy_logging._convert_mcp_to_llm_format.return_value = {} + proxy_logging.pre_call_hook = AsyncMock(side_effect=rate_limit_error) + proxy_logging.enforce_mcp_server_rate_limits = AsyncMock() + + with pytest.raises(ProxyRateLimitError) as rejected: + await MCPServerManager().pre_call_tool_check( + name="echo", + arguments={}, + server_name=server.name, + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + + assert rejected.value is rate_limit_error + proxy_logging.enforce_mcp_server_rate_limits.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_disallowed_tool_does_not_consume_mcp_server_rpm() -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + server: Final = MCPServer( + server_id="catalog-call-authorization", + name="catalog-call-authorization", + server_name="catalog-call-authorization", + transport=MCPTransport.http, + allowed_tools=["allowed"], + rpm=1, + ) + proxy_logging: Final = _mcp_rate_limited_proxy_logging() + manager: Final = MCPServerManager() + + with pytest.raises(MCPServerManagerHTTPException) as denied_call: + await manager.pre_call_tool_check( + name="disallowed", + arguments={}, + server_name=server.name, + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + + assert denied_call.value.status_code == 403 + await manager.pre_call_tool_check( + name="allowed", + arguments={}, + server_name=server.name, + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + + @pytest.mark.asyncio @pytest.mark.parametrize( "method", ["prompts/list", "prompts/get", "resources/list", "resources/templates/list", "resources/read"] @@ -692,6 +965,102 @@ async def test_discovery_lists_each_capability_with_the_same_caller(available): assert listing.await_args.args[0] is context +@pytest.mark.asyncio +async def test_discovery_shares_one_server_admission_across_catalog_listings() -> None: + from mcp import types + + from litellm.proxy._experimental.mcp_server import catalog + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListOk + + admitted: Final = MCPServer( + server_id="discover-admitted", + name="discover-admitted", + server_name="discover-admitted", + transport=MCPTransport.http, + ) + rejected: Final = MCPServer( + server_id="discover-rejected", + name="discover-rejected", + server_name="discover-rejected", + transport=MCPTransport.http, + ) + caller: Final = UserAPIKeyAuth(api_key="sk-discovery-admission") + proxy_logging: Final = _mcp_rate_limited_proxy_logging() + + async def enforce_server_rpm(_user_api_key_auth: UserAPIKeyAuth | None, server: MCPServer) -> None: + if server.server_id == rejected.server_id: + raise ProxyRateLimitError(detail="server RPM exceeded") + + async def fetch_tools(server: MCPServer, **_: object) -> tuple[types.ListToolsResult, ServerListOk]: + tools: Final = [types.Tool(name=f"{server.server_id}-tool", inputSchema={"type": "object"})] + return types.ListToolsResult(tools=tools), ServerListOk(tool_count=len(tools)) + + async def fetch_prompts(*, server: MCPServer, **_: object) -> types.ListPromptsResult: + return types.ListPromptsResult(prompts=[types.Prompt(name=f"{server.server_id}-prompt")]) + + async def fetch_resources(*, server: MCPServer, **_: object) -> types.ListResourcesResult: + return types.ListResourcesResult( + resources=[types.Resource(name=f"{server.server_id}-resource", uri=f"test://{server.server_id}")] + ) + + async def fetch_resource_templates(*, server: MCPServer, **_: object) -> types.ListResourceTemplatesResult: + return types.ListResourceTemplatesResult( + resource_templates=[ + types.ResourceTemplate( + name=f"{server.server_id}-template", + uri_template=f"test://{server.server_id}/{{name}}", + ) + ] + ) + + enforcement: Final = AsyncMock(side_effect=enforce_server_rpm) + upstream_calls: Final = ( + AsyncMock(side_effect=fetch_tools), + AsyncMock(side_effect=fetch_prompts), + AsyncMock(side_effect=fetch_resources), + AsyncMock(side_effect=fetch_resource_templates), + ) + async def fetch_optional_page( + _context: OperationContext, + request: types.ListPromptsRequest | types.ListResourcesRequest | types.ListResourceTemplatesRequest, + server: MCPServer, + _allowed: Sequence[MCPServer], + _cursor: str | None, + ) -> types.ListPromptsResult | types.ListResourcesResult | types.ListResourceTemplatesResult: + if isinstance(request, types.ListPromptsRequest): + return await upstream_calls[1](server=server) + if isinstance(request, types.ListResourcesRequest): + return await upstream_calls[2](server=server) + return await upstream_calls[3](server=server) + + optional_fetch: Final = AsyncMock(side_effect=fetch_optional_page) + with ( + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[admitted, rejected])), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + patch.object(proxy_logging, "enforce_mcp_server_rate_limits", enforcement), + patch.object(catalog, "get_filtered_server_tools", upstream_calls[0]), + patch.object(catalog, "fetch_optional_catalog_page", optional_fetch), + ): + result: Final = await GatewayOperations().execute( + types.DiscoverRequest(), + prepare_context(caller, mcp_servers=[admitted.server_id, rejected.server_id]), + ) + + assert enforcement.await_count == 2 + assert {call.args[1].server_id for call in enforcement.await_args_list} == { + admitted.server_id, + rejected.server_id, + } + assert tuple(call.args[0].server_id for call in upstream_calls[0].await_args_list) == (admitted.server_id,) + assert optional_fetch.await_count == 3 + assert all(call.args[2].server_id == admitted.server_id for call in optional_fetch.await_args_list) + assert all(upstream.await_count == 1 for upstream in upstream_calls) + assert result.capabilities.tools is not None + assert result.capabilities.prompts is not None + assert result.capabilities.resources is not None + + @pytest.mark.asyncio @pytest.mark.parametrize("outcome", ["success", "failure", "cancel"]) async def test_discovery_concurrent_listings_drain_on_failure_and_cancellation(outcome): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index 0e38d03ca3a..c748bdc6a6b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -1206,6 +1206,87 @@ class TestTestToolsList: class TestListToolsRestAPI: pytestmark = pytest.mark.asyncio + async def test_single_server_rate_limit_returns_429_without_fetching_tools( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + + server: Final = MCPServer( + server_id="rate-limited-server", + name="rate-limited-server", + server_name="rate-limited-server", + transport=MCPTransport.http, + ) + caller: Final = UserAPIKeyAuth() + manager: Final = MCPServerManager() + monkeypatch.setitem(manager.registry, server.server_id, server) + enforcement: Final = AsyncMock(side_effect=ProxyRateLimitError(detail="server RPM exceeded")) + proxy_logging: Final = MagicMock(enforce_mcp_server_rate_limits=enforcement) + upstream: Final = AsyncMock(return_value=[Tool(name="should-not-list", inputSchema={})]) + + async def allowed_servers(*_: object, **__: object) -> list[str]: + return [server.server_id] + + monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager) + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=[caller])) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) + monkeypatch.setattr(manager, "get_allowed_mcp_servers", allowed_servers) + monkeypatch.setattr(manager, "filter_server_ids_by_ip_with_info", lambda ids, _ip: (ids, 0)) + monkeypatch.setattr(manager, "_get_tools_from_server", upstream) + + with pytest.raises(HTTPException) as error: + await rest_endpoints.list_tool_rest_api( + _build_request(path="/mcp-rest/tools/list", method="GET"), + server_id=server.server_id, + user_api_key_dict=caller, + ) + + assert error.value.status_code == 429 + enforcement.assert_awaited_once_with(caller, server) + upstream.assert_not_awaited() + + async def test_admin_unfiltered_tools_list_does_not_enforce_server_rpm( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import LitellmUserRoles + + server: Final = MCPServer( + server_id="admin-unfiltered-server", + name="admin-unfiltered-server", + server_name="admin-unfiltered-server", + transport=MCPTransport.http, + allowed_tools=["enabled-tool"], + ) + caller: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + manager: Final = MCPServerManager() + monkeypatch.setitem(manager.registry, server.server_id, server) + enforcement: Final = AsyncMock() + proxy_logging: Final = MagicMock(enforce_mcp_server_rate_limits=enforcement) + upstream: Final = AsyncMock(return_value=[Tool(name="disabled-tool", inputSchema={})]) + + async def allowed_servers(*_: object, **__: object) -> list[str]: + return [server.server_id] + + monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager) + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=[caller])) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging) + monkeypatch.setattr(manager, "get_allowed_mcp_servers", allowed_servers) + monkeypatch.setattr(manager, "filter_server_ids_by_ip_with_info", lambda ids, _ip: (ids, 0)) + monkeypatch.setattr(manager, "_get_tools_from_server", upstream) + + result: Final = await rest_endpoints.list_tool_rest_api( + _build_request(path="/mcp-rest/tools/list", method="GET"), + server_id=server.server_id, + include_disabled_tools=True, + user_api_key_dict=caller, + ) + + assert [tool.name for tool in result["tools"]] == ["disabled-tool"] + enforcement.assert_not_awaited() + upstream.assert_awaited_once() + async def test_rejects_disallowed_server(self, monkeypatch): async def fake_contexts(user_api_key_auth): return [user_api_key_auth] diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index d9a9ecf807f..fecedd8c498 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -41,7 +41,8 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.mcp import MCPPreCallRequestObject +from litellm.types.mcp import MCPPreCallRequestObject, MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.utils import ( EmbeddingResponse, ModelResponse, @@ -3785,200 +3786,200 @@ async def test_failure_event_settles_project_itpm_otpm_at_recovered_partial_usag # ----------------------- Per-MCP-server rate limiting (v3) ----------------------- -def _make_mcp_handler(): - local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( +def _make_mcp_handler() -> tuple[_PROXY_MaxParallelRequestsHandler, DualCache]: + local_cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) return handler, local_cache -def _find_descriptor(descriptors, key): - return next((d for d in descriptors if d["key"] == key), None) - - -def _build_mcp_descriptors(handler, user_api_key_dict, data, call_type="call_mcp_tool"): - return handler._create_rate_limit_descriptors( - user_api_key_dict=user_api_key_dict, - data=data, - rpm_limit_type=None, - tpm_limit_type=None, - model_has_failures=False, - call_type=call_type, - ) - - -def test_mcp_per_key_descriptor_created_for_matching_server_v3(): - handler, _ = _make_mcp_handler() - api_key = hash_token("sk-mcp-key") - user_api_key_dict = UserAPIKeyAuth( - api_key=api_key, - metadata={"mcp_rpm_limit": {"github": 5}}, - ) - - descriptors = _build_mcp_descriptors( - handler, user_api_key_dict, {"mcp_server_name": "github"} - ) - - descriptor = _find_descriptor(descriptors, "mcp_per_key") - assert descriptor is not None - assert descriptor["value"] == f"{api_key}:github" - assert descriptor["rate_limit"]["requests_per_unit"] == 5 - # MCP tool calls have no token usage; tokens_per_unit must stay None so the - # TPM reservation path is never engaged (otherwise budget would leak). - assert descriptor["rate_limit"]["tokens_per_unit"] is None - - -def test_mcp_per_key_descriptor_skipped_for_non_matching_server_v3(): - handler, _ = _make_mcp_handler() - user_api_key_dict = UserAPIKeyAuth( +@pytest.mark.asyncio +async def test_mcp_per_key_rate_limit_uses_trusted_server_alias_v3() -> None: + handler, local_cache = _make_mcp_handler() + user_api_key_dict: Final = UserAPIKeyAuth( api_key=hash_token("sk-mcp-key"), - metadata={"mcp_rpm_limit": {"github": 5}}, + metadata={"mcp_rpm_limit": {"github-alias": 1}}, + ) + server: Final = MCPServer( + server_id="server-1", + name="github", + alias="github-alias", + server_name="github", + transport=MCPTransport.http, ) - descriptors = _build_mcp_descriptors( - handler, user_api_key_dict, {"mcp_server_name": "slack"} + await handler.enforce_mcp_server_rate_limits(user_api_key_dict, server) + + with pytest.raises(ProxyRateLimitError, match="mcp_per_key"): + await handler.enforce_mcp_server_rate_limits(user_api_key_dict, server) + + assert all( + value == 0 + for key, value in local_cache.in_memory_cache.cache_dict.items() + if key.endswith(":tokens") ) - assert _find_descriptor(descriptors, "mcp_per_key") is None - - -def test_mcp_descriptor_skipped_for_non_mcp_request_v3(): - """A non-MCP request must not create an MCP descriptor even if the caller - injects mcp_server_name in the body; otherwise an LLM call could consume a - target server's MCP quota and 429 legitimate tool calls.""" - handler, _ = _make_mcp_handler() - user_api_key_dict = UserAPIKeyAuth( - api_key=hash_token("sk-mcp-key"), - metadata={"mcp_rpm_limit": {"github": 5}}, - ) - - descriptors = _build_mcp_descriptors( - handler, - user_api_key_dict, - {"model": "gpt-4", "mcp_server_name": "github"}, - call_type="completion", - ) - - assert _find_descriptor(descriptors, "mcp_per_key") is None - - -def test_mcp_descriptor_skipped_for_raw_rest_body_v3(): - handler, _ = _make_mcp_handler() - user_api_key_dict = UserAPIKeyAuth( - api_key=hash_token("sk-mcp-key"), - team_id="team-1", - metadata={"mcp_rpm_limit": {"github": 5}}, - team_metadata={"mcp_rpm_limit": {"github": 3}}, - ) - - descriptors = _build_mcp_descriptors( - handler, - user_api_key_dict, - { - "server_id": "slack", - "name": "demo-tool", - "arguments": {}, - "mcp_server_name": "github", - }, - ) - - assert _find_descriptor(descriptors, "mcp_per_key") is None - assert _find_descriptor(descriptors, "mcp_per_team") is None - - -def test_mcp_per_team_descriptor_created_from_team_metadata_v3(): - handler, _ = _make_mcp_handler() - user_api_key_dict = UserAPIKeyAuth( - api_key=hash_token("sk-mcp-key"), - team_id="team-1", - team_metadata={"mcp_rpm_limit": {"github": 3}}, - ) - - descriptors = _build_mcp_descriptors( - handler, user_api_key_dict, {"mcp_server_name": "github"} - ) - - descriptor = _find_descriptor(descriptors, "mcp_per_team") - assert descriptor is not None - assert descriptor["value"] == "team-1:github" - assert descriptor["rate_limit"]["requests_per_unit"] == 3 - assert descriptor["rate_limit"]["tokens_per_unit"] is None - @pytest.mark.asyncio -async def test_mcp_per_key_rpm_enforced_v3(monkeypatch): - """ - A key configured with mcp_rpm_limit={"github": 2} must allow 2 calls to the - github MCP server within the window and reject the 3rd with a 429, while - calls to a different MCP server are unaffected. - """ - monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") - api_key = hash_token("sk-mcp-enforce") - local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) +async def test_mcp_per_key_rejection_does_not_consume_shared_server_rpm_v3() -> None: + handler, _ = _make_mcp_handler() + server: Final = MCPServer( + server_id="server-shared", + name="github", + server_name="github", + transport=MCPTransport.http, + rpm=2, + ) + first_key: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-limited"), + metadata={"mcp_rpm_limit": {"github": 1}}, + ) + second_key: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-unlimited")) + + await handler.enforce_mcp_server_rate_limits(first_key, server) + with pytest.raises(ProxyRateLimitError, match="mcp_per_key") as key_rejected: + await handler.enforce_mcp_server_rate_limits(first_key, server) + + assert key_rejected.value.headers is not None + assert key_rejected.value.headers["retry-after"] == str(handler.window_size) + assert key_rejected.value.headers["rate_limit_type"] == "requests" + assert key_rejected.value.headers["reset_at"] + await handler.enforce_mcp_server_rate_limits(second_key, server) + + with pytest.raises(ProxyRateLimitError, match="mcp_server") as server_rejected: + await handler.enforce_mcp_server_rate_limits(second_key, server) + + assert server_rejected.value.headers is not None + assert server_rejected.value.headers["retry-after"] == str(handler.window_size) + assert server_rejected.value.headers["rate_limit_type"] == "requests" + assert server_rejected.value.headers["reset_at"] + + +@pytest.mark.asyncio +async def test_mcp_per_key_rate_limit_is_scoped_to_server_identity_v3() -> None: + handler, _ = _make_mcp_handler() + user_api_key_dict: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 1}}, + ) + github: Final = MCPServer( + server_id="server-github", + name="github", + server_name="github", + transport=MCPTransport.http, + ) + slack: Final = MCPServer( + server_id="server-slack", + name="slack", + server_name="slack", + transport=MCPTransport.http, ) - window_starts: Dict[str, int] = {} - request_counts: Dict[str, int] = {} + await handler.enforce_mcp_server_rate_limits(user_api_key_dict, github) + await handler.enforce_mcp_server_rate_limits(user_api_key_dict, slack) - async def mock_batch_rate_limiter(*args, **kwargs): - keys = kwargs.get("keys") if kwargs else args[0] - args_list = kwargs.get("args") if kwargs else args[1] - now = args_list[0] - window_size = args_list[1] - results = [] - for i in range(0, len(keys), 2): - window_key = keys[i] - counter_key = keys[i + 1] - prev_window = window_starts.get(window_key) - prev_counter = request_counts.get(counter_key, 0) - if prev_window is None or (now - prev_window) >= window_size: - window_starts[window_key] = now - new_counter = 1 - else: - new_counter = prev_counter + 1 - request_counts[counter_key] = new_counter - results.append(now) - results.append(new_counter) - return results + with pytest.raises(ProxyRateLimitError, match="mcp_per_key"): + await handler.enforce_mcp_server_rate_limits(user_api_key_dict, github) - handler.batch_rate_limiter_script = mock_batch_rate_limiter - user_api_key_dict = UserAPIKeyAuth( - api_key=api_key, - metadata={"mcp_rpm_limit": {"github": 2}}, +@pytest.mark.asyncio +async def test_raw_mcp_server_name_does_not_create_mcp_descriptor_v3() -> None: + handler, _ = _make_mcp_handler() + user_api_key_dict: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 5}}, ) - for _ in range(2): - await handler.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=local_cache, - data={"mcp_server_name": "github"}, - call_type="call_mcp_tool", - ) + local_cache: Final = DualCache() + await handler.async_pre_call_hook( + user_api_key_dict, + local_cache, + {"model": "gpt-4o-mini", "mcp_server_name": "github"}, + "call_mcp_tool", + ) - with pytest.raises(HTTPException) as exc_info: - await handler.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=local_cache, - data={"mcp_server_name": "github"}, - call_type="call_mcp_tool", - ) - assert exc_info.value.status_code == 429 + assert not any("mcp_per_" in key for key in local_cache.in_memory_cache.cache_dict) - # A different server has no configured limit -> not rate limited. - for _ in range(5): - await handler.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=local_cache, - data={"mcp_server_name": "slack"}, - call_type="call_mcp_tool", - ) - # The TPM counter must never be created for an MCP descriptor. - assert not any(":tokens" in key and "github" in key for key in request_counts) +@pytest.mark.asyncio +async def test_mcp_per_team_rate_limit_is_enforced_from_team_metadata_v3() -> None: + handler, _ = _make_mcp_handler() + server: Final = MCPServer( + server_id="server-1", + name="github", + server_name="github", + transport=MCPTransport.http, + ) + first_key: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-first"), + team_id="team-1", + team_metadata={"mcp_rpm_limit": {"github": 1}}, + ) + second_key: Final = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-second"), + team_id="team-1", + team_metadata={"mcp_rpm_limit": {"github": 1}}, + ) + + await handler.enforce_mcp_server_rate_limits(first_key, server) + + with pytest.raises(ProxyRateLimitError, match="mcp_per_team"): + await handler.enforce_mcp_server_rate_limits(second_key, server) + + +@pytest.mark.asyncio +async def test_mcp_server_rpm_is_shared_across_keys_v3() -> None: + handler, _ = _make_mcp_handler() + server: Final = MCPServer( + server_id="server-1", + name="github", + server_name="github", + transport=MCPTransport.http, + rpm=2, + ) + first_key: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-first")) + second_key: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-second")) + + await handler.enforce_mcp_server_rate_limits(first_key, server) + await handler.enforce_mcp_server_rate_limits(second_key, server) + + with pytest.raises(ProxyRateLimitError, match="mcp_server"): + await handler.enforce_mcp_server_rate_limits(second_key, server) + + +@pytest.mark.asyncio +async def test_mcp_server_rpm_zero_rejects_first_request_v3() -> None: + handler, _ = _make_mcp_handler() + server: Final = MCPServer( + server_id="server-zero", + name="github", + server_name="github", + transport=MCPTransport.http, + rpm=0, + ) + + with pytest.raises(ProxyRateLimitError, match="mcp_server"): + await handler.enforce_mcp_server_rate_limits(None, server) + + +@pytest.mark.asyncio +async def test_mcp_server_without_any_rate_limits_skips_cache_v3() -> None: + from unittest.mock import AsyncMock, patch + + handler, local_cache = _make_mcp_handler() + server: Final = MCPServer( + server_id="server-unlimited", + name="github", + server_name="github", + transport=MCPTransport.http, + ) + + with patch.object(handler, "should_rate_limit", new_callable=AsyncMock) as should_rate_limit: + await handler.enforce_mcp_server_rate_limits(None, server) + + should_rate_limit.assert_not_awaited() + assert local_cache.in_memory_cache.cache_dict == {} def test_get_key_mcp_rpm_limit_precedence(): diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 790df66b909..3c2e099e06c 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -58,6 +58,7 @@ def generate_mock_mcp_server_db_record( url: str = "https://db-server.example.com/mcp", transport: str = "sse", auth_type: Optional[str] = None, + rpm: int | None = None, ) -> LiteLLM_MCPServerTable: """Generate a mock MCP server record from database""" now = datetime.now() @@ -71,6 +72,7 @@ def generate_mock_mcp_server_db_record( updated_at=now, created_by="test_user", updated_by="test_user", + rpm=rpm, ) @@ -4080,6 +4082,7 @@ class TestUpdateMCPServer: url="https://test.example.com/mcp", transport="http", ) + assert existing_server.rpm is None existing_server.extra_headers = [] # Initially empty # Create update request with extra_headers @@ -4087,6 +4090,7 @@ class TestUpdateMCPServer: server_id="test-server-1", alias="Updated Test Server", extra_headers=["X-Custom-Header", "X-Another-Header"], + rpm=5, ) # Mock the updated server with extra_headers @@ -4095,6 +4099,7 @@ class TestUpdateMCPServer: alias="Updated Test Server", url="https://test.example.com/mcp", transport="http", + rpm=5, ) updated_server.extra_headers = ["X-Custom-Header", "X-Another-Header"] @@ -4148,10 +4153,12 @@ class TestUpdateMCPServer: "X-Another-Header", ] assert called_payload.alias == "Updated Test Server" + assert called_payload.rpm == 5 # Verify the result includes extra_headers assert result.extra_headers == ["X-Custom-Header", "X-Another-Header"] assert result.alias == "Updated Test Server" + assert result.rpm == 5 class TestAddMCPServerAtomicity: @@ -4174,9 +4181,10 @@ class TestAddMCPServerAtomicity: alias="echo", url="https://echo.example.com/mcp", transport=MCPTransport.http, + rpm=5, ) admin = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") - created_server = generate_mock_mcp_server_db_record(server_id="created-1", alias="echo") + created_server = generate_mock_mcp_server_db_record(server_id="created-1", alias="echo", rpm=5) mock_manager = MagicMock() mock_manager.add_server = AsyncMock() @@ -4203,8 +4211,10 @@ class TestAddMCPServerAtomicity: result = await add_mcp_server(payload=payload, user_api_key_dict=admin) create_mock.assert_awaited_once() + assert create_mock.call_args.args[1].rpm == 5 mock_manager.reload_servers_from_database.assert_awaited_once() assert result.server_id == "created-1" + assert result.rpm == 5 @pytest.mark.asyncio async def test_create_500s_and_skips_registry_when_db_write_fails(self): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx index f6ad2e38658..65656a053ce 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx @@ -1044,6 +1044,8 @@ describe("CreateMCPServer", () => { const limitInput = screen.getByPlaceholderText("e.g. 10"); fireEvent.change(limitInput, { target: { value: "5" } }); + const rpmInput = screen.getByPlaceholderText("e.g. 60"); + fireEvent.change(rpmInput, { target: { value: "7" } }); vi.mocked(networking.createMCPServer).mockResolvedValue({ server_id: "new-server-1", @@ -1069,6 +1071,7 @@ describe("CreateMCPServer", () => { const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; expect(payload.max_concurrent_requests).toBe(5); + expect(payload.rpm).toBe(7); }); it("routes OAuth Token Exchange (OBO) config to the backend payload", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx index d53d11ddcfb..07618abb9a4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx @@ -816,6 +816,28 @@ const CreateMCPServer: React.FC = ({ )} + + RPM limit (all callers) + + + + + } + name="rpm" + > + {(control) => ( + + )} + + {/* Authentication - show for HTTP, SSE, and OpenAPI */} {transportType !== "stdio" && transportType !== "" && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.cases.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.cases.ts index ef8f728a609..7db2d0dda36 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.cases.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.cases.ts @@ -14,6 +14,7 @@ const SERVER: MCPServer = { updated_at: "2024-01-01T00:00:00Z", updated_by: "user-1", mcp_access_groups: [], + rpm: undefined, }; export const baseUi: EditServerUiState = { @@ -45,6 +46,7 @@ const ROOT = { url: "https://example.com/mcp", auth_type: "none", max_concurrent_requests: undefined, + rpm: undefined, mcp_access_groups: [], extra_headers: [], static_headers: [], diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.integration.test.tsx index f3cd37cc580..7aed6488b54 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.integration.test.tsx @@ -88,6 +88,7 @@ const EXPECTED_BASE: Readonly> = { tool_allowlist_enforced: false, }, oauth_passthrough: false, + rpm: undefined, server_id: "srv_1", server_name: "srv", static_headers: {}, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index 91db38db389..35e212faf37 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -2209,12 +2209,14 @@ describe("MCPServerEdit (max concurrent requests)", () => { ...interactiveOAuthServer, auth_type: "none", max_concurrent_requests: 5, + rpm: 5, }; it("prefills the existing limit and sends an updated value in the payload", async () => { vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...limitedServer, max_concurrent_requests: 2, + rpm: 2, }); render( @@ -2229,8 +2231,11 @@ describe("MCPServerEdit (max concurrent requests)", () => { const limitInput = screen.getByPlaceholderText("e.g. 10") as HTMLInputElement; expect(limitInput.value).toBe("5"); + const rpmInput = screen.getByPlaceholderText("e.g. 60") as HTMLInputElement; + expect(rpmInput.value).toBe("5"); fireEvent.change(limitInput, { target: { value: "2" } }); + fireEvent.change(rpmInput, { target: { value: "2" } }); const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); await act(async () => { @@ -2243,6 +2248,7 @@ describe("MCPServerEdit (max concurrent requests)", () => { const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; expect(payload.max_concurrent_requests).toBe(2); + expect(payload.rpm).toBe(2); }); it("sends null when the limit is cleared so the backend unsets it", async () => { @@ -2278,6 +2284,40 @@ describe("MCPServerEdit (max concurrent requests)", () => { const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; expect(payload.max_concurrent_requests).toBeNull(); }); + + it("sends null when the RPM limit is cleared", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...limitedServer, + rpm: null, + }); + + render( + , + ); + + const rpmInput = screen.getByPlaceholderText("e.g. 60") as HTMLInputElement; + expect(rpmInput.value).toBe("5"); + + fireEvent.change(rpmInput, { target: { value: "" } }); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.rpm).toBeNull(); + }); }); describe("MCPServerEdit (dcr_bridge toggle)", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 1b1e4c27bde..f35e3ab3d3a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -972,6 +972,28 @@ const MCPServerEdit: React.FC = ({ )} + + RPM limit (all callers) + + + + + } + name="rpm" + > + {(control) => ( + + )} + + {/* Authentication - for HTTP, SSE, and OpenAPI */} {!isStdioTransport && ( <> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.test.ts index 13aec81d9e8..a235069ece1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.test.ts @@ -40,6 +40,7 @@ describe("edit root: transport gates", () => { "description", "transport", "max_concurrent_requests", + "rpm", "command", "args", "env_json", @@ -210,7 +211,7 @@ describe("create root: where it diverges from edit", () => { }); }); -const ALWAYS = ["server_name", "alias", "description", "transport", "max_concurrent_requests"]; +const ALWAYS = ["server_name", "alias", "description", "transport", "max_concurrent_requests", "rpm"]; const PERMS = [ "allow_all_keys", "available_on_public_internet", @@ -478,6 +479,7 @@ describe("projection shape", () => { expect("description" in projected).toBe(true); expect(projected.description).toBeUndefined(); expect(Object.keys(projected)).toContain("max_concurrent_requests"); + expect(Object.keys(projected)).toContain("rpm"); }); it("emits mounted-but-unset CREDENTIAL keys as undefined rather than omitting them", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.ts index af9cbb58b2b..e980b69132c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mountedServerFields.ts @@ -8,7 +8,14 @@ export interface MountedFieldNames { const ENTRA_OBO_PROFILE = "entra_obo"; -const ALWAYS_MOUNTED_ROOT = ["server_name", "alias", "description", "transport", "max_concurrent_requests"] as const; +const ALWAYS_MOUNTED_ROOT = [ + "server_name", + "alias", + "description", + "transport", + "max_concurrent_requests", + "rpm", +] as const; const PERMISSION_SECTION_ROOT = [ "allow_all_keys", diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index afeff869b7a..2a6b3f2793b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -453,6 +453,7 @@ export interface MCPServer { oauth_passthrough?: boolean; dcr_bridge?: boolean | null; max_concurrent_requests?: number | null; + rpm?: number | null; /** Redacted to null in server responses; present when constructing a server locally. */ credentials?: Record | null; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index bce470a36f1..aa29164663b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -35069,6 +35069,8 @@ export interface components { review_notes?: string | null; /** Reviewed At */ reviewed_at?: string | null; + /** Rpm */ + rpm?: number | null; /** Server Id */ server_id: string; /** Server Name */ @@ -39132,6 +39134,8 @@ export interface components { per_server_oauth_discovery: boolean; /** Registration Url */ registration_url?: string | null; + /** Rpm */ + rpm?: number | null; /** Server Id */ server_id?: string | null; /** Server Name */ @@ -48539,6 +48543,8 @@ export interface components { per_server_oauth_discovery: boolean; /** Registration Url */ registration_url?: string | null; + /** Rpm */ + rpm?: number | null; /** Server Id */ server_id: string; /** Server Name */