feat(proxy): add per-MCP-server RPM rate limiting for keys and teams

Adds mcp_rpm_limit, a dict keyed by MCP server name (alias if set, else the
configured name) that caps requests per minute per server for a key or team.
The v3 rate limiter builds a per-server descriptor only when a limit is
configured for the server being called, so other servers stay uncapped and no
TPM reservation is engaged. Server identity is surfaced into the request data
via mcp_rate_limit_server_name so the limiter can resolve it.
This commit is contained in:
Sameer Kankute 2026-06-02 13:15:24 +05:30
parent f48a87ef12
commit f0b5f03bb6
No known key found for this signature in database
11 changed files with 675 additions and 1 deletions

View file

@ -2650,6 +2650,9 @@ class MCPServerManager:
"name": name,
"arguments": arguments,
"server_name": server_name,
"mcp_rate_limit_server_name": server.alias
or server.server_name
or server.name,
"user_api_key_auth": user_api_key_auth,
"user_api_key_user_id": (
getattr(user_api_key_auth, "user_id", None)

View file

@ -1050,6 +1050,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
prompts: Optional[List[str]] = None
@ -1851,6 +1852,7 @@ class NewTeamRequest(TeamBase):
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
model_tpm_limit: Optional[Dict[str, int]] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
@ -1920,6 +1922,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
prompts: Optional[List[str]] = None
model_rpm_limit: Optional[Dict[str, int]] = None
model_tpm_limit: Optional[Dict[str, int]] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
enforced_batch_output_expires_after: Optional[dict] = None
enforced_file_expires_after: Optional[dict] = None
@ -4285,6 +4288,7 @@ class PassThroughEndpointLoggingTypedDict(TypedDict):
LiteLLM_ManagementEndpoint_MetadataFields = [
"model_rpm_limit",
"model_tpm_limit",
"mcp_rpm_limit",
"rpm_limit_type",
"tpm_limit_type",
"enforced_params",

View file

@ -934,6 +934,40 @@ def get_team_model_tpm_limit(
return None
def get_key_mcp_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:
"""
Get the per-MCP-server rpm limit for a given api key.
Priority order (returns first found):
1. Key metadata (mcp_rpm_limit)
2. Team metadata (mcp_rpm_limit)
The returned dict is keyed by MCP server name (alias if set, else the
configured server name).
"""
if user_api_key_dict.metadata:
result = user_api_key_dict.metadata.get("mcp_rpm_limit")
if result:
return result
if user_api_key_dict.team_metadata:
team_limit = user_api_key_dict.team_metadata.get("mcp_rpm_limit")
if team_limit:
return team_limit
return None
def get_team_mcp_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:
if user_api_key_dict.team_metadata:
return user_api_key_dict.team_metadata.get("mcp_rpm_limit")
return None
def get_project_model_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:

View file

@ -1375,6 +1375,79 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
)
def _add_mcp_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
mcp_server_name: Optional[str],
descriptors: List[RateLimitDescriptor],
) -> None:
"""
Add a per-MCP-server rpm descriptor for the API key, if a limit is
configured for the server being called.
MCP tool calls have no token usage, so only requests_per_unit is set;
tokens_per_unit stays None so the TPM reservation path is never engaged.
"""
from litellm.proxy.auth.auth_utils import get_key_mcp_rpm_limit
if not mcp_server_name or not user_api_key_dict.api_key:
return
mcp_rpm_limit = get_key_mcp_rpm_limit(user_api_key_dict)
if not mcp_rpm_limit:
return
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
return
descriptors.append(
RateLimitDescriptor(
key="mcp_per_key",
value=f"{user_api_key_dict.api_key}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
def _add_mcp_per_team_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
mcp_server_name: Optional[str],
descriptors: List[RateLimitDescriptor],
) -> None:
"""
Add a per-MCP-server rpm descriptor for the team, if a limit is
configured for the server being called.
"""
from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit
if not mcp_server_name or not user_api_key_dict.team_id:
return
mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict)
if not mcp_rpm_limit:
return
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
return
descriptors.append(
RateLimitDescriptor(
key="mcp_per_team",
value=f"{user_api_key_dict.team_id}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
def _should_enforce_rate_limit(
self,
limit_type: Optional[str],
@ -1653,6 +1726,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
descriptors=descriptors,
)
# Per-MCP-server rate limits
mcp_server_name = 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,
)
if (
get_team_model_rpm_limit(user_api_key_dict) is not None
or get_team_model_tpm_limit(user_api_key_dict) is not None

View file

@ -1388,6 +1388,7 @@ async def generate_key_fn(
- model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget.
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
@ -1606,6 +1607,7 @@ async def generate_service_account_key_fn(
- model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget.
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request

View file

@ -863,8 +863,9 @@ async def new_team( # noqa: PLR0915
- members_with_roles: List[{"role": "admin" or "user", "user_id": "<user-id>"}] - A list of users and their roles in the team. Get user_id when making a new user via `/user/new`.
- team_member_permissions: Optional[List[str]] - A list of routes that non-admin team members can access. example: ["/key/generate", "/key/update", "/key/delete"]
- metadata: Optional[dict] - Metadata for team, store information for team. Example metadata = {"extra_info": "some info"}
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team.
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team.
- model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit for this team - applied across all keys for this team.
- mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team.
- tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit
- rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit
- rpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] - The type of RPM limit enforcement. Use "guaranteed_throughput" to raise an error if overallocating RPM, or "best_effort_throughput" for best effort enforcement.

View file

@ -643,6 +643,7 @@ class ProxyLogging:
"user_api_key_request_route": kwargs.get("user_api_key_request_route"),
"mcp_tool_name": request_obj.tool_name, # Keep original for reference
"mcp_arguments": request_obj.arguments, # Keep original for reference
"mcp_server_name": kwargs.get("mcp_rate_limit_server_name"),
# Raw Bearer token from the original HTTP request — allows guardrails
# (e.g. MCPJWTSigner) to independently verify the caller's identity
# before re-signing an outbound token (FR-5 verify+re-sign).

238
scripts/test_mcp_rpm_limit.sh Executable file
View file

@ -0,0 +1,238 @@
#!/usr/bin/env bash
#
# End-to-end manual test for per-MCP RPM rate limiting.
#
# What it does, with no other setup required:
# 1. Writes a throwaway proxy config with one stdio MCP server (the
# `uvx mcp-server-fetch` server, aliased "fetch_mcp").
# 2. Boots the proxy in the background and waits until it is ready.
# 3. Generates two keys, both with full access to fetch_mcp:
# - "limited" key: mcp_rpm_limit caps "fetch_mcp" at 2 req/min.
# - "control" key: mcp_rpm_limit caps a DIFFERENT server name
# ("other_mcp") at 2 req/min, so calls to fetch_mcp
# are uncapped.
# 4. Fires 4 fetch_mcp calls with each key. The limited key must trip at the
# 3rd call (429); the control key must never be rate limited. This proves
# the limit is keyed per MCP server name, not globally per key.
#
# Using one physical server with two keys (rather than two servers) keeps the
# test deterministic: a single server's access resolution is exercised, and the
# only variable between the two runs is which server name the key's limit
# targets.
#
# Usage:
# ./scripts/test_mcp_rpm_limit.sh
#
# Requirements: jq, curl, uvx (for the stdio fetch MCP server), and a reachable
# DATABASE_URL (read from .env).
set -uo pipefail
# --- locate repo root (this script lives in <repo>/scripts) -------------------
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/.." && pwd)"
cd "${REPO_ROOT}"
RPM_LIMIT=2
SERVER="fetch_mcp"
OTHER_SERVER="other_mcp"
WORKDIR="$(mktemp -d)"
CONFIG="${WORKDIR}/mcp_rpm_test_config.yaml"
PROXY_LOG="${WORKDIR}/proxy.log"
PROXY_PID=""
# --- load secrets (DATABASE_URL, provider keys, master key) -------------------
# Parse .env line-by-line and export each KEY=VALUE verbatim. We avoid
# `source`-ing it because some values contain characters (e.g. '#') that the
# shell would try to execute.
if [[ -f .env ]]; then
while IFS= read -r line; do
[[ "${line}" =~ ^[A-Za-z_][A-Za-z0-9_]*= ]] || continue
key="${line%%=*}"
val="${line#*=}"
val="${val%\"}"; val="${val#\"}" # strip surrounding double quotes
val="${val%\'}"; val="${val#\'}" # strip surrounding single quotes
export "${key}=${val}"
done < .env
fi
MASTER_KEY="${LITELLM_MASTER_KEY:-sk-1234}"
# --- pick a free TCP port (start at 4000) so we never collide with a proxy
# already running from a prior session -----------------------------------------
PORT=""
for candidate in $(seq 4000 4050); do
if ! lsof -iTCP:"${candidate}" -sTCP:LISTEN -n -P >/dev/null 2>&1; then
PORT="${candidate}"
break
fi
done
if [[ -z "${PORT}" ]]; then
echo "ERROR: no free port found in 4000-4050"
exit 1
fi
BASE="http://localhost:${PORT}"
echo ">> using port ${PORT}"
cleanup() {
if [[ -n "${PROXY_PID}" ]] && kill -0 "${PROXY_PID}" 2>/dev/null; then
echo ">> stopping proxy (pid ${PROXY_PID})"
kill "${PROXY_PID}" 2>/dev/null
# kill the whole process group in case uvicorn spawned children
pkill -P "${PROXY_PID}" 2>/dev/null
wait "${PROXY_PID}" 2>/dev/null
fi
echo ">> logs kept at: ${PROXY_LOG}"
}
trap cleanup EXIT
require() { command -v "$1" >/dev/null 2>&1 || { echo "ERROR: '$1' is required but not installed"; exit 1; }; }
require jq
require curl
# --- 1. write throwaway config ------------------------------------------------
cat > "${CONFIG}" <<YAML
model_list: []
mcp_servers:
${SERVER}:
transport: "stdio"
command: "uvx"
args: ["mcp-server-fetch", "--ignore-robots-txt"]
alias: "${SERVER}"
allow_all_keys: true
general_settings:
master_key: ${MASTER_KEY}
store_model_in_db: false
YAML
echo ">> config written to ${CONFIG}"
# --- 2. start proxy -----------------------------------------------------------
echo ">> starting proxy on :${PORT} (log: ${PROXY_LOG})"
# Put the repo root first on PYTHONPATH so the local litellm source shadows any
# stale `litellm` installed in site-packages (running the cli as a script puts
# litellm/proxy/ on sys.path instead of the repo root).
PYTHONPATH="${REPO_ROOT}:${PYTHONPATH:-}" python litellm/proxy/proxy_cli.py \
--config "${CONFIG}" \
--port "${PORT}" \
--detailed_debug \
--use_v2_migration_resolver > "${PROXY_LOG}" 2>&1 &
PROXY_PID=$!
echo -n ">> waiting for readiness"
ready=false
for _ in $(seq 1 90); do
if curl -sf "${BASE}/health/readiness" >/dev/null 2>&1; then
ready=true
break
fi
if ! kill -0 "${PROXY_PID}" 2>/dev/null; then
echo ""
echo "ERROR: proxy process died during startup. Tail of log:"
tail -n 40 "${PROXY_LOG}"
exit 1
fi
echo -n "."
sleep 1
done
echo ""
if [[ "${ready}" != "true" ]]; then
echo "ERROR: proxy did not become ready in time. Tail of log:"
tail -n 40 "${PROXY_LOG}"
exit 1
fi
echo ">> proxy is ready"
# --- 3. generate the two keys -------------------------------------------------
# Both keys get explicit access to fetch_mcp. They differ only in which server
# name their mcp_rpm_limit targets.
generate_key() {
local rpm_target="$1"
curl -sf -X POST "${BASE}/key/generate" \
-H "Authorization: Bearer ${MASTER_KEY}" \
-H "Content-Type: application/json" \
-d "{\"mcp_rpm_limit\": {\"${rpm_target}\": ${RPM_LIMIT}}, \"object_permission\": {\"mcp_servers\": [\"${SERVER}\"]}}" \
| jq -r '.key'
}
echo ">> generating limited key (mcp_rpm_limit {\"${SERVER}\": ${RPM_LIMIT}})"
LIMITED_KEY="$(generate_key "${SERVER}")"
echo ">> generating control key (mcp_rpm_limit {\"${OTHER_SERVER}\": ${RPM_LIMIT}})"
CONTROL_KEY="$(generate_key "${OTHER_SERVER}")"
for k in "${LIMITED_KEY}" "${CONTROL_KEY}"; do
if [[ -z "${k}" || "${k}" == "null" ]]; then
echo "ERROR: /key/generate failed. Is DATABASE_URL set and reachable?"
tail -n 40 "${PROXY_LOG}"
exit 1
fi
done
echo ">> limited key: ${LIMITED_KEY:0:12}... control key: ${CONTROL_KEY:0:12}..."
# --- discover a real tool name on the server ----------------------------------
TOOL_NAME="$(curl -sf "${BASE}/mcp-rest/tools/list?server_id=${SERVER}" \
-H "Authorization: Bearer ${LIMITED_KEY}" 2>/dev/null \
| jq -r '.tools[0].name // empty')"
if [[ -z "${TOOL_NAME}" ]]; then
echo ">> could not auto-discover a tool name; falling back to 'fetch'"
TOOL_NAME="fetch"
fi
echo ">> using tool: ${TOOL_NAME}"
# tool_call: the server alias is accepted directly as server_id. Point the fetch
# tool at the proxy's own health endpoint so the call is fast and always
# reachable; that way a non-429 response unambiguously means "the rate limiter
# let this through" rather than "the upstream fetch flaked".
FETCH_URL="${BASE}/health/readiness"
call_mcp() {
local key="$1"
curl -s -o /dev/null -w "%{http_code}" \
-X POST "${BASE}/mcp-rest/tools/call" \
-H "Authorization: Bearer ${key}" \
-H "Content-Type: application/json" \
-d "{\"server_id\": \"${SERVER}\", \"name\": \"${TOOL_NAME}\", \"arguments\": {\"url\": \"${FETCH_URL}\", \"max_length\": 100}}"
}
# --- 4a. limited key: expect 429 once the cap is exceeded ---------------------
echo ""
echo "=== limited key (caps ${SERVER} at ${RPM_LIMIT}/min) ==="
limited_codes=()
for i in 1 2 3 4; do
code="$(call_mcp "${LIMITED_KEY}")"
limited_codes+=("${code}")
echo " call ${i} -> HTTP ${code}"
done
# --- 4b. control key: caps a different server name, so fetch_mcp is uncapped --
echo ""
echo "=== control key (caps ${OTHER_SERVER}, so ${SERVER} is uncapped) ==="
control_codes=()
for i in 1 2 3 4; do
code="$(call_mcp "${CONTROL_KEY}")"
control_codes+=("${code}")
echo " call ${i} -> HTTP ${code}"
done
# --- evaluate -----------------------------------------------------------------
echo ""
echo "=== result ==="
pass=true
# limited: first two must NOT be 429, last two MUST be 429
[[ "${limited_codes[0]}" != "429" ]] || { echo "FAIL: limited call 1 was rate limited"; pass=false; }
[[ "${limited_codes[1]}" != "429" ]] || { echo "FAIL: limited call 2 was rate limited"; pass=false; }
[[ "${limited_codes[2]}" == "429" ]] || { echo "FAIL: limited call 3 was NOT rate limited (got ${limited_codes[2]})"; pass=false; }
[[ "${limited_codes[3]}" == "429" ]] || { echo "FAIL: limited call 4 was NOT rate limited (got ${limited_codes[3]})"; pass=false; }
# control: none may be 429
for c in "${control_codes[@]}"; do
[[ "${c}" != "429" ]] || { echo "FAIL: control key was rate limited on ${SERVER} (got ${c})"; pass=false; }
done
if [[ "${pass}" == "true" ]]; then
echo "PASS: ${SERVER} tripped at call 3 (429) under the limited key; the control key (which caps ${OTHER_SERVER}) was never rate limited on ${SERVER}."
exit 0
else
echo "See proxy log for detail: ${PROXY_LOG}"
exit 1
fi

View file

@ -826,3 +826,87 @@ class TestUserAPIKeyAuthJwtClaims:
auth.jwt_claims = claims
assert auth.jwt_claims == claims
assert auth.jwt_claims["groups"] == ["admin"]
class TestMcpRateLimitServerNameSurfacing:
"""
The per-MCP-server rate limiter only sees the request `data` dict, so the
server identity must be surfaced into it. These tests pin the contract
between pre_call_tool_check, _convert_mcp_to_llm_format, and the limiter.
"""
def setup_method(self):
self.proxy_logging = ProxyLogging(user_api_key_cache=MagicMock())
def test_convert_mcp_to_llm_format_surfaces_rate_limit_server_name(self):
request_obj = MagicMock()
request_obj.tool_name = "list_repos"
request_obj.arguments = {"org": "acme"}
result = self.proxy_logging._convert_mcp_to_llm_format(
request_obj, {"mcp_rate_limit_server_name": "github"}
)
assert result["mcp_server_name"] == "github"
def test_convert_mcp_to_llm_format_server_name_none_when_absent(self):
request_obj = MagicMock()
request_obj.tool_name = "list_repos"
request_obj.arguments = {}
result = self.proxy_logging._convert_mcp_to_llm_format(request_obj, {})
assert result["mcp_server_name"] is None
@pytest.mark.asyncio
async def test_pre_call_tool_check_resolves_alias_for_rate_limit(self):
"""
The rate-limit server key must be the alias when set (falling back to
server_name), matching how an admin keys mcp_rpm_limit in config.
"""
manager = MCPServerManager()
server = MCPServer(
server_id="test-id",
name="gh",
alias="gh",
server_name="github_full_name",
url="https://example.com",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
)
captured = {}
def capture_convert(request_obj, kwargs):
captured["kwargs"] = kwargs
return {"model": "fake"}
proxy_logging = MagicMock(spec=ProxyLogging)
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
return_value=MagicMock()
)
proxy_logging._convert_mcp_to_llm_format = MagicMock(
side_effect=capture_convert
)
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock(
return_value={"arguments": {}}
)
with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
with patch.object(
manager,
"check_tool_permission_for_key_team",
new_callable=AsyncMock,
):
with patch.object(manager, "validate_allowed_params"):
await manager.pre_call_tool_check(
name="list_repos",
arguments={},
server_name="github_full_name",
user_api_key_auth=None,
proxy_logging_obj=proxy_logging,
server=server,
)
assert captured["kwargs"]["mcp_rate_limit_server_name"] == "gh"

View file

@ -2893,3 +2893,197 @@ async def test_pre_call_hook_rejects_caller_supplied_stash_values():
):
leaked = [k for k in _LITELLM_STASH_KEYS if k in channel]
assert not leaked, f"caller-supplied stash survived in {channel!r}: {leaked}"
# ----------------------- Per-MCP-server rate limiting (v3) -----------------------
def _make_mcp_handler():
local_cache = DualCache()
handler = _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):
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,
)
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(
api_key=hash_token("sk-mcp-key"),
metadata={"mcp_rpm_limit": {"github": 5}},
)
descriptors = _build_mcp_descriptors(
handler, user_api_key_dict, {"mcp_server_name": "slack"}
)
assert _find_descriptor(descriptors, "mcp_per_key") is None
def test_mcp_descriptor_skipped_for_non_mcp_request_v3():
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"})
assert _find_descriptor(descriptors, "mcp_per_key") 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)
)
window_starts: Dict[str, int] = {}
request_counts: Dict[str, int] = {}
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
handler.batch_rate_limiter_script = mock_batch_rate_limiter
user_api_key_dict = UserAPIKeyAuth(
api_key=api_key,
metadata={"mcp_rpm_limit": {"github": 2}},
)
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",
)
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
# 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)
def test_get_key_mcp_rpm_limit_precedence():
from litellm.proxy.auth.auth_utils import (
get_key_mcp_rpm_limit,
get_team_mcp_rpm_limit,
)
# Key metadata takes precedence over team metadata.
key_first = UserAPIKeyAuth(
api_key=hash_token("sk-mcp-key"),
metadata={"mcp_rpm_limit": {"github": 10}},
team_metadata={"mcp_rpm_limit": {"github": 99}},
)
assert get_key_mcp_rpm_limit(key_first) == {"github": 10}
# Falls back to team metadata when key has none.
team_only = UserAPIKeyAuth(
api_key=hash_token("sk-mcp-key"),
team_metadata={"mcp_rpm_limit": {"github": 7}},
)
assert get_key_mcp_rpm_limit(team_only) == {"github": 7}
assert get_team_mcp_rpm_limit(team_only) == {"github": 7}
# No configuration anywhere.
none_set = UserAPIKeyAuth(api_key=hash_token("sk-mcp-key"))
assert get_key_mcp_rpm_limit(none_set) is None
assert get_team_mcp_rpm_limit(none_set) is None

