diff --git a/litellm/proxy/_experimental/mcp_server/server_resolution.py b/litellm/proxy/_experimental/mcp_server/server_resolution.py new file mode 100644 index 00000000000..8168fea9068 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/server_resolution.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass +from typing import Final, Literal, Protocol + +from fastapi import HTTPException, status + +from litellm.proxy._experimental.mcp_server.ui_session_utils import can_access_mcp_server +from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +class MCPServerRegistry(Protocol): + def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None: ... + + def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: ... + + def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: ... + + def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: ... + + async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth) -> list[str]: ... + + +ResolutionSource = Literal["temp", "db", "registry"] + + +@dataclass(frozen=True, slots=True) +class ResolvedMCPServer: + table: LiteLLM_MCPServerTable + runtime: MCPServer | None + source: ResolutionSource + + +async def resolve_mcp_server( + server_id: str, + *, + manager: MCPServerRegistry, + db_lookup: Callable[[str], Awaitable[LiteLLM_MCPServerTable | None]] | None = None, + temp_lookup: Callable[[str], Awaitable[MCPServer | None]] | None = None, + id_client_ip: str | None = None, + name_client_ip: str | None = None, + match_name: bool = False, +) -> ResolvedMCPServer | None: + if temp_lookup is not None: + temporary_server: Final[MCPServer | None] = await temp_lookup(server_id) + if temporary_server is not None: + return ResolvedMCPServer( + table=manager._build_mcp_server_table(temporary_server), + runtime=temporary_server, + source="temp", + ) + + if db_lookup is not None: + database_server: Final[LiteLLM_MCPServerTable | None] = await db_lookup(server_id) + if database_server is not None: + return ResolvedMCPServer(table=database_server, runtime=None, source="db") + + registry_candidate: Final[MCPServer | None] = manager.get_mcp_server_by_id(server_id) + registry_server: Final[MCPServer | None] = ( + registry_candidate + if registry_candidate is not None + and (id_client_ip is None or manager._is_server_accessible_from_ip(registry_candidate, id_client_ip)) + else None + ) + if registry_server is not None: + return ResolvedMCPServer( + table=manager._build_mcp_server_table(registry_server), + runtime=registry_server, + source="registry", + ) + + if match_name: + named_server: Final[MCPServer | None] = manager.get_mcp_server_by_name(server_id, client_ip=name_client_ip) + if named_server is not None: + return ResolvedMCPServer( + table=manager._build_mcp_server_table(named_server), + runtime=named_server, + source="registry", + ) + + return None + + +async def authorize_mcp_server( + resolved: ResolvedMCPServer | None, + user_api_key_dict: UserAPIKeyAuth, + *, + manager: MCPServerRegistry, + is_admin_view: bool, + not_found_detail: Mapping[str, str], + forbidden_detail: Mapping[str, str], + non_admin_missing: Literal["not_found", "forbidden"], + allow_catalog_view: bool = False, +) -> ResolvedMCPServer: + if resolved is None: + if is_admin_view or non_admin_missing == "not_found": + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=dict(not_found_detail), + ) + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=dict(forbidden_detail), + ) + + if is_admin_view: + return resolved + + if resolved.source == "temp" or ( + not allow_catalog_view + and not await can_access_mcp_server( + user_api_key_dict, resolved.table.server_id, manager.get_allowed_mcp_servers + ) + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=dict(forbidden_detail), + ) + + return resolved diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_server_resolution.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_server_resolution.py new file mode 100644 index 00000000000..f88088a4fd8 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_server_resolution.py @@ -0,0 +1,462 @@ +from __future__ import annotations + +import asyncio + +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass +from typing import Final, Literal +from unittest.mock import Mock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._experimental.mcp_server.server_resolution import ( + ResolutionSource, + ResolvedMCPServer, + authorize_mcp_server, + resolve_mcp_server, +) +from litellm.proxy._experimental.mcp_server.ui_session_utils import can_access_mcp_server +from litellm.proxy._types import LiteLLM_MCPServerTable, LitellmUserRoles, UserAPIKeyAuth +from litellm.types.mcp import MCPTransport +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +@dataclass(frozen=True) +class FakeMCPServerManager: + servers_by_id: Mapping[str, MCPServer] + servers_by_name: Mapping[str, MCPServer] + allowed_server_ids: tuple[str, ...] + id_lookup_spy: Mock + name_lookup_spy: Mock + ip_filter_spy: Mock + allowed_servers_spy: Mock + ip_accessible: bool + + def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None: + self.id_lookup_spy(server_id) + return self.servers_by_id.get(server_id) + + def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: + self.name_lookup_spy(server_name, client_ip) + return self.servers_by_name.get(server_name) + + def _is_server_accessible_from_ip(self, server: MCPServer, client_ip: str | None) -> bool: + self.ip_filter_spy(server, client_ip) + return self.ip_accessible + + def _build_mcp_server_table(self, server: MCPServer) -> LiteLLM_MCPServerTable: + return LiteLLM_MCPServerTable( + server_id=server.server_id, + alias=server.alias, + server_name=server.server_name, + url=server.url, + transport=server.transport, + ) + + async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth) -> list[str]: + self.allowed_servers_spy(user_api_key_auth) + return list(self.allowed_server_ids) + + +def _runtime_server(server_id: str = "canonical-server") -> MCPServer: + return MCPServer( + server_id=server_id, + name=server_id, + alias=f"{server_id}-alias", + server_name=f"{server_id}-name", + url="https://example.com/mcp", + transport=MCPTransport.http, + ) + + +def _table_server(server_id: str = "database-server") -> LiteLLM_MCPServerTable: + return LiteLLM_MCPServerTable( + server_id=server_id, + alias=f"{server_id}-alias", + server_name=f"{server_id}-name", + url="https://example.com/mcp", + transport=MCPTransport.http, + ) + + +def _manager( + *, + servers_by_id: Mapping[str, MCPServer] | None = None, + servers_by_name: Mapping[str, MCPServer] | None = None, + allowed_server_ids: tuple[str, ...] = (), + ip_accessible: bool = True, +) -> FakeMCPServerManager: + return FakeMCPServerManager( + servers_by_id={} if servers_by_id is None else servers_by_id, + servers_by_name={} if servers_by_name is None else servers_by_name, + allowed_server_ids=allowed_server_ids, + id_lookup_spy=Mock(), + name_lookup_spy=Mock(), + ip_filter_spy=Mock(), + allowed_servers_spy=Mock(), + ip_accessible=ip_accessible, + ) + + +def _auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="resolver-test-user", + api_key="resolver-test-key", + ) + + +@pytest.mark.asyncio +async def test_temp_resolution_precedes_db_and_registry() -> None: + temporary_server: Final = _runtime_server("temporary-server") + manager: Final = _manager(servers_by_id={temporary_server.server_id: temporary_server}) + temp_lookup: Final[Mock] = Mock() + db_lookup: Final[Mock] = Mock() + + async def lookup_temp(server_id: str) -> MCPServer | None: + temp_lookup(server_id) + return temporary_server + + async def lookup_db(server_id: str) -> LiteLLM_MCPServerTable | None: + db_lookup(server_id) + return _table_server(server_id) + + resolved: Final = await resolve_mcp_server( + "requested-id", + manager=manager, + temp_lookup=lookup_temp, + db_lookup=lookup_db, + ) + + assert resolved == ResolvedMCPServer( + table=manager._build_mcp_server_table(temporary_server), + runtime=temporary_server, + source="temp", + ) + temp_lookup.assert_called_once_with("requested-id") + db_lookup.assert_not_called() + manager.id_lookup_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_db_resolution_precedes_registry_id() -> None: + database_server: Final = _table_server("database-server") + registry_server: Final = _runtime_server(database_server.server_id) + manager: Final = _manager(servers_by_id={registry_server.server_id: registry_server}) + db_lookup: Final = Mock() + + async def lookup_db(server_id: str) -> LiteLLM_MCPServerTable | None: + db_lookup(server_id) + return database_server + + resolved: Final = await resolve_mcp_server( + database_server.server_id, + manager=manager, + db_lookup=lookup_db, + ) + + assert resolved == ResolvedMCPServer(table=database_server, runtime=None, source="db") + db_lookup.assert_called_once_with(database_server.server_id) + manager.id_lookup_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_registry_id_resolution_precedes_name() -> None: + server: Final = _runtime_server() + name_collision: Final = _runtime_server("other-server") + manager: Final = _manager( + servers_by_id={server.server_id: server}, + servers_by_name={server.server_id: name_collision}, + ) + + resolved: Final = await resolve_mcp_server( + server.server_id, + manager=manager, + match_name=True, + ) + + assert resolved == ResolvedMCPServer( + table=manager._build_mcp_server_table(server), + runtime=server, + source="registry", + ) + manager.id_lookup_spy.assert_called_once_with(server.server_id) + manager.name_lookup_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_lookup_ip_arguments_are_scoped_and_name_matching_can_be_disabled() -> None: + server: Final = _runtime_server() + manager: Final = _manager(servers_by_name={"server-alias": server}) + + resolved: Final = await resolve_mcp_server( + "server-alias", + manager=manager, + id_client_ip="id-client", + name_client_ip="name-client", + match_name=True, + ) + + assert resolved is not None + assert resolved.source == "registry" + assert resolved.runtime == server + manager.id_lookup_spy.assert_called_once_with("server-alias") + manager.ip_filter_spy.assert_not_called() + manager.name_lookup_spy.assert_called_once_with("server-alias", "name-client") + + disabled_manager: Final = _manager(servers_by_name={"server-alias": server}) + not_resolved: Final = await resolve_mcp_server( + "server-alias", + manager=disabled_manager, + name_client_ip="name-client", + ) + + assert not_resolved is None + disabled_manager.id_lookup_spy.assert_called_once_with("server-alias") + disabled_manager.name_lookup_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_id_lookup_applies_ip_filter_after_unfiltered_registry_lookup() -> None: + server: Final = _runtime_server() + manager: Final = _manager(servers_by_id={server.server_id: server}, ip_accessible=False) + + resolved: Final = await resolve_mcp_server( + server.server_id, + manager=manager, + id_client_ip="external-client", + ) + + assert resolved is None + manager.id_lookup_spy.assert_called_once_with(server.server_id) + manager.ip_filter_spy.assert_called_once_with(server, "external-client") + manager.name_lookup_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_db_lookup_none_skips_db_and_returns_registry_source() -> None: + server: Final = _runtime_server() + manager: Final = _manager(servers_by_id={server.server_id: server}) + + resolved: Final = await resolve_mcp_server(server.server_id, manager=manager, db_lookup=None) + + assert resolved == ResolvedMCPServer( + table=manager._build_mcp_server_table(server), + runtime=server, + source="registry", + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "is_admin_view,missing_policy,expected_status", + [ + pytest.param(True, "not_found", 404, id="admin-view-not-found"), + pytest.param(False, "not_found", 404, id="non-admin-not-found"), + pytest.param(False, "forbidden", 403, id="non-admin-forbidden"), + ], +) +async def test_authorize_missing_uses_caller_policy( + is_admin_view: bool, + missing_policy: Literal["not_found", "forbidden"], + expected_status: int, +) -> None: + manager: Final = _manager() + with pytest.raises(HTTPException) as exc_info: + await authorize_mcp_server( + None, + _auth(), + manager=manager, + is_admin_view=is_admin_view, + not_found_detail={"error": "not found"}, + forbidden_detail={"error": "forbidden"}, + non_admin_missing=missing_policy, + ) + + assert exc_info.value.status_code == expected_status + assert exc_info.value.detail == ({"error": "not found"} if expected_status == 404 else {"error": "forbidden"}) + manager.allowed_servers_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_non_admin_temp_resolution_is_denied_before_allowed_lookup() -> None: + server: Final = _runtime_server() + manager: Final = _manager(allowed_server_ids=(server.server_id,)) + resolved: Final = ResolvedMCPServer( + table=manager._build_mcp_server_table(server), + runtime=server, + source="temp", + ) + + with pytest.raises(HTTPException) as exc_info: + await authorize_mcp_server( + resolved, + _auth(), + manager=manager, + is_admin_view=False, + not_found_detail={"error": "not found"}, + forbidden_detail={"error": "forbidden"}, + non_admin_missing="not_found", + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == {"error": "forbidden"} + manager.allowed_servers_spy.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "allowed_server_ids,expected_status", + [ + pytest.param(("canonical-server",), None, id="allowed-canonical-id"), + pytest.param((), 403, id="denied-canonical-id"), + ], +) +async def test_authorize_uses_real_access_helper_for_canonical_id( + allowed_server_ids: tuple[str, ...], + expected_status: int | None, + monkeypatch: pytest.MonkeyPatch, +) -> None: + server: Final = _runtime_server("canonical-server") + manager: Final = _manager( + servers_by_name={"display-alias": server}, + allowed_server_ids=allowed_server_ids, + ) + resolved: Final = await resolve_mcp_server( + "display-alias", + manager=manager, + match_name=True, + ) + assert resolved is not None + access_spy: Final = Mock() + + async def spy_access( + user_api_key_auth: UserAPIKeyAuth, + requested_server_id: str, + allowed_servers: Callable[[UserAPIKeyAuth], Awaitable[list[str]]], + ) -> bool: + access_spy(requested_server_id) + return await can_access_mcp_server(user_api_key_auth, requested_server_id, allowed_servers) + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server_resolution.can_access_mcp_server", + spy_access, + ) + if expected_status is None: + authorized: Final = await authorize_mcp_server( + resolved, + _auth(), + manager=manager, + is_admin_view=False, + not_found_detail={"error": "not found"}, + forbidden_detail={"error": "forbidden"}, + non_admin_missing="not_found", + ) + assert authorized is resolved + else: + with pytest.raises(HTTPException) as exc_info: + await authorize_mcp_server( + resolved, + _auth(), + manager=manager, + is_admin_view=False, + not_found_detail={"error": "not found"}, + forbidden_detail={"error": "forbidden"}, + non_admin_missing="not_found", + ) + + assert exc_info.value.status_code == expected_status + assert exc_info.value.detail == {"error": "forbidden"} + + access_spy.assert_called_once_with(server.server_id) + manager.allowed_servers_spy.assert_called_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("source", ["db", "registry", "temp"]) +@pytest.mark.parametrize("admin", [False, True]) +async def test_catalog_visibility_never_opens_temporary_setup_to_non_admins( + source: ResolutionSource, + admin: bool, +) -> None: + server: Final = _runtime_server() + manager: Final = _manager() + resolved: Final = ResolvedMCPServer(manager._build_mcp_server_table(server), server, source) + operation: Final = authorize_mcp_server( + resolved, + _auth(), + manager=manager, + is_admin_view=admin, + not_found_detail={"error": "missing"}, + forbidden_detail={"error": "forbidden"}, + non_admin_missing="forbidden", + allow_catalog_view=True, + ) + if source == "temp" and not admin: + with pytest.raises(HTTPException) as error: + await operation + assert error.value.status_code == 403 + assert error.value.detail == {"error": "forbidden"} + else: + assert await operation is resolved + manager.allowed_servers_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_empty_temp_and_db_lookups_fall_through_to_ip_filtered_registry() -> None: + server: Final = _runtime_server() + manager: Final = _manager(servers_by_id={server.server_id: server}) + lookups: Final = Mock() + + async def temp_lookup(server_id: str) -> MCPServer | None: + lookups.temp(server_id) + return None + + async def db_lookup(server_id: str) -> LiteLLM_MCPServerTable | None: + lookups.db(server_id) + return None + + resolved: Final = await resolve_mcp_server( + server.server_id, + manager=manager, + temp_lookup=temp_lookup, + db_lookup=db_lookup, + id_client_ip="127.0.0.1", + ) + assert resolved is not None + assert resolved.runtime is server + assert resolved.source == "registry" + assert [call[0] for call in lookups.mock_calls] == ["temp", "db"] + manager.ip_filter_spy.assert_called_once_with(server, "127.0.0.1") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", [RuntimeError, asyncio.CancelledError]) +@pytest.mark.parametrize("source", ["db", "temp"]) +async def test_lookup_failure_or_cancellation_never_falls_back( + failure: type[RuntimeError] | type[asyncio.CancelledError], + source: str, +) -> None: + server: Final = _runtime_server() + manager: Final = _manager(servers_by_id={server.server_id: server}) + + async def lookup(server_id: str) -> None: + raise failure(server_id) + + with pytest.raises(failure, match=server.server_id): + await resolve_mcp_server( + server.server_id, + manager=manager, + db_lookup=lookup if source == "db" else None, + temp_lookup=lookup if source == "temp" else None, + ) + manager.id_lookup_spy.assert_not_called() + manager.name_lookup_spy.assert_not_called() + + +@pytest.mark.asyncio +async def test_missing_alias_does_not_produce_a_resolution() -> None: + manager: Final = _manager() + assert await resolve_mcp_server("missing", manager=manager, match_name=True) is None + manager.name_lookup_spy.assert_called_once_with("missing", None)