mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): keep worker MCP configurations consistent via a catalog revision (#42568)
* fix(mcp): refresh server catalog for each gateway operation * fix(mcp): reject configuration changes during scoped dispatch * fix(mcp): refresh shared catalog state for each operation * fix(mcp): refresh catalog before native alias routing * docs(mcp): clarify native route resolution order * test(mcp): align streaming fixture and generated API documentation * fix(mcp): preserve discovery published during catalog refresh * fix(mcp): coalesce queued catalog reads without a stale window * fix(mcp): reconcile concurrent route changes when publishing catalog * test(mcp): provide catalog scope in post-call hook fixtures * fix(mcp): retain valid live routes and handlers during refresh * fix(mcp): keep refreshed OpenAPI operation membership authoritative * test(mcp): preserve logging fixtures after catalog integration * fix(mcp): isolate transport test admission state and exhaust OAuth outcomes * refactor(mcp): narrow catalog consistency change to ticket scope Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(mcp): gate catalog refresh on a database revision marker Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): refresh waiters that observed a newer catalog revision Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): awaitable catalog revision doubles in prisma mocks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): satisfy catalog refresh lint and type gates * test(mcp): provide catalog scope in toolset fixtures * test(mcp): exercise temporary OAuth through catalog operations * fix(mcp): share active catalog snapshots across discovery tasks * fix(mcp): retain discovered routes across cached catalog operations * fix(mcp): preserve local tool ownership across discovery * refactor(mcp): satisfy tightened immutability lint ceiling Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): authorize local handlers by registered server ownership * fix(mcp): check registered handler freshness and isolate fixtures * fix(mcp): read catalog revisions and snapshots from the writer * fix(mcp): resolve access groups from the operation catalog * fix(mcp): preserve empty access group restrictions * fix(mcp): retain rediscovered routes across concurrent updates * fix(mcp): enforce registered ownership for local tool dispatch * fix(mcp): preserve issuer discovery across catalog snapshots --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
48c88bae04
commit
950da7c2ec
44 changed files with 3529 additions and 793 deletions
|
|
@ -0,0 +1,13 @@
|
|||
CREATE OR REPLACE FUNCTION litellm_bump_mcp_catalog_revision() RETURNS TRIGGER AS $$
|
||||
BEGIN
|
||||
INSERT INTO "LiteLLM_Config" ("param_name", "reload_revision")
|
||||
VALUES ('mcp_catalog', 1)
|
||||
ON CONFLICT ("param_name") DO UPDATE SET "reload_revision" = "LiteLLM_Config"."reload_revision" + 1;
|
||||
RETURN NULL;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
|
||||
DROP TRIGGER IF EXISTS "litellm_mcp_catalog_revision" ON "LiteLLM_MCPServerTable";
|
||||
CREATE TRIGGER "litellm_mcp_catalog_revision"
|
||||
AFTER INSERT OR UPDATE OR DELETE ON "LiteLLM_MCPServerTable"
|
||||
FOR EACH STATEMENT EXECUTE PROCEDURE litellm_bump_mcp_catalog_revision();
|
||||
|
|
@ -1,5 +1,6 @@
|
|||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
|
|
@ -15,6 +16,7 @@ import litellm
|
|||
from litellm._internal_context import with_service_target
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_ALL_TOOLS_WILDCARD
|
||||
from litellm.proxy._experimental.mcp_server.catalog import global_manager
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
get_passthrough_resource_metadata_url,
|
||||
get_passthrough_www_authenticate,
|
||||
|
|
@ -448,188 +450,191 @@ class MCPRequestHandler:
|
|||
Raises:
|
||||
HTTPException: If headers are invalid or missing required headers
|
||||
"""
|
||||
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
|
||||
async with global_manager().catalog.operation():
|
||||
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
|
||||
|
||||
# Check if there is an explicit LiteLLM API key (primary header)
|
||||
has_explicit_litellm_key: Final = headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) is not None
|
||||
# Check if there is an explicit LiteLLM API key (primary header)
|
||||
has_explicit_litellm_key: Final = (
|
||||
headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) is not None
|
||||
)
|
||||
|
||||
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
|
||||
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
|
||||
|
||||
# Get the old mcp_auth_header for backward compatibility
|
||||
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
|
||||
# Get the old mcp_auth_header for backward compatibility
|
||||
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
|
||||
|
||||
# Get the new server-specific auth headers
|
||||
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
|
||||
# Get the new server-specific auth headers
|
||||
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
|
||||
|
||||
# Get the oauth2 headers
|
||||
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
|
||||
# Get the oauth2 headers
|
||||
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
|
||||
|
||||
# Parse MCP servers from header
|
||||
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
|
||||
verbose_logger.debug("Raw MCP servers header: %s", mcp_servers_header)
|
||||
mcp_servers = None
|
||||
if mcp_servers_header is not None:
|
||||
try:
|
||||
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
|
||||
verbose_logger.debug("Parsed MCP servers: %s", mcp_servers)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error parsing mcp_servers header: %s", e)
|
||||
mcp_servers = None
|
||||
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
|
||||
mcp_servers = []
|
||||
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
|
||||
request: Final = Request(scope=scope)
|
||||
# Parse MCP servers from header
|
||||
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
|
||||
verbose_logger.debug("Raw MCP servers header: %s", mcp_servers_header)
|
||||
mcp_servers = None
|
||||
if mcp_servers_header is not None:
|
||||
try:
|
||||
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
|
||||
verbose_logger.debug("Parsed MCP servers: %s", mcp_servers)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error parsing mcp_servers header: %s", e)
|
||||
mcp_servers = None
|
||||
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
|
||||
mcp_servers = []
|
||||
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
|
||||
request: Final = Request(scope=scope)
|
||||
|
||||
async def mock_body():
|
||||
return b"{}"
|
||||
async def mock_body():
|
||||
return b"{}"
|
||||
|
||||
request.body = mock_body
|
||||
# Inline import — auth_utils participates in a proxy import cycle.
|
||||
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
|
||||
get_request_route,
|
||||
)
|
||||
request.body = mock_body
|
||||
# Inline import — auth_utils participates in a proxy import cycle.
|
||||
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
|
||||
get_request_route,
|
||||
)
|
||||
|
||||
request_route: Final = get_request_route(request)
|
||||
# Only OAuth metadata routes registered under /.well-known/ are public.
|
||||
if request_route.startswith("/.well-known/"):
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif (
|
||||
has_explicit_litellm_key
|
||||
and oauth2_headers
|
||||
and is_bridge_envelope_shaped(oauth2_headers["Authorization"])
|
||||
and (
|
||||
dual_bridge_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
request_route: Final = get_request_route(request)
|
||||
# Only OAuth metadata routes registered under /.well-known/ are public.
|
||||
if request_route.startswith("/.well-known/"):
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif (
|
||||
has_explicit_litellm_key
|
||||
and oauth2_headers
|
||||
and is_bridge_envelope_shaped(oauth2_headers["Authorization"])
|
||||
and (
|
||||
dual_bridge_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
)
|
||||
)
|
||||
is not None
|
||||
):
|
||||
(
|
||||
validated_user_api_key_auth,
|
||||
mcp_server_auth_headers,
|
||||
) = await MCPRequestHandler._admit_dcr_bridge_dual_credential(
|
||||
server=dual_bridge_target.server,
|
||||
requested_name=dual_bridge_target.requested_name,
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
litellm_api_key=litellm_api_key,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request=request,
|
||||
route=request_route,
|
||||
)
|
||||
elif has_explicit_litellm_key:
|
||||
# An explicit x-litellm-api-key is always a LiteLLM credential, even
|
||||
# for a delegated server, so validate it: identity / spend / rate
|
||||
# limits resolve and any stored upstream token can be forwarded.
|
||||
validated_user_api_key_auth = await user_api_key_auth(
|
||||
api_key=f"Bearer {_get_bearer_token_or_received_api_key(litellm_api_key)}",
|
||||
request=request,
|
||||
)
|
||||
elif MCPRequestHandler._target_servers_are_true_passthrough(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
) or (
|
||||
MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
)
|
||||
)
|
||||
is not None
|
||||
):
|
||||
(
|
||||
validated_user_api_key_auth,
|
||||
mcp_server_auth_headers,
|
||||
) = await MCPRequestHandler._admit_dcr_bridge_dual_credential(
|
||||
server=dual_bridge_target.server,
|
||||
requested_name=dual_bridge_target.requested_name,
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
litellm_api_key=litellm_api_key,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request=request,
|
||||
route=request_route,
|
||||
)
|
||||
elif has_explicit_litellm_key:
|
||||
# An explicit x-litellm-api-key is always a LiteLLM credential, even
|
||||
# for a delegated server, so validate it: identity / spend / rate
|
||||
# limits resolve and any stored upstream token can be forwarded.
|
||||
validated_user_api_key_auth = await user_api_key_auth(
|
||||
api_key=f"Bearer {_get_bearer_token_or_received_api_key(litellm_api_key)}",
|
||||
request=request,
|
||||
)
|
||||
elif MCPRequestHandler._target_servers_are_true_passthrough(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
) or (
|
||||
MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
)
|
||||
is not None
|
||||
and not oauth2_headers
|
||||
and not mcp_server_auth_headers
|
||||
and not mcp_auth_header
|
||||
):
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif (
|
||||
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
)
|
||||
) is not None and oauth2_headers:
|
||||
(
|
||||
validated_user_api_key_auth,
|
||||
mcp_server_auth_headers,
|
||||
) = await MCPRequestHandler._admit_dcr_bridge_authorization(
|
||||
server=bridge_delegate_target.server,
|
||||
requested_name=bridge_delegate_target.requested_name,
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
litellm_api_key=litellm_api_key,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request=request,
|
||||
route=request_route,
|
||||
)
|
||||
elif oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]):
|
||||
# A gateway DCR session bearer at any MCP scope: open the identity-only session
|
||||
# token and admit under the live litellm user; downstream grant resolution
|
||||
# intersects the admitted subject's servers with any path or header target, so a
|
||||
# per-server scope narrows and never broadens. One that does not open fails
|
||||
# closed with the scope's invalid_token challenge; a non-session bearer falls
|
||||
# through to the oauth2 arm.
|
||||
validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session(
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
request=request,
|
||||
route=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
)
|
||||
elif oauth2_headers:
|
||||
# Authorization on a non-delegated server: the bearer must be a real
|
||||
# LiteLLM credential, so a failed validation is a genuine 401/403 and
|
||||
# propagates unless a fallback in _admission_failure_fallback applies.
|
||||
try:
|
||||
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
|
||||
except (HTTPException, ProxyException) as e:
|
||||
validated_user_api_key_auth = _admission_failure_fallback(
|
||||
request=request,
|
||||
request_route=request_route,
|
||||
is not None
|
||||
and not oauth2_headers
|
||||
and not mcp_server_auth_headers
|
||||
and not mcp_auth_header
|
||||
):
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif (
|
||||
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
|
||||
path=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
exc=e,
|
||||
bearer_presented=True,
|
||||
client_ip=IPAddressUtils.get_mcp_client_ip(request),
|
||||
)
|
||||
else:
|
||||
try:
|
||||
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
|
||||
except (HTTPException, ProxyException) as exc:
|
||||
validated_user_api_key_auth = _admission_failure_fallback(
|
||||
) is not None and oauth2_headers:
|
||||
(
|
||||
validated_user_api_key_auth,
|
||||
mcp_server_auth_headers,
|
||||
) = await MCPRequestHandler._admit_dcr_bridge_authorization(
|
||||
server=bridge_delegate_target.server,
|
||||
requested_name=bridge_delegate_target.requested_name,
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
litellm_api_key=litellm_api_key,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request=request,
|
||||
request_route=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
exc=exc,
|
||||
bearer_presented=False,
|
||||
route=request_route,
|
||||
)
|
||||
elif oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]):
|
||||
# A gateway DCR session bearer at any MCP scope: open the identity-only session
|
||||
# token and admit under the live litellm user; downstream grant resolution
|
||||
# intersects the admitted subject's servers with any path or header target, so a
|
||||
# per-server scope narrows and never broadens. One that does not open fails
|
||||
# closed with the scope's invalid_token challenge; a non-session bearer falls
|
||||
# through to the oauth2 arm.
|
||||
validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session(
|
||||
authorization_value=oauth2_headers["Authorization"],
|
||||
request=request,
|
||||
route=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
)
|
||||
elif oauth2_headers:
|
||||
# Authorization on a non-delegated server: the bearer must be a real
|
||||
# LiteLLM credential, so a failed validation is a genuine 401/403 and
|
||||
# propagates unless a fallback in _admission_failure_fallback applies.
|
||||
try:
|
||||
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
|
||||
except (HTTPException, ProxyException) as e:
|
||||
validated_user_api_key_auth = _admission_failure_fallback(
|
||||
request=request,
|
||||
request_route=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
exc=e,
|
||||
bearer_presented=True,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
|
||||
except (HTTPException, ProxyException) as exc:
|
||||
validated_user_api_key_auth = _admission_failure_fallback(
|
||||
request=request,
|
||||
request_route=request_route,
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
exc=exc,
|
||||
bearer_presented=False,
|
||||
)
|
||||
|
||||
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
|
||||
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
|
||||
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
|
||||
# credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged.
|
||||
raw_headers = dict(headers)
|
||||
(
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
mcp_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
) = MCPRequestHandler._scrub_gateway_admission_credentials(
|
||||
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
)
|
||||
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
|
||||
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
|
||||
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
|
||||
# credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged.
|
||||
raw_headers = dict(headers)
|
||||
(
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
mcp_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
) = MCPRequestHandler._scrub_gateway_admission_credentials(
|
||||
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
)
|
||||
|
||||
return (
|
||||
validated_user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
)
|
||||
return (
|
||||
validated_user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
mcp_servers,
|
||||
mcp_server_auth_headers,
|
||||
oauth2_headers,
|
||||
raw_headers,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_gateway_admission_credential(value: str | None) -> bool:
|
||||
|
|
@ -1613,9 +1618,13 @@ class MCPRequestHandler:
|
|||
# Get allowed servers from key and team
|
||||
allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
|
||||
|
||||
# The key explicitly opted out of every MCP server. This overrides
|
||||
# team inheritance and additive grants (mirrors no-default-models).
|
||||
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key:
|
||||
# Only an explicit opt-out overrides additive grants. An empty group
|
||||
# scope uses the same marker to restrict inheritance below.
|
||||
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key and (
|
||||
user_api_key_auth is None
|
||||
or (permission := await MCPRequestHandler._key_object_permission_hydrated(user_api_key_auth)) is None
|
||||
or SpecialMCPServerNames.no_mcp_servers.value in (permission.mcp_servers or [])
|
||||
):
|
||||
return MCPServerAccess(server_ids=(), scope="scoped")
|
||||
|
||||
allowed_mcp_servers_for_team = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
|
||||
|
|
@ -1742,7 +1751,7 @@ class MCPRequestHandler:
|
|||
|
||||
declares_key_mcp_scope: Final = getattr(key_object_permission, "mcp_servers", None) is not None
|
||||
return MCPServerAccess(
|
||||
server_ids=tuple(set(allowed_mcp_servers)),
|
||||
server_ids=tuple(set(allowed_mcp_servers) - {SpecialMCPServerNames.no_mcp_servers.value}),
|
||||
scope=(
|
||||
"scoped"
|
||||
if has_lower_level_mcp_restrictions or org_restricts or declares_key_mcp_scope
|
||||
|
|
@ -2543,6 +2552,20 @@ class MCPRequestHandler:
|
|||
verbose_logger.warning("Failed to get key access group MCP server grants: %s", e)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _server_grants_with_group_scope(
|
||||
servers: AbstractSet[str],
|
||||
permission: LiteLLM_ObjectPermissionTable,
|
||||
) -> AbstractSet[str]:
|
||||
"""Keep a configured empty group scope distinct from an unrestricted level.
|
||||
|
||||
Combine same-level grants first; the existing zero-server marker prevents
|
||||
inheritance and ceiling substitution without discarding independent grants.
|
||||
"""
|
||||
if servers or not permission.mcp_access_groups:
|
||||
return servers
|
||||
return {SpecialMCPServerNames.no_mcp_servers.value}
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_key(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
|
|
@ -2625,7 +2648,7 @@ class MCPRequestHandler:
|
|||
|
||||
# Combine all lists
|
||||
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers + toolset_servers
|
||||
return list(set(all_servers))
|
||||
return list(MCPRequestHandler._server_grants_with_group_scope(set(all_servers), key_object_permission))
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for key: %s", e)
|
||||
return []
|
||||
|
|
@ -2698,7 +2721,7 @@ class MCPRequestHandler:
|
|||
team_access_group_servers: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> set[str]:
|
||||
) -> AbstractSet[str]:
|
||||
"""The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct
|
||||
``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups,
|
||||
tool-perm-referenced servers, toolset-referenced servers) unioned with its unified
|
||||
|
|
@ -2719,13 +2742,14 @@ class MCPRequestHandler:
|
|||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions, requires_fresh_policy=requires_fresh_policy
|
||||
)
|
||||
return (
|
||||
servers: Final = (
|
||||
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
|
||||
| set(legacy_access_group_servers)
|
||||
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
|
||||
| toolset_grants.keys()
|
||||
| set(team_access_group_servers)
|
||||
)
|
||||
return MCPRequestHandler._server_grants_with_group_scope(servers, object_permissions)
|
||||
|
||||
@staticmethod
|
||||
async def _allowed_mcp_servers_for_single_team(
|
||||
|
|
@ -2935,7 +2959,7 @@ class MCPRequestHandler:
|
|||
all_servers: Final = tuple(
|
||||
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
|
||||
)
|
||||
return list(set(all_servers))
|
||||
return list(MCPRequestHandler._server_grants_with_group_scope(set(all_servers), object_permissions))
|
||||
except Exception as e:
|
||||
# None = ceiling UNRESOLVED, distinct from [] = org places no restriction. Collapsing them
|
||||
# let a DB fault silently drop a ceiling; the caller picks fail-open/closed from this signal.
|
||||
|
|
@ -3033,7 +3057,7 @@ class MCPRequestHandler:
|
|||
|
||||
# Combine all lists
|
||||
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers
|
||||
return list(set(all_servers))
|
||||
return list(MCPRequestHandler._server_grants_with_group_scope(set(all_servers), object_permission))
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for end_user: %s", e)
|
||||
return []
|
||||
|
|
@ -3168,7 +3192,12 @@ class MCPRequestHandler:
|
|||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions, requires_fresh_policy=fresh
|
||||
)
|
||||
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
|
||||
return tuple(
|
||||
MCPRequestHandler._server_grants_with_group_scope(
|
||||
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants},
|
||||
object_permissions,
|
||||
)
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
|
||||
return None
|
||||
|
|
@ -3515,7 +3544,12 @@ class MCPRequestHandler:
|
|||
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
inline_tools: Final = global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions)
|
||||
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools})
|
||||
return list(
|
||||
MCPRequestHandler._server_grants_with_group_scope(
|
||||
{*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools},
|
||||
obj_perm,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
|
|
@ -3638,7 +3672,7 @@ class MCPRequestHandler:
|
|||
requires_fresh_policy: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers.
|
||||
Resolve MCP access groups against the operation snapshot, or database and config outside an operation.
|
||||
``requires_fresh_policy`` reads the writer and propagates a read fault instead of resolving to no servers.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
@ -3649,6 +3683,10 @@ class MCPRequestHandler:
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
snapshot: Final = global_mcp_server_manager.catalog.current()
|
||||
if snapshot is not None:
|
||||
return list(MCPRequestHandler._get_config_server_ids_for_access_groups(snapshot.servers, access_groups))
|
||||
|
||||
# Use the new helper for config-loaded servers
|
||||
server_ids: Final = MCPRequestHandler._get_config_server_ids_for_access_groups(
|
||||
global_mcp_server_manager.config_mcp_servers, access_groups
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from fastapi import APIRouter, Depends, Form, HTTPException, Request
|
|||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
|
||||
from litellm.proxy._experimental.mcp_server.db import store_user_credential
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
BYOK_RESOURCE_METADATA_PATH,
|
||||
|
|
@ -738,6 +739,7 @@ async def byok_protected_resource_metadata(request: Request) -> JSONResponse:
|
|||
|
||||
|
||||
@router.get("/v1/mcp/oauth/authorize", include_in_schema=False)
|
||||
@public_catalog_operation
|
||||
async def byok_authorize_get(
|
||||
request: Request,
|
||||
client_id: str | None = None,
|
||||
|
|
@ -778,7 +780,7 @@ async def byok_authorize_get(
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
registry: Final = global_mcp_server_manager.get_registry()
|
||||
registry: Final = await global_mcp_server_manager.catalog.list()
|
||||
if server_id in registry:
|
||||
srv: Final = registry[server_id]
|
||||
server_name = srv.server_name or srv.name
|
||||
|
|
|
|||
578
litellm/proxy/_experimental/mcp_server/catalog.py
Normal file
578
litellm/proxy/_experimental/mcp_server/catalog.py
Normal file
|
|
@ -0,0 +1,578 @@
|
|||
"""Authoritative MCP catalog snapshots shared by legacy lookup adapters."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
from collections import UserDict
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass, replace
|
||||
from functools import wraps
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pydantic import BaseModel
|
||||
|
||||
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")
|
||||
|
||||
|
||||
class _OperationRoutes(UserDict[str, str]):
|
||||
def __init__(self, initial: Mapping[str, str]) -> None:
|
||||
super().__init__()
|
||||
self.data = dict(initial)
|
||||
self.written_names: set[str] = set() # mutable-ok: record route writes without copying the journal per tool
|
||||
|
||||
def __setitem__(self, name: str, owner: str) -> None:
|
||||
self.data[name] = owner
|
||||
self.written_names.add(name)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CatalogSnapshot:
|
||||
servers: Mapping[str, MCPServer]
|
||||
identity: str
|
||||
tools: Mapping[str, MCPTool]
|
||||
routing: MutableMapping[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",
|
||||
"authorization_response_iss_parameter_supported",
|
||||
)
|
||||
)
|
||||
| (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._arrival_ticket = 0
|
||||
self._completed_ticket = 0
|
||||
self._shared_snapshot: CatalogSnapshot | None = None
|
||||
self._applied_revision: int | None = None
|
||||
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
|
||||
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
|
||||
self._operation: ContextVar[tuple[CatalogSnapshot, asyncio.Event] | 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) -> MutableMapping[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) -> Mapping[str, MCPServer]:
|
||||
snapshot: Final = self.current()
|
||||
return snapshot.servers if snapshot is not None else self.manager.config_mcp_servers | self.manager.registry
|
||||
|
||||
async def _fresh_snapshot(self) -> CatalogSnapshot:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
return _snapshot(self.manager, self._database_identity)
|
||||
from litellm.proxy.proxy_server import should_load_db_object
|
||||
|
||||
if not should_load_db_object("mcp"):
|
||||
return _snapshot(self.manager, self._database_identity)
|
||||
from litellm.proxy._experimental.mcp_server.db import get_mcp_catalog_revision
|
||||
|
||||
try:
|
||||
revision: Final = await get_mcp_catalog_revision(prisma_client)
|
||||
if revision is not None and revision == self._applied_revision and self._shared_snapshot is not None:
|
||||
return self._shared_snapshot
|
||||
self._arrival_ticket += 1
|
||||
arrival: Final = self._arrival_ticket
|
||||
async with self._refresh_lock:
|
||||
if arrival > self._completed_ticket or revision != self._applied_revision:
|
||||
await self._publish_refresh(revision, reuse_unchanged=True)
|
||||
if self._shared_snapshot is None:
|
||||
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed")
|
||||
return self._shared_snapshot
|
||||
except Exception as exc:
|
||||
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed") from exc
|
||||
|
||||
async def list(self) -> Mapping[str, MCPServer]:
|
||||
async with self.operation() as snapshot:
|
||||
return snapshot.servers
|
||||
|
||||
def assert_current(self, server: MCPServer) -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
snapshot: Final = self.current()
|
||||
if snapshot is None or server.server_id not in snapshot.servers:
|
||||
return
|
||||
expected: Final = snapshot.servers[server.server_id]
|
||||
registered: Final = self.manager.registry.get(server.server_id) or self.manager.config_mcp_servers.get(
|
||||
server.server_id
|
||||
)
|
||||
if (
|
||||
registered is None
|
||||
or registered.updated_at != expected.updated_at
|
||||
or server.updated_at != expected.updated_at
|
||||
):
|
||||
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
|
||||
|
||||
@asynccontextmanager
|
||||
async def operation(self) -> AsyncGenerator[CatalogSnapshot]:
|
||||
current: Final = self.current()
|
||||
if current is not None:
|
||||
yield current
|
||||
return
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
shared: Final = await self._fresh_snapshot()
|
||||
initial_routing: Final = self._unchanged_routing(shared.servers, self.manager.published_tool_routes)
|
||||
routing: Final = _OperationRoutes(initial_routing)
|
||||
snapshot: Final = replace(
|
||||
shared,
|
||||
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
|
||||
routing=routing,
|
||||
)
|
||||
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, frozenset(routing.written_names))
|
||||
|
||||
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, written_names: frozenset[str]) -> None:
|
||||
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
|
||||
snapshot.servers,
|
||||
{name: owner for name, owner in snapshot.routing.items() if name in written_names},
|
||||
)
|
||||
|
||||
def _unchanged_routing(
|
||||
self, servers: Mapping[str, MCPServer], routing: Mapping[str, str]
|
||||
) -> MappingProxyType[str, str]:
|
||||
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
|
||||
|
||||
current: Final = self.manager.config_mcp_servers | self.manager.registry
|
||||
unchanged: Final = (
|
||||
server
|
||||
for key, server in servers.items()
|
||||
if (candidate := current.get(key)) is not None
|
||||
and _configuration_identity(candidate) == _configuration_identity(server)
|
||||
)
|
||||
unchanged_owners: Final = frozenset(chain.from_iterable(map(self.manager.owned_mapping_values, unchanged)))
|
||||
return MappingProxyType(
|
||||
{name: owner for name, owner in routing.items() if normalize_server_name(owner) in unchanged_owners}
|
||||
)
|
||||
|
||||
async def resolve(self, identifier: str, client_ip: str | None = None) -> MCPServer | None:
|
||||
async with self.operation():
|
||||
return self.manager.get_mcp_server_by_name(
|
||||
identifier, client_ip=client_ip
|
||||
) or self.manager.get_mcp_server_by_id(identifier, client_ip=client_ip)
|
||||
|
||||
async def resolve_oauth_metadata(
|
||||
self,
|
||||
server: MCPServer,
|
||||
resolve: Callable[[MCPServer], Awaitable[MCPServer]],
|
||||
) -> MCPServer:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import oauth_endpoints_unresolved
|
||||
|
||||
snapshot: Final = self.current()
|
||||
selected: Final = snapshot.servers.get(server.server_id) if snapshot is not None else None
|
||||
if selected is None:
|
||||
return await resolve(server)
|
||||
if not oauth_endpoints_unresolved(selected):
|
||||
return selected
|
||||
registered: Final = self.manager.registry.get(server.server_id) or self.manager.config_mcp_servers.get(
|
||||
server.server_id
|
||||
)
|
||||
self.assert_current(server)
|
||||
_check_oauth_revision(selected, registered)
|
||||
resolved: Final = await resolve(selected)
|
||||
_check_oauth_revision(selected, resolved)
|
||||
return resolved
|
||||
|
||||
async def reload(self) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.db import get_mcp_catalog_revision
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
revision: Final = await get_mcp_catalog_revision(prisma_client) if prisma_client is not None else None
|
||||
async with self._refresh_lock:
|
||||
await self._publish_refresh(revision)
|
||||
|
||||
async def _publish_refresh(self, revision: int | None, *, reuse_unchanged: bool = False) -> None:
|
||||
covered: Final = self._arrival_ticket
|
||||
token: Final = self._operation.set(None)
|
||||
try:
|
||||
await self._reload_and_publish(reuse_unchanged=reuse_unchanged)
|
||||
except Exception:
|
||||
self._shared_snapshot = None
|
||||
self._completed_ticket = covered
|
||||
raise
|
||||
else:
|
||||
self._shared_snapshot = _snapshot(self.manager, self._database_identity)
|
||||
self._applied_revision = revision
|
||||
self._completed_ticket = covered
|
||||
finally:
|
||||
self._operation.reset(token)
|
||||
|
||||
async def _reload_and_publish(self, *, reuse_unchanged: bool) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
|
||||
|
||||
previous_config: Final = MappingProxyType(
|
||||
{key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()}
|
||||
)
|
||||
live_registry: Final = self.manager.registry
|
||||
previous_servers: Final = previous_config | MappingProxyType(
|
||||
{key: value.model_copy(deep=True) for key, value in live_registry.items()}
|
||||
)
|
||||
config_identities: Final = MappingProxyType(
|
||||
{key: _configuration_identity(value) for key, value in previous_config.items()}
|
||||
)
|
||||
staged_config: Final = MappingProxyType(
|
||||
{key: value.model_copy(deep=True) for key, value in previous_config.items()}
|
||||
)
|
||||
await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values()))
|
||||
initial_routing: Final = MappingProxyType(dict(self.manager.published_tool_routes))
|
||||
staged_routing: Final = dict(initial_routing)
|
||||
closed: Final = asyncio.Event()
|
||||
routing_token: Final = self._staged_routing.set((staged_routing, closed))
|
||||
initial_tools: Final = MappingProxyType(dict(global_mcp_tool_registry.published_tools))
|
||||
try:
|
||||
with global_mcp_tool_registry.catalog_scope(initial_tools) as staged_tools:
|
||||
await self._reload(reuse_unchanged=reuse_unchanged)
|
||||
refreshed_openapi: Final = (
|
||||
server
|
||||
for server in self.manager.registry.values()
|
||||
if server.spec_path and server is not live_registry.get(server.server_id)
|
||||
)
|
||||
refreshed_openapi_owners: Final = frozenset(
|
||||
chain.from_iterable(map(self.manager.owned_mapping_values, refreshed_openapi))
|
||||
)
|
||||
live_routes: Final = self._unchanged_routing(
|
||||
previous_servers,
|
||||
MappingProxyType(
|
||||
{
|
||||
name: owner
|
||||
for name, owner in self.manager.published_tool_routes.items()
|
||||
if normalize_server_name(owner) not in refreshed_openapi_owners
|
||||
}
|
||||
),
|
||||
)
|
||||
concurrent_routes: Final = self._unchanged_routing(
|
||||
self.manager.config_mcp_servers | live_registry,
|
||||
MappingProxyType(
|
||||
{
|
||||
name: owner
|
||||
for name, owner in self.manager.published_tool_routes.items()
|
||||
if initial_routing.get(name) != owner
|
||||
and (name not in staged_routing or staged_routing.get(name) == initial_routing.get(name))
|
||||
}
|
||||
),
|
||||
)
|
||||
removed_routes: Final = initial_routing.keys() - self.manager.published_tool_routes.keys()
|
||||
retained_staged_routes: Final = self._unchanged_routing(
|
||||
previous_config | self.manager.registry,
|
||||
MappingProxyType(
|
||||
{
|
||||
name: owner
|
||||
for name, owner in staged_routing.items()
|
||||
if name not in removed_routes or owner != initial_routing[name]
|
||||
}
|
||||
),
|
||||
)
|
||||
concurrent_tools: Final = MappingProxyType(
|
||||
{
|
||||
name: tool
|
||||
for name, tool in global_mcp_tool_registry.published_tools.items()
|
||||
if (
|
||||
name in live_routes
|
||||
or name in concurrent_routes
|
||||
or (name not in initial_routing and name not in self.manager.published_tool_routes)
|
||||
)
|
||||
and initial_tools.get(name) is staged_tools.get(name)
|
||||
}
|
||||
)
|
||||
self.manager.config_mcp_servers = {
|
||||
key: value.model_copy(
|
||||
update=staged_config[key].model_dump(
|
||||
include=frozenset(("client_id", "client_secret", "token_endpoint_auth_method"))
|
||||
)
|
||||
)
|
||||
if config_identities.get(key) == _configuration_identity(value)
|
||||
else value
|
||||
for key, value in self.manager.config_mcp_servers.items()
|
||||
}
|
||||
global_mcp_tool_registry.tools = (
|
||||
MappingProxyType(
|
||||
{
|
||||
name: tool
|
||||
for name, tool in staged_tools.items()
|
||||
if name in global_mcp_tool_registry.published_tools or tool is not initial_tools.get(name)
|
||||
}
|
||||
)
|
||||
| concurrent_tools
|
||||
)
|
||||
self.manager.published_tool_routes = live_routes | retained_staged_routes | concurrent_routes
|
||||
finally:
|
||||
closed.set()
|
||||
self._staged_routing.reset(routing_token)
|
||||
|
||||
async def _stage_servers(
|
||||
self, rows: Sequence[BaseModel], *, reuse_unchanged: bool
|
||||
) -> dict[str, MCPServer]: # mutable-ok: assign_unique_short_prefix requires a dict registry
|
||||
from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
carry_forward_resolved_oauth_endpoints,
|
||||
oauth_endpoints_unresolved,
|
||||
warn_on_server_name_fields,
|
||||
)
|
||||
|
||||
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 rows:
|
||||
try:
|
||||
server = LiteLLM_MCPServerTable.model_validate(row.model_dump())
|
||||
existing_server = previous_registry.get(server.server_id)
|
||||
|
||||
if (
|
||||
existing_server is not None
|
||||
and (reuse_unchanged or not existing_server.spec_path)
|
||||
and existing_server.updated_at is not None
|
||||
and server.updated_at is not None
|
||||
and existing_server.updated_at == server.updated_at
|
||||
and (
|
||||
self.manager.oauth_discovery_slot(server.server_id) is not None
|
||||
or not oauth_endpoints_unresolved(existing_server)
|
||||
)
|
||||
):
|
||||
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
|
||||
# which can perform network discovery for OAuth2 servers.
|
||||
new_registry[server.server_id] = existing_server
|
||||
continue
|
||||
|
||||
warn_on_server_name_fields(
|
||||
server_id=server.server_id,
|
||||
alias=getattr(server, "alias", None),
|
||||
server_name=getattr(server, "server_name", None),
|
||||
)
|
||||
self.manager.warn_if_newly_blocked_stdio(server, existing_server)
|
||||
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,
|
||||
)
|
||||
|
||||
return new_registry
|
||||
|
||||
async def _reload(self, *, reuse_unchanged: bool) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
config_ids_capturing_db_identifiers,
|
||||
warn_on_shared_identifier_prefixes,
|
||||
)
|
||||
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 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 = await self._stage_servers(raw_rows, reuse_unchanged=reuse_unchanged)
|
||||
|
||||
# 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_server_definition_caches(server_id)
|
||||
self.manager.invalidate_oauth_discovery_state(server_id)
|
||||
self._database_identity = database_identity
|
||||
self.manager.registry = registered_registry
|
||||
if not reuse_unchanged:
|
||||
self.manager.clear_initialize_instructions()
|
||||
warn_on_shared_identifier_prefixes(registered_registry.values())
|
||||
# A discovery task may have published into ``previous_registry`` while
|
||||
# this replacement was being staged. Reconcile every published entry
|
||||
# synchronously after the swap so a lost publication cannot also leave
|
||||
# the replacement unresolved with no retry slot.
|
||||
registered_servers: Final = tuple(registered_registry.values())
|
||||
self.manager.reconcile_oauth_discovery_slots_for_servers(registered_servers)
|
||||
self.manager.prime_oauth_metadata_discovery_for_servers(registered_servers)
|
||||
|
||||
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
|
||||
|
||||
# get_registry() is ``config_mcp_servers | registry``, so a database row sharing an id with a
|
||||
# config.yaml server hides that server everywhere. Only reachable once an operator pins
|
||||
# ``server_id`` in config.yaml; say so rather than letting the server disappear silently.
|
||||
shadowed_config_server_ids: Final = frozenset(
|
||||
self.manager.config_mcp_servers.keys() & registered_registry.keys()
|
||||
)
|
||||
if shadowed_config_server_ids and shadowed_config_server_ids != self._warned_shadowed_config_server_ids:
|
||||
verbose_logger.warning(
|
||||
"config.yaml MCP server_id(s) %s are also database-backed MCP servers. The database "
|
||||
"entry takes precedence, so the config.yaml server is unreachable. Give the config "
|
||||
"entry a different server_id.",
|
||||
", ".join(sorted(shadowed_config_server_ids)),
|
||||
)
|
||||
self._warned_shadowed_config_server_ids = shadowed_config_server_ids
|
||||
|
||||
# The mirror image of the block above: a config server_id that is a database server's name
|
||||
# answers that server's grants instead, because ids are matched before names.
|
||||
capturing_config_server_ids: Final = config_ids_capturing_db_identifiers(
|
||||
self.manager.config_mcp_servers.keys(), registered_registry.values()
|
||||
)
|
||||
if capturing_config_server_ids and capturing_config_server_ids != self._warned_capturing_config_server_ids:
|
||||
verbose_logger.warning(
|
||||
"config.yaml MCP server_id(s) %s are the name or alias of a database-backed MCP "
|
||||
"server. Permission entries naming them resolve to the config.yaml server, not the "
|
||||
"database one. Give the config entry a different server_id.",
|
||||
", ".join(sorted(capturing_config_server_ids)),
|
||||
)
|
||||
self._warned_capturing_config_server_ids = capturing_config_server_ids
|
||||
|
||||
|
||||
def catalog_operation(
|
||||
manager: Callable[[], MCPServerManager],
|
||||
) -> Callable[[Callable[_P, Awaitable[_R]]], Callable[_P, Awaitable[_R]]]:
|
||||
def decorate(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
|
||||
@wraps(function)
|
||||
async def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: # kwargs-ok: preserves ParamSpec
|
||||
|
||||
async with manager().catalog.operation():
|
||||
return await function(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
||||
return decorate
|
||||
|
||||
|
||||
def global_manager() -> MCPServerManager:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
return global_mcp_server_manager
|
||||
|
||||
|
||||
def public_catalog_operation(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
|
||||
return catalog_operation(global_manager)(function)
|
||||
|
|
@ -42,6 +42,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import (
|
||||
|
|
@ -723,6 +724,20 @@ async def _db_find_mcp_server_row(
|
|||
return await _mcp_server_table_actions(prisma_client).find_unique(where={"server_id": server_id})
|
||||
|
||||
|
||||
MCP_CATALOG_REVISION_PARAM_NAME: Final = "mcp_catalog"
|
||||
|
||||
|
||||
async def get_mcp_catalog_revision(prisma_client: PrismaClient) -> int | None:
|
||||
"""The ``LiteLLM_Config`` revision the MCP server table trigger last published, or None when
|
||||
no write has ever bumped it (or the trigger is not installed), meaning always reload."""
|
||||
row: Final = await ConfigRepository(prisma_client, use_writer=True).table.find_unique(
|
||||
where={"param_name": MCP_CATALOG_REVISION_PARAM_NAME}
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
return int(row.reload_revision or 0)
|
||||
|
||||
|
||||
async def _db_update_mcp_server_row(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
|
|
@ -853,6 +868,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 MCPServerRepository(prisma_client, use_writer=True).table.find_many(where=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
|
||||
|
|
|
|||
|
|
@ -37,6 +37,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 public_catalog_operation
|
||||
from litellm.proxy._experimental.mcp_server.faults import (
|
||||
CallerRejected,
|
||||
CredentialSource,
|
||||
|
|
@ -1894,7 +1895,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:
|
||||
|
|
@ -2020,6 +2021,7 @@ async def register_client_with_server(
|
|||
|
||||
|
||||
@router.get("/authorize/mcp-session")
|
||||
@public_catalog_operation
|
||||
async def authorize_mcp_session(
|
||||
request: Request,
|
||||
redirect_uri: str,
|
||||
|
|
@ -2045,6 +2047,7 @@ async def authorize_mcp_session(
|
|||
|
||||
@router.get("/{mcp_server_name}/authorize")
|
||||
@router.get("/authorize")
|
||||
@public_catalog_operation
|
||||
async def authorize(
|
||||
request: Request,
|
||||
redirect_uri: str,
|
||||
|
|
@ -2083,9 +2086,11 @@ async def authorize(
|
|||
resource=resource,
|
||||
)
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
lookup_name: Final[str | None] = mcp_server_name or client_id
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip) if lookup_name else None
|
||||
mcp_server = await global_mcp_server_manager.catalog.resolve(lookup_name, client_ip) if lookup_name else None
|
||||
if mcp_server is None and mcp_server_name is None:
|
||||
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if mcp_server is None:
|
||||
|
|
@ -2119,6 +2124,7 @@ async def authorize(
|
|||
|
||||
@router.post("/{mcp_server_name}/token")
|
||||
@router.post("/token")
|
||||
@public_catalog_operation
|
||||
async def token_endpoint(
|
||||
request: Request,
|
||||
grant_type: str = Form(...),
|
||||
|
|
@ -2212,6 +2218,7 @@ async def _vendor_credential_state(user_id: str, server_id: str) -> VendorCreden
|
|||
|
||||
|
||||
@router.get("/authorize/flow")
|
||||
@public_catalog_operation
|
||||
async def authorize_flow(request: Request, flow: str) -> Response:
|
||||
return await describe_connect_flow(
|
||||
request=request,
|
||||
|
|
@ -2223,6 +2230,7 @@ async def authorize_flow(request: Request, flow: str) -> Response:
|
|||
|
||||
|
||||
@router.post("/authorize/complete")
|
||||
@public_catalog_operation
|
||||
async def authorize_complete(
|
||||
request: Request,
|
||||
flow: str = Form(...),
|
||||
|
|
@ -2610,6 +2618,7 @@ def is_network_error(exc: Exception) -> bool:
|
|||
return isinstance(exc, httpx.TransportError)
|
||||
|
||||
|
||||
@public_catalog_operation
|
||||
async def _build_oauth_protected_resource_response(
|
||||
request: Request,
|
||||
mcp_server_name: str | None,
|
||||
|
|
@ -2873,6 +2882,7 @@ async def oauth_authorization_server_aggregate(request: Request):
|
|||
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
|
||||
# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot)
|
||||
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
|
||||
@public_catalog_operation
|
||||
async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_name: str):
|
||||
"""
|
||||
OAuth protected resource discovery endpoint using standard MCP URL pattern.
|
||||
|
|
@ -2893,6 +2903,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam
|
|||
# LiteLLM legacy pattern: /.well-known/oauth-protected-resource/{server_name}/mcp
|
||||
# Kept for backward compatibility with existing deployments
|
||||
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
|
||||
@public_catalog_operation
|
||||
async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str | None = None):
|
||||
"""
|
||||
OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern.
|
||||
|
|
@ -2916,12 +2927,7 @@ def _build_oauth_authorization_server_response(
|
|||
*,
|
||||
issuer_path: str | None = None,
|
||||
) -> dict:
|
||||
"""Build OAuth authorization server metadata response (gateway-as-AS shape).
|
||||
|
||||
Synchronous because the body only does dict construction and synchronous
|
||||
registry lookups; unlike :func:`_build_oauth_protected_resource_response`
|
||||
it does not need to await any upstream IO.
|
||||
"""
|
||||
"""Build OAuth authorization server metadata response (gateway-as-AS shape)."""
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
explicitly_named: Final = mcp_server_name is not None
|
||||
|
|
@ -2969,6 +2975,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}}")
|
||||
@public_catalog_operation
|
||||
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
|
||||
"""
|
||||
OAuth authorization server discovery endpoint using standard MCP URL pattern.
|
||||
|
|
@ -2986,6 +2993,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")
|
||||
@public_catalog_operation
|
||||
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None):
|
||||
"""
|
||||
OAuth authorization server discovery endpoint.
|
||||
|
|
@ -3059,6 +3067,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")
|
||||
@public_catalog_operation
|
||||
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
|
||||
"""
|
||||
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
|
||||
|
|
@ -3072,6 +3081,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
|
|||
|
||||
@router.post("/{mcp_server_name}/register")
|
||||
@router.post("/register")
|
||||
@public_catalog_operation
|
||||
async def register_client(request: Request, mcp_server_name: str | None = None):
|
||||
# Get the correct base URL considering X-Forwarded-* headers
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
|
|
@ -3080,46 +3090,44 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
data: Final[dict] = {**request_data}
|
||||
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
|
||||
|
||||
dummy_return: Final = {
|
||||
"client_id": mcp_server_name or "dummy_client",
|
||||
"client_secret": "dummy",
|
||||
"redirect_uris": client_redirect_uris or [f"{request_base_url}/callback"],
|
||||
}
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
if not mcp_server_name:
|
||||
# A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR
|
||||
# endpoint the aggregate authorization-server metadata advertises. A single-server
|
||||
# deployment registers at /{server}/register instead (its bare-origin discovery
|
||||
# advertises that), so this does not affect it. A request without redirect_uris is not
|
||||
# a DCR request, so the legacy single-server-or-dummy fallback is kept for it.
|
||||
if data.get("redirect_uris"):
|
||||
return await register_aggregate_client(
|
||||
request=request, request_body=data, token_exchange_available=token_exchange_available()
|
||||
)
|
||||
resolved: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if resolved:
|
||||
return await register_client_with_server(
|
||||
request=request,
|
||||
mcp_server=resolved,
|
||||
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=resolved.server_name or resolved.name,
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
)
|
||||
return dummy_return
|
||||
if not mcp_server_name and data.get("redirect_uris"):
|
||||
return await register_aggregate_client(
|
||||
request=request, request_body=data, token_exchange_available=token_exchange_available()
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
if mcp_server is None:
|
||||
return dummy_return
|
||||
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=mcp_server_name,
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
)
|
||||
async with global_mcp_server_manager.catalog.operation():
|
||||
dummy_return: Final = {
|
||||
"client_id": mcp_server_name or "dummy_client",
|
||||
"client_secret": "dummy",
|
||||
"redirect_uris": client_redirect_uris or [f"{request_base_url}/callback"],
|
||||
}
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
if not mcp_server_name:
|
||||
resolved: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if resolved:
|
||||
return await register_client_with_server(
|
||||
request=request,
|
||||
mcp_server=resolved,
|
||||
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=resolved.server_name or resolved.name,
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
)
|
||||
return dummy_return
|
||||
|
||||
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
if mcp_server is None:
|
||||
return dummy_return
|
||||
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=mcp_server_name,
|
||||
client_redirect_uris=client_redirect_uris,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
|
|||
from litellm._internal_context import with_service_target
|
||||
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,
|
||||
|
|
@ -774,6 +775,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]:
|
||||
|
|
|
|||
|
|
@ -49,7 +49,7 @@ from mcp.types import (
|
|||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl, BaseModel, Field, TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
from typing_extensions import ReadOnly, assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import with_service_target
|
||||
|
|
@ -195,7 +195,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.integrations.slack_alerting import AlertType
|
||||
from litellm.types.llms.base import LiteLLMBaseModel
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
|
@ -318,7 +317,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]
|
||||
|
|
@ -573,7 +572,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]:
|
||||
|
|
@ -753,7 +752,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
|
||||
|
|
@ -812,7 +811,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
|
||||
|
|
@ -1511,7 +1510,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,
|
||||
|
|
@ -1537,7 +1536,7 @@ def _warn_on_server_name_fields(
|
|||
_warn("server_name", server_name)
|
||||
|
||||
|
||||
def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None:
|
||||
def warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None:
|
||||
"""Warn once per identifier that several servers share.
|
||||
|
||||
``get_server_prefix`` resolves alias first, so two servers sharing a
|
||||
|
|
@ -1971,6 +1970,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] = {}
|
||||
|
|
@ -1997,7 +1999,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",
|
||||
|
|
@ -2013,14 +2015,12 @@ 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._catalog_alert_signatures: Mapping[tuple[str, AlertType], str] = MappingProxyType({})
|
||||
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:
|
||||
|
|
@ -2033,7 +2033,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()
|
||||
|
|
@ -2046,8 +2046,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()
|
||||
|
|
@ -2129,7 +2129,7 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
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:
|
||||
|
|
@ -2166,7 +2166,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)
|
||||
|
|
@ -2177,7 +2177,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 (
|
||||
|
|
@ -2224,7 +2224,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
|
||||
|
|
@ -2240,7 +2240,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:
|
||||
|
|
@ -2269,19 +2269,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
|
||||
|
|
@ -2300,6 +2305,7 @@ class MCPServerManager:
|
|||
incomplete metadata for a server whose OAuth flow the gateway
|
||||
runs itself.
|
||||
"""
|
||||
self.catalog.assert_current(server)
|
||||
acquisition: Final = self._get_or_start_oauth_discovery_task(server)
|
||||
if acquisition is None:
|
||||
return self._registered_server(server)
|
||||
|
|
@ -2312,11 +2318,13 @@ class MCPServerManager:
|
|||
raise
|
||||
match outcome:
|
||||
case _OAuthDiscoveryResolved(resolved_server):
|
||||
self.catalog.assert_current(resolved_server)
|
||||
return resolved_server
|
||||
case _OAuthDiscoveryStale():
|
||||
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
|
||||
case _OAuthDiscoveryFailed(timed_out=timed_out):
|
||||
current: Final = self._registered_server(server)
|
||||
self.catalog.assert_current(current)
|
||||
if current.is_client_forwarded_token:
|
||||
return current
|
||||
server_ref: Final = current.alias or current.server_name or current.name or current.server_id
|
||||
|
|
@ -2326,11 +2334,13 @@ class MCPServerManager:
|
|||
detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}",
|
||||
)
|
||||
|
||||
return assert_never(outcome)
|
||||
|
||||
async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer:
|
||||
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")
|
||||
|
||||
|
|
@ -2408,11 +2418,21 @@ class MCPServerManager:
|
|||
e,
|
||||
)
|
||||
|
||||
def get_registry(self) -> dict[str, MCPServer]:
|
||||
def get_registry(self) -> Mapping[str, MCPServer]:
|
||||
"""
|
||||
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,
|
||||
) -> MutableMapping[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).
|
||||
|
|
@ -2499,7 +2519,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,
|
||||
|
|
@ -2715,10 +2735,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_server_definition_caches(server_id)
|
||||
self.invalidate_server_definition_caches(server_id)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
self._set_oauth_discovery_deferred(
|
||||
server_id,
|
||||
|
|
@ -2739,13 +2759,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."""
|
||||
|
|
@ -2753,7 +2773,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(
|
||||
|
|
@ -2894,6 +2914,7 @@ class MCPServerManager:
|
|||
description=description,
|
||||
input_schema=input_schema,
|
||||
handler=tool_func,
|
||||
server_id=server.server_id,
|
||||
)
|
||||
|
||||
# Update tool name to server name mapping (for both prefixed and base names)
|
||||
|
|
@ -2918,17 +2939,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_server_definition_caches(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_server_definition_caches(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
|
||||
|
|
@ -2939,13 +2961,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)
|
||||
|
|
@ -2976,7 +2998,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)
|
||||
|
||||
|
|
@ -3067,6 +3089,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))
|
||||
|
|
@ -3300,13 +3323,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)
|
||||
|
|
@ -3333,12 +3357,12 @@ class MCPServerManager:
|
|||
# `credentials` field is the only one still encrypted here).
|
||||
# Re-decrypting plaintext would zero the values, so build with
|
||||
# env_vars_are_encrypted=False.
|
||||
self._warn_if_newly_blocked_stdio(mcp_server, None)
|
||||
self.warn_if_newly_blocked_stdio(mcp_server, None)
|
||||
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_server_definition_caches(mcp_server.server_id)
|
||||
self.assign_unique_short_prefix(new_server)
|
||||
self.invalidate_server_definition_caches(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)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -3358,7 +3382,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:
|
||||
|
|
@ -3370,14 +3394,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_server_definition_caches(mcp_server.server_id)
|
||||
self.assign_unique_short_prefix(new_server)
|
||||
self.invalidate_server_definition_caches(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)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
|
|
@ -4200,6 +4224,7 @@ class MCPServerManager:
|
|||
client_ip: str | None = None,
|
||||
protocol_version_override: MCPUpstreamProtocol | None = None,
|
||||
) -> MCPClient:
|
||||
self.catalog.assert_current(server)
|
||||
record_auth_resolution(server.server_id, AuthResolution.unresolved)
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
return await prepare_upstream_client(
|
||||
|
|
@ -4433,17 +4458,23 @@ 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 clear_initialize_instructions(self) -> None:
|
||||
self._upstream_initialize_instructions_by_server_id.clear()
|
||||
self._upstream_initialize_instructions_probed_at.clear()
|
||||
|
||||
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)
|
||||
|
||||
def _invalidate_server_definition_caches(self, server_id: str) -> None:
|
||||
def invalidate_server_definition_caches(self, server_id: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton
|
||||
invalidate_oauth_metadata_cache,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self.invalidate_discovery_lists(server_id)
|
||||
self._drop_listed_tools(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
|
|
@ -4559,7 +4590,7 @@ class MCPServerManager:
|
|||
return server.server_id, hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None:
|
||||
def warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None:
|
||||
if previous is None or previous.transport != row.transport:
|
||||
warn_if_mcp_stdio_blocked(row.alias or row.server_name, row.transport)
|
||||
|
||||
|
|
@ -5336,7 +5367,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,
|
||||
|
|
@ -5467,6 +5498,8 @@ class MCPServerManager:
|
|||
Returns:
|
||||
List of tools with prefixed names
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
prefixed_tools: Final = []
|
||||
prefix: Final = get_server_prefix(server)
|
||||
|
||||
|
|
@ -5479,6 +5512,11 @@ class MCPServerManager:
|
|||
# short ID) so call_tool can resolve regardless of which form a
|
||||
# caller / cached client is using.
|
||||
for spelling in iter_known_tool_name_spellings(original_name, server):
|
||||
namespace_owner = self.server_owning_tool_name_prefix(spelling)
|
||||
if namespace_owner is not None and namespace_owner.server_id != server.server_id:
|
||||
continue
|
||||
if namespace_owner is None and global_mcp_tool_registry.get_tool(spelling) is not None:
|
||||
continue
|
||||
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
|
||||
|
||||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
|
|
@ -5682,6 +5720,7 @@ class MCPServerManager:
|
|||
# Registration used add_server_prefix_to_name(base, get_server_prefix(server)),
|
||||
# and tool_name is the bare base name by the time call_tool reaches here, so
|
||||
# rebuilding the key the same way reproduces it exactly
|
||||
self.catalog.assert_current(server)
|
||||
registry_key: Final = add_server_prefix_to_name(tool_name, get_server_prefix(server))
|
||||
tool: Final = global_mcp_tool_registry.get_tool(registry_key)
|
||||
if tool is None:
|
||||
|
|
@ -6300,7 +6339,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
|
||||
|
|
@ -6612,7 +6651,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
|
||||
|
|
@ -6649,6 +6688,14 @@ class MCPServerManager:
|
|||
Returns:
|
||||
MCPServer if found, None otherwise
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
registered_tool: Final = global_mcp_tool_registry.get_tool(tool_name)
|
||||
if registered_tool is not None:
|
||||
if registered_tool.server_id is not None:
|
||||
return self.get_mcp_server_by_id(registered_tool.server_id)
|
||||
return self.server_owning_tool_name_prefix(tool_name)
|
||||
|
||||
registry_servers: Final = list(self.get_registry().values())
|
||||
prefix_to_server: Final = self._known_prefix_to_server()
|
||||
|
||||
|
|
@ -6677,154 +6724,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),
|
||||
)
|
||||
self._warn_if_newly_blocked_stdio(server, existing_server)
|
||||
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_server_definition_caches(server_id)
|
||||
self.registry = registered_registry
|
||||
_warn_on_shared_identifier_prefixes(registered_registry.values())
|
||||
# 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 = []
|
||||
|
|
@ -7015,7 +6915,7 @@ class MCPServerManager:
|
|||
return server
|
||||
return None
|
||||
|
||||
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
|
||||
def get_filtered_registry(self, client_ip: str | None = None) -> Mapping[str, MCPServer]:
|
||||
"""
|
||||
Get registry filtered by client IP access control.
|
||||
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ from litellm.proxy._experimental.mcp_server.capabilities import (
|
|||
build_discovery,
|
||||
configured_versions,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
|
||||
from litellm.proxy._experimental.mcp_server.contracts import (
|
||||
AuthorizedToolCall,
|
||||
OperationContext,
|
||||
|
|
@ -627,6 +628,7 @@ def apply_display_name_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,
|
||||
|
|
@ -943,6 +945,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,
|
||||
|
|
@ -1474,6 +1477,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,
|
||||
|
|
@ -1529,6 +1533,7 @@ async def _list_mcp_tools(
|
|||
return AggregateToolListing(tools=[], outcomes={})
|
||||
|
||||
|
||||
@catalog_operation(global_manager)
|
||||
async def _list_mcp_prompts(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -1570,6 +1575,7 @@ async def _list_mcp_prompts(
|
|||
return managed_prompts
|
||||
|
||||
|
||||
@catalog_operation(global_manager)
|
||||
async def _list_mcp_resources(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -1599,6 +1605,7 @@ async def _list_mcp_resources(
|
|||
return managed_resources
|
||||
|
||||
|
||||
@catalog_operation(global_manager)
|
||||
async def _list_mcp_resource_templates(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -2087,6 +2094,9 @@ async def _execute_mcp_tool(
|
|||
),
|
||||
)
|
||||
|
||||
if local_tool.server_id is not None and local_tool.server_id != mcp_server.server_id:
|
||||
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
|
||||
|
||||
# `pre_call_tool_check` calls into `proxy_logging_obj` for the
|
||||
# pre-call guardrail hooks, so source it from the canonical
|
||||
# `proxy_server` module the same way `_handle_managed_mcp_tool`
|
||||
|
|
@ -2217,6 +2227,12 @@ async def _execute_mcp_tool(
|
|||
),
|
||||
)
|
||||
|
||||
if (
|
||||
registered_local_tool.server_id is not None
|
||||
and registered_local_tool.server_id != prefix_server.server_id
|
||||
):
|
||||
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
|
||||
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
hook_result = await global_mcp_server_manager.pre_call_tool_check(
|
||||
|
|
@ -2398,6 +2414,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,
|
||||
|
|
@ -2694,6 +2711,15 @@ async def _handle_local_mcp_tool(
|
|||
tool: Final = global_mcp_tool_registry.get_tool(name)
|
||||
if not tool:
|
||||
raise HTTPException(status_code=404, detail=f"Tool '{name}' not found")
|
||||
server: Final = (
|
||||
global_mcp_server_manager.get_mcp_server_by_id(tool.server_id)
|
||||
if tool.server_id is not None
|
||||
else global_mcp_server_manager.server_owning_tool_name_prefix(name)
|
||||
)
|
||||
if tool.server_id is not None and server is None:
|
||||
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
|
||||
if server is not None:
|
||||
global_mcp_server_manager.catalog.assert_current(server)
|
||||
|
||||
try:
|
||||
if inspect.iscoroutinefunction(tool.handler):
|
||||
|
|
@ -3217,6 +3243,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 DiscoverRequest():
|
||||
|
|
|
|||
|
|
@ -23,6 +23,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,
|
||||
|
|
@ -950,6 +951,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"),
|
||||
|
|
@ -1174,6 +1176,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 (
|
||||
|
|
|
|||
|
|
@ -1599,6 +1599,9 @@ if MCP_AVAILABLE:
|
|||
await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[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, Generator, 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,32 @@ class MCPToolRegistry:
|
|||
|
||||
def __init__(self):
|
||||
# Registry to store all registered tools
|
||||
self.tools: dict[str, MCPTool] = {}
|
||||
self.published_tools: dict[str, MCPTool] = {} # mutable-ok: register_tool writes entries through .tools
|
||||
self._catalog_tools: ContextVar[tuple[dict[str, MCPTool], asyncio.Event] | None] = ContextVar(
|
||||
"mcp_catalog_tools", default=None
|
||||
)
|
||||
|
||||
@property
|
||||
def tools(self) -> dict[str, MCPTool]: # mutable-ok: register_tool mutates the returned mapping
|
||||
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: # mutable-ok: stored dict is mutated by register_tool
|
||||
self.published_tools = tools
|
||||
|
||||
@contextmanager
|
||||
def catalog_scope(
|
||||
self, tools: Mapping[str, MCPTool]
|
||||
) -> Generator[dict[str, MCPTool]]: # mutable-ok: yields the mutable staged copy
|
||||
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,
|
||||
|
|
@ -30,6 +58,8 @@ class MCPToolRegistry:
|
|||
description: str,
|
||||
input_schema: dict[str, Any],
|
||||
handler: Callable,
|
||||
*,
|
||||
server_id: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Register a new tool in the registry
|
||||
|
|
@ -39,6 +69,7 @@ class MCPToolRegistry:
|
|||
description=description,
|
||||
input_schema=input_schema,
|
||||
handler=handler,
|
||||
server_id=server_id,
|
||||
)
|
||||
verbose_logger.debug("Registered tool: %s", name)
|
||||
|
||||
|
|
|
|||
|
|
@ -76,7 +76,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.
|
||||
|
|
@ -109,7 +109,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.
|
||||
"""
|
||||
|
|
@ -309,7 +309,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:
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Validates that MCP servers referenced in request tools are registered
|
|||
on the LiteLLM gateway. Blocks or alerts when unregistered servers are found.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -46,7 +47,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,21 +91,22 @@ class MCPSecurityGuardrail(CustomGuardrail):
|
|||
return server_names
|
||||
|
||||
@staticmethod
|
||||
def _find_unregistered_mcp_servers(data: dict) -> set[str]:
|
||||
async def _find_unregistered_mcp_servers(data: Mapping[str, object]) -> frozenset[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):
|
||||
return set()
|
||||
return frozenset()
|
||||
|
||||
requested_servers: Final = MCPSecurityGuardrail._extract_mcp_server_names_from_tools(tools)
|
||||
if not requested_servers:
|
||||
return set()
|
||||
return frozenset()
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
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 frozenset(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 AsyncGenerator, Iterable, Mapping, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import (
|
||||
|
|
@ -47,6 +48,8 @@ from fastapi.responses import JSONResponse
|
|||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, public_catalog_operation
|
||||
|
||||
try:
|
||||
from prisma.errors import RecordNotFoundError, UniqueViolationError
|
||||
except ImportError:
|
||||
|
|
@ -1139,6 +1142,7 @@ if MCP_AVAILABLE:
|
|||
tags=["mcp"],
|
||||
description="MCP registry endpoint. Spec: https://github.com/modelcontextprotocol/registry",
|
||||
)
|
||||
@public_catalog_operation
|
||||
async def get_mcp_registry(request: Request):
|
||||
if not _is_public_registry_enabled():
|
||||
raise HTTPException(
|
||||
|
|
@ -1157,7 +1161,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)
|
||||
|
||||
|
|
@ -1187,6 +1192,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]:
|
||||
|
|
@ -1225,6 +1231,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]:
|
||||
|
|
@ -1246,6 +1253,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
|
||||
|
|
@ -1262,6 +1270,7 @@ if MCP_AVAILABLE:
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=list[LiteLLM_MCPServerTable],
|
||||
)
|
||||
@catalog_operation(lambda: global_mcp_server_manager)
|
||||
async def fetch_all_mcp_servers(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
team_id: str | None = Query(
|
||||
|
|
@ -1378,6 +1387,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,
|
||||
|
|
@ -1784,6 +1794,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,
|
||||
|
|
@ -2145,6 +2156,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
return _redact_mcp_credentials(temp_record)
|
||||
|
||||
@public_catalog_operation
|
||||
async def _mcp_oauth_user_api_key_auth(request: Request) -> UserAPIKeyAuth:
|
||||
"""
|
||||
Auth dependency for MCP OAuth browser-navigation endpoints (/authorize, /token).
|
||||
|
|
@ -2202,9 +2214,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
|
||||
|
|
@ -2271,6 +2281,18 @@ if MCP_AVAILABLE:
|
|||
assert authorized.runtime is not None
|
||||
return authorized.runtime
|
||||
|
||||
@asynccontextmanager
|
||||
async def _oauth_server_operation(
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request: Request | None = None,
|
||||
) -> AsyncGenerator[MCPServer]:
|
||||
if await get_cached_temporary_mcp_server(server_id) is not None:
|
||||
yield await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request)
|
||||
return
|
||||
async with global_mcp_server_manager.catalog.operation():
|
||||
yield await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request)
|
||||
|
||||
@router.get(
|
||||
"/server/oauth/{server_id}/authorize",
|
||||
include_in_schema=False,
|
||||
|
|
@ -2288,47 +2310,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",
|
||||
|
|
@ -2348,47 +2370,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",
|
||||
|
|
@ -2400,22 +2422,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[Mapping[str, object]] = {**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}",
|
||||
|
|
@ -2791,6 +2813,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,
|
||||
|
|
|
|||
|
|
@ -336,6 +336,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
|||
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache
|
||||
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
|
||||
from litellm.proxy._experimental.mcp_server.stdio_gate import MCP_STDIO_ENABLED_ENV_VAR, is_mcp_stdio_flag_key
|
||||
from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot
|
||||
from litellm.proxy._types import *
|
||||
|
|
@ -20368,30 +20369,26 @@ 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
|
||||
deduped: Final = tuple(
|
||||
token for token in dict.fromkeys(raw.strip() for raw in csv_segment.split(",")) if token
|
||||
)[:DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS]
|
||||
|
||||
resolved: Final[list[str]] = [] # mutable-ok: sequential await per token cannot live in a comprehension
|
||||
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:
|
||||
|
|
@ -20428,12 +20425,13 @@ async def _is_mcp_access_group_cached(name: str) -> bool:
|
|||
"/{mcp_server_name}/mcp",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
|
||||
)
|
||||
@public_catalog_operation
|
||||
async def dynamic_mcp_route(mcp_server_name: str, request: Request):
|
||||
"""Handle /{name}/mcp for MCP server aliases, toolsets, MCP access group tags, and comma-separated lists.
|
||||
|
||||
Resolution order:
|
||||
1. Registered MCP server alias / name
|
||||
2. Comma-separated list (short-circuits before any DB call)
|
||||
2. Comma-separated list
|
||||
3. Toolset name (DB lookup, cached)
|
||||
4. MCP access group tag (DB lookup, cached)
|
||||
"""
|
||||
|
|
@ -20446,7 +20444,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
|
||||
|
|
|
|||
|
|
@ -22,6 +22,9 @@ class _ConfigRow(Protocol):
|
|||
@property
|
||||
def param_value(self) -> object: ...
|
||||
|
||||
@property
|
||||
def reload_revision(self) -> int | None: ...
|
||||
|
||||
|
||||
class _ConfigTable(Protocol):
|
||||
async def find_unique(self, *, where: Mapping[str, str]) -> _ConfigRow | None: ...
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -112,6 +113,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,
|
||||
|
|
@ -185,6 +187,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
_split_mcp_tools = split_mcp_tools
|
||||
|
||||
@staticmethod
|
||||
@catalog_operation(global_manager)
|
||||
async def routes_through_gateway(
|
||||
tools: Iterable[Mapping[str, object]] | None,
|
||||
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
|
||||
|
|
@ -238,6 +241,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,
|
||||
|
|
@ -704,6 +708,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],
|
||||
|
|
|
|||
|
|
@ -238,7 +238,7 @@ class MCPServer(LiteLLMBaseModel):
|
|||
# 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.
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from collections.abc import Callable
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from pydantic import ConfigDict
|
||||
from pydantic import ConfigDict, Field
|
||||
|
||||
from litellm.types.llms.base import LiteLLMBaseModel
|
||||
|
||||
|
|
@ -12,6 +12,7 @@ class MCPTool(LiteLLMBaseModel):
|
|||
description: str
|
||||
input_schema: dict[str, Any]
|
||||
handler: Callable
|
||||
server_id: str | None = Field(default=None, frozen=True)
|
||||
|
||||
|
||||
class ToolSchema(LiteLLMBaseModel):
|
||||
|
|
|
|||
|
|
@ -283,38 +283,40 @@ def test_access_group_membership_follows_edits(gateway: Gateway) -> None:
|
|||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_peer_worker_observes_create_edit_and_delete_without_restart(gateway: Gateway, peer: Gateway) -> None:
|
||||
with mcp_peer() as first, mcp_peer() as second, gateway.scenario() as scenario:
|
||||
def test_peer_worker_observes_create_edit_and_delete_without_restart(gateway: Gateway, tmp_path: Path) -> None:
|
||||
config: Final = tmp_path / "peer.yaml"
|
||||
config.write_text(yaml.safe_dump({
|
||||
"model_list": [],
|
||||
"general_settings": {"master_key": gateway.key, "store_model_in_db": True, "proxy_config_reload_interval_seconds": 3600},
|
||||
}))
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, database_setup=()) as peer,
|
||||
mcp_peer() as first,
|
||||
mcp_peer() as second,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "mgmt" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, first, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
eventually(
|
||||
lambda: peer.client.get(
|
||||
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
|
||||
),
|
||||
lambda value: value.status_code == 200 and value.json() != [],
|
||||
seconds=40,
|
||||
listing: Final = peer.client.get(
|
||||
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
|
||||
)
|
||||
assert listing.status_code == 200 and listing.json() != [], listing.text
|
||||
names: Final = tool_names(peer, key, identity)
|
||||
assert call_tool(peer, key, identity, names["add"], ADD).status_code == 200
|
||||
assert len(tool_calls(first.drain())) == 1
|
||||
moved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "url": second.url})
|
||||
assert moved.status_code == 202, moved.text
|
||||
eventually(
|
||||
lambda: call_tool(peer, key, identity, names["add"], ADD),
|
||||
lambda value: value.status_code == 200 and len(tool_calls(second.drain())) == 1,
|
||||
seconds=40,
|
||||
)
|
||||
called: Final = call_tool(peer, key, identity, names["add"], ADD)
|
||||
assert called.status_code == 200, called.text
|
||||
assert tool_calls(first.drain()) == ()
|
||||
assert len(tool_calls(second.drain())) == 1
|
||||
delete_mcp(gateway, identity)
|
||||
eventually(
|
||||
lambda: peer.client.get(
|
||||
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
|
||||
),
|
||||
lambda value: value.status_code >= 400 or value.json() == [],
|
||||
seconds=40,
|
||||
deleted: Final = peer.client.get(
|
||||
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
|
||||
)
|
||||
second.drain()
|
||||
assert call_tool(peer, key, identity, names["add"], ADD).status_code >= 400
|
||||
assert deleted.status_code in (403, 404) or (deleted.status_code == 200 and deleted.json() == []), deleted.text
|
||||
assert call_tool(peer, key, identity, names["add"], ADD).status_code in (403, 404)
|
||||
assert tool_calls(second.drain()) == ()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -235,9 +235,11 @@ def _clear_proxy_database_env() -> typing.Iterator[None]:
|
|||
|
||||
|
||||
async def _initialize_proxy(config_path: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
cleanup_router_config_variables()
|
||||
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
|
||||
await initialize(config=config_path, debug=True)
|
||||
for server_id, upstream in tuple(global_mcp_server_manager.registry.items()):
|
||||
if upstream.server_name != "math_restricted":
|
||||
|
|
|
|||
|
|
@ -7990,9 +7990,13 @@ class TestGatewaySessionAdmission:
|
|||
rpm_limit=rpm_limit,
|
||||
)
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
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
|
||||
|
|
@ -10143,10 +10147,14 @@ class TestScopedSessionAdmission:
|
|||
rpm_limit=None,
|
||||
)
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
|
||||
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", 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)
|
||||
|
|
@ -10200,3 +10208,21 @@ async def test_managed_agent_permission_resolution_outage_is_not_an_unrestricted
|
|||
)
|
||||
with pytest.raises(RuntimeError, match="policy unavailable"):
|
||||
await resolution
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unreadable_empty_key_scope_cannot_gain_additive_grants(monkeypatch):
|
||||
auth = UserAPIKeyAuth(api_key="test-key", object_permission_id="key-scope")
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_get_allowed_mcp_servers_for_key",
|
||||
AsyncMock(return_value=[SpecialMCPServerNames.no_mcp_servers.value]),
|
||||
)
|
||||
monkeypatch.setattr(MCPRequestHandler, "_key_object_permission_hydrated", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[]))
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=["unrelated-server"])
|
||||
)
|
||||
access = await MCPRequestHandler.get_mcp_server_access(auth)
|
||||
assert access.server_ids == ()
|
||||
assert access.scope == "scoped"
|
||||
|
|
|
|||
|
|
@ -81,10 +81,18 @@ def config_only_mcp_manager_factory():
|
|||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _hermetic_mcp_server_registry():
|
||||
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
saved_tools = global_mcp_tool_registry.published_tools
|
||||
global_mcp_tool_registry.published_tools = {}
|
||||
saved_catalog = global_mcp_server_manager.catalog
|
||||
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
|
||||
saved_registry = dict(global_mcp_server_manager.registry)
|
||||
saved_config_servers = dict(global_mcp_server_manager.config_mcp_servers)
|
||||
saved_tool_mapping = dict(global_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
|
||||
|
|
@ -103,6 +111,8 @@ def _hermetic_mcp_server_registry():
|
|||
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear()
|
||||
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.update(saved_tool_mapping)
|
||||
global_mcp_server_manager._oauth_discovery_slots = saved_oauth_slots
|
||||
global_mcp_server_manager.catalog = saved_catalog
|
||||
global_mcp_tool_registry.published_tools = saved_tools
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
|
|||
|
|
@ -655,7 +655,10 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
|
|||
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy")
|
||||
mcp_operations.byok_credential_cache.flush_cache()
|
||||
server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True)
|
||||
monkeypatch.setattr(proxy_server, "should_load_db_object", lambda _kind: False)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
|
|||
|
|
@ -507,7 +507,7 @@ class _MapTable:
|
|||
@pytest.mark.parametrize("quoted", [False, True])
|
||||
async def test_secret_maps_create_update_round_trip(map_algorithm: str, field: str, quoted: bool) -> None:
|
||||
table: Final = _MapTable(quoted=quoted)
|
||||
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
|
||||
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
|
||||
original: Final = {"TOKEN": " sensitive-secret\n", "PREFIX": "v2:gcm:literal", "TEMPLATE": "Bearer ${TOKEN}"}
|
||||
create: Final = NewMCPServerRequest.model_validate({
|
||||
"server_id": "srv-map", "transport": "http", "url": "https://up.example.com/mcp", field: original,
|
||||
|
|
@ -580,7 +580,7 @@ async def test_secret_map_rotation_migrates_rekeys_and_preserves_corrupt(
|
|||
{"server_id": "encrypted", field: old, other: None},
|
||||
)
|
||||
prisma: Final = SimpleNamespace(db=SimpleNamespace(
|
||||
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[]))
|
||||
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[])), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))
|
||||
))
|
||||
await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key="rotated-map-key")
|
||||
assert table.rows["broken"][field] == corrupt
|
||||
|
|
@ -608,7 +608,7 @@ async def test_bulk_reads_isolate_corrupt_secret_maps(reader, field, map_algorit
|
|||
]
|
||||
snapshot = [row.model_dump() for row in rows]
|
||||
table = SimpleNamespace(find_many=AsyncMock(return_value=rows))
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
|
||||
result = await reader(prisma, ["broken", "healthy"]) if reader is get_mcp_servers else await reader(prisma)
|
||||
items = result.items if reader is get_mcp_submissions else result
|
||||
assert [row.server_id for row in items] == ["healthy"]
|
||||
|
|
@ -627,7 +627,7 @@ async def test_bulk_reads_do_not_swallow_unrelated_validation_errors(reader):
|
|||
|
||||
row = _prisma_map_row({"server_id": "invalid", "transport": "unsupported"})
|
||||
table = SimpleNamespace(find_many=AsyncMock(return_value=[row]))
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
|
||||
request = reader(prisma, ["invalid"]) if reader is get_mcp_servers else reader(prisma)
|
||||
with pytest.raises(ValidationError, match="transport"):
|
||||
await request
|
||||
|
|
@ -1545,6 +1545,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch):
|
|||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(server_id="config_faros", credentials=blob_old)]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Final
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
|
@ -92,7 +92,7 @@ def mock_mcp_client_ip():
|
|||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolate_global_mcp_registry():
|
||||
def isolate_global_mcp_registry(monkeypatch):
|
||||
"""Restore the module-global MCP server registry after each test.
|
||||
|
||||
Tests here register servers on ``global_mcp_server_manager`` directly; without a
|
||||
|
|
@ -101,6 +101,8 @@ def isolate_global_mcp_registry():
|
|||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
|
||||
monkeypatch.setattr(global_mcp_server_manager, "catalog", TargetCatalog(global_mcp_server_manager))
|
||||
snapshot = dict(global_mcp_server_manager.registry)
|
||||
yield
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
|
@ -115,7 +117,8 @@ def _mock_callback_request(base_url: str = "http://localhost:3000/"):
|
|||
and trusted ``X-Forwarded-*`` headers). A simple MagicMock with the
|
||||
right attributes is sufficient.
|
||||
"""
|
||||
req = MagicMock()
|
||||
req = MagicMock(spec=Request)
|
||||
req.client = None
|
||||
req.base_url = base_url
|
||||
req.headers = {}
|
||||
req.cookies = {}
|
||||
|
|
@ -739,7 +742,8 @@ async def test_token_endpoint_forwards_code_verifier():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_client_without_mcp_server_name_returns_dummy():
|
||||
@pytest.mark.parametrize("server_name", [None, "missing"])
|
||||
async def test_register_client_without_mcp_server_name_returns_dummy(server_name):
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
|
|
@ -762,17 +766,18 @@ async def test_register_client_without_mcp_server_name_returns_dummy():
|
|||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
):
|
||||
result = await register_client(request=mock_request)
|
||||
result = await register_client(request=mock_request, mcp_server_name=server_name)
|
||||
|
||||
assert result == {
|
||||
"client_id": "dummy_client",
|
||||
"client_id": server_name or "dummy_client",
|
||||
"client_secret": "dummy",
|
||||
"redirect_uris": ["https://proxy.litellm.example/callback"],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_client_returns_existing_server_credentials():
|
||||
@pytest.mark.parametrize("use_root", [False, True])
|
||||
async def test_register_client_returns_existing_server_credentials(use_root):
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
|
|
@ -812,7 +817,9 @@ async def test_register_client_returns_existing_server_credentials():
|
|||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
):
|
||||
result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name)
|
||||
result = await register_client(
|
||||
request=mock_request, mcp_server_name=None if use_root else oauth2_server.server_name
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
|
@ -4885,7 +4892,7 @@ async def test_token_endpoint_authorization_code_missing_code():
|
|||
)
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8245,7 +8252,7 @@ async def test_authorize_endpoint_rejects_non_oauth2_server():
|
|||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8284,7 +8291,7 @@ async def test_token_endpoint_rejects_non_oauth2_server():
|
|||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8326,7 +8333,7 @@ async def test_register_client_rejects_non_oauth2_server():
|
|||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8399,7 +8406,7 @@ async def test_oauth_authorization_server_404_for_non_oauth2_server():
|
|||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8447,7 +8454,7 @@ async def test_oauth_protected_resource_passthrough_none_auth_not_404():
|
|||
)
|
||||
global_mcp_server_manager.registry[passthrough_server.server_id] = passthrough_server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8481,7 +8488,7 @@ async def test_oauth_protected_resource_404_for_unknown_server_name():
|
|||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -8509,7 +8516,7 @@ async def test_oauth_authorization_server_404_for_unknown_server_name():
|
|||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -9518,7 +9525,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()
|
||||
|
|
@ -9535,7 +9542,9 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
|
|||
)
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.writer_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
hydrate_spy = AsyncMock()
|
||||
with (
|
||||
|
|
@ -9543,7 +9552,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()
|
||||
|
||||
|
|
@ -9814,7 +9823,7 @@ async def test_authorize_wall_names_the_fix_for_urlless_servers():
|
|||
auth_type=MCPAuth.oauth2,
|
||||
spec_path="https://example.com/openapi.yaml",
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -9852,7 +9861,7 @@ async def test_token_wall_names_the_fix_for_urlless_servers():
|
|||
spec_path="https://example.com/openapi.yaml",
|
||||
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -9892,7 +9901,7 @@ async def test_register_wall_names_the_fix_for_urlless_servers():
|
|||
auth_type=MCPAuth.oauth2,
|
||||
spec_path="https://example.com/openapi.yaml",
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -9931,7 +9940,7 @@ async def test_authorize_wall_points_at_discovery_failure_for_url_servers():
|
|||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -9967,7 +9976,7 @@ async def test_token_wall_points_at_discovery_failure_for_url_servers():
|
|||
auth_type=MCPAuth.oauth2,
|
||||
authorization_url="https://idp.example.com/authorize",
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -10007,7 +10016,7 @@ async def test_authorize_wall_names_the_issuer_for_anchored_servers():
|
|||
issuer="https://idp.example.com",
|
||||
issuer_is_anchored=True,
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -10054,7 +10063,7 @@ async def test_authorize_uses_admin_entered_github_oauth_urls_after_issuer_yield
|
|||
configured_authorization_url="https://github.com/login/oauth/authorize",
|
||||
configured_token_url="https://github.com/login/oauth/access_token",
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -10076,7 +10085,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
|
||||
|
|
@ -10093,7 +10102,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
|
||||
|
|
@ -10146,7 +10155,7 @@ async def test_token_exchange_with_configured_token_url_never_joins_discovery(mo
|
|||
"get_async_httpx_client",
|
||||
lambda llm_provider: fake_http_client,
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -10271,7 +10280,7 @@ async def test_bridge_authorize_relays_with_registration_url_resolved_by_deferre
|
|||
"ensure_oauth_metadata_discovered",
|
||||
resolve_discovery,
|
||||
)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
|
|
@ -11593,7 +11602,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 []
|
||||
|
|
@ -12807,6 +12816,65 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
|
|||
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ["create", "update", "delete"])
|
||||
async def test_authorize_observes_committed_peer_server_changes(monkeypatch, change):
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-consistency-test-key")
|
||||
stamp = datetime.now(timezone.utc)
|
||||
old_server = _create_id_lookup_oauth2_server()
|
||||
old_server.updated_at = stamp
|
||||
row = LiteLLM_MCPServerTable(
|
||||
server_id=old_server.server_id,
|
||||
server_name=old_server.server_name,
|
||||
alias=old_server.alias,
|
||||
url="https://upstream.example/mcp",
|
||||
transport="http",
|
||||
auth_type="oauth2",
|
||||
authorization_url="https://new-provider.example/authorize",
|
||||
token_url="https://new-provider.example/token",
|
||||
scopes=["read"],
|
||||
credentials={"client_id": "current-client", "client_secret": "current-secret"},
|
||||
created_at=stamp,
|
||||
updated_at=stamp + timedelta(seconds=1),
|
||||
)
|
||||
read_rows = AsyncMock(return_value=[] if change == "delete" else [row])
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
|
||||
prisma.writer_db = prisma.db
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(
|
||||
global_mcp_server_manager, "registry", {} if change == "create" else {old_server.server_id: old_server}
|
||||
)
|
||||
monkeypatch.setattr(global_mcp_server_manager, "config_mcp_servers", {})
|
||||
request = Request(
|
||||
{"type": "http", "scheme": "https", "server": ("gateway.example", 443), "path": "/authorize", "headers": []}
|
||||
)
|
||||
|
||||
if change == "delete":
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await discoverable_endpoints.authorize(
|
||||
request, "http://localhost/callback", mcp_server_name=old_server.server_id
|
||||
)
|
||||
assert exc.value.status_code == 404
|
||||
else:
|
||||
response = await discoverable_endpoints.authorize(
|
||||
request, "http://localhost/callback", mcp_server_name=old_server.server_id
|
||||
)
|
||||
assert response.status_code == 307
|
||||
assert response.headers["location"].startswith("https://new-provider.example/authorize?")
|
||||
assert "client_id=current-client" in response.headers["location"]
|
||||
read_rows.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_server_drops_cached_upstream_oauth_metadata():
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
|
|
|
|||
|
|
@ -66,8 +66,7 @@ async def _call_block(logging_obj, order: list, *, user_api_key_auth=mock.sentin
|
|||
proxy_logging_obj = mock.MagicMock()
|
||||
proxy_logging_obj.post_call_failure_hook.side_effect = _record_post_call_failure_hook
|
||||
|
||||
fake_proxy_server = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy_server.proxy_logging_obj = proxy_logging_obj # pyright: ignore[reportAttributeAccessIssue]
|
||||
fake_proxy_server = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None)
|
||||
|
||||
with mock.patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}):
|
||||
with contextlib.suppress(HTTPException):
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
import os
|
||||
import pytest
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
from contextlib import asynccontextmanager
|
||||
from contextlib import asynccontextmanager, nullcontext
|
||||
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
|
|
@ -918,6 +918,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
|
||||
# Create a mock manager
|
||||
mock_manager = AsyncMock()
|
||||
mock_manager.catalog.operation = nullcontext
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["server1_id", "server2_id"]
|
||||
)
|
||||
|
|
@ -946,6 +947,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
# Test Case 2: Without specific MCP servers
|
||||
# Create a different mock manager for the second test case
|
||||
mock_manager_2 = AsyncMock()
|
||||
mock_manager_2.catalog.operation = nullcontext
|
||||
mock_manager_2.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["server1_id", "server2_id"]
|
||||
)
|
||||
|
|
@ -996,6 +998,7 @@ async def test_get_tools_from_mcp_servers():
|
|||
# Test Case 3: With specific MCP servers and access groups
|
||||
# Create a mock manager
|
||||
mock_manager = AsyncMock()
|
||||
mock_manager.catalog.operation = nullcontext
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["server1_id", "server2_id", "server3_id"]
|
||||
)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -5405,7 +5405,9 @@ class TestMCPServerManagerReload:
|
|||
db_row = _make_db_mcp_server("server-1", timestamp)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.writer_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
|
|
@ -5447,7 +5449,9 @@ class TestMCPServerManagerReload:
|
|||
)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.writer_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
|
|
@ -5461,7 +5465,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
|
||||
|
|
@ -5500,9 +5504,11 @@ class TestMCPServerManagerReload:
|
|||
return another_healthy_server
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.writer_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
|
||||
return_value=[healthy_row, bad_row, another_healthy_row]
|
||||
)
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
|
|
@ -5513,7 +5519,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()
|
||||
|
|
@ -5572,7 +5578,9 @@ class TestMCPServerManagerReload:
|
|||
raise RuntimeError("blocked address")
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.writer_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[healthy_row, bad_openapi_row])
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
|
|
@ -5585,7 +5593,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"),
|
||||
|
|
@ -7754,8 +7762,8 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
|
|||
),
|
||||
patch.object(
|
||||
mcp_operations.global_mcp_server_manager,
|
||||
"get_registry",
|
||||
return_value={
|
||||
"registry",
|
||||
{
|
||||
requested_server.server_id: requested_server,
|
||||
collision_server.server_id: collision_server,
|
||||
},
|
||||
|
|
@ -8019,6 +8027,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
fake_tool.server_id = fake_server.server_id
|
||||
|
||||
start_time = datetime.now(timezone.utc)
|
||||
litellm_logging_obj, _ = function_setup(
|
||||
|
|
@ -9000,6 +9009,7 @@ async def test_stateful_mcp_tool_call_uses_current_requests_otel_destinations(_m
|
|||
server = MCPServer(
|
||||
server_id="otel-context-test",
|
||||
name="otelcontext",
|
||||
server_name="otelcontext",
|
||||
transport=MCPTransport.http,
|
||||
allow_all_keys=True,
|
||||
)
|
||||
|
|
@ -9091,6 +9101,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_queries_active_rows(
|
|||
row.server_id = "submitted-1"
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
|
||||
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
result = await get_active_submitted_mcp_server_ids_for_user(prisma_client, "submitter-user")
|
||||
|
||||
|
|
@ -9111,6 +9122,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock()
|
||||
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
assert await get_active_submitted_mcp_server_ids_for_user(prisma_client, "") == []
|
||||
prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited()
|
||||
|
|
@ -10026,7 +10038,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
|
||||
|
|
|
|||
|
|
@ -995,6 +995,7 @@ class TestRotateCredentials:
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
|
||||
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
|
||||
|
||||
|
|
@ -1043,6 +1044,7 @@ class TestRotateCredentials:
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
|
||||
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
|
||||
|
||||
|
|
|
|||
|
|
@ -134,6 +134,7 @@ def _byok_key_row(server_id):
|
|||
def _mock_prisma(null_rows, token_rows):
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=null_rows)
|
||||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_mcpservertable.update_many = AsyncMock(return_value=MagicMock())
|
||||
mock_prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=token_rows)
|
||||
return mock_prisma
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ def _row(**overrides):
|
|||
def _prisma(rows):
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows)
|
||||
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_mcpservertable.update = AsyncMock()
|
||||
return prisma_client
|
||||
|
||||
|
|
|
|||
|
|
@ -622,6 +622,7 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c
|
|||
|
||||
body = '["a","b"]'
|
||||
tool = MagicMock()
|
||||
tool.server_id = None
|
||||
tool.handler = AsyncMock(return_value=parse_http_body(body))
|
||||
with patch.object(global_mcp_tool_registry, "get_tool", return_value=tool):
|
||||
result = await operations._handle_local_mcp_tool("reports-list_tags", {}, WireCompat(compat))
|
||||
|
|
@ -778,3 +779,172 @@ async def test_list_mcp_tools_records_the_catalog_only_when_asked(
|
|||
manager._drop_listed_tools(server.server_id)
|
||||
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
|
||||
assert (listed is not None) is recorded
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discovery_keeps_one_catalog_revision_across_concurrent_listings(monkeypatch):
|
||||
from mcp.types import (
|
||||
DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult,
|
||||
ListResourceTemplatesResult,
|
||||
)
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
manager = MCPServerManager()
|
||||
original = MCPServer(server_id="catalog-server", name="before", transport=MCPTransport.http)
|
||||
updated = original.model_copy(update={"name": "after"})
|
||||
manager.registry = {original.server_id: original}
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
|
||||
observed = []
|
||||
responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]),
|
||||
ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[]))
|
||||
|
||||
def listing(index):
|
||||
async def run(*args, **kwargs):
|
||||
async with manager.catalog.operation():
|
||||
observed.append(manager.get_mcp_server_by_id(original.server_id).name)
|
||||
manager.registry = {updated.server_id: updated}
|
||||
return responses[index]
|
||||
return run
|
||||
|
||||
with (
|
||||
patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)),
|
||||
patch.object(operations, "_execute_list_prompts", side_effect=listing(1)),
|
||||
patch.object(operations, "_execute_list_resources", side_effect=listing(2)),
|
||||
patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)),
|
||||
):
|
||||
result = await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped")))
|
||||
assert result.capabilities.model_dump(exclude_none=True) == {}
|
||||
assert observed == ["before"] * 4
|
||||
async with manager.catalog.operation():
|
||||
assert manager.get_mcp_server_by_id(original.server_id).name == "after"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("changed_server_id", ["private", "allowed"])
|
||||
async def test_local_handler_freshness_tracks_registered_owner_with_overlapping_alias(monkeypatch, changed_server_id):
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
manager = MCPServerManager()
|
||||
now = datetime.now(timezone.utc)
|
||||
private = MCPServer(server_id="private", name="billing", alias="billing", transport=MCPTransport.http, updated_at=now)
|
||||
allowed = MCPServer(server_id="allowed", name="billing_admin", alias="billing-admin", transport=MCPTransport.http, updated_at=now)
|
||||
manager.config_mcp_servers = {server.server_id: server for server in (private, allowed)}
|
||||
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
|
||||
handler = AsyncMock(return_value="private result")
|
||||
global_mcp_tool_registry.register_tool("billing-admin-export", "export", {}, handler, server_id=private.server_id)
|
||||
|
||||
async with manager.catalog.operation():
|
||||
manager.config_mcp_servers[changed_server_id] = manager.config_mcp_servers[changed_server_id].model_copy(update={"updated_at": now + timedelta(seconds=1)})
|
||||
if changed_server_id == "private":
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await operations._handle_local_mcp_tool("billing-admin-export", {})
|
||||
assert denied.value.status_code == 503
|
||||
handler.assert_not_awaited()
|
||||
else:
|
||||
result = await operations._handle_local_mcp_tool("billing-admin-export", {})
|
||||
assert result.is_error is False
|
||||
handler.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("owner_present", [False, True])
|
||||
@pytest.mark.parametrize("requested_server", [False, True])
|
||||
async def test_local_call_cannot_borrow_another_servers_authority(
|
||||
monkeypatch: pytest.MonkeyPatch, owner_present: bool, requested_server: bool
|
||||
) -> None:
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
allowed: Final = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http)
|
||||
owner: Final = MCPServer(server_id="private", name="private", transport=MCPTransport.http)
|
||||
manager.config_mcp_servers = {
|
||||
server.server_id: server for server in ((allowed, owner) if owner_present else (allowed,))
|
||||
}
|
||||
handler: Final = AsyncMock(return_value="private result")
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
|
||||
tool_name: Final = "export" if requested_server else "private-export"
|
||||
global_mcp_tool_registry.register_tool(tool_name, "export", {}, handler, server_id=owner.server_id)
|
||||
|
||||
async with manager.catalog.operation():
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await operations._execute_mcp_tool(
|
||||
name=tool_name if requested_server else "allowed-private-export",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[allowed],
|
||||
start_time=datetime.now(),
|
||||
requested_server_id=allowed.server_id if requested_server else None,
|
||||
)
|
||||
assert denied.value.status_code == 403
|
||||
handler.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("server_owned", [False, True])
|
||||
@pytest.mark.parametrize("requested_server", [False, True])
|
||||
async def test_local_call_preserves_matching_and_legacy_handlers(
|
||||
monkeypatch: pytest.MonkeyPatch, server_owned: bool, requested_server: bool
|
||||
) -> None:
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http)
|
||||
manager.config_mcp_servers = {server.server_id: server}
|
||||
handler: Final = AsyncMock(return_value="allowed result")
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
|
||||
global_mcp_tool_registry.register_tool(
|
||||
"export", "export", {}, handler, server_id=server.server_id if server_owned else None
|
||||
)
|
||||
|
||||
async with manager.catalog.operation():
|
||||
result: Final = await operations._execute_mcp_tool(
|
||||
name="export" if requested_server else "allowed-export",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
requested_server_id=server.server_id if requested_server else None,
|
||||
)
|
||||
assert result.is_error is False
|
||||
handler.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_handler_rejects_an_owner_absent_from_the_catalog(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
manager: Final = MCPServerManager()
|
||||
handler: Final = AsyncMock(return_value="private result")
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
|
||||
global_mcp_tool_registry.register_tool("private-export", "export", {}, handler, server_id="private")
|
||||
|
||||
async with manager.catalog.operation():
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await operations._handle_local_mcp_tool("private-export", {})
|
||||
assert denied.value.status_code == 503
|
||||
handler.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -1509,6 +1509,7 @@ class TestListToolsRestAPI:
|
|||
mcp_info={"server_name": "stub"},
|
||||
)
|
||||
stub_server.available_on_public_internet = True
|
||||
monkeypatch.setattr(rest_endpoints.global_mcp_server_manager, "registry", {"server-1": stub_server})
|
||||
|
||||
mock_transport_ctx = AsyncMock()
|
||||
mock_transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock()))
|
||||
|
|
@ -1577,9 +1578,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,32 @@ 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.writer_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], []))
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
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,6 @@
|
|||
from contextlib import nullcontext
|
||||
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -16,7 +19,7 @@ import httpx
|
|||
import pytest
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
from respx import MockRouter
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm
|
||||
|
|
@ -152,7 +155,7 @@ def create_mcp_router_test_client() -> TestClient:
|
|||
|
||||
|
||||
def patch_proxy_general_settings(settings: dict):
|
||||
fake_proxy_server_module = types.SimpleNamespace(general_settings=settings)
|
||||
fake_proxy_server_module = types.SimpleNamespace(general_settings=settings, prisma_client=None)
|
||||
return patch.dict(
|
||||
sys.modules,
|
||||
{"litellm.proxy.proxy_server": fake_proxy_server_module},
|
||||
|
|
@ -2732,14 +2735,15 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
algorithm="HS256",
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.client = None
|
||||
mock_request.headers = {}
|
||||
mock_request.cookies = {"token": token_cookie}
|
||||
|
||||
expected_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key=api_key_in_cookie
|
||||
)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=master_key)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=master_key, prisma_client=None, general_settings={})
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2772,7 +2776,8 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
)
|
||||
|
||||
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.client = None
|
||||
mock_request.headers = {"Authorization": "Bearer sk-header-key"}
|
||||
mock_request.cookies = {}
|
||||
|
||||
|
|
@ -2806,7 +2811,8 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
)
|
||||
|
||||
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.client = None
|
||||
mock_request.headers = {}
|
||||
mock_request.cookies = {}
|
||||
mock_request.path_params = {"server_id": "server-1"}
|
||||
|
|
@ -2816,7 +2822,8 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = non_oauth_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None)
|
||||
mock_manager.catalog.resolve = AsyncMock(return_value=non_oauth_server)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None, general_settings={})
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2854,7 +2861,8 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
)
|
||||
|
||||
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
mock_request = MagicMock()
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.client = None
|
||||
mock_request.headers = {}
|
||||
mock_request.cookies = {}
|
||||
mock_request.path_params = {"server_id": "server-1"}
|
||||
|
|
@ -2868,7 +2876,8 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = internal_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None)
|
||||
mock_manager.catalog.resolve = AsyncMock(return_value=internal_server)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None, general_settings={})
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2925,6 +2934,65 @@ 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()
|
||||
prisma.writer_db = prisma.db
|
||||
|
||||
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)
|
||||
prisma.db.litellm_config.find_unique = 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 (
|
||||
|
|
@ -2941,8 +3009,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",
|
||||
|
|
@ -2963,7 +3031,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,
|
||||
|
|
@ -2991,8 +3059,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",
|
||||
|
|
@ -3096,8 +3164,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",
|
||||
|
|
@ -3242,8 +3310,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",
|
||||
|
|
@ -3301,8 +3369,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",
|
||||
|
|
@ -3355,8 +3423,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",
|
||||
|
|
@ -3405,8 +3473,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",
|
||||
|
|
@ -3453,8 +3521,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",
|
||||
|
|
@ -3493,8 +3561,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",
|
||||
|
|
@ -3538,8 +3606,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",
|
||||
|
|
@ -3561,7 +3629,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,
|
||||
|
|
@ -3592,8 +3660,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",
|
||||
|
|
@ -3615,7 +3683,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,
|
||||
|
|
@ -3652,8 +3720,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",
|
||||
|
|
@ -3671,7 +3739,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,
|
||||
|
|
@ -3715,8 +3783,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",
|
||||
|
|
@ -3760,8 +3828,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",
|
||||
|
|
@ -8307,6 +8375,38 @@ async def test_config_server_edit_preserves_api_contract_without_creating_rows(r
|
|||
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.writer_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
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
|
||||
|
||||
|
||||
class TestDuplicateIdentifierRejection:
|
||||
"""server_name/alias must be unique across live servers, case-insensitive.
|
||||
|
||||
|
|
@ -8801,11 +8901,15 @@ def _mock_mcp_resolution_prisma_client(
|
|||
object_permission: LiteLLM_ObjectPermissionTable | None = None,
|
||||
) -> MagicMock:
|
||||
prisma: Final = MagicMock()
|
||||
prisma.writer_db = prisma.db
|
||||
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=SimpleNamespace(object_permission=key_permission)
|
||||
)
|
||||
|
||||
def matches_server_filter(name: str, condition: object) -> bool:
|
||||
if name == "OR" and condition == [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]:
|
||||
return server.approval_status in (None, "active", "approved")
|
||||
if name == "submitted_by":
|
||||
return server.submitted_by == condition
|
||||
if name == "server_id":
|
||||
|
|
@ -9213,7 +9317,7 @@ class TestMCPServerResolutionRegressions:
|
|||
allowed_id: Final = "lit3974-allowed-config"
|
||||
denied_id: Final = "lit3974-denied-config"
|
||||
prisma: Final = _mock_mcp_resolution_prisma_client(
|
||||
generate_mock_mcp_server_db_record(server_id=denied_id),
|
||||
generate_mock_mcp_server_db_record(server_id=denied_id, alias="denied_alias"),
|
||||
LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="lit3974-alias-permission",
|
||||
mcp_servers=[allowed_id],
|
||||
|
|
@ -9376,6 +9480,15 @@ class TestMCPServerResolutionCharacterization:
|
|||
}
|
||||
)
|
||||
prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=hidden)
|
||||
eligible: Final = approval_status in (None, "active", "approved")
|
||||
existing_find_many: Final = prisma.db.litellm_mcpservertable.find_many
|
||||
|
||||
async def current_rows(**kwargs: object) -> list[LiteLLM_MCPServerTable]:
|
||||
if kwargs.get("where") == {"OR": [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]}:
|
||||
return [hidden] if eligible else []
|
||||
return await existing_find_many(**kwargs)
|
||||
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=current_rows)
|
||||
if not registered:
|
||||
manager.config_mcp_servers = {}
|
||||
health: Final = AsyncMock(return_value=hidden)
|
||||
|
|
@ -9390,7 +9503,13 @@ class TestMCPServerResolutionCharacterization:
|
|||
patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "view_all"}),
|
||||
):
|
||||
listed: Final = await mgmt_endpoints.fetch_all_mcp_servers(auth, team_id=None)
|
||||
assert (server_id in {item.server_id for item in listed}) is registered
|
||||
assert (server_id in {item.server_id for item in listed}) is (registered or eligible)
|
||||
if caller != "admin" and eligible:
|
||||
public_detail: Final = await mgmt_endpoints.fetch_mcp_server(_make_mock_request(), server_id, auth)
|
||||
assert public_detail.credentials is None
|
||||
assert public_detail.url is None
|
||||
assert public_detail.static_headers is None
|
||||
return
|
||||
if caller != "admin":
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await mgmt_endpoints.fetch_mcp_server(_make_mock_request(), server_id, auth)
|
||||
|
|
@ -10491,7 +10610,7 @@ class TestMCPServerResolutionCharacterization:
|
|||
effects: Final = _ResolutionEffects()
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_cache_temporary_mcp_server_in_redis,
|
||||
_get_cached_temporary_mcp_server_or_404,
|
||||
_oauth_server_operation,
|
||||
_TemporaryMCPServerEntry,
|
||||
)
|
||||
|
||||
|
|
@ -10576,11 +10695,8 @@ class TestMCPServerResolutionCharacterization:
|
|||
):
|
||||
if expected_status in (403, 404):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id,
|
||||
auth,
|
||||
request=_make_mock_request(),
|
||||
)
|
||||
async with _oauth_server_operation(server_id, auth, request=_make_mock_request()):
|
||||
pytest.fail("Denied OAuth resolution must not enter the operation")
|
||||
|
||||
expected_detail: Final = (
|
||||
{"error": f"MCP server {server_id} not found"}
|
||||
|
|
@ -10594,20 +10710,16 @@ class TestMCPServerResolutionCharacterization:
|
|||
effects.assert_no_writes()
|
||||
assert httpx_mock.calls.call_count == 0, f"{source}/{caller}: no upstream HTTP"
|
||||
else:
|
||||
resolved: Final = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id,
|
||||
auth,
|
||||
request=_make_mock_request(),
|
||||
)
|
||||
expected_alias: Final = (
|
||||
db_server.alias
|
||||
if source == "temp_draft"
|
||||
else "LIT3974 OAuth"
|
||||
if source == "config"
|
||||
else temp_server.alias
|
||||
)
|
||||
assert resolved.server_id == server_id, f"{source}/{caller}: resolved OAuth server"
|
||||
assert resolved.alias == expected_alias, f"{source}/{caller}: resolved OAuth display name"
|
||||
async with _oauth_server_operation(server_id, auth, request=_make_mock_request()) as resolved:
|
||||
expected_alias: Final = (
|
||||
db_server.alias
|
||||
if source == "temp_draft"
|
||||
else "LIT3974 OAuth"
|
||||
if source == "config"
|
||||
else temp_server.alias
|
||||
)
|
||||
assert resolved.server_id == server_id, f"{source}/{caller}: resolved OAuth server"
|
||||
assert resolved.alias == expected_alias, f"{source}/{caller}: resolved OAuth display name"
|
||||
finally:
|
||||
mgmt_endpoints.litellm.cache = original_cache
|
||||
|
||||
|
|
|
|||
|
|
@ -3,8 +3,7 @@ Tests for the dynamic_mcp_route handler in proxy_server.py.
|
|||
|
||||
Covers the resolution order:
|
||||
1. Registered MCP server alias → forwards to /mcp/{name}
|
||||
2. Comma-separated list → short-circuits before any DB call;
|
||||
forwarded to /mcp/{segment}
|
||||
2. Comma-separated list → forwards to /mcp/{segment}
|
||||
3. Toolset name (cached) → sets toolset scope, forwards to /mcp
|
||||
4. MCP access group tag (cached) → forwards to /mcp/{name} when the group
|
||||
resolves to at least one server
|
||||
|
|
@ -611,3 +610,66 @@ def test_aggregate_mcp_route_returns_404_when_mcp_unavailable():
|
|||
|
||||
assert response.status_code == 404
|
||||
assert handler_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ["create", "rename", "delete"])
|
||||
@pytest.mark.parametrize("csv", [False, True])
|
||||
async def test_dynamic_route_observes_committed_peer_catalog_changes(monkeypatch, change, csv):
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
from starlette.responses import Response
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable
|
||||
|
||||
old_row = LiteLLM_MCPServerTable(
|
||||
server_id="catalog-route", server_name="previous", alias="previous", transport="http",
|
||||
url="https://previous.example/mcp", updated_at=datetime(2026, 1, 1),
|
||||
)
|
||||
new_row = old_row.model_copy(update={
|
||||
"server_name": "current", "alias": "current", "url": "https://current.example/mcp",
|
||||
"updated_at": datetime(2026, 1, 2),
|
||||
})
|
||||
|
||||
async def find_rows(*, where):
|
||||
if "mcp_access_groups" in where or change == "delete":
|
||||
return []
|
||||
return [new_row]
|
||||
|
||||
read_rows = AsyncMock(side_effect=find_rows)
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(
|
||||
litellm_mcpservertable=SimpleNamespace(find_many=read_rows),
|
||||
litellm_mcptoolsettable=SimpleNamespace(find_first=AsyncMock(return_value=None)),
|
||||
litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
|
||||
))
|
||||
prisma.writer_db = prisma.db
|
||||
manager = mcp_server_manager.MCPServerManager()
|
||||
if change != "create":
|
||||
manager.registry = {old_row.server_id: await manager.build_mcp_server_from_table(old_row)}
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
|
||||
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
|
||||
name = "previous" if change == "delete" else "current"
|
||||
segment = f"{name},missing" if csv else name
|
||||
request = _make_request(f"/{segment}/mcp")
|
||||
|
||||
async def forwarded(path_segment, request):
|
||||
if change != "delete":
|
||||
assert manager.get_mcp_server_by_name(path_segment).url == new_row.url
|
||||
return Response(content=b"forwarded", status_code=200)
|
||||
|
||||
with patch(_FORWARD, new=AsyncMock(side_effect=forwarded)) as forward:
|
||||
if change == "delete":
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await proxy_server.dynamic_mcp_route(segment, request)
|
||||
assert exc.value.status_code == 404
|
||||
forward.assert_not_awaited()
|
||||
else:
|
||||
response = await proxy_server.dynamic_mcp_route(segment, request)
|
||||
assert response.status_code == 200
|
||||
forward.assert_awaited_once_with(name, request)
|
||||
assert sum("OR" in call.kwargs["where"] for call in read_rows.await_args_list) == 1
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import subprocess
|
|||
import sys
|
||||
import textwrap
|
||||
import types
|
||||
from contextlib import nullcontext
|
||||
from typing import Any, Final, Literal, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -42,10 +43,11 @@ class _DummyMCPResult:
|
|||
|
||||
def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
|
||||
"""Patch MCP globals so _execute_tool_calls can run in tests."""
|
||||
proxy_module = types.SimpleNamespace(proxy_logging_obj=object())
|
||||
proxy_module = types.SimpleNamespace(proxy_logging_obj=object(), prisma_client=None)
|
||||
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
|
||||
|
|
@ -63,7 +65,7 @@ def _setup_proxy_logging(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
|
|||
"""Patch proxy_logging_obj so failure hook can be asserted."""
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||||
proxy_module = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj)
|
||||
proxy_module = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None)
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
|
||||
return proxy_logging_obj.post_call_failure_hook
|
||||
|
||||
|
|
@ -386,6 +388,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")),
|
||||
)
|
||||
|
|
@ -504,6 +507,7 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch):
|
|||
isError=False,
|
||||
)
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(return_value=result),
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
|
|
@ -545,6 +549,7 @@ async def test_execute_tool_calls_returns_proxy_result_without_logging(monkeypat
|
|||
)
|
||||
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(return_value=result),
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
|
|
@ -578,6 +583,7 @@ async def test_execute_tool_calls_passes_logging_details_to_proxy_hook(monkeypat
|
|||
)
|
||||
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(return_value=result),
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
|
|
@ -613,6 +619,7 @@ async def test_execute_tool_calls_continues_when_post_call_logging_fails(monkeyp
|
|||
|
||||
result = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(return_value=result),
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
|
|
@ -668,6 +675,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=[]),
|
||||
|
|
@ -725,6 +733,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=[]),
|
||||
|
|
@ -1287,6 +1296,7 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
"authorization": "proxy-sentinel",
|
||||
}
|
||||
manager: Final = 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=[]),
|
||||
|
|
@ -1347,6 +1357,7 @@ async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
|
|||
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
|
||||
]
|
||||
fake_manager: Final = 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=[]),
|
||||
|
|
@ -1485,6 +1496,7 @@ async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch
|
|||
|
||||
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
|
||||
return 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=[]),
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import sys
|
||||
import types
|
||||
from contextlib import nullcontext
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
|
|
@ -77,6 +78,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
|
|||
"""Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests."""
|
||||
call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False))
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=call_tool,
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
|
|
@ -89,7 +91,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
|
|||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"litellm.proxy.proxy_server",
|
||||
types.SimpleNamespace(proxy_logging_obj=MagicMock()),
|
||||
types.SimpleNamespace(proxy_logging_obj=MagicMock(), prisma_client=None),
|
||||
)
|
||||
return call_tool
|
||||
|
||||
|
|
|
|||
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -24697,7 +24697,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
@ -24708,7 +24708,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
@ -24719,7 +24719,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
@ -24730,7 +24730,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
@ -24741,7 +24741,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
@ -24752,7 +24752,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
@ -24763,7 +24763,7 @@ export interface paths {
|
|||
*
|
||||
* Resolution order:
|
||||
* 1. Registered MCP server alias / name
|
||||
* 2. Comma-separated list (short-circuits before any DB call)
|
||||
* 2. Comma-separated list
|
||||
* 3. Toolset name (DB lookup, cached)
|
||||
* 4. MCP access group tag (DB lookup, cached)
|
||||
*/
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue