mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): refresh server catalog for each gateway operation
This commit is contained in:
parent
2ea43214c9
commit
c6b5f6da7f
16 changed files with 478 additions and 69 deletions
|
|
@ -13,6 +13,7 @@ from typing_extensions import assert_never
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
get_passthrough_resource_metadata_url,
|
||||
get_passthrough_www_authenticate,
|
||||
|
|
@ -398,6 +399,7 @@ class MCPRequestHandler:
|
|||
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
|
||||
|
||||
@staticmethod
|
||||
@with_mcp_catalog
|
||||
async def process_mcp_request(
|
||||
scope: Scope,
|
||||
) -> tuple[
|
||||
|
|
|
|||
|
|
@ -689,7 +689,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
|
||||
|
|
|
|||
94
litellm/proxy/_experimental/mcp_server/catalog.py
Normal file
94
litellm/proxy/_experimental/mcp_server/catalog.py
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
|
||||
from contextlib import asynccontextmanager
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from functools import wraps
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
_P: Final = ParamSpec("_P")
|
||||
_R: Final = TypeVar("_R")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _CatalogScope:
|
||||
servers: Mapping[str, MCPServer]
|
||||
owner_task_id: int
|
||||
active: bool = True
|
||||
|
||||
|
||||
class TargetCatalog:
|
||||
def __init__(self, manager: MCPServerManager) -> None:
|
||||
self._manager = manager
|
||||
self._reload_lock = asyncio.Lock()
|
||||
self._scope: ContextVar[_CatalogScope | None] = ContextVar("mcp_catalog_scope", default=None)
|
||||
|
||||
@property
|
||||
def current(self) -> Mapping[str, MCPServer] | None:
|
||||
scope: Final = self._scope.get()
|
||||
return scope.servers if scope is not None and scope.active else None
|
||||
|
||||
async def refresh(self) -> None:
|
||||
async with self._reload_lock:
|
||||
await self._manager._reload_servers_from_database() # pyright: ignore[reportPrivateUsage] # existing staged loader
|
||||
|
||||
async def list(self) -> Mapping[str, MCPServer]:
|
||||
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime proxy dependency
|
||||
|
||||
scope: Final = self._scope.get()
|
||||
if scope is not None and scope.active and scope.owner_task_id == id(asyncio.current_task()):
|
||||
return scope.servers
|
||||
async with self._reload_lock:
|
||||
if prisma_client is not None:
|
||||
try:
|
||||
await self._manager._reload_servers_from_database(reuse_unchanged=True) # pyright: ignore[reportPrivateUsage] # existing staged loader
|
||||
except Exception as exc: # noqa: BLE001 # never serve an unverified database snapshot
|
||||
from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency
|
||||
|
||||
raise HTTPException(
|
||||
status_code=503, detail="MCP server configuration could not be refreshed"
|
||||
) from exc
|
||||
return MappingProxyType(self._manager.config_mcp_servers | self._manager.registry)
|
||||
|
||||
@asynccontextmanager
|
||||
async def operation(self) -> AsyncIterator[Mapping[str, MCPServer]]:
|
||||
existing: Final = self._scope.get()
|
||||
if existing is not None and existing.active and existing.owner_task_id == id(asyncio.current_task()):
|
||||
yield existing.servers
|
||||
return
|
||||
scope: Final = _CatalogScope(await self.list(), id(asyncio.current_task()))
|
||||
token: Final = self._scope.set(scope)
|
||||
try:
|
||||
yield scope.servers
|
||||
finally:
|
||||
scope.active = False
|
||||
self._scope.reset(token)
|
||||
|
||||
async def resolve(self, lookup: str, client_ip: str | None = None) -> MCPServer | None:
|
||||
async with self.operation():
|
||||
return self._manager.get_mcp_server_by_name(
|
||||
lookup, client_ip=client_ip
|
||||
) or self._manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
|
||||
|
||||
|
||||
def with_mcp_catalog(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
|
||||
@wraps(function)
|
||||
async def wrapped(
|
||||
*args: _P.args,
|
||||
**kwargs: _P.kwargs, # kwargs-ok: ParamSpec preserves each wrapped endpoint keyword contract
|
||||
) -> _R:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # manager imports catalog
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
async with global_mcp_server_manager.catalog.operation():
|
||||
return await function(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
|
@ -36,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
|
|||
can_store_oauth_credential,
|
||||
oauth_authorization_uses_gateway_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
|
||||
from litellm.proxy._experimental.mcp_server.faults import (
|
||||
CallerRejected,
|
||||
CredentialSource,
|
||||
|
|
@ -1875,6 +1876,7 @@ async def register_client_with_server(
|
|||
|
||||
|
||||
@router.get("/authorize/mcp-session")
|
||||
@with_mcp_catalog
|
||||
async def authorize_mcp_session(
|
||||
request: Request,
|
||||
redirect_uri: str,
|
||||
|
|
@ -1900,6 +1902,7 @@ async def authorize_mcp_session(
|
|||
|
||||
@router.get("/{mcp_server_name}/authorize")
|
||||
@router.get("/authorize")
|
||||
@with_mcp_catalog
|
||||
async def authorize(
|
||||
request: Request,
|
||||
redirect_uri: str,
|
||||
|
|
@ -1938,9 +1941,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:
|
||||
|
|
@ -1974,6 +1979,7 @@ async def authorize(
|
|||
|
||||
@router.post("/{mcp_server_name}/token")
|
||||
@router.post("/token")
|
||||
@with_mcp_catalog
|
||||
async def token_endpoint(
|
||||
request: Request,
|
||||
grant_type: str = Form(...),
|
||||
|
|
@ -2067,6 +2073,7 @@ async def _vendor_credential_state(user_id: str, server_id: str) -> VendorCreden
|
|||
|
||||
|
||||
@router.get("/authorize/flow")
|
||||
@with_mcp_catalog
|
||||
async def authorize_flow(request: Request, flow: str) -> Response:
|
||||
return await describe_connect_flow(
|
||||
request=request,
|
||||
|
|
@ -2078,6 +2085,7 @@ async def authorize_flow(request: Request, flow: str) -> Response:
|
|||
|
||||
|
||||
@router.post("/authorize/complete")
|
||||
@with_mcp_catalog
|
||||
async def authorize_complete(
|
||||
request: Request,
|
||||
flow: str = Form(...),
|
||||
|
|
@ -2435,6 +2443,7 @@ def is_network_error(exc: Exception) -> bool:
|
|||
return isinstance(exc, httpx.TransportError)
|
||||
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _build_oauth_protected_resource_response(
|
||||
request: Request,
|
||||
mcp_server_name: str | None,
|
||||
|
|
@ -2739,12 +2748,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
|
||||
|
|
@ -2792,6 +2796,7 @@ def _build_oauth_authorization_server_response(
|
|||
|
||||
# Standard MCP pattern: /.well-known/oauth-authorization-server/mcp/{server_name}
|
||||
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
|
||||
@with_mcp_catalog
|
||||
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
|
||||
"""
|
||||
OAuth authorization server discovery endpoint using standard MCP URL pattern.
|
||||
|
|
@ -2809,6 +2814,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n
|
|||
# LiteLLM legacy pattern and root endpoint
|
||||
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}")
|
||||
@router.get("/.well-known/oauth-authorization-server")
|
||||
@with_mcp_catalog
|
||||
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None):
|
||||
"""
|
||||
OAuth authorization server discovery endpoint.
|
||||
|
|
@ -2882,6 +2888,7 @@ async def jwks_json(request: Request):
|
|||
|
||||
# Additional legacy pattern support
|
||||
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
|
||||
@with_mcp_catalog
|
||||
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
|
||||
"""
|
||||
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
|
||||
|
|
@ -2903,46 +2910,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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -73,6 +73,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
|||
MCPServerAccess,
|
||||
_is_mcp_admitted_user_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
|
||||
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
|
||||
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
||||
MCP_ELICITATION_AVAILABLE,
|
||||
|
|
@ -1860,6 +1861,7 @@ class MCPServerManager:
|
|||
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
|
||||
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
|
||||
)
|
||||
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] = {}
|
||||
|
|
@ -2287,11 +2289,12 @@ 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
|
||||
snapshot: Final = self.catalog.current
|
||||
return snapshot if snapshot is not None else self.config_mcp_servers | self.registry
|
||||
|
||||
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).
|
||||
|
|
@ -6507,14 +6510,15 @@ class MCPServerManager:
|
|||
return None
|
||||
|
||||
async def reload_servers_from_database(self):
|
||||
await self.catalog.refresh()
|
||||
|
||||
async def _reload_servers_from_database(self, *, reuse_unchanged: bool = False):
|
||||
"""Re-synchronize the in-memory MCP server registry with the database."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
get_prisma_client_or_throw,
|
||||
)
|
||||
|
||||
verbose_logger.debug("Loading MCP servers from database into registry...")
|
||||
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")
|
||||
|
|
@ -6593,7 +6597,8 @@ class MCPServerManager:
|
|||
# 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)
|
||||
if not reuse_unchanged or new_server is not previous_registry.get(server_id):
|
||||
await self._maybe_register_openapi_tools(new_server, initialize_mapping=False)
|
||||
registered_registry[server_id] = new_server
|
||||
if new_server.spec_path:
|
||||
registered_openapi_tools = True
|
||||
|
|
@ -6607,10 +6612,13 @@ class MCPServerManager:
|
|||
|
||||
dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys()
|
||||
for registry_key in dropped_registry_keys:
|
||||
self._cleanup_server_tool_routing_artifacts(previous_registry[registry_key])
|
||||
self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
|
||||
|
||||
for server_id in previous_registry.keys() | registered_registry.keys():
|
||||
if previous_registry.get(server_id) != registered_registry.get(server_id):
|
||||
self._upstream_initialize_instructions_by_server_id.pop(server_id, None)
|
||||
self._upstream_initialize_instructions_probed_at.pop(server_id, None)
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self.registry = registered_registry
|
||||
# A discovery task may have published into ``previous_registry`` while
|
||||
|
|
@ -6834,7 +6842,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.
|
||||
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
|
|||
cache_byok_credential,
|
||||
get_cached_byok_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
|
||||
from litellm.proxy._experimental.mcp_server.contracts import (
|
||||
AuthorizedToolCall,
|
||||
OperationContext,
|
||||
|
|
@ -614,6 +615,7 @@ def apply_tool_overrides(
|
|||
return tools
|
||||
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _get_allowed_mcp_servers(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
mcp_servers: Sequence[str] | None,
|
||||
|
|
@ -1435,6 +1437,7 @@ async def filter_tools_by_key_team_permissions(
|
|||
]
|
||||
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _list_mcp_tools(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -1485,6 +1488,7 @@ async def _list_mcp_tools(
|
|||
return AggregateToolListing(tools=[], outcomes={})
|
||||
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _list_mcp_prompts(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -1526,6 +1530,7 @@ async def _list_mcp_prompts(
|
|||
return managed_prompts
|
||||
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _list_mcp_resources(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -1555,6 +1560,7 @@ async def _list_mcp_resources(
|
|||
return managed_resources
|
||||
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _list_mcp_resource_templates(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
mcp_auth_header: str | None = None,
|
||||
|
|
@ -3055,6 +3061,7 @@ class GatewayOperations:
|
|||
@overload
|
||||
async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ...
|
||||
|
||||
@with_mcp_catalog
|
||||
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
|
||||
match operation:
|
||||
case AuthorizedToolCall():
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from litellm.exceptions import (
|
|||
GuardrailRaisedException,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPServerListError,
|
||||
MCPServerURLCredentialsError,
|
||||
|
|
@ -861,6 +862,7 @@ if MCP_AVAILABLE:
|
|||
return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id)
|
||||
|
||||
@router.get("/tools/list", dependencies=[Depends(user_api_key_auth)])
|
||||
@with_mcp_catalog
|
||||
async def list_tool_rest_api(
|
||||
request: Request,
|
||||
server_id: str | None = Query(None, description="The server id to list tools for"),
|
||||
|
|
@ -1085,6 +1087,7 @@ if MCP_AVAILABLE:
|
|||
}
|
||||
|
||||
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
|
||||
@with_mcp_catalog
|
||||
async def call_tool_rest_api(
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
|
|
|
|||
|
|
@ -45,6 +45,8 @@ from fastapi import (
|
|||
from fastapi.responses import JSONResponse
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
|
||||
|
||||
try:
|
||||
from prisma.errors import RecordNotFoundError, UniqueViolationError
|
||||
except ImportError:
|
||||
|
|
@ -1043,6 +1045,7 @@ if MCP_AVAILABLE:
|
|||
tags=["mcp"],
|
||||
description="MCP registry endpoint. Spec: https://github.com/modelcontextprotocol/registry",
|
||||
)
|
||||
@with_mcp_catalog
|
||||
async def get_mcp_registry(request: Request):
|
||||
if not _is_public_registry_enabled():
|
||||
raise HTTPException(
|
||||
|
|
@ -1166,6 +1169,7 @@ if MCP_AVAILABLE:
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=list[LiteLLM_MCPServerTable],
|
||||
)
|
||||
@with_mcp_catalog
|
||||
async def fetch_all_mcp_servers(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
team_id: str | None = Query(
|
||||
|
|
@ -1284,6 +1288,7 @@ if MCP_AVAILABLE:
|
|||
description="Health check for MCP servers",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@with_mcp_catalog
|
||||
async def health_check_servers(
|
||||
server_ids: list[str] | None = Query(
|
||||
None,
|
||||
|
|
@ -1959,6 +1964,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
return _redact_mcp_credentials(temp_record)
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _mcp_oauth_user_api_key_auth(request: Request) -> UserAPIKeyAuth:
|
||||
"""
|
||||
Auth dependency for MCP OAuth browser-navigation endpoints (/authorize, /token).
|
||||
|
|
@ -2059,6 +2065,7 @@ if MCP_AVAILABLE:
|
|||
request_data=request_data,
|
||||
)
|
||||
|
||||
@with_mcp_catalog
|
||||
async def _get_cached_temporary_mcp_server_or_404(
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
|
|||
|
|
@ -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 with_mcp_catalog
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
iter_known_server_prefixes,
|
||||
logging_safe_mcp_headers,
|
||||
|
|
@ -178,6 +179,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
@with_mcp_catalog
|
||||
async def routes_through_gateway(
|
||||
tools: Iterable[Mapping[str, object]] | None,
|
||||
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
|
||||
|
|
@ -229,6 +231,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return user_api_key_auth
|
||||
|
||||
@staticmethod
|
||||
@with_mcp_catalog
|
||||
async def _get_mcp_tools_from_manager(
|
||||
user_api_key_auth: "UserAPIKeyAuth | None",
|
||||
mcp_tools_with_litellm_proxy: Iterable[Mapping[str, object]] | None,
|
||||
|
|
@ -680,6 +683,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
return result_text or "Tool executed successfully"
|
||||
|
||||
@staticmethod
|
||||
@with_mcp_catalog
|
||||
async def _execute_tool_calls(
|
||||
tool_server_map: dict[str, str],
|
||||
tool_calls: Sequence[object],
|
||||
|
|
|
|||
|
|
@ -7482,9 +7482,11 @@ class TestGatewaySessionAdmission:
|
|||
rpm_limit=rpm_limit,
|
||||
)
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
with (
|
||||
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
|
||||
):
|
||||
yield get_user_object
|
||||
|
|
@ -9571,10 +9573,12 @@ class TestScopedSessionAdmission:
|
|||
rpm_limit=None,
|
||||
)
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
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.prisma_client", prisma),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
|
||||
):
|
||||
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict)
|
||||
|
|
|
|||
|
|
@ -632,6 +632,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
|
|||
mcp_operations.byok_credential_cache.flush_cache()
|
||||
server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
|
|||
|
|
@ -738,7 +738,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
|
||||
|
||||
|
|
@ -761,17 +762,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
|
||||
|
||||
|
|
@ -811,7 +813,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()
|
||||
|
||||
|
|
@ -12347,3 +12351,61 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
|
|||
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
|
||||
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)))
|
||||
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()
|
||||
|
|
|
|||
|
|
@ -14516,3 +14516,209 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
|
|||
assert captured["client_ip"] is None
|
||||
finally:
|
||||
auth_context_var.reset(token)
|
||||
|
||||
|
||||
def _catalog_row(name="initial"):
|
||||
return LiteLLM_MCPServerTable(
|
||||
server_id="catalog-server",
|
||||
server_name=name,
|
||||
transport="http",
|
||||
url="https://upstream.example/mcp",
|
||||
created_at=datetime(2026, 1, 1),
|
||||
updated_at=datetime(2026, 1, 1 if name == "initial" else 2),
|
||||
)
|
||||
|
||||
|
||||
def _catalog_database(monkeypatch, read_rows):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
client = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows)))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_pins_one_snapshot_and_next_operation_reads_current_rows(monkeypatch):
|
||||
read_rows = AsyncMock(return_value=[_catalog_row()])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
async with manager.catalog.operation():
|
||||
initial = manager.get_mcp_server_by_id("catalog-server")
|
||||
read_rows.return_value = [_catalog_row("updated")]
|
||||
async with manager.catalog.operation():
|
||||
assert await manager.catalog.resolve("initial") is initial
|
||||
assert tuple((await manager.catalog.list()).values()) == (initial,)
|
||||
read_rows.assert_awaited_once()
|
||||
async with manager.catalog.operation():
|
||||
assert manager.get_mcp_server_by_id("catalog-server").name == "updated"
|
||||
assert await manager.catalog.resolve("initial") is None
|
||||
assert manager.catalog.current is None
|
||||
assert read_rows.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_database_failure_keeps_registry_but_rejects_operation(monkeypatch):
|
||||
read_rows = AsyncMock(return_value=[_catalog_row()])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
snapshot = await manager.catalog.list()
|
||||
read_rows.side_effect = RuntimeError("private connection detail")
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
async with manager.catalog.operation():
|
||||
pytest.fail("An unverified database snapshot must not execute")
|
||||
assert exc.value.status_code == 503
|
||||
assert exc.value.detail == "MCP server configuration could not be refreshed"
|
||||
assert manager.registry == dict(snapshot)
|
||||
assert manager.catalog.current is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_serializes_background_and_operation_refreshes(monkeypatch):
|
||||
entered = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def first_read(**kwargs):
|
||||
entered.set()
|
||||
await release.wait()
|
||||
return [_catalog_row()]
|
||||
|
||||
read_rows = AsyncMock(side_effect=first_read)
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
first = asyncio.create_task(manager.reload_servers_from_database())
|
||||
await entered.wait()
|
||||
read_rows.side_effect = None
|
||||
read_rows.return_value = [_catalog_row("updated")]
|
||||
second = asyncio.create_task(manager.catalog.list())
|
||||
await asyncio.sleep(0)
|
||||
assert read_rows.await_count == 1
|
||||
release.set()
|
||||
await first
|
||||
latest = await second
|
||||
assert latest["catalog-server"].name == "updated"
|
||||
assert manager.registry["catalog-server"].name == "updated"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_cancellation_releases_refresh_lock(monkeypatch):
|
||||
entered = asyncio.Event()
|
||||
blocked = asyncio.Event()
|
||||
|
||||
async def blocked_read(**kwargs):
|
||||
entered.set()
|
||||
await blocked.wait()
|
||||
return []
|
||||
|
||||
read_rows = AsyncMock(side_effect=blocked_read)
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
task = asyncio.create_task(manager.catalog.list())
|
||||
await entered.wait()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
read_rows.side_effect = None
|
||||
read_rows.return_value = [_catalog_row()]
|
||||
snapshot = await asyncio.wait_for(manager.catalog.list(), timeout=1)
|
||||
assert tuple(snapshot) == ("catalog-server",)
|
||||
assert manager.catalog.current is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_child_operation_does_not_inherit_stale_session_snapshot(monkeypatch):
|
||||
read_rows = AsyncMock(return_value=[_catalog_row()])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def child_operation():
|
||||
await release.wait()
|
||||
async with manager.catalog.operation():
|
||||
return manager.get_mcp_server_by_id("catalog-server")
|
||||
|
||||
async with manager.catalog.operation():
|
||||
child = asyncio.create_task(child_operation())
|
||||
read_rows.return_value = [_catalog_row("updated")]
|
||||
release.set()
|
||||
assert (await child).name == "updated"
|
||||
assert read_rows.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_config_only_snapshot_cleans_up_after_error(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
manager = MCPServerManager()
|
||||
configured = MCPServer(server_id="configured", name="configured", transport="stdio", command="echo")
|
||||
manager.config_mcp_servers = {configured.server_id: configured}
|
||||
async def fail_operation():
|
||||
async with manager.catalog.operation():
|
||||
assert await manager.catalog.resolve("configured") is configured
|
||||
raise ValueError("stop operation")
|
||||
|
||||
with pytest.raises(ValueError, match="stop operation"):
|
||||
await fail_operation()
|
||||
assert manager.catalog.current is None
|
||||
assert dict(await manager.catalog.list()) == {"configured": configured}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_delete_drops_derived_tool_mapping(monkeypatch):
|
||||
read_rows = AsyncMock(return_value=[_catalog_row()])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
await manager.catalog.list()
|
||||
manager.tool_name_to_mcp_server_name_mapping = {"initial-echo": "initial"}
|
||||
read_rows.return_value = []
|
||||
assert not await manager.catalog.list()
|
||||
assert manager.tool_name_to_mcp_server_name_mapping == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_unchanged_read_preserves_derived_initialize_instructions(monkeypatch):
|
||||
read_rows = AsyncMock(return_value=[_catalog_row()])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
await manager.catalog.list()
|
||||
manager._upstream_initialize_instructions_by_server_id = {"catalog-server": "cached instructions"}
|
||||
manager._upstream_initialize_instructions_probed_at = {"catalog-server": 123.0}
|
||||
await manager.catalog.list()
|
||||
assert manager._upstream_initialize_instructions_by_server_id == {"catalog-server": "cached instructions"}
|
||||
assert manager._upstream_initialize_instructions_probed_at == {"catalog-server": 123.0}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catalog_reuses_openapi_tools_until_configuration_or_background_refresh(monkeypatch, tmp_path):
|
||||
from litellm.proxy._experimental.mcp_server import openapi_to_mcp_generator, tool_registry
|
||||
|
||||
registry = tool_registry.MCPToolRegistry()
|
||||
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
|
||||
spec_path = tmp_path / "catalog.json"
|
||||
spec_path.write_text(json.dumps({
|
||||
"openapi": "3.0.0", "info": {"title": "Catalog", "version": "1"},
|
||||
"paths": {"/echo": {"get": {"operationId": "echo"}}},
|
||||
}))
|
||||
row = _catalog_row().model_copy(update={"spec_path": str(spec_path)})
|
||||
read_rows = AsyncMock(return_value=[row])
|
||||
_catalog_database(monkeypatch, read_rows)
|
||||
manager = MCPServerManager()
|
||||
with patch.object(
|
||||
openapi_to_mcp_generator, "load_openapi_spec_async",
|
||||
wraps=openapi_to_mcp_generator.load_openapi_spec_async,
|
||||
) as load_spec:
|
||||
await manager.catalog.list()
|
||||
first_tools = tuple(registry.list_tools())
|
||||
assert len(first_tools) == 1
|
||||
await manager.catalog.list()
|
||||
assert load_spec.await_count == 1
|
||||
assert [tool.name for tool in registry.list_tools()] == [tool.name for tool in first_tools]
|
||||
await manager.reload_servers_from_database()
|
||||
assert load_spec.await_count == 2
|
||||
read_rows.return_value = [row.model_copy(update={"updated_at": datetime(2026, 1, 2)})]
|
||||
await manager.catalog.list()
|
||||
assert load_spec.await_count == 3
|
||||
read_rows.return_value = []
|
||||
await manager.catalog.list()
|
||||
assert registry.list_tools() == []
|
||||
|
|
|
|||
|
|
@ -1495,6 +1495,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()))
|
||||
|
|
|
|||
|
|
@ -136,7 +136,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},
|
||||
|
|
@ -2552,7 +2552,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
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)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2629,7 +2629,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = non_oauth_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2681,7 +2681,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = internal_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from contextlib import nullcontext
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
|
|
@ -28,10 +29,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
|
||||
|
|
@ -49,7 +51,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
|
||||
|
||||
|
|
@ -378,6 +380,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
|
|||
post_call_failure_hook = _setup_proxy_logging(monkeypatch)
|
||||
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom"))
|
||||
)
|
||||
|
|
@ -506,6 +509,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
|
|||
|
||||
# Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields.
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
|
|
@ -559,6 +563,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
|
|||
mock_get_tools,
|
||||
)
|
||||
fake_manager = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue