fix(mcp): bound public catalog work and consolidate worker snapshots

This commit is contained in:
Joshua Valluru 2026-09-22 17:08:58 -07:00
commit 6bca95bdb0
31 changed files with 2183 additions and 806 deletions

View file

@ -0,0 +1,97 @@
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

@ -13,7 +13,7 @@ from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import global_manager
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
@ -399,7 +399,6 @@ class MCPRequestHandler:
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
@staticmethod
@with_mcp_catalog
async def process_mcp_request(
scope: Scope,
) -> tuple[
@ -432,163 +431,166 @@ class MCPRequestHandler:
Raises:
HTTPException: If headers are invalid or missing required headers
"""
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
async with global_manager().catalog.operation(request=Request(scope)):
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
# Check if there is an explicit LiteLLM API key (primary header)
has_explicit_litellm_key: Final = headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) is not None
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
# Get the old mcp_auth_header for backward compatibility
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
# Get the new server-specific auth headers
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
# Get the oauth2 headers
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
# Parse MCP servers from header
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
verbose_logger.debug("Raw MCP servers header: %s", mcp_servers_header)
mcp_servers = None
if mcp_servers_header is not None:
try:
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
verbose_logger.debug("Parsed MCP servers: %s", mcp_servers)
except Exception as e:
verbose_logger.debug("Error parsing mcp_servers header: %s", e)
mcp_servers = None
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
mcp_servers = []
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
request: Final = Request(scope=scope)
async def mock_body():
return b"{}"
request.body = mock_body
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
request_route: Final = get_request_route(request)
# Only OAuth metadata routes registered under /.well-known/ are public.
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif has_explicit_litellm_key:
# An explicit x-litellm-api-key is always a LiteLLM credential, even
# for a delegated server, so validate it: identity / spend / rate
# limits resolve and any stored upstream token can be forwarded.
validated_user_api_key_auth = await user_api_key_auth(
api_key=f"Bearer {_get_bearer_token_or_received_api_key(litellm_api_key)}",
request=request,
# Check if there is an explicit LiteLLM API key (primary header)
has_explicit_litellm_key: Final = (
headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) is not None
)
elif MCPRequestHandler._target_servers_are_true_passthrough(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
) or (
MCPRequestHandler._single_dcr_bridge_delegate_target(
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
# Get the old mcp_auth_header for backward compatibility
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
# Get the new server-specific auth headers
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
# Get the oauth2 headers
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
# Parse MCP servers from header
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
verbose_logger.debug("Raw MCP servers header: %s", mcp_servers_header)
mcp_servers = None
if mcp_servers_header is not None:
try:
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
verbose_logger.debug("Parsed MCP servers: %s", mcp_servers)
except Exception as e:
verbose_logger.debug("Error parsing mcp_servers header: %s", e)
mcp_servers = None
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
mcp_servers = []
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
request: Final = Request(scope=scope)
async def mock_body():
return b"{}"
request.body = mock_body
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
request_route: Final = get_request_route(request)
# Only OAuth metadata routes registered under /.well-known/ are public.
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif has_explicit_litellm_key:
# An explicit x-litellm-api-key is always a LiteLLM credential, even
# for a delegated server, so validate it: identity / spend / rate
# limits resolve and any stored upstream token can be forwarded.
validated_user_api_key_auth = await user_api_key_auth(
api_key=f"Bearer {_get_bearer_token_or_received_api_key(litellm_api_key)}",
request=request,
)
elif MCPRequestHandler._target_servers_are_true_passthrough(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
is not None
and not oauth2_headers
and not mcp_server_auth_headers
and not mcp_auth_header
):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
) is not None and oauth2_headers:
) or (
MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
is not None
and not oauth2_headers
and not mcp_server_auth_headers
and not mcp_auth_header
):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
) is not None and oauth2_headers:
(
validated_user_api_key_auth,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_authorization(
server=bridge_delegate_target.server,
requested_name=bridge_delegate_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
route=request_route,
)
elif oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]):
# A gateway DCR session bearer at any MCP scope: open the identity-only session
# token and admit under the live litellm user; downstream grant resolution
# intersects the admitted subject's servers with any path or header target, so a
# per-server scope narrows and never broadens. One that does not open fails
# closed with the scope's invalid_token challenge; a non-session bearer falls
# through to the oauth2 arm.
validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session(
authorization_value=oauth2_headers["Authorization"],
request=request,
route=request_route,
mcp_servers=mcp_servers,
)
elif oauth2_headers:
# Authorization on a non-delegated server: the bearer must be a real
# LiteLLM credential, so a failed validation is a genuine 401/403 and
# propagates unless a fallback in _admission_failure_fallback applies.
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as e:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=e,
bearer_presented=True,
)
else:
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as exc:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=exc,
bearer_presented=False,
)
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
# credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged.
raw_headers = dict(headers)
(
validated_user_api_key_auth,
oauth2_headers,
raw_headers,
mcp_auth_header,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_authorization(
server=bridge_delegate_target.server,
requested_name=bridge_delegate_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
) = MCPRequestHandler._scrub_gateway_admission_credentials(
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
route=request_route,
)
elif oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]):
# A gateway DCR session bearer at any MCP scope: open the identity-only session
# token and admit under the live litellm user; downstream grant resolution
# intersects the admitted subject's servers with any path or header target, so a
# per-server scope narrows and never broadens. One that does not open fails
# closed with the scope's invalid_token challenge; a non-session bearer falls
# through to the oauth2 arm.
validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session(
authorization_value=oauth2_headers["Authorization"],
request=request,
route=request_route,
mcp_servers=mcp_servers,
return (
validated_user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
)
elif oauth2_headers:
# Authorization on a non-delegated server: the bearer must be a real
# LiteLLM credential, so a failed validation is a genuine 401/403 and
# propagates unless a fallback in _admission_failure_fallback applies.
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as e:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=e,
bearer_presented=True,
)
else:
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as exc:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=exc,
bearer_presented=False,
)
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
# credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged.
raw_headers = dict(headers)
(
oauth2_headers,
raw_headers,
mcp_auth_header,
mcp_server_auth_headers,
) = MCPRequestHandler._scrub_gateway_admission_credentials(
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
)
return (
validated_user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
)
@staticmethod
def _is_gateway_admission_credential(value: str | None) -> bool:

View file

@ -25,6 +25,7 @@ from fastapi import APIRouter, Depends, Form, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._experimental.mcp_server.db import store_user_credential
from litellm.proxy._experimental.mcp_server.oauth_utils import (
BYOK_RESOURCE_METADATA_PATH,
@ -649,6 +650,7 @@ async def byok_protected_resource_metadata(request: Request) -> JSONResponse:
@router.get("/v1/mcp/oauth/authorize", include_in_schema=False)
@public_catalog_operation
async def byok_authorize_get(
request: Request,
client_id: str | None = None,

View file

@ -1,70 +1,163 @@
"""Authoritative MCP catalog snapshots shared by legacy lookup adapters."""
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
import hashlib
import json
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence
from contextlib import asynccontextmanager
from contextvars import ContextVar
from dataclasses import dataclass
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
_P: Final = ParamSpec("_P")
_R: Final = TypeVar("_R")
_REFRESH_FAILURE: Final = "MCP server configuration could not be refreshed"
_P = ParamSpec("_P")
_R = TypeVar("_R")
@dataclass(slots=True)
class _CatalogScope:
@dataclass(frozen=True, slots=True)
class CatalogSnapshot:
servers: Mapping[str, MCPServer]
owner_task_id: int
active: bool = True
identity: str
tools: Mapping[str, MCPTool]
routing: dict[str, str]
def _configuration_identity(server: MCPServer) -> str:
return json.dumps(
server.model_dump(
mode="json",
exclude=frozenset(("short_prefix", "scopes", "authorization_url", "token_url", "registration_url"))
| (frozenset() if server.issuer_is_anchored else frozenset(("issuer",))),
),
sort_keys=True,
)
def _check_oauth_revision(selected: MCPServer, candidate: MCPServer | None) -> None:
from fastapi import HTTPException
if candidate is None or _configuration_identity(selected) != _configuration_identity(candidate):
raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly")
def _snapshot(manager: MCPServerManager, database_identity: str) -> CatalogSnapshot:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
servers: Final = manager.config_mcp_servers | manager.registry
detached: Final = MappingProxyType({key: value.model_copy(deep=True) for key, value in servers.items()})
serialized: Final = json.dumps(
(
database_identity,
tuple(sorted(manager.registry)),
tuple((key, _configuration_identity(value)) for key, value in sorted(manager.config_mcp_servers.items())),
),
sort_keys=True,
)
return CatalogSnapshot(
detached,
hashlib.sha256(serialized.encode()).hexdigest(),
MappingProxyType(dict(global_mcp_tool_registry.published_tools)),
dict(manager.published_tool_routes),
)
class TargetCatalog:
def __init__(self, manager: MCPServerManager) -> None:
self._manager = manager
self._reload_lock = asyncio.Lock()
self._scope: ContextVar[_CatalogScope | None] = ContextVar("mcp_catalog_scope", default=None)
self.manager = manager
self._refresh_lock = asyncio.Lock()
self._database_identity = ""
self._arrival_ticket = 0
self._completed_ticket = 0
self._snapshot: Mapping[str, MCPServer] | None = None
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(
"mcp_catalog_snapshot", default=None
)
self._staged_routing: ContextVar[tuple[dict[str, str], asyncio.Event] | None] = ContextVar(
"mcp_catalog_routing", default=None
)
@property
def current(self) -> Mapping[str, MCPServer] | None:
scope: Final = self._scope.get()
return scope.servers if scope is not None and scope.active else None
def current(self) -> CatalogSnapshot | None:
scoped: Final = self._operation.get()
return scoped[0] if scoped is not None and not scoped[1].is_set() else None
async def refresh(self) -> None:
async with self._reload_lock:
await self._reload()
def routing(self) -> dict[str, str]:
staged: Final = self._staged_routing.get()
if staged is not None and not staged[1].is_set():
return staged[0]
snapshot: Final = self.current()
return snapshot.routing if snapshot is not None else self.manager.published_tool_routes
async def _reload(self, *, reuse_unchanged: bool = False) -> None:
token: Final = self._scope.set(None)
# A shared read must start after each covered operation arrives.
covered_ticket: Final = self._arrival_ticket
try:
await self._manager._reload_servers_from_database(reuse_unchanged=reuse_unchanged) # pyright: ignore[reportPrivateUsage] # existing staged loader
except Exception:
self._snapshot = None
self._completed_ticket = covered_ticket
raise
else:
self._snapshot = MappingProxyType(self._manager.config_mcp_servers | self._manager.registry)
self._completed_ticket = covered_ticket
finally:
self._scope.reset(token)
def registry(self) -> Mapping[str, MCPServer]:
snapshot: Final = self.current()
return snapshot.servers if snapshot is not None else self.manager.config_mcp_servers | self.manager.registry
async def _fresh_snapshot(self) -> CatalogSnapshot:
from fastapi import HTTPException
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return _snapshot(self.manager, self._database_identity)
from litellm.proxy.proxy_server import should_load_db_object
if not should_load_db_object("mcp"):
return _snapshot(self.manager, self._database_identity)
self._arrival_ticket += 1
arrival: Final = self._arrival_ticket
async with self._refresh_lock:
if arrival > self._completed_ticket:
try:
await self._publish_refresh(reuse_unchanged=True)
except Exception as exc:
raise HTTPException(
status_code=503, detail="MCP server configuration could not be refreshed"
) from exc
if self._shared_snapshot is None:
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:
return snapshot.servers
def assert_current(self, server: MCPServer) -> None:
snapshot: Final = self.current
if snapshot is None or server.server_id not in snapshot:
from fastapi import HTTPException
snapshot: Final = self.current()
if snapshot is None or server.server_id not in snapshot.servers:
return
expected: Final = snapshot[server.server_id]
registered: Final = self._manager.registry.get(server.server_id) or self._manager.config_mcp_servers.get(
expected: Final = snapshot.servers[server.server_id]
registered: Final = self.manager.registry.get(server.server_id) or self.manager.config_mcp_servers.get(
server.server_id
)
if (
@ -72,64 +165,412 @@ class TargetCatalog:
or registered.updated_at != expected.updated_at
or server.updated_at != expected.updated_at
):
from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
async def list(self) -> Mapping[str, MCPServer]:
from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime proxy dependency
scope: Final = self._scope.get()
if scope is not None and scope.active and scope.owner_task_id == id(asyncio.current_task()):
return scope.servers
self._arrival_ticket += 1
arrival_ticket: Final = self._arrival_ticket
async with self._reload_lock:
if prisma_client is None:
return MappingProxyType(self._manager.config_mcp_servers | self._manager.registry)
if arrival_ticket > self._completed_ticket:
try:
await self._reload(reuse_unchanged=True)
except Exception as exc: # noqa: BLE001 # never serve an unverified database snapshot
raise HTTPException(status_code=503, detail=_REFRESH_FAILURE) from exc
if self._snapshot is None:
raise HTTPException(status_code=503, detail=_REFRESH_FAILURE)
return self._snapshot
@asynccontextmanager
async def operation(self) -> AsyncIterator[Mapping[str, MCPServer]]:
existing: Final = self._scope.get()
if existing is not None and existing.active and existing.owner_task_id == id(asyncio.current_task()):
yield existing.servers
async def operation(self, *, request: Request | None = None) -> 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()):
yield current
return
scope: Final = _CatalogScope(await self.list(), id(asyncio.current_task()))
token: Final = self._scope.set(scope)
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
shared: Final = await self._acquire_snapshot(request)
snapshot: Final = replace(
shared,
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
routing=dict(shared.routing),
)
closed: Final = asyncio.Event()
token: Final = self._operation.set((snapshot, closed, id(asyncio.current_task())))
try:
yield scope.servers
with global_mcp_tool_registry.catalog_scope(snapshot.tools):
yield snapshot
finally:
scope.active = False
self._scope.reset(token)
closed.set()
self._operation.reset(token)
self._retain_discovered_routing(snapshot)
async def resolve(self, lookup: str, client_ip: str | None = None) -> MCPServer | None:
async with self.operation():
return self._manager.get_mcp_server_by_name(
lookup, client_ip=client_ip
) or self._manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
def with_mcp_catalog(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
@wraps(function)
async def wrapped(
*args: _P.args,
**kwargs: _P.kwargs, # kwargs-ok: ParamSpec preserves each wrapped endpoint keyword contract
) -> _R:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # manager imports catalog
global_mcp_server_manager,
def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None:
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
snapshot.servers, snapshot.routing
)
async with global_mcp_server_manager.catalog.operation():
def _unchanged_routing(
self, servers: Mapping[str, MCPServer], routing: Mapping[str, str]
) -> MappingProxyType[str, str]:
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
current: Final = self.manager.config_mcp_servers | self.manager.registry
unchanged_owners: Final = frozenset(
owner
for key, server in servers.items()
if (candidate := current.get(key)) is not None
and _configuration_identity(candidate) == _configuration_identity(server)
for owner in self.manager.owned_mapping_values(server)
)
return MappingProxyType(
{name: owner for name, owner in routing.items() if normalize_server_name(owner) in unchanged_owners}
)
async def resolve(self, identifier: str, client_ip: str | None = None) -> MCPServer | None:
async with self.operation():
return self.manager.get_mcp_server_by_name(
identifier, client_ip=client_ip
) or self.manager.get_mcp_server_by_id(identifier, client_ip=client_ip)
async def resolve_oauth_metadata(
self,
server: MCPServer,
resolve: Callable[[MCPServer], Awaitable[MCPServer]],
) -> MCPServer:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import oauth_endpoints_unresolved
snapshot: Final = self.current()
selected: Final = snapshot.servers.get(server.server_id) if snapshot is not None else None
if selected is None:
return await resolve(server)
if not oauth_endpoints_unresolved(selected):
return selected
registered: Final = self.manager.registry.get(server.server_id) or self.manager.config_mcp_servers.get(
server.server_id
)
self.assert_current(server)
_check_oauth_revision(selected, registered)
resolved: Final = await resolve(selected)
_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()
async def _publish_refresh(self, *, reuse_unchanged: bool = False) -> None:
covered: Final = self._arrival_ticket
token: Final = self._operation.set(None)
try:
await self._reload_and_publish(reuse_unchanged=reuse_unchanged)
except Exception:
self._shared_snapshot = None
self._completed_ticket = covered
raise
else:
self._shared_snapshot = _snapshot(self.manager, self._database_identity)
self._completed_ticket = covered
finally:
self._operation.reset(token)
async def _reload_and_publish(self, *, reuse_unchanged: bool) -> None:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
previous_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()}
)
live_registry: Final = self.manager.registry
previous_servers: Final = previous_config | MappingProxyType(
{key: value.model_copy(deep=True) for key, value in live_registry.items()}
)
config_identities: Final = MappingProxyType(
{key: _configuration_identity(value) for key, value in previous_config.items()}
)
staged_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in previous_config.items()}
)
await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values()))
initial_routing: Final = MappingProxyType(dict(self.manager.published_tool_routes))
staged_routing: Final = dict(initial_routing)
closed: Final = asyncio.Event()
routing_token: Final = self._staged_routing.set((staged_routing, closed))
initial_tools: Final = MappingProxyType(dict(global_mcp_tool_registry.published_tools))
try:
with global_mcp_tool_registry.catalog_scope(initial_tools) as staged_tools:
await self._reload(reuse_unchanged=reuse_unchanged)
refreshed_openapi_owners: Final = frozenset(
owner
for server in self.manager.registry.values()
if server.spec_path and server is not live_registry.get(server.server_id)
for owner in self.manager.owned_mapping_values(server)
)
live_routes: Final = self._unchanged_routing(
previous_servers,
MappingProxyType(
{
name: owner
for name, owner in self.manager.published_tool_routes.items()
if normalize_server_name(owner) not in refreshed_openapi_owners
}
),
)
concurrent_routes: Final = self._unchanged_routing(
self.manager.config_mcp_servers | live_registry,
MappingProxyType(
{
name: owner
for name, owner in self.manager.published_tool_routes.items()
if initial_routing.get(name) != owner
and (name not in staged_routing or staged_routing.get(name) == initial_routing.get(name))
}
),
)
removed_routes: Final = initial_routing.keys() - self.manager.published_tool_routes.keys()
retained_staged_routes: Final = self._unchanged_routing(
previous_config | self.manager.registry,
MappingProxyType(
{
name: owner
for name, owner in staged_routing.items()
if name not in removed_routes or owner != initial_routing[name]
}
),
)
concurrent_tools: Final = MappingProxyType(
{
name: tool
for name, tool in global_mcp_tool_registry.published_tools.items()
if (
name in live_routes
or name in concurrent_routes
or (name not in initial_routing and name not in self.manager.published_tool_routes)
)
and initial_tools.get(name) is staged_tools.get(name)
}
)
self.manager.config_mcp_servers = {
key: value.model_copy(
update=staged_config[key].model_dump(
include=frozenset(("client_id", "client_secret", "token_endpoint_auth_method"))
)
)
if config_identities.get(key) == _configuration_identity(value)
else value
for key, value in self.manager.config_mcp_servers.items()
}
global_mcp_tool_registry.tools = (
MappingProxyType(
{
name: tool
for name, tool in staged_tools.items()
if name in global_mcp_tool_registry.published_tools or tool is not initial_tools.get(name)
}
)
| concurrent_tools
)
self.manager.published_tool_routes = live_routes | retained_staged_routes | concurrent_routes
finally:
closed.set()
self._staged_routing.reset(routing_token)
async def _reload(self, *, reuse_unchanged: bool) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
carry_forward_resolved_oauth_endpoints,
config_ids_capturing_db_identifiers,
oauth_endpoints_unresolved,
warn_on_server_name_fields,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_prisma_client_or_throw,
)
verbose_logger.debug("Loading MCP servers from database into registry...")
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
# Load only "active", legacy "approved", and NULL (no approval workflow) rows.
# Pending/rejected servers are excluded at the DB level so we never load them.
from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable, get_runtime_mcp_server_rows
raw_rows: Final[Sequence[BaseModel]] = await get_runtime_mcp_server_rows(prisma_client)
database_identity: Final = hashlib.sha256(
json.dumps(
tuple(sorted(json.dumps(row.model_dump(mode="json"), sort_keys=True, default=str) for row in raw_rows))
).encode()
).hexdigest()
verbose_logger.info("Found %s MCP servers in database", len(raw_rows))
previous_registry: Final = self.manager.registry
new_registry: Final[dict[str, MCPServer]] = {}
# Stage one: build every server. Stage two assigns short prefixes
# against the *full* set so dedup is deterministic regardless of
# iteration order.
for row in raw_rows:
try:
server = LiteLLM_MCPServerTable.model_validate(row.model_dump())
existing_server = previous_registry.get(server.server_id)
if (
existing_server is not None
and (reuse_unchanged or not existing_server.spec_path)
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
and (
self.manager.oauth_discovery_slot(server.server_id) is not None
or not oauth_endpoints_unresolved(existing_server)
)
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
warn_on_server_name_fields(
server_id=server.server_id,
alias=getattr(server, "alias", None),
server_name=getattr(server, "server_name", None),
)
verbose_logger.debug("Building server from DB: %s (%s)", server.server_id, server.server_name)
# raw_rows come straight from the DB, so their global env var
# values (like credentials) are still encrypted here, unlike the
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.manager.build_mcp_server_from_table(
server, env_vars_are_encrypted=True, register_oauth_discovery=False
)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
new_server.short_prefix = existing_server.short_prefix
carry_forward_resolved_oauth_endpoints(new_server=new_server, previous_server=existing_server)
new_registry[server.server_id] = new_server
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
getattr(row, "server_id", None),
getattr(row, "alias", None),
e,
)
# Assign short prefixes against the full candidate set without
# publishing the staged registry to concurrent callers.
registered_registry: Final[dict[str, MCPServer]] = {}
for server_id, new_server in new_registry.items():
try:
if new_server is not previous_registry.get(server_id):
self.manager.assign_unique_short_prefix(new_server, registry=new_registry)
# Register OpenAPI tools *after* the final short prefix is assigned
# so the tools are stored in the global registry under the same
# prefix that lookups will use.
if new_server is not previous_registry.get(server_id):
if previous_server := previous_registry.get(server_id):
self.manager.remove_server_tool_routing(previous_server)
await self.manager.maybe_register_openapi_tools(new_server, initialize_mapping=False)
registered_registry[server_id] = new_server
except Exception as e:
self.manager.remove_server_tool_routing(new_server)
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
new_server.server_id,
getattr(new_server, "alias", None),
e,
)
dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys()
for registry_key in dropped_registry_keys:
self.manager.remove_server_tool_routing(previous_registry[registry_key])
self.manager.invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
for server_id in previous_registry.keys() | registered_registry.keys():
if previous_registry.get(server_id) != registered_registry.get(server_id):
self.manager.invalidate_discovery_lists(server_id)
self.manager.invalidate_oauth_discovery_state(server_id)
self._database_identity = database_identity
self.manager.registry = registered_registry
# A discovery task may have published into ``previous_registry`` while
# this replacement was being staged. Reconcile every published entry
# synchronously after the swap so a lost publication cannot also leave
# the replacement unresolved with no retry slot.
registered_servers: Final = tuple(registered_registry.values())
self.manager.reconcile_oauth_discovery_slots_for_servers(registered_servers)
self.manager.prime_oauth_metadata_discovery_for_servers(registered_servers)
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
# get_registry() is ``config_mcp_servers | registry``, so a database row sharing an id with a
# config.yaml server hides that server everywhere. Only reachable once an operator pins
# ``server_id`` in config.yaml; say so rather than letting the server disappear silently.
shadowed_config_server_ids: Final = frozenset(
self.manager.config_mcp_servers.keys() & registered_registry.keys()
)
if shadowed_config_server_ids and shadowed_config_server_ids != self._warned_shadowed_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are also database-backed MCP servers. The database "
"entry takes precedence, so the config.yaml server is unreachable. Give the config "
"entry a different server_id.",
", ".join(sorted(shadowed_config_server_ids)),
)
self._warned_shadowed_config_server_ids = shadowed_config_server_ids
# The mirror image of the block above: a config server_id that is a database server's name
# answers that server's grants instead, because ids are matched before names.
capturing_config_server_ids: Final = config_ids_capturing_db_identifiers(
self.manager.config_mcp_servers.keys(), registered_registry.values()
)
if capturing_config_server_ids and capturing_config_server_ids != self._warned_capturing_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are the name or alias of a database-backed MCP "
"server. Permission entries naming them resolve to the config.yaml server, not the "
"database one. Give the config entry a different server_id.",
", ".join(sorted(capturing_config_server_ids)),
)
self._warned_capturing_config_server_ids = capturing_config_server_ids
def catalog_operation(
manager: Callable[[], MCPServerManager],
) -> Callable[[Callable[_P, Awaitable[_R]]], Callable[_P, Awaitable[_R]]]:
def decorate(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
@wraps(function)
async def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: # kwargs-ok: preserves ParamSpec
async with manager().catalog.operation():
return await function(*args, **kwargs)
return wrapped
return decorate
def global_manager() -> MCPServerManager:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
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

View file

@ -645,6 +645,15 @@ async def get_all_mcp_servers(
return list(_readable_mcp_servers(mcp_servers))
async def get_runtime_mcp_server_rows(
prisma_client: PrismaClient,
) -> Sequence["prisma_db_models.LiteLLM_MCPServerTable"]:
where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = {
"OR": [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]
}
return await _db_find_mcp_server_rows(prisma_client, where)
async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> LiteLLM_MCPServerTable | None:
"""
Returns the matching mcp server from the db iff exists

View file

@ -36,7 +36,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
can_store_oauth_credential,
oauth_authorization_uses_gateway_credential,
)
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
CredentialSource,
@ -1770,7 +1770,7 @@ async def resolve_ephemeral_dcr_client(
def _register_flow_needed_endpoint(mcp_server: MCPServer) -> str | None:
"""The register flow's deferred-discovery join gate. A DCR bridge with no admin-configured
client can only register callers through the upstream's registration endpoint
(``_oauth_endpoints_unresolved`` keeps its discovery slot armed for exactly this shape), so
(``oauth_endpoints_unresolved`` keeps its discovery slot armed for exactly this shape), so
the flow must keep joining discovery while registration is still missing instead of silently
degrading to the dummy short-circuit. Every other shape only needs the authorization url."""
if mcp_server.is_dcr_bridge and not mcp_server.client_id and mcp_server.effective_registration_url is None:
@ -1876,7 +1876,7 @@ async def register_client_with_server(
@router.get("/authorize/mcp-session")
@with_mcp_catalog
@public_catalog_operation
async def authorize_mcp_session(
request: Request,
redirect_uri: str,
@ -1902,7 +1902,7 @@ async def authorize_mcp_session(
@router.get("/{mcp_server_name}/authorize")
@router.get("/authorize")
@with_mcp_catalog
@public_catalog_operation
async def authorize(
request: Request,
redirect_uri: str,
@ -1979,7 +1979,7 @@ async def authorize(
@router.post("/{mcp_server_name}/token")
@router.post("/token")
@with_mcp_catalog
@public_catalog_operation
async def token_endpoint(
request: Request,
grant_type: str = Form(...),
@ -2073,7 +2073,7 @@ async def _vendor_credential_state(user_id: str, server_id: str) -> VendorCreden
@router.get("/authorize/flow")
@with_mcp_catalog
@public_catalog_operation
async def authorize_flow(request: Request, flow: str) -> Response:
return await describe_connect_flow(
request=request,
@ -2085,7 +2085,7 @@ async def authorize_flow(request: Request, flow: str) -> Response:
@router.post("/authorize/complete")
@with_mcp_catalog
@public_catalog_operation
async def authorize_complete(
request: Request,
flow: str = Form(...),
@ -2443,7 +2443,7 @@ def is_network_error(exc: Exception) -> bool:
return isinstance(exc, httpx.TransportError)
@with_mcp_catalog
@public_catalog_operation
async def _build_oauth_protected_resource_response(
request: Request,
mcp_server_name: str | None,
@ -2705,6 +2705,7 @@ async def oauth_authorization_server_aggregate(request: Request):
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot)
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
@public_catalog_operation
async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth protected resource discovery endpoint using standard MCP URL pattern.
@ -2725,6 +2726,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam
# LiteLLM legacy pattern: /.well-known/oauth-protected-resource/{server_name}/mcp
# Kept for backward compatibility with existing deployments
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
@public_catalog_operation
async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern.
@ -2796,7 +2798,7 @@ def _build_oauth_authorization_server_response(
# Standard MCP pattern: /.well-known/oauth-authorization-server/mcp/{server_name}
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
@with_mcp_catalog
@public_catalog_operation
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery endpoint using standard MCP URL pattern.
@ -2814,7 +2816,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n
# LiteLLM legacy pattern and root endpoint
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}")
@router.get("/.well-known/oauth-authorization-server")
@with_mcp_catalog
@public_catalog_operation
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth authorization server discovery endpoint.
@ -2888,7 +2890,7 @@ async def jwks_json(request: Request):
# Additional legacy pattern support
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
@with_mcp_catalog
@public_catalog_operation
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
@ -2902,6 +2904,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
@router.post("/{mcp_server_name}/register")
@router.post("/register")
@public_catalog_operation
async def register_client(request: Request, mcp_server_name: str | None = None):
# Get the correct base URL considering X-Forwarded-* headers
request_base_url: Final = get_request_base_url(request)
@ -2916,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():
async with global_mcp_server_manager.catalog.operation(request=request):
dummy_return: Final = {
"client_id": mcp_server_name or "dummy_client",
"client_secret": "dummy",

View file

@ -55,6 +55,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
canonical_resource_uri,
@ -770,6 +771,7 @@ def _open_flow_for(
return flow
@catalog_operation(global_manager)
async def _flow_target(
flow: _ConnectFlow, lookup_server_reachability: LookupServerReachability
) -> tuple[Literal["unscoped", "interactive", "m2m", "stale"], MCPServer | None]:

View file

@ -73,7 +73,6 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPServerAccess,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
@ -179,7 +178,6 @@ from litellm.proxy.middleware.per_request_root_path_middleware import (
get_request_root_path,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.table_repositories import MCPServerRepository
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import (
DEFAULT_SUBJECT_TOKEN_TYPE,
@ -288,7 +286,7 @@ def _requires_oauth_discovery(
use_issuer_anchor: bool,
server: MCPServer,
) -> bool:
return _has_oauth_discovery_source(server_url, use_issuer_anchor) and _oauth_endpoints_unresolved(server)
return _has_oauth_discovery_source(server_url, use_issuer_anchor) and oauth_endpoints_unresolved(server)
_StringList: TypeAlias = list[str]
@ -542,7 +540,7 @@ def _config_identifier_owners(
)
def _config_ids_capturing_db_identifiers(
def config_ids_capturing_db_identifiers(
config_server_ids: Container[str],
db_servers: Iterable[MCPServer],
) -> frozenset[str]:
@ -723,7 +721,7 @@ def _flow_endpoints_missing(
return authorization_url is None or token_url is None
def _oauth_endpoints_unresolved(server: MCPServer) -> bool:
def oauth_endpoints_unresolved(server: MCPServer) -> bool:
"""``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
@ -782,7 +780,7 @@ def _endpoints_corroborate_authorization_url(
) == _normalized_authorize_endpoint(trusted_authorization_url)
def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_server: MCPServer | None) -> None:
def carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_server: MCPServer | None) -> None:
"""Keep the last known good OAuth endpoints when a rebuild's re-discovery comes back empty.
A rebuild wholesale-replaces the registry entry, so without this a transient upstream outage
@ -1428,7 +1426,7 @@ def _obo_retry_applies(server: MCPServer, subject_token: str | None) -> bool:
return server.auth_type == MCPAuth.oauth2_token_exchange and bool(subject_token)
def _warn_on_server_name_fields(
def warn_on_server_name_fields(
*,
server_id: str,
alias: str | None,
@ -1861,6 +1859,8 @@ class MCPServerManager:
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
)
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
self.catalog = TargetCatalog(self)
self.registry: dict[str, MCPServer] = {}
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
@ -1888,7 +1888,7 @@ class MCPServerManager:
# semaphore so an edited limit rebuilds it instead of keeping the old cap
# until restart.
self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {}
self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {}
self.published_tool_routes: dict[str, str] = {}
"""
{
"gmail_send_email": "zapier_mcp_server",
@ -1902,13 +1902,11 @@ class MCPServerManager:
# Last set of config server ids found shadowed by database rows. reload_servers_from_database
# runs on the config-reload timer, so this keeps a standing misconfiguration from re-logging
# the same warning every interval; a change in the set logs again.
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
self._oauth_discovery_on_startup = _mcp_oauth_discovery_on_startup_enabled()
self._oauth_discovery_generation_counter = 0
self._oauth_discovery_slots: tuple[_OAuthDiscoverySlot, ...] = ()
def _oauth_discovery_slot(self, server_id: str) -> _OAuthDiscoverySlot | None:
def oauth_discovery_slot(self, server_id: str) -> _OAuthDiscoverySlot | None:
return next((slot for slot in self._oauth_discovery_slots if slot.server_id == server_id), None)
def _remove_oauth_discovery_slot(self, server_id: str) -> None:
@ -1921,7 +1919,7 @@ class MCPServerManager:
)
def _set_oauth_discovery_deferred(self, server_id: str, discovery_deferred: bool) -> None:
previous: Final = self._oauth_discovery_slot(server_id)
previous: Final = self.oauth_discovery_slot(server_id)
self._remove_oauth_discovery_slot(server_id)
if previous is not None and previous.task is not None and not previous.task.done():
previous.task.cancel()
@ -1934,8 +1932,8 @@ class MCPServerManager:
)
)
def _invalidate_oauth_discovery_state(self, server_id: str) -> None:
previous: Final = self._oauth_discovery_slot(server_id)
def invalidate_oauth_discovery_state(self, server_id: str) -> None:
previous: Final = self.oauth_discovery_slot(server_id)
self._remove_oauth_discovery_slot(server_id)
if previous is not None and previous.task is not None and not previous.task.done():
previous.task.cancel()
@ -2010,7 +2008,7 @@ class MCPServerManager:
return resolved
def _oauth_discovery_slot_is_current(self, server_id: str, generation: int) -> bool:
slot: Final = self._oauth_discovery_slot(server_id)
slot: Final = self.oauth_discovery_slot(server_id)
return slot is not None and slot.generation == generation
def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None:
@ -2047,7 +2045,7 @@ class MCPServerManager:
if not self._oauth_discovery_slot_is_current(server.server_id, generation):
return _OAuthDiscoveryStale(server_id=server.server_id)
current: Final = self._registered_server(server)
if not _oauth_endpoints_unresolved(current):
if not oauth_endpoints_unresolved(current):
published: Final = self._publish_resolved_oauth_server(current, generation)
return (
_OAuthDiscoveryResolved(server=published)
@ -2058,7 +2056,7 @@ class MCPServerManager:
if not self._oauth_discovery_slot_is_current(server.server_id, generation):
return _OAuthDiscoveryStale(server_id=server.server_id)
candidate: Final = self._merge_discovered_oauth_metadata(self._registered_server(server), metadata)
if _oauth_endpoints_unresolved(candidate):
if oauth_endpoints_unresolved(candidate):
return None
published_candidate: Final = self._publish_resolved_oauth_server(candidate, generation)
return (
@ -2105,7 +2103,7 @@ class MCPServerManager:
return outcome
def _record_oauth_discovery_failure(self, server_id: str, generation: int) -> None:
slot: Final = self._oauth_discovery_slot(server_id)
slot: Final = self.oauth_discovery_slot(server_id)
if slot is None or slot.generation != generation:
return
consecutive_failures: Final = slot.consecutive_failures + 1
@ -2121,7 +2119,7 @@ class MCPServerManager:
self,
server: MCPServer,
) -> tuple[asyncio.Task[_OAuthDiscoveryOutcome], int] | None:
slot: Final = self._oauth_discovery_slot(server.server_id)
slot: Final = self.oauth_discovery_slot(server.server_id)
if slot is None:
return None
if slot.task is not None:
@ -2150,19 +2148,24 @@ class MCPServerManager:
"""
self._get_or_start_oauth_discovery_task(server)
def _prime_oauth_metadata_discovery_for_servers(self, servers: Sequence[MCPServer]) -> None:
def prime_oauth_metadata_discovery_for_servers(self, servers: Sequence[MCPServer]) -> None:
for server in servers:
self.prime_oauth_metadata_discovery(server)
def _reconcile_oauth_discovery_slots_for_servers(self, servers: Sequence[MCPServer]) -> None:
def reconcile_oauth_discovery_slots_for_servers(self, servers: Sequence[MCPServer]) -> None:
"""Align retry slots after an atomic registry replacement."""
for server in servers:
should_defer = _requires_oauth_discovery(server.url, server.issuer_is_anchored, server)
has_slot = self._oauth_discovery_slot(server.server_id) is not None
has_slot = self.oauth_discovery_slot(server.server_id) is not None
if should_defer != has_slot:
self._set_oauth_discovery_deferred(server.server_id, should_defer)
async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
return await self.catalog.resolve_oauth_metadata(
server, lambda selected: self._ensure_oauth_metadata_discovered(selected, _retry_stale=_retry_stale)
)
async def _ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
"""Join the bounded discovery task and return the resolved server.
Concurrent callers share one task per server. A failed attempt remains
@ -2214,7 +2217,7 @@ class MCPServerManager:
if retry_stale:
return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False)
current: Final = self._registered_server(server)
if not _oauth_endpoints_unresolved(current) or current.is_client_forwarded_token:
if not oauth_endpoints_unresolved(current) or current.is_client_forwarded_token:
return current
raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly")
@ -2296,8 +2299,15 @@ class MCPServerManager:
"""
Get the registered MCP Servers from the registry and union with the config MCP Servers
"""
snapshot: Final = self.catalog.current
return snapshot if snapshot is not None else self.config_mcp_servers | self.registry
return self.catalog.registry()
@property
def tool_name_to_mcp_server_name_mapping(self) -> dict[str, str]:
return self.catalog.routing()
@tool_name_to_mcp_server_name_mapping.setter
def tool_name_to_mcp_server_name_mapping(self, mapping: dict[str, str]) -> None:
self.published_tool_routes = mapping
def is_config_declared_server(self, server_id: str) -> bool:
"""True when server_id was declared in config.yaml (present in the in-memory config map).
@ -2384,7 +2394,7 @@ class MCPServerManager:
)
assigned_server_ids[server_id] = server_name
_warn_on_server_name_fields(
warn_on_server_name_fields(
server_id=server_id,
alias=alias,
server_name=server_name,
@ -2590,10 +2600,10 @@ class MCPServerManager:
token_validation=server_config.get("token_validation", None),
oauth_identity_binding=server_config.get("oauth_identity_binding", None),
)
self._assign_unique_short_prefix(new_server)
self.assign_unique_short_prefix(new_server)
_warn_legacy_delegate_auth_if_applicable(new_server, source="config")
_warn_config_id_jag_server_outruns_sso(new_server)
self._invalidate_discovery_lists(server_id)
self.invalidate_discovery_lists(server_id)
self.config_mcp_servers[server_id] = new_server
self._set_oauth_discovery_deferred(
server_id,
@ -2614,13 +2624,13 @@ class MCPServerManager:
"Loaded MCP Servers: %s", json.dumps(_redacted_registry_dump(self.config_mcp_servers), indent=4)
)
await self._hydrate_config_servers_dcr_clients()
await self.hydrate_config_servers_dcr_clients()
self._prime_oauth_metadata_discovery_for_servers(tuple(self.config_mcp_servers.values()))
self.prime_oauth_metadata_discovery_for_servers(tuple(self.config_mcp_servers.values()))
self.initialize_tool_name_to_mcp_server_name_mapping()
async def _hydrate_config_servers_dcr_clients(self) -> None:
async def hydrate_config_servers_dcr_clients(self, servers: Sequence[MCPServer] | None = None) -> None:
"""Overlay each config-declared server's persisted DCR client (from the server-scoped
store) onto its in-memory object so token refresh authenticates after a restart. A
best-effort no-op when the DB is unreachable at config-load time."""
@ -2628,7 +2638,7 @@ class MCPServerManager:
hydrate_config_server_dcr_client,
)
for server in self.config_mcp_servers.values():
for server in servers if servers is not None else self.config_mcp_servers.values():
try:
if await hydrate_config_server_dcr_client(server):
verbose_logger.debug(
@ -2793,17 +2803,18 @@ class MCPServerManager:
mappings make ``_get_mcp_server_from_tool_name`` resolve to a prefix that
no longer exists in the live registry.
"""
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
self.invalidate_discovery_lists(server.server_id)
self.remove_server_tool_routing(server)
def remove_server_tool_routing(self, server: MCPServer) -> None:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
self._invalidate_discovery_lists(server.server_id)
prefix_root: Final = normalize_server_name(get_server_prefix(server))
if server.spec_path and prefix_root:
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
global_mcp_tool_registry.unregister_tools_with_prefix(openapi_key_prefix)
owned_normalized: Final = self._owned_mapping_values(server)
owned_normalized: Final = self.owned_mapping_values(server)
stale_mapping_keys: Final = tuple(
tool_name
@ -2814,13 +2825,13 @@ class MCPServerManager:
for key in stale_mapping_keys:
del self.tool_name_to_mcp_server_name_mapping[key]
def _owned_mapping_values(self, server: MCPServer) -> frozenset[str]:
def owned_mapping_values(self, server: MCPServer) -> frozenset[str]:
return frozenset(
normalize_server_name(value) for value in (*iter_known_server_prefixes(server), server.name) if value
)
def server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool:
owned: Final = self._owned_mapping_values(server)
owned: Final = self.owned_mapping_values(server)
mapped_owners: Final = (
self.tool_name_to_mcp_server_name_mapping.get(spelling)
for spelling in iter_known_tool_name_spellings(tool_name, server)
@ -2851,7 +2862,7 @@ class MCPServerManager:
if evicted is not None:
verbose_logger.debug("Removed MCP Server: %s", mcp_server.server_id or mcp_server.server_name)
self._cleanup_server_tool_routing_artifacts(evicted)
self._invalidate_oauth_discovery_state(evicted.server_id)
self.invalidate_oauth_discovery_state(evicted.server_id)
else:
verbose_logger.warning("Server ID %s not found in registry", mcp_server.server_id)
@ -2942,6 +2953,7 @@ class MCPServerManager:
*,
credentials_are_encrypted: bool = True,
env_vars_are_encrypted: bool | None = None,
register_oauth_discovery: bool = True,
) -> MCPServer:
_mcp_info: Final[MCPInfo] = mcp_server.mcp_info or {}
env_dict: Final = _deserialize_json_dict(getattr(mcp_server, "env", None))
@ -3166,13 +3178,14 @@ class MCPServerManager:
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
)
_warn_legacy_delegate_auth_if_applicable(new_server, source="database")
self._set_oauth_discovery_deferred(
new_server.server_id,
_requires_oauth_discovery(server_url, use_issuer_anchor, new_server),
)
if register_oauth_discovery:
self._set_oauth_discovery_deferred(
new_server.server_id,
_requires_oauth_discovery(server_url, use_issuer_anchor, new_server),
)
return new_server
async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
async def maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
"""Register OpenAPI tools if the server has a spec_path configured."""
if server.spec_path:
verbose_logger.info("Loading OpenAPI spec from %s for server %s", server.spec_path, server.name)
@ -3200,10 +3213,10 @@ class MCPServerManager:
# Re-decrypting plaintext would zero the values, so build with
# env_vars_are_encrypted=False.
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
self._assign_unique_short_prefix(new_server)
self._invalidate_discovery_lists(mcp_server.server_id)
self.assign_unique_short_prefix(new_server)
self.invalidate_discovery_lists(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
await self.maybe_register_openapi_tools(new_server)
self.prime_oauth_metadata_discovery(new_server)
verbose_logger.debug("Added MCP Server: %s", new_server.name)
@ -3221,7 +3234,7 @@ class MCPServerManager:
evicted = self.registry.pop(mcp_server.server_name, None)
if evicted is not None:
self._cleanup_server_tool_routing_artifacts(evicted)
self._invalidate_oauth_discovery_state(evicted.server_id)
self.invalidate_oauth_discovery_state(evicted.server_id)
return
try:
if mcp_server.server_id in self.registry:
@ -3233,14 +3246,14 @@ class MCPServerManager:
existing_prefix: Final = self.registry[mcp_server.server_id].short_prefix
if existing_prefix and not new_server.short_prefix:
new_server.short_prefix = existing_prefix
_carry_forward_resolved_oauth_endpoints(
carry_forward_resolved_oauth_endpoints(
new_server=new_server,
previous_server=self.registry[mcp_server.server_id],
)
self._assign_unique_short_prefix(new_server)
self._invalidate_discovery_lists(mcp_server.server_id)
self.assign_unique_short_prefix(new_server)
self.invalidate_discovery_lists(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
await self.maybe_register_openapi_tools(new_server)
self.prime_oauth_metadata_discovery(new_server)
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
@ -4135,6 +4148,7 @@ class MCPServerManager:
Returns:
Configured MCP client instance.
"""
self.catalog.assert_current(server)
record_auth_resolution(server.server_id, AuthResolution.unresolved)
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
transport: Final = resolved_server.transport or MCPTransport.sse
@ -4472,7 +4486,9 @@ class MCPServerManager:
)
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
def _invalidate_discovery_lists(self, server_id: str) -> None:
def invalidate_discovery_lists(self, server_id: str) -> None:
self._upstream_initialize_instructions_by_server_id.pop(server_id, None)
self._upstream_initialize_instructions_probed_at.pop(server_id, None)
self._prompt_discovery_cache.invalidate(server_id)
self._resource_discovery_cache.invalidate(server_id)
self._template_discovery_cache.invalidate(server_id)
@ -5262,7 +5278,7 @@ class MCPServerManager:
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024
def _assign_unique_short_prefix(
def assign_unique_short_prefix(
self,
server: MCPServer,
registry: dict[str, MCPServer] | None = None,
@ -6153,7 +6169,7 @@ class MCPServerManager:
failure is logged, never raised, because the DB write already succeeded and the TTL remains
the backstop.
"""
self._invalidate_discovery_lists(server_id)
self.invalidate_discovery_lists(server_id)
try:
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
@ -6452,7 +6468,7 @@ class MCPServerManager:
Note: This now handles prefixed tool names
"""
for server in self.get_registry().values():
if self._oauth_discovery_slot(server.server_id) is not None:
if self.oauth_discovery_slot(server.server_id) is not None:
continue
if server.needs_user_oauth_token:
# Skip OAuth2 servers that rely on user-provided tokens
@ -6514,157 +6530,7 @@ class MCPServerManager:
return None
async def reload_servers_from_database(self):
await self.catalog.refresh()
async def _reload_servers_from_database(self, *, reuse_unchanged: bool = False):
"""Re-synchronize the in-memory MCP server registry with the database."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_prisma_client_or_throw,
)
verbose_logger.debug("Loading MCP servers from database into registry...")
# perform authz check to filter the mcp servers user has access to
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
# Load only "active", legacy "approved", and NULL (no approval workflow) rows.
# Pending/rejected servers are excluded at the DB level so we never load them.
from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable
raw_rows: Final[Sequence[BaseModel]] = await MCPServerRepository(prisma_client).table.find_many(
where={
"OR": [
{"approval_status": None},
{"approval_status": {"in": ["active", "approved"]}},
]
}
)
verbose_logger.info("Found %s MCP servers in database", len(raw_rows))
previous_registry: Final = self.registry
new_registry: Final[dict[str, MCPServer]] = {}
# Stage one: build every server. Stage two assigns short prefixes
# against the *full* set so dedup is deterministic regardless of
# iteration order.
for row in raw_rows:
try:
server = LiteLLM_MCPServerTable.model_validate(row.model_dump())
existing_server = previous_registry.get(server.server_id)
if (
existing_server is not None
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
and (
self._oauth_discovery_slot(server.server_id) is not None
or not _oauth_endpoints_unresolved(existing_server)
)
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
_warn_on_server_name_fields(
server_id=server.server_id,
alias=getattr(server, "alias", None),
server_name=getattr(server, "server_name", None),
)
verbose_logger.debug("Building server from DB: %s (%s)", server.server_id, server.server_name)
# raw_rows come straight from the DB, so their global env var
# values (like credentials) are still encrypted here, unlike the
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
new_server.short_prefix = existing_server.short_prefix
_carry_forward_resolved_oauth_endpoints(new_server=new_server, previous_server=existing_server)
new_registry[server.server_id] = new_server
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
getattr(row, "server_id", None),
getattr(row, "alias", None),
e,
)
# Assign short prefixes against the full candidate set without
# publishing the staged registry to concurrent callers.
registered_registry: Final[dict[str, MCPServer]] = {}
registered_openapi_tools = False
for server_id, new_server in new_registry.items():
try:
self._assign_unique_short_prefix(new_server, registry=new_registry)
# Register OpenAPI tools *after* the final short prefix is assigned
# so the tools are stored in the global registry under the same
# prefix that lookups will use.
if not reuse_unchanged or new_server is not previous_registry.get(server_id):
await self._maybe_register_openapi_tools(new_server, initialize_mapping=False)
if new_server.spec_path:
registered_openapi_tools = True
registered_registry[server_id] = new_server
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
new_server.server_id,
getattr(new_server, "alias", None),
e,
)
dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys()
for registry_key in dropped_registry_keys:
self._cleanup_server_tool_routing_artifacts(previous_registry[registry_key])
self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
for server_id in previous_registry.keys() | registered_registry.keys():
if previous_registry.get(server_id) != registered_registry.get(server_id):
self._upstream_initialize_instructions_by_server_id.pop(server_id, None)
self._upstream_initialize_instructions_probed_at.pop(server_id, None)
self._invalidate_discovery_lists(server_id)
self.registry = registered_registry
# A discovery task may have published into ``previous_registry`` while
# this replacement was being staged. Reconcile every published entry
# synchronously after the swap so a lost publication cannot also leave
# the replacement unresolved with no retry slot.
registered_servers: Final = tuple(registered_registry.values())
self._reconcile_oauth_discovery_slots_for_servers(registered_servers)
self._prime_oauth_metadata_discovery_for_servers(registered_servers)
if registered_openapi_tools:
self.initialize_tool_name_to_mcp_server_name_mapping()
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
# get_registry() is ``config_mcp_servers | registry``, so a database row sharing an id with a
# config.yaml server hides that server everywhere. Only reachable once an operator pins
# ``server_id`` in config.yaml; say so rather than letting the server disappear silently.
shadowed_config_server_ids: Final = frozenset(self.config_mcp_servers.keys() & registered_registry.keys())
if shadowed_config_server_ids and shadowed_config_server_ids != self._warned_shadowed_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are also database-backed MCP servers. The database "
"entry takes precedence, so the config.yaml server is unreachable. Give the config "
"entry a different server_id.",
", ".join(sorted(shadowed_config_server_ids)),
)
self._warned_shadowed_config_server_ids = shadowed_config_server_ids
# The mirror image of the block above: a config server_id that is a database server's name
# answers that server's grants instead, because ids are matched before names.
capturing_config_server_ids: Final = _config_ids_capturing_db_identifiers(
self.config_mcp_servers.keys(), registered_registry.values()
)
if capturing_config_server_ids and capturing_config_server_ids != self._warned_capturing_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are the name or alias of a database-backed MCP "
"server. Permission entries naming them resolve to the config.yaml server, not the "
"database one. Give the config entry a different server_id.",
", ".join(sorted(capturing_config_server_ids)),
)
self._warned_capturing_config_server_ids = capturing_config_server_ids
await self._hydrate_config_servers_dcr_clients()
await self.catalog.reload()
def get_mcp_servers_from_ids(self, server_ids: list[str]) -> list[MCPServer]:
servers: Final = []

View file

@ -1,6 +1,5 @@
"""Shared MCP operation policy and dispatch."""
import asyncio
import traceback
import types
import uuid
@ -50,7 +49,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 with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog, catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -615,7 +614,7 @@ def apply_tool_overrides(
return tools
@with_mcp_catalog
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_allowed_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_servers: Sequence[str] | None,
@ -932,6 +931,7 @@ def _aggregate_server_key(server: MCPServer) -> str:
return get_server_prefix(server) or "unknown"
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_tools_from_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | None,
@ -1153,24 +1153,17 @@ 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)
# 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
}
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
# If logging is enabled, enrich spend_logs_metadata with counts
if litellm_logging_obj:
per_server_tool_counts: Final[dict[str, int]] = {
_aggregate_server_key(server): len(server_tools)
for server, (server_tools, _) in zip(allowed_mcp_servers, results)
if server is not None
key: outcome.tool_count if isinstance(outcome, ServerListOk) else 0
for key, outcome in server_outcomes.items()
}
metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata")
@ -1437,7 +1430,7 @@ async def filter_tools_by_key_team_permissions(
]
@with_mcp_catalog
@catalog_operation(lambda: global_mcp_server_manager)
async def _list_mcp_tools(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1488,7 +1481,7 @@ async def _list_mcp_tools(
return AggregateToolListing(tools=[], outcomes={})
@with_mcp_catalog
@catalog_operation(global_manager)
async def _list_mcp_prompts(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1530,7 +1523,7 @@ async def _list_mcp_prompts(
return managed_prompts
@with_mcp_catalog
@catalog_operation(global_manager)
async def _list_mcp_resources(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1560,7 +1553,7 @@ async def _list_mcp_resources(
return managed_resources
@with_mcp_catalog
@catalog_operation(global_manager)
async def _list_mcp_resource_templates(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -2277,6 +2270,7 @@ async def fire_mcp_tool_call_failure_logging(
@client
@catalog_operation(lambda: global_mcp_server_manager)
async def call_mcp_tool(
name: str,
arguments: dict[str, object] | None = None,
@ -3064,7 +3058,7 @@ class GatewayOperations:
@overload
async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ...
@with_mcp_catalog
@catalog_operation(lambda: global_mcp_server_manager)
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case AuthorizedToolCall():

View file

@ -22,7 +22,7 @@ from litellm.exceptions import (
GuardrailRaisedException,
ModifyResponseException,
)
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPServerURLCredentialsError,
@ -862,7 +862,7 @@ if MCP_AVAILABLE:
return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id)
@router.get("/tools/list", dependencies=[Depends(user_api_key_auth)])
@with_mcp_catalog
@catalog_operation(global_manager)
async def list_tool_rest_api(
request: Request,
server_id: str | None = Query(None, description="The server id to list tools for"),
@ -1087,7 +1087,7 @@ if MCP_AVAILABLE:
}
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
@with_mcp_catalog
@catalog_operation(global_manager)
async def call_tool_rest_api(
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),

View file

@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
from litellm.exceptions import ContextWindowExceededError
from litellm.litellm_core_utils.exception_mapping_utils import ExceptionCheckers
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
from litellm.proxy._experimental.mcp_server.utils import MCP_TOOL_PREFIX_SEPARATOR
@ -78,6 +79,7 @@ class SemanticMCPToolFilter:
self._tool_map: dict[str, object] = {} # MCPTool objects or OpenAI function dicts
self._index_sync_lock = asyncio.Lock()
@catalog_operation(global_manager)
async def build_router_from_mcp_registry(self) -> None:
"""Build semantic router from all MCP tools in the registry (no auth checks)."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (

View file

@ -1548,6 +1548,9 @@ if MCP_AVAILABLE:
)
return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id})
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation
@catalog_operation(lambda: operations.global_mcp_server_manager)
async def _raise_preemptive_401_for_unauthenticated_servers(
scope: Scope,
mcp_servers: list[str] | None,

View file

@ -1,5 +1,8 @@
import asyncio
import json
from collections.abc import Callable
from collections.abc import Callable, Iterator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
@ -22,7 +25,30 @@ class MCPToolRegistry:
def __init__(self):
# Registry to store all registered tools
self.tools: dict[str, MCPTool] = {}
self.published_tools: dict[str, MCPTool] = {}
self._catalog_tools: ContextVar[tuple[dict[str, MCPTool], asyncio.Event] | None] = ContextVar(
"mcp_catalog_tools", default=None
)
@property
def tools(self) -> dict[str, MCPTool]:
scoped: Final = self._catalog_tools.get()
return scoped[0] if scoped is not None and not scoped[1].is_set() else self.published_tools
@tools.setter
def tools(self, tools: dict[str, MCPTool]) -> None:
self.published_tools = tools
@contextmanager
def catalog_scope(self, tools: Mapping[str, MCPTool]) -> Iterator[dict[str, MCPTool]]:
detached: Final = dict(tools)
closed: Final = asyncio.Event()
token: Final = self._catalog_tools.set((detached, closed))
try:
yield detached
finally:
closed.set()
self._catalog_tools.reset(token)
def register_tool(
self,

View file

@ -81,7 +81,7 @@ MCP_TOOL_PREFIX_FORMAT: Final = "{server_name}{separator}{tool_name}"
# principle hash to the same three chars; that natural-hash collision
# IS a routing-correctness issue (the second registrant would otherwise
# have its tools misrouted to the first), so registration goes through
# ``MCPServerManager._assign_unique_short_prefix`` which rehashes with
# ``MCPServerManager.assign_unique_short_prefix`` which rehashes with
# a deterministic attempt counter until it finds an unused prefix and
# caches the result on ``MCPServer.short_prefix``. A collision is
# logged at INFO when it happens.
@ -114,7 +114,7 @@ def compute_short_server_prefix(server_id: str, attempt: int = 0) -> str:
and whose remaining characters are drawn from the full base62
alphabet. Pass ``attempt > 0`` to rehash to a different prefix when
the natural hash collides with a prefix already assigned to another
server (see ``MCPServerManager._assign_unique_short_prefix``). An
server (see ``MCPServerManager.assign_unique_short_prefix``). An
empty ``server_id`` raises ``ValueError`` — short prefixes require a
stable identifier to be deterministic.
"""
@ -314,7 +314,7 @@ def get_server_prefix(server: object) -> str:
When the short-prefix mode is enabled (``LITELLM_USE_SHORT_MCP_TOOL_PREFIX``)
a three-character base62 ID is returned. We prefer the cached
``server.short_prefix`` value when set — that field is populated at
registration time by ``MCPServerManager._assign_unique_short_prefix``
registration time by ``MCPServerManager.assign_unique_short_prefix``
and resolves natural-hash collisions deterministically — and only fall
back to the natural hash for ad-hoc / temp-server objects without a
cached value. In default mode the historical behaviour is preserved:

View file

@ -46,7 +46,7 @@ class MCPSecurityGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True:
return data
unregistered: Final = self._find_unregistered_mcp_servers(data)
unregistered: Final = await self._find_unregistered_mcp_servers(data)
if not unregistered:
return data
@ -90,7 +90,7 @@ class MCPSecurityGuardrail(CustomGuardrail):
return server_names
@staticmethod
def _find_unregistered_mcp_servers(data: dict) -> set[str]:
async def _find_unregistered_mcp_servers(data: dict) -> set[str]:
"""Check tools in data against the MCP server registry. Returns set of unregistered server names."""
tools: Final = data.get("tools")
if not tools or not isinstance(tools, list):
@ -104,7 +104,8 @@ class MCPSecurityGuardrail(CustomGuardrail):
global_mcp_server_manager,
)
registry: Final = global_mcp_server_manager.get_registry()
registered_names: Final = set(registry.keys())
async with global_mcp_server_manager.catalog.operation():
registry: Final = global_mcp_server_manager.get_registry()
registered_names: Final = set(registry.keys())
return requested_servers - registered_names
return requested_servers - registered_names

View file

@ -19,7 +19,8 @@ import functools
import importlib
import json
import os
from collections.abc import Iterable, Mapping, Sequence
from collections.abc import AsyncIterator, Iterable, Mapping, Sequence
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import (
@ -45,7 +46,7 @@ from fastapi import (
from fastapi.responses import JSONResponse
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, public_catalog_operation
try:
from prisma.errors import RecordNotFoundError, UniqueViolationError
@ -1045,7 +1046,7 @@ if MCP_AVAILABLE:
tags=["mcp"],
description="MCP registry endpoint. Spec: https://github.com/modelcontextprotocol/registry",
)
@with_mcp_catalog
@public_catalog_operation
async def get_mcp_registry(request: Request):
if not _is_public_registry_enabled():
raise HTTPException(
@ -1064,7 +1065,8 @@ if MCP_AVAILABLE:
registry_servers.append({"server": _build_builtin_registry_entry(base_url)})
# Centralized IP-based filtering: external callers only see public servers
registered_servers: Final = list(global_mcp_server_manager.get_filtered_registry(client_ip).values())
async with global_mcp_server_manager.catalog.operation():
registered_servers: Final = list(global_mcp_server_manager.get_filtered_registry(client_ip).values())
registered_servers.sort(key=_build_mcp_registry_server_name)
@ -1094,6 +1096,7 @@ if MCP_AVAILABLE:
return "view_all"
return "restricted"
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_team_scoped_mcp_server_list(
team_id: str,
) -> list[LiteLLM_MCPServerTable]:
@ -1132,6 +1135,7 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials_list(servers)
@catalog_operation(lambda: global_mcp_server_manager)
async def _resolve_accessible_mcp_servers(
user_api_key_dict: UserAPIKeyAuth,
) -> list[LiteLLM_MCPServerTable]:
@ -1153,6 +1157,7 @@ if MCP_AVAILABLE:
aggregated.setdefault(server.server_id, server)
return list(aggregated.values())
@catalog_operation(lambda: global_mcp_server_manager)
async def _connected_app_reachable_server_ids(user_api_key_dict: UserAPIKeyAuth) -> frozenset[str]:
"""Server ids a connected app authorized by this dashboard user is served on the aggregate
MCP endpoint, resolved through the one owner of the admitted subject so the page and the
@ -1169,7 +1174,7 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
response_model=list[LiteLLM_MCPServerTable],
)
@with_mcp_catalog
@catalog_operation(lambda: global_mcp_server_manager)
async def fetch_all_mcp_servers(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
team_id: str | None = Query(
@ -1288,7 +1293,7 @@ if MCP_AVAILABLE:
description="Health check for MCP servers",
dependencies=[Depends(user_api_key_auth)],
)
@with_mcp_catalog
@catalog_operation(lambda: global_mcp_server_manager)
async def health_check_servers(
server_ids: list[str] | None = Query(
None,
@ -1597,6 +1602,7 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
response_model=LiteLLM_MCPServerTable,
)
@catalog_operation(lambda: global_mcp_server_manager)
async def fetch_mcp_server(
request: Request,
server_id: str,
@ -1964,7 +1970,7 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials(temp_record)
@with_mcp_catalog
@public_catalog_operation
async def _mcp_oauth_user_api_key_auth(request: Request) -> UserAPIKeyAuth:
"""
Auth dependency for MCP OAuth browser-navigation endpoints (/authorize, /token).
@ -2022,9 +2028,7 @@ if MCP_AVAILABLE:
server_id: Final[str] = request.path_params.get("server_id", "")
if server_id:
_s = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if not _s:
_s = global_mcp_server_manager.get_mcp_server_by_name(server_id)
_s = await global_mcp_server_manager.catalog.resolve(server_id)
if (
_s
and getattr(_s, "auth_type", None) == MCPAuth.oauth2
@ -2065,48 +2069,50 @@ if MCP_AVAILABLE:
request_data=request_data,
)
@with_mcp_catalog
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_cached_temporary_mcp_server_or_404(
server_id: str,
user_api_key_dict: UserAPIKeyAuth,
request: Request | None = None,
) -> MCPServer:
server = await get_cached_temporary_mcp_server(server_id)
resolved_from_temp_cache: Final = server is not None
if server is None:
# Fall back to real DB/config server (e.g. for the user-side OAuth flow
# which calls these endpoints with a real server_id, not a temp session id).
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as server:
return server
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) if request else None
server = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(server_id, client_ip=client_ip)
if server is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail={"error": f"MCP server {server_id} not found"},
)
@asynccontextmanager
async def _oauth_server_operation(
server_id: str,
user_api_key_dict: UserAPIKeyAuth,
request: Request | None = None,
) -> AsyncIterator[MCPServer]:
temporary: Final = await get_cached_temporary_mcp_server(server_id)
if temporary is not None:
if not _user_has_admin_view(user_api_key_dict):
raise HTTPException(status_code=403, detail={"error": f"Access denied to MCP server {server_id}"})
yield temporary
return
async with global_mcp_server_manager.catalog.operation():
yield await _resolve_saved_oauth_server(server_id, user_api_key_dict, request)
# Per-server access policy mirrors `fetch_mcp_server`: admin-view
# callers are unrestricted; non-admins must have the server in their
# allowed-servers set. Temporary cached servers come from the
# admin-only `/server/oauth/session` setup flow and are not exposed
# to non-admins.
@catalog_operation(lambda: global_mcp_server_manager)
async def _resolve_saved_oauth_server(
server_id: str,
user_api_key_dict: UserAPIKeyAuth,
request: Request | None,
) -> MCPServer:
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) if request else None
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(server_id, client_ip=client_ip)
if server is None:
raise HTTPException(status_code=404, detail={"error": f"MCP server {server_id} not found"})
if not _user_has_admin_view(user_api_key_dict):
if resolved_from_temp_cache:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": f"Access denied to MCP server {server_id}"},
)
allowed_server_ids: Final[set[str]] = set()
for auth_context in await build_effective_auth_contexts(user_api_key_dict):
allowed_server_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(auth_context))
if server.server_id not in allowed_server_ids:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": f"Access denied to MCP server {server_id}"},
)
allowed_ids: Final[set[str]] = set()
for context in await build_effective_auth_contexts(user_api_key_dict):
allowed_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(context))
if server.server_id not in allowed_ids:
raise HTTPException(status_code=403, detail={"error": f"Access denied to MCP server {server_id}"})
return server
@router.get(
@ -2126,47 +2132,47 @@ if MCP_AVAILABLE:
response_type: str | None = None,
scope: str | None = None,
):
mcp_server: Final = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
_raise_if_not_oauth2(mcp_server)
# Use the server's stored client_id when the caller doesn't supply one
stored_or_supplied_client_id: Final = mcp_server.client_id or client_id or ""
ephemeral_dcr_client: Final = (
await resolve_ephemeral_dcr_client(
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server:
_raise_if_not_oauth2(mcp_server)
# Use the server's stored client_id when the caller doesn't supply one
stored_or_supplied_client_id: Final = mcp_server.client_id or client_id or ""
ephemeral_dcr_client: Final = (
await resolve_ephemeral_dcr_client(
request=request,
mcp_server=mcp_server,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
redirect_uri=redirect_uri,
)
if not stored_or_supplied_client_id
else None
)
resolved_client_id: Final = stored_or_supplied_client_id or (
ephemeral_dcr_client.client_id if ephemeral_dcr_client else ""
)
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
)
return await authorize_with_server(
request=request,
mcp_server=mcp_server,
client_id=resolved_client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
redirect_uri=redirect_uri,
response_type=response_type,
scope=scope,
ephemeral_dcr_client=ephemeral_dcr_client,
)
if not stored_or_supplied_client_id
else None
)
resolved_client_id: Final = stored_or_supplied_client_id or (
ephemeral_dcr_client.client_id if ephemeral_dcr_client else ""
)
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
)
return await authorize_with_server(
request=request,
mcp_server=mcp_server,
client_id=resolved_client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
response_type=response_type,
scope=scope,
ephemeral_dcr_client=ephemeral_dcr_client,
)
@router.post(
"/server/oauth/{server_id}/token",
@ -2186,47 +2192,47 @@ if MCP_AVAILABLE:
refresh_token: str | None = Form(None),
scope: str | None = Form(None),
):
mcp_server: Final = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
_raise_if_not_oauth2(mcp_server)
# Sealed passthrough codes exist only for the authorization_code grant. A refresh_token
# grant must never open one: the minted client is unrecoverable after the single flow by
# contract, so an expired browser-held token re-runs authorize instead.
sealed_code: Final = (
redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier)
if grant_type == "authorization_code"
else None
)
resolved_code: Final = sealed_code.upstream_code if sealed_code else code
# A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit
# or plain flow alike), so the exchange must present that binding, not the browser page.
resolved_redirect_uri: Final = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri
caller_client_id: Final = sealed_code.client_id if sealed_code else client_id
caller_client_secret: Final = sealed_code.client_secret if sealed_code else client_secret
resolved_client_id: Final = mcp_server.client_id or caller_client_id or ""
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server:
_raise_if_not_oauth2(mcp_server)
# Sealed passthrough codes exist only for the authorization_code grant. A refresh_token
# grant must never open one: the minted client is unrecoverable after the single flow by
# contract, so an expired browser-held token re-runs authorize instead.
sealed_code: Final = (
redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier)
if grant_type == "authorization_code"
else None
)
resolved_code: Final = sealed_code.upstream_code if sealed_code else code
# A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit
# or plain flow alike), so the exchange must present that binding, not the browser page.
resolved_redirect_uri: Final = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri
caller_client_id: Final = sealed_code.client_id if sealed_code else client_id
caller_client_secret: Final = sealed_code.client_secret if sealed_code else client_secret
resolved_client_id: Final = mcp_server.client_id or caller_client_id or ""
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
)
return await exchange_token_with_server(
request=request,
mcp_server=mcp_server,
grant_type=grant_type,
code=resolved_code,
redirect_uri=resolved_redirect_uri,
client_id=resolved_client_id,
client_secret=caller_client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None,
)
return await exchange_token_with_server(
request=request,
mcp_server=mcp_server,
grant_type=grant_type,
code=resolved_code,
redirect_uri=resolved_redirect_uri,
client_id=resolved_client_id,
client_secret=caller_client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None,
)
@router.post(
"/server/oauth/{server_id}/register",
@ -2238,22 +2244,22 @@ if MCP_AVAILABLE:
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
mcp_server: Final = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
request_data: Final = await _read_request_body(request=request)
data: Final[dict] = {**request_data}
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server:
request_data: Final = await _read_request_body(request=request)
data: Final[dict] = {**request_data}
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
return await register_client_with_server(
request=request,
mcp_server=mcp_server,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=server_id,
persist_credentials=_user_is_full_admin(user_api_key_dict),
client_redirect_uris=client_redirect_uris,
)
return await register_client_with_server(
request=request,
mcp_server=mcp_server,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=server_id,
persist_credentials=_user_is_full_admin(user_api_key_dict),
client_redirect_uris=client_redirect_uris,
)
@router.delete(
"/server/{server_id}",
@ -2604,6 +2610,7 @@ if MCP_AVAILABLE:
# ── Per-user MCP env var endpoints ────────────────────────────────────────
@catalog_operation(lambda: global_mcp_server_manager)
async def _authorize_and_fetch_mcp_server(
prisma_client,
user_api_key_dict: UserAPIKeyAuth,

View file

@ -322,7 +322,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot
from litellm.proxy._types import *
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
@ -19700,30 +19700,33 @@ async def _resolve_mcp_csv_tokens(csv_segment: str, client_ip: str | None) -> li
all-unmatched server filter falls back to the full ``allowed_mcp_servers``
list and silently broadens the request scope).
"""
from litellm.constants import DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.catalog import global_manager
seen: Final[set] = set()
deduped: Final[list[str]] = []
for raw in csv_segment.split(","):
token = raw.strip()
if not token or token in seen:
continue
seen.add(token)
deduped.append(token)
if len(deduped) >= DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS:
break
async with global_manager().catalog.operation():
from litellm.constants import DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
resolved: Final[list[str]] = []
for token in deduped:
if global_mcp_server_manager.get_mcp_server_by_name(token, client_ip=client_ip):
resolved.append(token)
continue
if await _is_mcp_access_group_cached(token):
resolved.append(token)
return resolved
seen: Final[set] = set()
deduped: Final[list[str]] = []
for raw in csv_segment.split(","):
token = raw.strip()
if not token or token in seen:
continue
seen.add(token)
deduped.append(token)
if len(deduped) >= DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS:
break
resolved: Final[list[str]] = []
for token in deduped:
if global_mcp_server_manager.get_mcp_server_by_name(token, client_ip=client_ip):
resolved.append(token)
continue
if await _is_mcp_access_group_cached(token):
resolved.append(token)
return resolved
async def _is_mcp_access_group_cached(name: str) -> bool:
@ -19760,7 +19763,7 @@ async def _is_mcp_access_group_cached(name: str) -> bool:
"/{mcp_server_name}/mcp",
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
)
@with_mcp_catalog
@public_catalog_operation
async def dynamic_mcp_route(mcp_server_name: str, request: Request):
"""Handle /{name}/mcp for MCP server aliases, toolsets, MCP access group tags, and comma-separated lists.
@ -19779,7 +19782,9 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request):
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
# 1. Registered MCP server alias
if global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip):
async with global_mcp_server_manager.catalog.operation():
server: Final = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
if server is not None:
return await _mcp_forward_as_path(mcp_server_name, request)
# 2. Comma-separated list — validate every token resolves to a known

View file

@ -10,7 +10,7 @@ from openai.types.responses.function_tool_param import FunctionToolParam
from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.utils import (
iter_known_server_prefixes,
logging_safe_mcp_headers,
@ -112,6 +112,7 @@ async def _toolset_exists(name: str) -> bool:
return False
@catalog_operation(global_manager)
async def _gateway_served_names(
names: Collection[str],
servers: Callable[[], Collection[MCPServer]] = _registered_mcp_servers,
@ -179,7 +180,7 @@ class LiteLLM_Proxy_MCP_Handler:
)
@staticmethod
@with_mcp_catalog
@catalog_operation(global_manager)
async def routes_through_gateway(
tools: Iterable[Mapping[str, object]] | None,
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
@ -231,7 +232,7 @@ class LiteLLM_Proxy_MCP_Handler:
return user_api_key_auth
@staticmethod
@with_mcp_catalog
@catalog_operation(global_manager)
async def _get_mcp_tools_from_manager(
user_api_key_auth: "UserAPIKeyAuth | None",
mcp_tools_with_litellm_proxy: Iterable[Mapping[str, object]] | None,
@ -687,7 +688,7 @@ class LiteLLM_Proxy_MCP_Handler:
return result_text or "Tool executed successfully"
@staticmethod
@with_mcp_catalog
@catalog_operation(global_manager)
async def _execute_tool_calls(
tool_server_map: dict[str, str],
tool_calls: Sequence[object],

View file

@ -204,7 +204,7 @@ class MCPServer(BaseModel):
# None or a value <= 0 means unlimited.
max_concurrent_requests: 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
# enabled. Set by ``MCPServerManager.assign_unique_short_prefix`` at
# registration time so that natural-hash collisions between two
# different ``server_id`` values are bumped deterministically. Left
# ``None`` in default-prefix mode.

View file

@ -2,7 +2,7 @@
import os
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from contextlib import asynccontextmanager
from contextlib import asynccontextmanager, nullcontext
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -918,6 +918,7 @@ async def test_get_tools_from_mcp_servers():
# Create a mock manager
mock_manager = AsyncMock()
mock_manager.catalog.operation = nullcontext
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
@ -946,6 +947,7 @@ async def test_get_tools_from_mcp_servers():
# Test Case 2: Without specific MCP servers
# Create a different mock manager for the second test case
mock_manager_2 = AsyncMock()
mock_manager_2.catalog.operation = nullcontext
mock_manager_2.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
@ -993,6 +995,7 @@ async def test_get_tools_from_mcp_servers():
# Test Case 3: With specific MCP servers and access groups
# Create a mock manager
mock_manager = AsyncMock()
mock_manager.catalog.operation = nullcontext
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id", "server3_id"]
)

View file

@ -0,0 +1,89 @@
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

@ -7491,7 +7491,8 @@ class TestGatewaySessionAdmission:
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
with (
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.should_load_db_object", return_value=False),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
yield get_user_object
@ -9583,7 +9584,8 @@ class TestScopedSessionAdmission:
with (
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.should_load_db_object", return_value=False),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict)

View file

@ -631,6 +631,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy")
mcp_operations.byok_credential_cache.flush_cache()
server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True)
monkeypatch.setattr(proxy_server, "should_load_db_object", lambda _kind: False)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None)

View file

@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from fastapi import HTTPException, Request
from litellm.types.mcp import MCPAuth
@ -91,7 +91,7 @@ def mock_mcp_client_ip():
@pytest.fixture(autouse=True)
def isolate_global_mcp_registry():
def isolate_global_mcp_registry(monkeypatch):
"""Restore the module-global MCP server registry after each test.
Tests here register servers on ``global_mcp_server_manager`` directly; without a
@ -100,6 +100,8 @@ def isolate_global_mcp_registry():
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
monkeypatch.setattr(global_mcp_server_manager, "catalog", TargetCatalog(global_mcp_server_manager))
snapshot = dict(global_mcp_server_manager.registry)
yield
global_mcp_server_manager.registry.clear()
@ -114,7 +116,8 @@ def _mock_callback_request(base_url: str = "http://localhost:3000/"):
and trusted ``X-Forwarded-*`` headers). A simple MagicMock with the
right attributes is sufficient.
"""
req = MagicMock()
req = MagicMock(spec=Request)
req.client = None
req.base_url = base_url
req.headers = {}
req.cookies = {}
@ -4700,7 +4703,7 @@ async def test_token_endpoint_authorization_code_missing_code():
)
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.example/"
mock_request.headers = {}
@ -7796,7 +7799,7 @@ async def test_authorize_endpoint_rejects_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -7835,7 +7838,7 @@ async def test_token_endpoint_rejects_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -7877,7 +7880,7 @@ async def test_register_client_rejects_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -7950,7 +7953,7 @@ async def test_oauth_authorization_server_404_for_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -7998,7 +8001,7 @@ async def test_oauth_protected_resource_passthrough_none_auth_not_404():
)
global_mcp_server_manager.registry[passthrough_server.server_id] = passthrough_server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8032,7 +8035,7 @@ async def test_oauth_protected_resource_404_for_unknown_server_name():
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8060,7 +8063,7 @@ async def test_oauth_authorization_server_404_for_unknown_server_name():
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9064,7 +9067,7 @@ async def test_load_servers_from_config_hydrates_dcr_clients():
)
hydrate_spy = AsyncMock()
with patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy):
with patch.object(global_mcp_server_manager, "hydrate_config_servers_dcr_clients", new=hydrate_spy):
await global_mcp_server_manager.load_servers_from_config({})
hydrate_spy.assert_awaited_once()
@ -9089,7 +9092,7 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=prisma,
),
patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy),
patch.object(global_mcp_server_manager, "hydrate_config_servers_dcr_clients", new=hydrate_spy),
):
await global_mcp_server_manager.reload_servers_from_database()
@ -9360,7 +9363,7 @@ async def test_authorize_wall_names_the_fix_for_urlless_servers():
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9398,7 +9401,7 @@ async def test_token_wall_names_the_fix_for_urlless_servers():
spec_path="https://example.com/openapi.yaml",
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9438,7 +9441,7 @@ async def test_register_wall_names_the_fix_for_urlless_servers():
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9477,7 +9480,7 @@ async def test_authorize_wall_points_at_discovery_failure_for_url_servers():
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9513,7 +9516,7 @@ async def test_token_wall_points_at_discovery_failure_for_url_servers():
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9553,7 +9556,7 @@ async def test_authorize_wall_names_the_issuer_for_anchored_servers():
issuer="https://idp.example.com",
issuer_is_anchored=True,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9600,7 +9603,7 @@ async def test_authorize_uses_admin_entered_github_oauth_urls_after_issuer_yield
configured_authorization_url="https://github.com/login/oauth/authorize",
configured_token_url="https://github.com/login/oauth/access_token",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9622,7 +9625,7 @@ def test_oauth_endpoints_count_admin_entered_urls_as_resolved():
"""A leftover issuer empties the resolved authorize/token fields but must not keep the
server on the deferred-discovery retry path when the admin already stored those URLs."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_oauth_endpoints_unresolved,
oauth_endpoints_unresolved,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -9639,7 +9642,7 @@ def test_oauth_endpoints_count_admin_entered_urls_as_resolved():
configured_authorization_url="https://github.com/login/oauth/authorize",
configured_token_url="https://github.com/login/oauth/access_token",
)
assert _oauth_endpoints_unresolved(server) is False
assert oauth_endpoints_unresolved(server) is False
@pytest.mark.asyncio
@ -9692,7 +9695,7 @@ async def test_token_exchange_with_configured_token_url_never_joins_discovery(mo
"get_async_httpx_client",
lambda llm_provider: fake_http_client,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9817,7 +9820,7 @@ async def test_bridge_authorize_relays_with_registration_url_resolved_by_deferre
"ensure_oauth_metadata_discovered",
resolve_discovery,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -11139,7 +11142,7 @@ def test_discovery_advertises_the_exchange_grant_only_where_the_gateway_can_serv
litellm_jwtauth=LiteLLM_JWTAuth(virtual_key_claim_field=virtual_key_claim_field),
)
monkeypatch.setattr("litellm.proxy.proxy_server.jwt_handler", handler)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": jwt_auth_enabled})
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": jwt_auth_enabled, "supported_db_objects": []})
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
exchange_grant = ["urn:ietf:params:oauth:grant-type:token-exchange"] if exchange_servable else []

View file

@ -66,8 +66,7 @@ async def _call_block(logging_obj, order: list, *, user_api_key_auth=mock.sentin
proxy_logging_obj = mock.MagicMock()
proxy_logging_obj.post_call_failure_hook.side_effect = _record_post_call_failure_hook
fake_proxy_server = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy_server.proxy_logging_obj = proxy_logging_obj # pyright: ignore[reportAttributeAccessIssue]
fake_proxy_server = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None)
with mock.patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}):
with contextlib.suppress(HTTPException):

View file

@ -5445,7 +5445,7 @@ class TestMCPServerManagerReload:
):
await manager.reload_servers_from_database()
mock_build.assert_awaited_once_with(db_row, env_vars_are_encrypted=True)
mock_build.assert_awaited_once_with(db_row, env_vars_are_encrypted=True, register_oauth_discovery=False)
assert manager.registry["server-1"] is rebuilt_server
@pytest.mark.asyncio
@ -5497,7 +5497,7 @@ class TestMCPServerManagerReload:
"build_mcp_server_from_table",
AsyncMock(side_effect=build_server),
),
patch.object(manager, "_maybe_register_openapi_tools", AsyncMock()),
patch.object(manager, "maybe_register_openapi_tools", AsyncMock()),
caplog.at_level("ERROR", logger="LiteLLM"),
):
await manager.reload_servers_from_database()
@ -5569,7 +5569,7 @@ class TestMCPServerManagerReload:
),
patch.object(
manager,
"_maybe_register_openapi_tools",
"maybe_register_openapi_tools",
AsyncMock(side_effect=register_openapi_tools),
),
caplog.at_level("ERROR", logger="LiteLLM"),
@ -9507,7 +9507,7 @@ class TestPreemptive401ModeAware:
assert resolved.authorization_url == "https://idp.example.com/authorize"
assert resolved.token_url == "https://idp.example.com/token"
assert resolved.registration_url == "https://idp.example.com/register"
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
assert exc.value.status_code == 401
@pytest.mark.asyncio

View file

@ -1560,9 +1560,8 @@ class TestListToolsRestAPI:
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: stub_server if server_id == "server-1" else None,
raising=False,
"registry",
{stub_server.server_id: stub_server},
)
request = _build_request(path="/mcp-rest/tools/list", method="GET")

View file

@ -345,7 +345,7 @@ class TestManagerShortPrefix:
class TestShortPrefixCollisionResolution:
"""``_assign_unique_short_prefix`` must rehash on collision.
"""``assign_unique_short_prefix`` must rehash on collision.
The dedup path is exercised by forcing two distinct ``server_id``
values to both hash to the same natural prefix via a monkeypatched
@ -355,7 +355,7 @@ class TestShortPrefixCollisionResolution:
def test_no_op_when_flag_off(self):
manager = MCPServerManager()
server = _make_server(server_id="abc")
manager._assign_unique_short_prefix(server)
manager.assign_unique_short_prefix(server)
assert server.short_prefix is None
def test_assigns_natural_hash_when_no_collision(self, monkeypatch):
@ -364,7 +364,7 @@ class TestShortPrefixCollisionResolution:
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
manager = MCPServerManager()
server = _make_server(server_id="abc")
manager._assign_unique_short_prefix(server)
manager.assign_unique_short_prefix(server)
assert server.short_prefix == mcp_utils.compute_short_server_prefix("abc")
@ -394,9 +394,9 @@ class TestShortPrefixCollisionResolution:
# Pretend both are already in the registry so dedup sees both.
manager.registry[first.server_id] = first
manager._assign_unique_short_prefix(first)
manager.assign_unique_short_prefix(first)
manager.registry[second.server_id] = second
manager._assign_unique_short_prefix(second)
manager.assign_unique_short_prefix(second)
assert first.short_prefix == "AAA"
assert second.short_prefix == "AAB"
@ -408,7 +408,7 @@ class TestShortPrefixCollisionResolution:
server = _make_server(server_id="abc")
server.short_prefix = "ZZZ" # pretend a previous registration set this
manager._assign_unique_short_prefix(server)
manager.assign_unique_short_prefix(server)
assert server.short_prefix == "ZZZ"

View file

@ -205,3 +205,30 @@ class TestInitializeGuardrail:
assert isinstance(result, MCPSecurityGuardrail)
assert result.on_violation == expected
assert result in litellm.callbacks
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["acompletion", "aresponses"])
async def test_guardrail_observes_saved_server_creation_and_deletion_on_another_worker(guardrail, call_type):
from unittest.mock import AsyncMock
from litellm.proxy._types import LiteLLM_MCPServerTable
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(server_id="peer-server", alias="peer_server", transport="http",
url="https://upstream.example/mcp")
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], []))
data = {"tools": [{"type": "mcp", "server_url": "litellm_proxy/mcp/peer-server"}],
"guardrails": ["test-mcp-security"]}
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager),
):
result = await guardrail.async_pre_call_hook(UserAPIKeyAuth(), MagicMock(), data, call_type)
assert result == data
with pytest.raises(HTTPException) as exc:
await guardrail.async_pre_call_hook(UserAPIKeyAuth(), MagicMock(), data, call_type)
assert exc.value.status_code == 400
assert exc.value.detail["unregistered_servers"] == ["peer-server"]

View file

@ -1,3 +1,4 @@
from contextlib import nullcontext
import os
import sys
import types
@ -11,7 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from respx import MockRouter
from fastapi import FastAPI, HTTPException
from fastapi import FastAPI, HTTPException, Request
from fastapi.testclient import TestClient
from litellm._uuid import uuid
@ -2545,14 +2546,15 @@ class TestTemporaryMCPSessionEndpoints:
algorithm="HS256",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {}
mock_request.cookies = {"token": token_cookie}
expected_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key=api_key_in_cookie
)
fake_proxy_server = types.SimpleNamespace(master_key=master_key, prisma_client=None)
fake_proxy_server = types.SimpleNamespace(master_key=master_key, prisma_client=None, general_settings={})
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2585,7 +2587,8 @@ class TestTemporaryMCPSessionEndpoints:
)
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {"Authorization": "Bearer sk-header-key"}
mock_request.cookies = {}
@ -2619,7 +2622,8 @@ class TestTemporaryMCPSessionEndpoints:
)
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {}
mock_request.cookies = {}
mock_request.path_params = {"server_id": "server-1"}
@ -2629,7 +2633,8 @@ class TestTemporaryMCPSessionEndpoints:
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = non_oauth_server
mock_manager.get_mcp_server_by_name.return_value = None
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None)
mock_manager.catalog.resolve = AsyncMock(return_value=non_oauth_server)
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None, general_settings={})
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2667,7 +2672,8 @@ class TestTemporaryMCPSessionEndpoints:
)
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {}
mock_request.cookies = {}
mock_request.path_params = {"server_id": "server-1"}
@ -2681,7 +2687,8 @@ class TestTemporaryMCPSessionEndpoints:
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = internal_server
mock_manager.get_mcp_server_by_name.return_value = None
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None)
mock_manager.catalog.resolve = AsyncMock(return_value=internal_server)
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None, general_settings={})
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2738,6 +2745,63 @@ class TestTemporaryMCPSessionEndpoints:
}
assert dependency_names == {None, "user_api_key_dict"}
@pytest.mark.asyncio
async def test_authorize_saved_server_on_cold_worker_without_oauth_session(self):
from collections.abc import Mapping
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy.management_endpoints.mcp_management_endpoints import mcp_authorize
row: Final = LiteLLM_MCPServerTable(
server_id="saved-server",
server_name="saved_server",
alias="saved_server",
transport=MCPTransport.http,
url="https://upstream.example.com/mcp",
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
authorization_url="https://upstream.example.com/authorize",
token_url="https://upstream.example.com/token",
approval_status="active",
)
manager: Final = MCPServerManager()
prisma: Final = MagicMock()
async def persisted_rows(*, where: Mapping[str, object]) -> list[LiteLLM_MCPServerTable]:
return [] if where.get("approval_status") == "draft" else [row]
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=persisted_rows)
prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=None)
request: Final = Request(
{"type": "http", "method": "GET", "scheme": "http", "server": ("localhost", 4000),
"path": "/v1/mcp/server/oauth/saved-server/authorize", "headers": [],
"query_string": b"", "client": ("127.0.0.1", 1234)}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.master_key", "sk-unit-test-catalog"),
patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
):
response = await mcp_authorize(
request=request,
server_id=row.server_id,
user_api_key_dict=generate_mock_user_api_key_auth(),
client_id="client-id",
redirect_uri="http://localhost:9876/callback",
state="saved-server-test",
code_challenge=None,
code_challenge_method=None,
response_type="code",
scope=None,
)
assert response.status_code == 307
assert response.headers["location"].startswith("https://upstream.example.com/authorize?")
@pytest.mark.asyncio
async def test_mcp_authorize_proxies_to_discoverable_endpoint(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
@ -2754,8 +2818,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
@ -2776,7 +2840,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is authorize_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
authorize_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -2804,8 +2868,8 @@ class TestTemporaryMCPSessionEndpoints:
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
patches = [
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
@ -2909,8 +2973,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client",
@ -3055,8 +3119,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3114,8 +3178,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3168,8 +3232,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3218,8 +3282,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3266,8 +3330,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
@ -3306,8 +3370,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3351,8 +3415,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3374,7 +3438,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -3405,8 +3469,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3428,7 +3492,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -3465,8 +3529,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
@ -3484,7 +3548,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is register_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
read_body.assert_awaited_once_with(request=request)
register_mock.assert_awaited_once_with(
request=request,
@ -3528,8 +3592,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
@ -3573,8 +3637,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
@ -7940,3 +8004,33 @@ async def test_config_server_edit_preserves_api_contract_without_creating_rows(r
prisma.tx.assert_not_called()
assert server.model_dump() == original
assert manager.registry == {}
@pytest.mark.asyncio
async def test_saved_server_authorize_denial_does_not_dispatch_upstream():
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager: Final = MCPServerManager()
row: Final = LiteLLM_MCPServerTable(server_id="denied-peer", alias="denied_peer", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, url="https://upstream.example/mcp",
authorization_url="https://upstream.example/authorize", token_url="https://upstream.example/token")
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
user: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.INTERNAL_USER)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
patch.object(mgmt_endpoints, "get_cached_temporary_mcp_server", AsyncMock(return_value=None)),
patch.object(mgmt_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=(user,))),
patch.object(manager, "get_allowed_mcp_servers", AsyncMock(return_value=[])),
patch.object(mgmt_endpoints, "authorize_with_server", new_callable=AsyncMock) as authorize,
patch.object(mgmt_endpoints, "resolve_ephemeral_dcr_client", new_callable=AsyncMock) as register,
):
with pytest.raises(HTTPException) as exc:
await mgmt_endpoints.mcp_authorize(request=None, server_id=row.server_id, user_api_key_dict=user,
client_id="client", redirect_uri="http://localhost/callback")
assert exc.value.status_code == 403
authorize.assert_not_awaited()
register.assert_not_awaited()
assert prisma.db.litellm_mcpservertable.find_many.await_count == 1