fix(mcp): refresh shared catalog state for each operation

This commit is contained in:
Joshua Valluru 2026-09-22 12:40:45 -07:00
parent 8a9305fa27
commit 74a470410d
28 changed files with 1288 additions and 490 deletions

View file

@ -13,6 +13,7 @@ from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
@ -398,6 +399,7 @@ class MCPRequestHandler:
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
@staticmethod
@catalog_operation(global_manager)
async def process_mcp_request(
scope: Scope,
) -> tuple[

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 catalog_operation, global_manager
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)
@catalog_operation(global_manager)
async def byok_authorize_get(
request: Request,
client_id: str | None = None,

View file

@ -0,0 +1,394 @@
"""Authoritative MCP catalog snapshots shared by legacy lookup adapters."""
from __future__ import annotations
import asyncio
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 functools import wraps
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 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 = ParamSpec("_P")
_R = TypeVar("_R")
@dataclass(frozen=True, slots=True)
class CatalogSnapshot:
servers: Mapping[str, MCPServer]
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._refresh_lock = asyncio.Lock()
self._database_identity = ""
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] | None] = ContextVar(
"mcp_catalog_snapshot", default=None
)
self._staged_routing: ContextVar[tuple[dict[str, str], asyncio.Event] | None] = ContextVar(
"mcp_catalog_routing", default=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
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
def registry(self) -> dict[str, MCPServer]:
snapshot: Final = self.current()
return (
dict(snapshot.servers) if snapshot is not None else self.manager.config_mcp_servers | self.manager.registry
)
@asynccontextmanager
async def operation(self) -> AsyncIterator[CatalogSnapshot]:
current: Final = self.current()
if current is not None:
yield current
return
from litellm.proxy.proxy_server import prisma_client
if prisma_client is not None:
from litellm.proxy.proxy_server import should_load_db_object
if should_load_db_object("mcp"):
await self.reload()
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
snapshot: Final = _snapshot(self.manager, self._database_identity)
closed: Final = asyncio.Event()
token: Final = self._operation.set((snapshot, closed))
try:
with global_mcp_tool_registry.catalog_scope(snapshot.tools):
yield snapshot
finally:
closed.set()
self._operation.reset(token)
self._retain_discovered_routing(snapshot)
def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None:
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 snapshot.servers.items()
if current.get(key) == server
for owner in self.manager.owned_mapping_values(server)
)
self.manager.published_tool_routes = self.manager.published_tool_routes | MappingProxyType(
{
name: owner
for name, owner in snapshot.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_id(identifier, client_ip) or self.manager.get_mcp_server_by_name(
identifier, 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
)
_check_oauth_revision(selected, registered)
resolved: Final = await resolve(selected)
_check_oauth_revision(selected, resolved)
return resolved
@staticmethod
async def 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:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
async with self._refresh_lock:
staged_config: Final = {
key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()
}
await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values()))
staged_routing: Final = dict(self.manager.published_tool_routes)
closed: Final = asyncio.Event()
routing_token: Final = self._staged_routing.set((staged_routing, closed))
try:
with global_mcp_tool_registry.catalog_scope(global_mcp_tool_registry.published_tools) as staged_tools:
await self._reload()
self.manager.config_mcp_servers = staged_config
global_mcp_tool_registry.tools = staged_tools
self.manager.published_tool_routes = staged_routing
finally:
closed.set()
self._staged_routing.reset(routing_token)
async def _reload(self) -> 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 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

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,6 +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 catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
CredentialSource,
@ -1769,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:
@ -1900,6 +1901,7 @@ async def authorize_mcp_session(
@router.get("/{mcp_server_name}/authorize")
@router.get("/authorize")
@catalog_operation(global_manager)
async def authorize(
request: Request,
redirect_uri: str,
@ -1974,6 +1976,7 @@ async def authorize(
@router.post("/{mcp_server_name}/token")
@router.post("/token")
@catalog_operation(global_manager)
async def token_endpoint(
request: Request,
grant_type: str = Form(...),
@ -2078,6 +2081,7 @@ async def authorize_flow(request: Request, flow: str) -> Response:
@router.post("/authorize/complete")
@catalog_operation(global_manager)
async def authorize_complete(
request: Request,
flow: str = Form(...),
@ -2696,6 +2700,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}}")
@catalog_operation(global_manager)
async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth protected resource discovery endpoint using standard MCP URL pattern.
@ -2716,6 +2721,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")
@catalog_operation(global_manager)
async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern.
@ -2792,6 +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}}")
@catalog_operation(global_manager)
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery endpoint using standard MCP URL pattern.
@ -2809,6 +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")
@catalog_operation(global_manager)
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth authorization server discovery endpoint.
@ -2882,6 +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")
@catalog_operation(global_manager)
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
@ -2895,6 +2904,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
@router.post("/{mcp_server_name}/register")
@router.post("/register")
@catalog_operation(global_manager)
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)

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

@ -178,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,
@ -287,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]
@ -541,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]:
@ -722,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
@ -781,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
@ -1427,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,
@ -1860,6 +1859,9 @@ 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)
self.config_mcp_servers: dict[str, MCPServer] = {}
@ -1886,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",
@ -1900,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:
@ -1919,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()
@ -1932,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()
@ -2008,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:
@ -2045,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)
@ -2056,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 (
@ -2103,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
@ -2119,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:
@ -2148,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
@ -2209,7 +2214,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")
@ -2291,7 +2296,15 @@ class MCPServerManager:
"""
Get the registered MCP Servers from the registry and union with the config MCP Servers
"""
return 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).
@ -2378,7 +2391,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,
@ -2584,10 +2597,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,
@ -2608,13 +2621,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."""
@ -2622,7 +2635,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(
@ -2787,17 +2800,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
@ -2808,13 +2822,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)
@ -2845,7 +2859,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)
@ -2936,6 +2950,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))
@ -3160,13 +3175,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)
@ -3194,10 +3210,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)
@ -3215,7 +3231,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:
@ -3227,14 +3243,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)
@ -4466,7 +4482,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)
@ -5256,7 +5274,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,
@ -6146,7 +6164,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
@ -6445,7 +6463,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
@ -6507,152 +6525,7 @@ class MCPServerManager:
return None
async def reload_servers_from_database(self):
"""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...")
self._upstream_initialize_instructions_by_server_id.clear()
self._upstream_initialize_instructions_probed_at.clear()
# 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.
await self._maybe_register_openapi_tools(new_server, initialize_mapping=False)
registered_registry[server_id] = new_server
if new_server.spec_path:
registered_openapi_tools = True
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._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._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,6 +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 TargetCatalog, catalog_operation
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -614,6 +614,7 @@ def apply_tool_overrides(
return tools
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_allowed_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_servers: Sequence[str] | None,
@ -930,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,
@ -1151,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.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")
@ -1435,6 +1430,7 @@ async def filter_tools_by_key_team_permissions(
]
@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,
@ -2271,6 +2267,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,
@ -3055,6 +3052,7 @@ class GatewayOperations:
@overload
async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ...
@catalog_operation(lambda: global_mcp_server_manager)
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case AuthorizedToolCall():

View file

@ -22,6 +22,7 @@ from litellm.exceptions import (
GuardrailRaisedException,
ModifyResponseException,
)
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPServerURLCredentialsError,
@ -861,6 +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)])
@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"),
@ -1085,6 +1087,7 @@ if MCP_AVAILABLE:
}
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
@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

