refactor(mcp): add shared server resolver without changing callers (#43262)

* test(mcp): characterize server resolution and authorization

* refactor(mcp): extract shared server resolution

* test(mcp): pin catalog isolation and batched credential permissions

* test(mcp): enforce identity isolation in database fixtures

* test(mcp): name resolution tests by behavior

* test(mcp): describe detail access assertion failures

* chore: keep agent naming discipline local

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-09-26 20:01:37 +00:00 • committed by GitHub
parent 69ad004015
commit 40297e6268
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 584 additions and 0 deletions

View file

@ -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

View file

@ -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)