refactor(mcp): narrow catalog consistency change to ticket scope

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-23 05:18:27 +00:00
parent b0f1eb656c
commit 51892184a3
7 changed files with 25 additions and 356 deletions

View file

@ -1,97 +0,0 @@
from __future__ import annotations
import math
import os
import time
from collections.abc import Callable, Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Final
from fastapi import HTTPException, Request
from litellm.proxy.auth.login_throttle import source_group
from litellm.proxy.auth.network import TrustedProxyConfig, resolve_client_ip
@dataclass(slots=True)
class _Window:
started: float
requests: int = 0
active: int = 0
def _limit(name: str, default: int) -> int:
value: Final = int(os.environ.get(name, str(default)))
if value < 1:
raise ValueError(f"{name} must be positive")
return value
class MCPAdmissionLimiter:
def __init__(self, clock: Callable[[], float] = time.monotonic) -> None:
self._clock = clock
self._client_rpm = _limit("LITELLM_MCP_PUBLIC_RPM", 120)
self._worker_rpm = _limit("LITELLM_MCP_PUBLIC_WORKER_RPM", 600)
self._client_active = _limit("LITELLM_MCP_PUBLIC_MAX_IN_FLIGHT", 64)
self._worker_active = _limit("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", 128)
self._max_sources = _limit("LITELLM_MCP_PUBLIC_MAX_SOURCES", 4096)
self._sources: dict[str, _Window] = {}
self._worker = _Window(clock())
@staticmethod
def _reject(retry_after: int) -> None:
raise HTTPException(
status_code=429,
detail="MCP configuration request limit exceeded; retry later",
headers={"Retry-After": str(max(1, retry_after))},
)
@contextmanager
def admit(self, source: str) -> Iterator[None]:
now: Final = self._clock()
if now - self._worker.started >= 60:
self._worker = _Window(now, active=self._worker.active)
self._sources = {
key: value for key, value in self._sources.items() if value.active or now - value.started < 60
}
if self._worker.active >= self._worker_active:
self._reject(1)
if self._worker.requests >= self._worker_rpm:
self._reject(math.ceil(60 - (now - self._worker.started)))
previous: Final = self._sources.get(source)
if previous is None and len(self._sources) >= self._max_sources:
self._reject(math.ceil(60 - (now - self._worker.started)))
client: Final = (
previous
if previous is not None and now - previous.started < 60
else _Window(now, active=previous.active if previous is not None else 0)
)
if client.active >= self._client_active:
self._reject(1)
if client.requests >= self._client_rpm:
self._reject(math.ceil(60 - (now - client.started)))
self._sources[source] = client
client.requests += 1
client.active += 1
self._worker.requests += 1
self._worker.active += 1
try:
yield
finally:
self._sources[source].active -= 1
self._worker.active -= 1
def admission_source(request: Request) -> str:
from litellm.proxy.proxy_server import general_settings # noqa: PLC0415 # runtime proxy configuration
settings: Final = general_settings or {}
config: Final = TrustedProxyConfig.model_validate(
{
"use_forwarded_for": settings.get("use_x_forwarded_for", False),
"trusted_proxy_cidrs": settings.get("mcp_trusted_proxy_ranges") or (),
}
)
client_ip, _ = resolve_client_ip(request, config)
return source_group(client_ip or "unknown")

View file