@ -1559,6 +1559,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,6 +46,8 @@ from fastapi import (
from fastapi.responses import JSONResponse
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation
try:
from prisma.errors import RecordNotFoundError, UniqueViolationError
except ImportError:
@ -1061,7 +1064,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)
@ -1091,6 +1095,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]:
@ -1129,6 +1134,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]:
@ -1150,6 +1156,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
@ -1284,6 +1291,7 @@ if MCP_AVAILABLE:
description="Health check for MCP servers",
dependencies=[Depends(user_api_key_auth)],
)
@catalog_operation(lambda: global_mcp_server_manager)
async def health_check_servers(
server_ids: list[str] | None = Query(
None,
@ -1592,6 +1600,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,
@ -2016,9 +2025,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
@ -2064,42 +2071,44 @@ if MCP_AVAILABLE:
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(
@ -2119,47 +2128,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",
@ -2179,47 +2188,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",
@ -2231,22 +2240,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}",
@ -2597,6 +2606,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

@ -19677,30 +19677,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:
@ -19755,7 +19758,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,6 +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 catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.utils import (
iter_known_server_prefixes,
logging_safe_mcp_headers,
@ -111,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,
@ -229,6 +231,7 @@ class LiteLLM_Proxy_MCP_Handler:
return user_api_key_auth
@staticmethod
@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,
@ -680,6 +683,7 @@ class LiteLLM_Proxy_MCP_Handler:
return result_text or "Tool executed successfully"
@staticmethod
@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

@ -7485,6 +7485,7 @@ class TestGatewaySessionAdmission:
with (
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
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
@ -9575,6 +9576,7 @@ class TestScopedSessionAdmission:
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", 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_mcpusercredentials.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)

View file

@ -9060,7 +9060,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()
@ -9085,7 +9085,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()
@ -9618,7 +9618,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
@ -9635,7 +9635,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
@ -11135,7 +11135,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

@ -5535,7 +5535,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
@ -5587,7 +5587,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()
@ -5659,7 +5659,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"),
@ -9534,7 +9534,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

@ -44,7 +44,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_deserialize_json_dict,
_flow_endpoints_missing,
_mcp_oauth_discovery_on_startup_enabled,
_oauth_endpoints_unresolved,
oauth_endpoints_unresolved,
_deserialize_json_list,
_normalize_mcp_server_cost_info,
_obo_retry_applies,
@ -694,7 +694,7 @@ class TestMCPServerManager:
assert resolved[0].scopes == ["mcp.read"]
assert manager.config_mcp_servers[server.server_id] is resolved[0]
assert server.authorization_url is None
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
@pytest.mark.asyncio
async def test_table_oauth_discovery_can_be_deferred_until_first_use(self):
@ -720,7 +720,7 @@ class TestMCPServerManager:
server = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
discovery.assert_not_awaited()
assert manager._oauth_discovery_slot(server.server_id) is not None
assert manager.oauth_discovery_slot(server.server_id) is not None
manager.registry[server.server_id] = server
with patch.object(manager, "_descovery_metadata", new=discovery):
@ -777,7 +777,7 @@ class TestMCPServerManager:
assert len({id(resolution) for resolution in resolutions}) == 1
assert resolutions[0].authorization_url == "https://idp.example.com/authorize"
assert resolutions[0].token_url == "https://idp.example.com/token"
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
@pytest.mark.asyncio
async def test_lazy_oauth_discovery_timeout_is_bounded(self):
@ -809,7 +809,7 @@ class TestMCPServerManager:
assert exc.value.status_code == 503
assert "timed out" in str(exc.value.detail)
discovery.assert_awaited_once_with(server)
assert manager._oauth_discovery_slot(server.server_id) is not None
assert manager.oauth_discovery_slot(server.server_id) is not None
@pytest.mark.asyncio
async def test_cancelling_one_waiter_does_not_cancel_shared_discovery(self):
@ -890,7 +890,7 @@ class TestMCPServerManager:
assert replacement.token_url is None
assert resolved.authorization_url == "https://idp.example.com/authorize"
assert resolved.token_url == "https://idp.example.com/token"
assert manager._oauth_discovery_slot(replacement.server_id) is None
assert manager.oauth_discovery_slot(replacement.server_id) is None
def test_registry_swap_reconcile_keeps_slot_for_issuer_anchored_server_without_url(self):
manager = MCPServerManager()
@ -907,9 +907,9 @@ class TestMCPServerManager:
manager.registry[server.server_id] = server
manager._set_oauth_discovery_deferred(server.server_id, True)
manager._reconcile_oauth_discovery_slots_for_servers([server])
manager.reconcile_oauth_discovery_slots_for_servers([server])
assert manager._oauth_discovery_slot(server.server_id) is not None
assert manager.oauth_discovery_slot(server.server_id) is not None
resolved = server.model_copy(
update={
@ -918,12 +918,12 @@ class TestMCPServerManager:
}
)
manager.registry[resolved.server_id] = resolved
manager._reconcile_oauth_discovery_slots_for_servers([resolved])
manager.reconcile_oauth_discovery_slots_for_servers([resolved])
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
def _assert_oauth_discovery_state_removed(self, manager, server_id):
assert manager._oauth_discovery_slot(server_id) is None
assert manager.oauth_discovery_slot(server_id) is None
@pytest.mark.asyncio
async def test_deactivated_server_clears_lazy_oauth_discovery_state(self):
@ -965,7 +965,7 @@ class TestMCPServerManager:
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repository,
),
patch(
@ -998,7 +998,7 @@ class TestMCPServerManager:
manager.registry[server.server_id] = server
previous_registry = manager.registry
manager._set_oauth_discovery_deferred(server.server_id, True)
old_generation = manager._oauth_discovery_slot(server.server_id).generation
old_generation = manager.oauth_discovery_slot(server.server_id).generation
resolved = server.model_copy(
update={
"authorization_url": "https://idp.example.com/authorize",
@ -1017,7 +1017,8 @@ class TestMCPServerManager:
raw_row = MagicMock()
raw_row.model_dump.return_value = row.model_dump()
repository = MagicMock()
repository.table.find_many = AsyncMock(return_value=[raw_row])
other_row: Final = row.model_copy(update={"server_id": "changed-other", "server_name": "changed_other", "auth_type": MCPAuth.none})
repository.table.find_many = AsyncMock(return_value=[raw_row, other_row])
async def publish_while_staged(*_args, **_kwargs):
assert manager.registry is previous_registry
@ -1025,7 +1026,7 @@ class TestMCPServerManager:
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repository,
),
patch(
@ -1034,16 +1035,16 @@ class TestMCPServerManager:
),
patch.object(
manager,
"_maybe_register_openapi_tools",
"maybe_register_openapi_tools",
new=AsyncMock(side_effect=publish_while_staged),
),
patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"),
patch.object(manager, "prime_oauth_metadata_discovery_for_servers"),
):
await manager.reload_servers_from_database()
assert previous_registry[server.server_id] is resolved
assert manager.registry[server.server_id] is server
retry_slot = manager._oauth_discovery_slot(server.server_id)
retry_slot = manager.oauth_discovery_slot(server.server_id)
assert retry_slot is not None
assert retry_slot.generation > old_generation
@ -1143,7 +1144,7 @@ class TestMCPServerManager:
assert manager.config_mcp_servers[server.server_id].authorization_url == "https://idp.example.com/authorize"
assert manager.config_mcp_servers[server.server_id].token_url is None
assert manager.config_mcp_servers[server.server_id].scopes is None
assert manager._oauth_discovery_slot(server.server_id) is not None
assert manager.oauth_discovery_slot(server.server_id) is not None
@pytest.mark.asyncio
async def test_create_mcp_client_triggers_deferred_oauth_discovery(self):
@ -1370,7 +1371,7 @@ class TestMCPServerManager:
manager = MCPServerManager()
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(self._id_jag_config())
@ -1386,7 +1387,7 @@ class TestMCPServerManager:
manager = MCPServerManager()
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(self._id_jag_config())
@ -1408,7 +1409,7 @@ class TestMCPServerManager:
}
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(config)
@ -1422,7 +1423,7 @@ class TestMCPServerManager:
manager = MCPServerManager()
with (
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
await manager.load_servers_from_config(self._id_jag_config())
@ -7377,7 +7378,7 @@ class TestMCPServerTimestamps:
with (
patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
@ -7463,10 +7464,10 @@ class TestMCPServerTimestamps:
client_id="cid",
client_secret="csec",
)
assert _oauth_endpoints_unresolved(m2m_shaped) is False
assert oauth_endpoints_unresolved(m2m_shaped) is False
interactive_unresolved = m2m_shaped.model_copy(update={"client_id": None, "client_secret": None})
assert _oauth_endpoints_unresolved(interactive_unresolved) is True
assert oauth_endpoints_unresolved(interactive_unresolved) is True
def test_dcr_bridge_relay_arm_needs_its_registration_endpoint(self):
"""A dcr_bridge server with no admin-configured client can only register callers through the
@ -7486,11 +7487,11 @@ class TestMCPServerTimestamps:
token_url="https://idp.example.com/token",
registration_url=None,
)
assert _oauth_endpoints_unresolved(relay_arm) is True
assert oauth_endpoints_unresolved(relay_arm) is True
assert (
_oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False
oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False
)
assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False
assert oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False
def test_entra_obo_without_scopes_is_unresolved(self):
"""entra_obo token exchange fails closed without a scope, and scopes can come from resource
@ -7506,9 +7507,9 @@ class TestMCPServerTimestamps:
token_url="https://idp.example.com/token",
scopes=None,
)
assert _oauth_endpoints_unresolved(entra) is True
assert _oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False
assert _oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False
assert oauth_endpoints_unresolved(entra) is True
assert oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False
assert oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False
@pytest.mark.asyncio
async def test_reload_fast_path_retries_unresolved_oauth_servers(self):
@ -7553,7 +7554,7 @@ class TestMCPServerTimestamps:
build_mock = AsyncMock(return_value=previous_entry)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
),
patch(
@ -7646,7 +7647,7 @@ class TestMCPServerTimestamps:
def test_carry_forward_skips_when_url_or_auth_type_changed(self):
"""Stale endpoints from a different upstream or auth mode must not carry forward."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
carry_forward_resolved_oauth_endpoints,
)
def make_server(url: str, auth_type: MCPAuth, authorization_url: Optional[str]) -> MCPServer:
@ -7662,19 +7663,19 @@ class TestMCPServerTimestamps:
previous = make_server("https://old.example.com/mcp", MCPAuth.oauth2, "https://idp.example.com/authorize")
url_changed = make_server("https://new.example.com/mcp", MCPAuth.oauth2, None)
_carry_forward_resolved_oauth_endpoints(new_server=url_changed, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=url_changed, previous_server=previous)
assert url_changed.authorization_url is None
auth_changed = make_server("https://old.example.com/mcp", MCPAuth.true_passthrough, None)
_carry_forward_resolved_oauth_endpoints(new_server=auth_changed, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=auth_changed, previous_server=previous)
assert auth_changed.authorization_url is None
same = make_server("https://old.example.com/mcp", MCPAuth.oauth2, None)
_carry_forward_resolved_oauth_endpoints(new_server=same, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=same, previous_server=previous)
assert same.authorization_url == "https://idp.example.com/authorize"
explicit = make_server("https://old.example.com/mcp", MCPAuth.oauth2, "https://configured.example.com/auth")
_carry_forward_resolved_oauth_endpoints(new_server=explicit, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=explicit, previous_server=previous)
assert explicit.authorization_url == "https://configured.example.com/auth"
def test_carry_forward_does_not_revive_token_url_across_authorization_url_change(self):
@ -7685,7 +7686,7 @@ class TestMCPServerTimestamps:
endpoint recreates the RFC 9700 mix-up, durably, and the discovery gate alone cannot catch
it because the stale endpoint comes from the registry, not from discovery."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
carry_forward_resolved_oauth_endpoints,
)
previous = MCPServer(
@ -7707,7 +7708,7 @@ class TestMCPServerTimestamps:
authorization_url="https://idp-b.example.com/authorize",
)
_carry_forward_resolved_oauth_endpoints(new_server=repointed, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=repointed, previous_server=previous)
assert repointed.authorization_url == "https://idp-b.example.com/authorize"
assert repointed.token_url is None
@ -7719,7 +7720,7 @@ class TestMCPServerTimestamps:
consistent group, and a rebuild that re-pins the same authorize endpoint (formatting aside)
keeps carrying the corroborated token endpoint."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
carry_forward_resolved_oauth_endpoints,
)
def previous() -> MCPServer:
@ -7742,7 +7743,7 @@ class TestMCPServerTimestamps:
auth_type=MCPAuth.oauth2,
authorization_url=None,
)
_carry_forward_resolved_oauth_endpoints(new_server=blipped, previous_server=previous())
carry_forward_resolved_oauth_endpoints(new_server=blipped, previous_server=previous())
assert blipped.authorization_url == "https://idp.example.com/authorize"
assert blipped.token_url == "https://idp.example.com/token"
assert blipped.registration_url == "https://idp.example.com/register"
@ -7755,7 +7756,7 @@ class TestMCPServerTimestamps:
auth_type=MCPAuth.oauth2,
authorization_url="https://IDP.example.com:443/authorize/",
)
_carry_forward_resolved_oauth_endpoints(new_server=same_authorize, previous_server=previous())
carry_forward_resolved_oauth_endpoints(new_server=same_authorize, previous_server=previous())
assert same_authorize.token_url == "https://idp.example.com/token"
assert same_authorize.registration_url == "https://idp.example.com/register"
@ -7767,7 +7768,7 @@ class TestMCPServerTimestamps:
scopes still carry as last-known-good. Anchoring is keyed on the explicit issuer_is_anchored
flag, not on issuer truthiness, so a discovered issuer does not trip this fail-closed branch."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
carry_forward_resolved_oauth_endpoints,
)
previous = MCPServer(
@ -7793,7 +7794,7 @@ class TestMCPServerTimestamps:
issuer_is_anchored=True,
)
_carry_forward_resolved_oauth_endpoints(new_server=failed_rebuild, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=failed_rebuild, previous_server=previous)
assert failed_rebuild.authorization_url is None
assert failed_rebuild.token_url is None
@ -7807,7 +7808,7 @@ class TestMCPServerTimestamps:
regression the explicit issuer_is_anchored flag prevents: keying fail-closed on issuer truthiness
alone would drop the working endpoints the moment the server learned its issuer."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_carry_forward_resolved_oauth_endpoints,
carry_forward_resolved_oauth_endpoints,
)
previous = MCPServer(
@ -7834,7 +7835,7 @@ class TestMCPServerTimestamps:
authorization_url=None,
)
_carry_forward_resolved_oauth_endpoints(new_server=blipped_rebuild, previous_server=previous)
carry_forward_resolved_oauth_endpoints(new_server=blipped_rebuild, previous_server=previous)
assert blipped_rebuild.authorization_url == "https://idp.example.com/authorize"
assert blipped_rebuild.token_url == "https://idp.example.com/token"
@ -11975,7 +11976,7 @@ class TestClientForwardedDiscoveryFailureIsNotFatal:
assert resolved is manager.config_mcp_servers[server.server_id]
assert resolved.authorization_url is None
assert resolved.token_url is None
assert manager._oauth_discovery_slot(server.server_id) is not None
assert manager.oauth_discovery_slot(server.server_id) is not None
@pytest.mark.parametrize(
"auth_type, serves_the_listing",
@ -12037,7 +12038,7 @@ class TestClientForwardedDiscoveryFailureIsNotFatal:
assert resolved.token_url == "https://idp.example.com/token"
assert resolved.registration_url == "https://idp.example.com/register"
assert manager.config_mcp_servers[server.server_id].authorization_url == "https://idp.example.com/authorize"
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
class TestResolveOpenapiToolAuth:
@ -12395,7 +12396,7 @@ class TestConfigServerIdPinning:
)
with (
patch( # test-quality-ok: the db reload path has no seam but its own repository
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository",
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repository,
),
patch( # test-quality-ok: same, the prisma client is fetched inside the reload
@ -13316,7 +13317,7 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None:
transport=MCPTransport.http, auth_type=MCPAuth.oauth2,
)
manager._set_oauth_discovery_deferred(original.server_id, True)
original_slot: Final = manager._oauth_discovery_slot(original.server_id)
original_slot: Final = manager.oauth_discovery_slot(original.server_id)
assert original_slot is not None
replacement: Final = original.model_copy(update={"url": "https://new.example.com/mcp"})
manager.registry[original.server_id] = replacement
@ -13335,30 +13336,30 @@ async def test_temporary_oauth_discovery_expires_without_more_requests() -> None
)
manager._set_oauth_discovery_deferred(server.server_id, True)
resolved: Final = await manager.ensure_oauth_metadata_discovered(server)
assert manager._oauth_discovery_slot(server.server_id) is not None
assert manager.oauth_discovery_slot(server.server_id) is not None
loop: Final = asyncio.get_running_loop()
expired: Final = loop.create_future()
with patch.object(loop, "time", return_value=loop.time() + 301):
loop.call_later(0, expired.set_result, None)
await expired
assert resolved.authorization_url == server.authorization_url
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
def test_old_temporary_discovery_expiry_preserves_replacement() -> None:
manager: Final = MCPServerManager()
manager._set_oauth_discovery_deferred("reused-session", True)
old_slot: Final = manager._oauth_discovery_slot("reused-session")
old_slot: Final = manager.oauth_discovery_slot("reused-session")
assert old_slot is not None
manager._set_oauth_discovery_deferred("reused-session", True)
replacement: Final = manager._oauth_discovery_slot("reused-session")
replacement: Final = manager.oauth_discovery_slot("reused-session")
manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation)
assert manager._oauth_discovery_slot("reused-session") is replacement
assert manager.oauth_discovery_slot("reused-session") is replacement
assert replacement is not None
manager._expire_temporary_oauth_discovery("reused-session", replacement.generation)
assert manager._oauth_discovery_slot("reused-session") is None
assert manager.oauth_discovery_slot("reused-session") is None
manager._expire_temporary_oauth_discovery("reused-session", replacement.generation)
assert manager._oauth_discovery_slot("reused-session") is None
assert manager.oauth_discovery_slot("reused-session") is None
@pytest.mark.asyncio
@ -13719,12 +13720,12 @@ async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old
with _mcp_upstream(upstream.respond):
task: Final = asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None))
await asyncio.wait_for(upstream.entered.wait(), timeout=5)
manager._invalidate_discovery_lists("discovery")
manager.invalidate_discovery_lists("discovery")
upstream.release.set()
assert (await task)[0].name == "discovery-example"
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
assert upstream.initializes == 2
manager._invalidate_discovery_lists("discovery")
manager.invalidate_discovery_lists("discovery")
assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1
assert upstream.initializes == 3
@ -14516,3 +14517,331 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
assert captured["client_ip"] is None
finally:
auth_context_var.reset(token)
@pytest.mark.asyncio
async def test_catalog_observes_committed_update_and_delete_without_background_reload():
from datetime import timedelta
timestamp: Final = datetime.now()
row: Final = LiteLLM_MCPServerTable(
server_id="catalog-server", server_name="catalog_server", alias="catalog_server",
transport=MCPTransport.http, url="https://first.example.com/mcp", updated_at=timestamp,
)
updated: Final = row.model_copy(update={"url": "https://second.example.com/mcp", "updated_at": timestamp + timedelta(seconds=1)})
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], []))
manager: Final = MCPServerManager()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
first = await manager.catalog.resolve(row.server_id)
manager.registry[row.server_id].short_prefix = "a12"
second = await manager.catalog.resolve(row.server_id)
deleted = await manager.catalog.resolve(row.server_id)
assert first is not None and first.url == row.url
assert second is not None and second.url == updated.url
assert second.short_prefix == "a12"
assert deleted is None
assert prisma.db.litellm_mcpservertable.find_many.await_count == 3
@pytest.mark.asyncio
async def test_catalog_lookup_uses_one_snapshot_until_operation_finishes():
from datetime import timedelta
row: Final = LiteLLM_MCPServerTable(
server_id="snapshot-server", alias="snapshot_server", transport=MCPTransport.http,
url="https://first.example.com/mcp", updated_at=datetime.now(),
)
updated: Final = row.model_copy(update={"url": "https://second.example.com/mcp", "updated_at": row.updated_at + timedelta(seconds=1)})
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], [updated]))
manager: Final = MCPServerManager()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
async with manager.catalog.operation() as before:
await manager.reload_servers_from_database()
during = await manager.catalog.resolve(row.server_id)
assert during is not None and during.url == row.url
assert manager.registry[row.server_id].url == updated.url
async with manager.catalog.operation() as after:
current = manager.get_mcp_server_by_id(row.server_id)
assert current is not None and current.url == updated.url
assert before.identity != after.identity
assert prisma.db.litellm_mcpservertable.find_many.await_count == 3
@pytest.mark.asyncio
async def test_catalog_failed_reload_preserves_published_discovery_state():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="healthy", name="healthy", transport=MCPTransport.http)
manager.registry = {server.server_id: server}
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions"
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("database unavailable"))
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(RuntimeError, match="database unavailable"):
await manager.reload_servers_from_database()
assert manager.get_mcp_server_by_id(server.server_id) is server
assert manager._upstream_initialize_instructions_by_server_id == {server.server_id: "healthy instructions"}
@pytest.mark.asyncio
async def test_catalog_cancellation_retains_state_and_releases_refresh_lock():
entered: Final = asyncio.Event()
release: Final = asyncio.Event()
async def blocked_read(**kwargs: object) -> list[LiteLLM_MCPServerTable]:
entered.set()
await release.wait()
return []
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="healthy", name="healthy", transport=MCPTransport.http)
manager.registry = {server.server_id: server}
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions"
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=blocked_read)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
pending = asyncio.create_task(manager.reload_servers_from_database())
await asyncio.wait_for(entered.wait(), timeout=2)
pending.cancel()
with pytest.raises(asyncio.CancelledError):
await pending
assert manager.registry == {server.server_id: server}
assert manager._upstream_initialize_instructions_by_server_id == {server.server_id: "healthy instructions"}
release.set()
await asyncio.wait_for(manager.reload_servers_from_database(), timeout=2)
assert manager.registry == {}
@pytest.mark.asyncio
async def test_catalog_background_lookup_after_operation_exit_observes_deletion():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="expired-snapshot", name="expired_snapshot", transport=MCPTransport.http)
manager.registry = {server.server_id: server}
release: Final = asyncio.Event()
async def lookup_after_exit() -> MCPServer | None:
await release.wait()
return await manager.catalog.resolve(server.server_id)
with patch("litellm.proxy.proxy_server.prisma_client", None):
async with manager.catalog.operation():
pending: Final = asyncio.create_task(lookup_after_exit())
manager.registry = {}
release.set()
assert await pending is None
@pytest.mark.asyncio
async def test_catalog_cancelled_openapi_refresh_retains_tools_and_discovery(monkeypatch):
from datetime import timedelta
from litellm.proxy._experimental.mcp_server import tool_registry
manager: Final = MCPServerManager()
stamp: Final = datetime.now()
server: Final = MCPServer(server_id="staged", name="staged", transport=MCPTransport.http,
url="https://before.example/mcp", spec_path="before.json", updated_at=stamp, auth_type=MCPAuth.oauth2)
manager.registry = {server.server_id: server}
manager._set_oauth_discovery_deferred(server.server_id, True)
original_slot: Final = manager.oauth_discovery_slot(server.server_id)
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "keep instructions"
manager.tool_name_to_mcp_server_name_mapping = {"staged-existing": "staged"}
registry: Final = tool_registry.MCPToolRegistry()
registry.register_tool("staged-existing", "existing", {}, lambda: "existing")
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
row: Final = LiteLLM_MCPServerTable(server_id=server.server_id, alias="staged", transport=MCPTransport.http,
url="https://after.example/mcp", spec_path="after.json", updated_at=stamp + timedelta(seconds=1), auth_type=MCPAuth.oauth2)
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
async def cancelled_registration(*args: object, **kwargs: object) -> None:
registry.register_tool("staged-new", "new", {}, lambda: "new")
raise asyncio.CancelledError()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch.object(manager, "maybe_register_openapi_tools", side_effect=cancelled_registration),
patch.object(manager, "invalidate_discovery_lists") as invalidate,
):
with pytest.raises(asyncio.CancelledError):
await manager.reload_servers_from_database()
invalidate.assert_not_called()
assert manager.registry == {server.server_id: server}
assert [tool.name for tool in registry.list_tools()] == ["staged-existing"]
assert manager.tool_name_to_mcp_server_name_mapping == {"staged-existing": "staged"}
assert manager._upstream_initialize_instructions_by_server_id == {server.server_id: "keep instructions"}
assert manager.oauth_discovery_slot(server.server_id) is original_slot
@pytest.mark.asyncio
async def test_catalog_snapshot_identity_is_independent_of_worker_oauth_discovery():
row: Final = LiteLLM_MCPServerTable(server_id="identity-server", alias="identity_server", transport=MCPTransport.http,
url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2, updated_at=datetime.now())
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
manager: Final = MCPServerManager()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch.object(manager, "prime_oauth_metadata_discovery_for_servers"),
):
async with manager.catalog.operation() as unresolved:
assert manager.get_mcp_server_by_id(row.server_id).authorization_url is None
resolved: Final = manager.registry[row.server_id].model_copy(update={
"authorization_url": "https://idp.example/authorize", "token_url": "https://idp.example/token"})
manager.registry[row.server_id] = resolved
async with manager.catalog.operation() as discovered:
assert manager.get_mcp_server_by_id(row.server_id).authorization_url == resolved.authorization_url
assert unresolved.identity == discovered.identity
@pytest.mark.asyncio
async def test_catalog_failed_openapi_row_does_not_publish_partial_handlers(monkeypatch):
from litellm.proxy._experimental.mcp_server import tool_registry
manager: Final = MCPServerManager()
registry: Final = tool_registry.MCPToolRegistry()
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
row: Final = LiteLLM_MCPServerTable(server_id="broken", alias="broken", transport=MCPTransport.http,
url="https://upstream.example/mcp", spec_path="broken.json")
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
async def broken_registration(server: MCPServer, **kwargs: object) -> None:
registry.register_tool("broken-partial", "partial", {}, lambda: "must not run")
manager.tool_name_to_mcp_server_name_mapping["broken-partial"] = "broken"
raise ValueError("invalid remaining operation")
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch.object(manager, "maybe_register_openapi_tools", side_effect=broken_registration),
):
await manager.reload_servers_from_database()
assert manager.registry == {}
assert registry.list_tools() == []
assert manager.tool_name_to_mcp_server_name_mapping == {}
@pytest.mark.asyncio
async def test_catalog_cancelled_config_hydration_preserves_published_credentials():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="config-hydration", name="config_hydration", transport=MCPTransport.http,
client_id="previous-client")
manager.config_mcp_servers = {server.server_id: server}
async def cancelled_hydration(target: MCPServer) -> bool:
target.client_id = "unpublished-client"
raise asyncio.CancelledError()
with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.hydrate_config_server_dcr_client", side_effect=cancelled_hydration):
with pytest.raises(asyncio.CancelledError):
await manager.reload_servers_from_database()
assert manager.get_mcp_server_by_id(server.server_id).client_id == "previous-client"
@pytest.mark.asyncio
async def test_catalog_oauth_resolution_cannot_replace_an_operation_target_after_update():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="pinned-oauth", name="pinned_oauth", transport=MCPTransport.http,
url="https://before.example/mcp", auth_type=MCPAuth.oauth2,
authorization_url="https://before.example/authorize", token_url="https://before.example/token")
manager.registry = {server.server_id: server}
with patch("litellm.proxy.proxy_server.prisma_client", None):
async with manager.catalog.operation():
selected: Final = manager.get_mcp_server_by_id(server.server_id)
manager.registry[server.server_id] = server.model_copy(update={
"url": "https://after.example/mcp", "authorization_url": "https://after.example/authorize"})
resolved: Final = await manager.ensure_oauth_metadata_discovered(selected)
assert resolved.url == "https://before.example/mcp"
assert resolved.authorization_url == "https://before.example/authorize"
assert manager.registry[server.server_id].url == "https://after.example/mcp"
@pytest.mark.asyncio
async def test_catalog_unresolved_oauth_snapshot_fails_closed_if_target_was_replaced():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="replaced-oauth", name="replaced_oauth", transport=MCPTransport.http,
url="https://before.example/mcp", auth_type=MCPAuth.oauth2)
manager.registry = {server.server_id: server}
with (
patch("litellm.proxy.proxy_server.prisma_client", None),
patch.object(manager, "_discover_oauth_metadata_for_server", new_callable=AsyncMock) as discovery,
):
async with manager.catalog.operation():
selected: Final = manager.get_mcp_server_by_id(server.server_id)
manager.registry[server.server_id] = server.model_copy(update={
"url": "https://after.example/mcp", "authorization_url": "https://after.example/authorize",
"token_url": "https://after.example/token"})
with pytest.raises(HTTPException) as exc:
await manager.ensure_oauth_metadata_discovered(selected)
assert exc.value.status_code == 503
discovery.assert_not_awaited()
@pytest.mark.asyncio
async def test_catalog_unresolved_oauth_snapshot_accepts_discovery_for_the_same_target():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="same-oauth", name="same_oauth", transport=MCPTransport.http,
url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2)
manager.registry = {server.server_id: server}
discovered: Final = server.model_copy(update={"authorization_url": "https://issuer.example/authorize",
"token_url": "https://issuer.example/token", "scopes": ["read"]})
with (
patch("litellm.proxy.proxy_server.prisma_client", None),
patch.object(manager, "_ensure_oauth_metadata_discovered", return_value=discovered) as discovery,
):
async with manager.catalog.operation():
resolved: Final = await manager.ensure_oauth_metadata_discovered(server)
assert resolved.authorization_url == "https://issuer.example/authorize"
assert resolved.token_url == "https://issuer.example/token"
assert resolved.scopes == ["read"]
discovery.assert_awaited_once()
@pytest.mark.asyncio
async def test_catalog_fresh_lookup_does_not_fall_back_to_stale_grants_when_database_fails():
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="stale-grant", name="stale_grant", transport=MCPTransport.http,
allow_all_keys=True)
manager.registry = {server.server_id: server}
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("unavailable"))
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
):
with pytest.raises(RuntimeError, match="unavailable"):
await manager.catalog.resolve(server.server_id)
assert manager.registry == {server.server_id: server}
@pytest.mark.asyncio
async def test_catalog_list_failure_cancels_and_joins_other_upstream_fetches():
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
started: Final = asyncio.Event()
stopped: Final = asyncio.Event()
pending: Final = MCPServer(server_id="pending", name="pending", transport=MCPTransport.http)
broken: Final = MCPServer(server_id="broken", name="broken", transport=MCPTransport.http)
async def fetch(server: MCPServer):
if server.server_id == broken.server_id:
await started.wait()
raise RuntimeError("unexpected fetch failure")
started.set()
try:
await asyncio.Event().wait()
finally:
stopped.set()
with pytest.raises(RuntimeError, match="unexpected fetch failure"):
await TargetCatalog.list((pending, broken), fetch, lambda server: server.server_id)
assert stopped.is_set()

View file

@ -1559,9 +1559,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
@ -2629,6 +2630,7 @@ 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
mock_manager.catalog.resolve = AsyncMock(return_value=non_oauth_server)
fake_proxy_server = types.SimpleNamespace(master_key=None)
with (
@ -2681,6 +2683,7 @@ 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
mock_manager.catalog.resolve = AsyncMock(return_value=internal_server)
fake_proxy_server = types.SimpleNamespace(master_key=None)
with (
@ -2738,6 +2741,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 +2814,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 +2836,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 +2864,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 +2969,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 +3115,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 +3174,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 +3228,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 +3278,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 +3326,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 +3366,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 +3411,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 +3434,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 +3465,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 +3488,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 +3525,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 +3544,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 +3588,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 +3633,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 +8000,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

View file

@ -1,3 +1,5 @@
from contextlib import nullcontext
import subprocess
import sys
import textwrap
@ -32,6 +34,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(return_value=_DummyMCPResult()),
# Newer logging path calls this to enrich spend logs metadata
@ -378,6 +381,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
post_call_failure_hook = _setup_proxy_logging(monkeypatch)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom"))
)
@ -506,6 +510,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
# Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields.
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
@ -559,6 +564,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
mock_get_tools,
)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),