View file

@ -482,6 +482,33 @@ class TestSetObjectMetadataField:
_set_object_metadata_field(team, "model_rpm_limit", {"x": 1})
assert team.metadata == {"model_rpm_limit": {"x": 1}}
def test_mcp_rpm_limit_is_hoisted_into_metadata(self):
"""
Per-MCP-server rpm limits are stored in the metadata JSON column, not a
dedicated DB column. The key/team management endpoints rely on
LiteLLM_ManagementEndpoint_MetadataFields to move the request field into
metadata; this regression guards that mcp_rpm_limit is in that list and
round-trips through the same loop the endpoints use.
"""
from litellm.proxy._types import LiteLLM_ManagementEndpoint_MetadataFields
assert "mcp_rpm_limit" in LiteLLM_ManagementEndpoint_MetadataFields
from types import SimpleNamespace
team = LiteLLM_TeamTable(team_id="t1", metadata={})
mcp_rpm_limit = {"github": 100}
data = SimpleNamespace(mcp_rpm_limit=mcp_rpm_limit)
with patch(
"litellm.proxy.management_endpoints.common_utils._premium_user_check"
):
for field in LiteLLM_ManagementEndpoint_MetadataFields:
if getattr(data, field, None) is not None:
_set_object_metadata_field(team, field, getattr(data, field))
assert team.metadata["mcp_rpm_limit"] == mcp_rpm_limit
class TestRequireCallerUserIdForNonAdmin:
"""