mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
e66dbfc366
commit
e48f8d928d
36 changed files with 1636 additions and 221 deletions
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "rpm" INTEGER;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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'",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
271
tests/integration/mcp/test_mcp_rate_limits.py
Normal file
271
tests/integration/mcp/test_mcp_rate_limits.py
Normal 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()) == ()
|
||||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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: {
|
||||
|
|
|
|||
|
|
@ -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={})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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: [],
|
||||
|
|
|
|||
|
|
@ -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: {},
|
||||
|
|
|
|||
|
|
@ -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)", () => {
|
||||
|
|
|
|||
|
|
@ -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 && (
|
||||
<>
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue