From 51892184a36937a1e982f9d6ad973f0279756e90 Mon Sep 17 00:00:00 2001 From: joshua Date: Wed, 23 Sep 2026 05:18:27 +0000 Subject: [PATCH] refactor(mcp): narrow catalog consistency change to ticket scope Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/auth/admission.py | 97 ---------------- .../mcp_server/auth/user_api_key_auth_mcp.py | 2 +- .../proxy/_experimental/mcp_server/catalog.py | 63 ++--------- .../mcp_server/discoverable_endpoints.py | 2 +- .../_experimental/mcp_server/operations.py | 24 ++-- .../mcp_server/auth/test_admission.py | 89 --------------- .../mcp_server/test_mcp_server_manager.py | 104 ------------------ 7 files changed, 25 insertions(+), 356 deletions(-) delete mode 100644 litellm/proxy/_experimental/mcp_server/auth/admission.py delete mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/auth/test_admission.py diff --git a/litellm/proxy/_experimental/mcp_server/auth/admission.py b/litellm/proxy/_experimental/mcp_server/auth/admission.py deleted file mode 100644 index e5d1167d4d3..00000000000 --- a/litellm/proxy/_experimental/mcp_server/auth/admission.py +++ /dev/null @@ -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") diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 0313bbc6fce..d2284cc901c 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index bbbecacccdb..ecf71089731 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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) +) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 6299068ae6c..f1e19b17d1e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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", diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 4ca704d9afc..9bb28793e1e 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_admission.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_admission.py deleted file mode 100644 index cc689b4db87..00000000000 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_admission.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 270a26cb5ff..ba26d4024a4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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