@ -446,7 +446,7 @@ class MCPRequestHandler:
Raises:
HTTPException: If headers are invalid or missing required headers
"""
async with global_manager().catalog.operation(request=Request(scope)):
async with global_manager().catalog.operation():
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
# Check if there is an explicit LiteLLM API key (primary header)

View file

@ -10,19 +10,14 @@ from contextlib import asynccontextmanager
from contextvars import ContextVar
from dataclasses import dataclass, replace
from functools import wraps
from inspect import signature
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
from litellm._logging import verbose_logger
if TYPE_CHECKING:
from mcp.types import Tool as SDKTool
from pydantic import BaseModel
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.auth.admission import MCPAdmissionLimiter
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerOutcome
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp_server.tool_registry import MCPTool
@ -86,7 +81,6 @@ class TargetCatalog:
self._arrival_ticket = 0
self._completed_ticket = 0
self._shared_snapshot: CatalogSnapshot | None = None
self._admission: MCPAdmissionLimiter | None = None
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
self._operation: ContextVar[tuple[CatalogSnapshot, asyncio.Event, int] | None] = ContextVar(
@ -136,18 +130,8 @@ class TargetCatalog:
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed")
return self._shared_snapshot
async def _acquire_snapshot(self, request: Request | None) -> CatalogSnapshot:
if request is None:
return await self._fresh_snapshot()
from litellm.proxy._experimental.mcp_server.auth.admission import MCPAdmissionLimiter, admission_source
if self._admission is None:
self._admission = MCPAdmissionLimiter()
with self._admission.admit(admission_source(request)):
return await self._fresh_snapshot()
async def list(self, *, request: Request | None = None) -> Mapping[str, MCPServer]:
async with self.operation(request=request) as snapshot:
async def list(self) -> Mapping[str, MCPServer]:
async with self.operation() as snapshot:
return snapshot.servers
def assert_current(self, server: MCPServer) -> None:
@ -168,7 +152,7 @@ class TargetCatalog:
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
@asynccontextmanager
async def operation(self, *, request: Request | None = None) -> AsyncIterator[CatalogSnapshot]:
async def operation(self) -> AsyncIterator[CatalogSnapshot]:
current: Final = self.current()
scoped: Final = self._operation.get()
if current is not None and scoped is not None and scoped[2] == id(asyncio.current_task()):
@ -176,7 +160,7 @@ class TargetCatalog:
return
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
shared: Final = await self._acquire_snapshot(request)
shared: Final = await self._fresh_snapshot()
snapshot: Final = replace(
shared,
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
@ -242,27 +226,6 @@ class TargetCatalog:
_check_oauth_revision(selected, resolved)
return resolved
@staticmethod
async def aggregate_list(
servers: Sequence[MCPServer],
fetch: Callable[[MCPServer], Awaitable[tuple[list[SDKTool], ServerOutcome]]],
server_key: Callable[[MCPServer], str],
) -> AggregateToolListing:
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
tasks: Final = tuple(asyncio.ensure_future(fetch(server)) for server in servers)
try:
results: Final = await asyncio.gather(*tasks)
return AggregateToolListing(
tools=[tool for tools, _ in results for tool in tools],
outcomes={server_key(server): outcome for server, (_, outcome) in zip(servers, results)},
)
finally:
for task in tasks:
if not task.done():
task.cancel()
await asyncio.gather(*tasks, return_exceptions=True)
async def reload(self) -> None:
async with self._refresh_lock:
await self._publish_refresh()
@ -559,18 +522,6 @@ def global_manager() -> MCPServerManager:
return global_mcp_server_manager
def public_catalog_operation(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
parameters: Final = signature(function)
@wraps(function)
async def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: # kwargs-ok: preserves ParamSpec
from starlette.requests import Request
arguments: Final = parameters.bind(*args, **kwargs).arguments
request: Final = arguments.get("request")
if not isinstance(request, Request):
raise TypeError("Public MCP operations require a Request")
async with global_manager().catalog.operation(request=request):
return await function(*args, **kwargs)
return wrapped
public_catalog_operation: Final[Callable[[Callable[_P, Awaitable[_R]]], Callable[_P, Awaitable[_R]]]] = (
catalog_operation(global_manager)
)

View file

@ -2919,7 +2919,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
async with global_mcp_server_manager.catalog.operation(request=request):
async with global_mcp_server_manager.catalog.operation():
dummy_return: Final = {
"client_id": mcp_server_name or "dummy_client",
"client_secret": "dummy",

View file

@ -1,5 +1,6 @@
"""Shared MCP operation policy and dispatch."""
import asyncio
import traceback
import types
import uuid
@ -49,7 +50,7 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog, catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -1153,17 +1154,24 @@ async def _get_tools_from_mcp_servers(
verbose_logger.exception("Error getting tools from server %s: %s", server.name, e)
return [], classify_list_exception(e)
listing: Final = await TargetCatalog.aggregate_list(
allowed_mcp_servers, _fetch_and_filter_server_tools, _aggregate_server_key
)
all_tools: Final = listing.tools
server_outcomes: Final = listing.outcomes
# Fetch tools from all servers in parallel
tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers]
results: Final = await asyncio.gather(*tasks)
# Flatten results into single list
all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools]
server_outcomes: Final[dict[str, ServerOutcome]] = {
_aggregate_server_key(server): outcome
for server, (_, outcome) in zip(allowed_mcp_servers, results)
if server is not None
}
# If logging is enabled, enrich spend_logs_metadata with counts
if litellm_logging_obj:
per_server_tool_counts: Final[dict[str, int]] = {
key: outcome.tool_count if isinstance(outcome, ServerListOk) else 0
for key, outcome in server_outcomes.items()
_aggregate_server_key(server): len(server_tools)
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
if server is not None
}
metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata")

View file

@ -1,89 +0,0 @@
from unittest.mock import Mock
import pytest
from fastapi import HTTPException, Request
from litellm.proxy._experimental.mcp_server.auth.admission import MCPAdmissionLimiter, admission_source
def test_client_and_worker_rate_budgets_recover_after_window(monkeypatch):
monkeypatch.setenv("LITELLM_MCP_PUBLIC_RPM", "2")
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_RPM", "3")
clock = Mock(return_value=0.0)
limiter = MCPAdmissionLimiter(clock)
for source in ("a", "a"):
with limiter.admit(source):
pass
with pytest.raises(HTTPException) as client_error, limiter.admit("a"):
pytest.fail("client budget must reject")
assert client_error.value.status_code == 429
assert client_error.value.headers == {"Retry-After": "60"}
with limiter.admit("b"):
pass
clock.return_value = 59.1
with pytest.raises(HTTPException) as worker_error, limiter.admit("c"):
pytest.fail("worker budget must reject rotated sources")
assert worker_error.value.headers == {"Retry-After": "1"}
clock.return_value = 60.0
with limiter.admit("a"), limiter.admit("c"):
pass
def test_inflight_limits_release_on_failure_across_window_boundary(monkeypatch):
monkeypatch.setenv("LITELLM_MCP_PUBLIC_MAX_IN_FLIGHT", "1")
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", "2")
clock = Mock(return_value=0.0)
limiter = MCPAdmissionLimiter(clock)
with limiter.admit("a"):
with pytest.raises(HTTPException) as client_error, limiter.admit("a"):
pytest.fail("one client cannot occupy another permit")
assert client_error.value.headers == {"Retry-After": "1"}
with limiter.admit("b"):
clock.return_value = 60.0
with pytest.raises(HTTPException) as worker_error, limiter.admit("c"):
pytest.fail("rotating sources cannot exceed active work")
assert worker_error.value.status_code == 429
with pytest.raises(RuntimeError, match="upstream"), limiter.admit("c"):
raise RuntimeError("upstream")
with limiter.admit("a"), limiter.admit("c"):
pass
def test_source_capacity_never_evicts_live_budget(monkeypatch):
monkeypatch.setenv("LITELLM_MCP_PUBLIC_MAX_SOURCES", "2")
clock = Mock(return_value=0.0)
limiter = MCPAdmissionLimiter(clock)
with limiter.admit("a"):
with limiter.admit("b"):
pass
for index in range(100):
with pytest.raises(HTTPException) as error, limiter.admit(str(index)):
pytest.fail("source spray must be bounded")
assert error.value.status_code == 429
clock.return_value = 60.0
with limiter.admit("c"):
pass
with limiter.admit("a"):
pass
@pytest.mark.parametrize("value", ["0", "-1", "invalid"])
def test_invalid_limits_are_not_silently_disabled(monkeypatch, value):
monkeypatch.setenv("LITELLM_MCP_PUBLIC_RPM", value)
with pytest.raises(ValueError, match=r"must be positive|invalid literal"):
MCPAdmissionLimiter()
@pytest.mark.parametrize("peer,trusted,xff,expected", [
("198.51.100.1", [], "203.0.113.1", "198.51.100.1"),
("10.0.0.1", ["10.0.0.0/8"], "203.0.113.1, 198.51.100.1", "198.51.100.1"),
("2001:db8::1", [], "", "2001:db8::/64"),
(None, [], "203.0.113.1", "unknown"),
])
def test_source_identity_ignores_untrusted_forwarded_addresses(monkeypatch, peer, trusted, xff, expected):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "general_settings", {"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": trusted})
request = Request({"type": "http", "client": (peer, 1234) if peer else None,
"headers": [(b"x-forwarded-for", xff.encode())]})
assert admission_source(request) == expected

View file

@ -15489,55 +15489,6 @@ async def test_catalog_fresh_lookup_does_not_fall_back_to_stale_grants_when_data
assert manager.registry == {server.server_id: server}
@pytest.mark.asyncio
async def test_catalog_list_failure_cancels_and_joins_other_upstream_fetches():
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
started: Final = asyncio.Event()
stopped: Final = asyncio.Event()
pending: Final = MCPServer(server_id="pending", name="pending", transport=MCPTransport.http)
broken: Final = MCPServer(server_id="broken", name="broken", transport=MCPTransport.http)
async def fetch(server: MCPServer):
if server.server_id == broken.server_id:
await started.wait()
raise RuntimeError("unexpected fetch failure")
started.set()
try:
await asyncio.Event().wait()
finally:
stopped.set()
with pytest.raises(RuntimeError, match="unexpected fetch failure"):
await TargetCatalog.aggregate_list((pending, broken), fetch, lambda server: server.server_id)
assert stopped.is_set()
@pytest.mark.asyncio
@pytest.mark.parametrize("headers", [[], [(b"authorization", b"Bearer arbitrary")]])
async def test_public_catalog_budget_rejects_before_database_read(monkeypatch, headers):
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import oauth_authorization_server_mcp
monkeypatch.setenv("LITELLM_MCP_PUBLIC_RPM", "2")
read_rows = AsyncMock(return_value=[])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
request = Request({"type": "http", "method": "GET", "path": "/.well-known/oauth-authorization-server/missing",
"headers": headers, "client": ("198.51.100.1", 1234), "scheme": "https",
"server": ("gateway.example", 443), "query_string": b""})
for expected in (404, 404, 429):
with pytest.raises(HTTPException) as exc:
await oauth_authorization_server_mcp(request=request, mcp_server_name="missing")
assert exc.value.status_code == expected
assert int(exc.value.headers["Retry-After"]) > 0
assert read_rows.await_count == 2
await manager.catalog.list()
assert read_rows.await_count == 3
@pytest.mark.asyncio
async def test_catalog_cancelled_registration_does_not_publish_partial_handlers(monkeypatch):
from litellm.proxy._experimental.mcp_server import tool_registry
@ -15558,58 +15509,3 @@ async def test_catalog_cancelled_registration_does_not_publish_partial_handlers(
assert [tool.name for tool in registry.list_tools()] == ["existing"]
assert manager.registry == {}
@pytest.mark.asyncio
@pytest.mark.parametrize("cancel", [False, True])
async def test_public_catalog_inflight_rejection_precedes_read_and_releases_permit(monkeypatch, cancel):
from starlette.requests import Request
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", "1")
started = asyncio.Event()
release = asyncio.Event()
async def blocked_rows(**kwargs):
started.set()
await release.wait()
return [_catalog_row()]
read_rows = AsyncMock(side_effect=blocked_rows)
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
request = Request({"type": "http", "client": ("198.51.100.1", 1234), "headers": []})
first = asyncio.create_task(manager.catalog.list(request=request))
await asyncio.wait_for(started.wait(), 1)
with pytest.raises(HTTPException) as error:
await manager.catalog.list(request=request)
assert error.value.status_code == 429
assert error.value.headers == {"Retry-After": "1"}
assert read_rows.await_count == 1
authenticated = asyncio.create_task(manager.catalog.list())
if cancel:
first.cancel()
release.set()
if cancel:
with pytest.raises(asyncio.CancelledError):
await first
else:
assert "catalog-server" in await first
assert "catalog-server" in await asyncio.wait_for(authenticated, 1)
assert "catalog-server" in await manager.catalog.list(request=request)
assert read_rows.await_count == 3
@pytest.mark.asyncio
async def test_public_catalog_releases_permit_before_operation_and_nested_use(monkeypatch):
from starlette.requests import Request
monkeypatch.setenv("LITELLM_MCP_PUBLIC_WORKER_MAX_IN_FLIGHT", "1")
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
request = Request({"type": "http", "client": ("198.51.100.1", 1234), "headers": []})
async with manager.catalog.operation(request=request):
async with manager.catalog.operation(request=request):
assert read_rows.await_count == 1
other = await asyncio.create_task(manager.catalog.list(request=request))
assert "catalog-server" in other
assert read_rows.await_count == 2