mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
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:
parent
69ad004015
commit
40297e6268
2 changed files with 584 additions and 0 deletions
122
litellm/proxy/_experimental/mcp_server/server_resolution.py
Normal file
122
litellm/proxy/_experimental/mcp_server/server_resolution.py
Normal 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
|
||||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue