feat(mcp): rate limit all MCP operations and add server-level rpm (#44600)

* feat(mcp): rate limit all MCP operations and add server-level rpm

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

* refactor(mcp): drop comments copied onto list fallbacks

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

* fix(mcp): count discovery once per server and rate limit REST tools listing

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

* test(mcp): cover rate-limited catalog error propagation

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

* test(mcp): avoid fastapi import in mcp operations test

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

* fix(mcp): apply server rate limits to paginated catalog listings

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

* refactor(mcp): restore main's unused prompt and resource listing helpers

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

* test(mcp): add server rate limit coverage

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

* test(mcp): harden rate-limit integration setup

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

* fix(mcp): prevent rejected calls from consuming shared limits

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

---------

Co-authored-by: joshua <joshua@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 16:30:52 -07:00 • committed by GitHub
parent e66dbfc366
commit e48f8d928d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
36 changed files with 1636 additions and 221 deletions

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "rpm" INTEGER;

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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": [
{

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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={})

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -816,6 +816,28 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
)}
</MountedFormField>
<MountedFormField
label={
<span className="text-sm font-medium text-foreground flex items-center">
RPM limit (all callers)
<SimpleTooltip content="Max requests per minute to this server across all keys and teams. Leave empty for no limit">
<Info className="ml-2 size-4 text-info hover:text-info/80 cursor-help" />
</SimpleTooltip>
</span>
}
name="rpm"
>
{(control) => (
<Input
{...numberControl(control, 0)}
min={0}
step={1}
placeholder="e.g. 60"
className="w-full rounded-lg"
/>
)}
</MountedFormField>
{/* Authentication - show for HTTP, SSE, and OpenAPI */}
{transportType !== "stdio" && transportType !== "" && (
<Collapsible defaultOpen className="mb-4">

View file

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

View file

@ -88,6 +88,7 @@ const EXPECTED_BASE: Readonly<Record<string, unknown>> = {
tool_allowlist_enforced: false,
},
oauth_passthrough: false,
rpm: undefined,
server_id: "srv_1",
server_name: "srv",
static_headers: {},

View file

@ -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(
<MCPServerEdit
mcpServer={limitedServer}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
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)", () => {

View file

@ -972,6 +972,28 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
)}
</MountedFormField>
<MountedFormField
label={
<span className="text-sm font-medium text-foreground flex items-center">
RPM limit (all callers)
<SimpleTooltip content="Max requests per minute to this server across all keys and teams. Leave empty for no limit">
<Info className="ml-2 size-4 text-info hover:text-info/80 cursor-help" />
</SimpleTooltip>
</span>
}
name="rpm"
>
{(control) => (
<Input
{...numberControl(control, 0)}
min={0}
step={1}
placeholder="e.g. 60"
className="w-full rounded-lg"
/>
)}
</MountedFormField>
{/* Authentication - for HTTP, SSE, and OpenAPI */}
{!isStdioTransport && (
<>

View file

@ -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", () => {

View file

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

View file

@ -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<string, unknown> | null;

View file

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