mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): refresh shared catalog state for each operation
This commit is contained in:
parent
8a9305fa27
commit
74a470410d
28 changed files with 1288 additions and 490 deletions
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
394
litellm/proxy/_experimental/mcp_server/catalog.py
Normal file
394
litellm/proxy/_experimental/mcp_server/catalog.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=[]),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue