mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-14 16:25:29 +00:00
Fix 404 when fetching config-based MCP servers by ID (#22711)
* fixed mcp api * added non-admin test * resolved greptile comemnt * fix: add IP filtering to get_mcp_server_by_id path in fetch_mcp_server Apply _is_server_accessible_from_ip check after get_mcp_server_by_id lookup to prevent external callers from accessing MCP servers configured with available_on_public_internet=False when they know the server_id. Made-with: Cursor
This commit is contained in:
@@ -1051,6 +1051,7 @@ if MCP_AVAILABLE:
|
||||
response_model=LiteLLM_MCPServerTable,
|
||||
)
|
||||
async def fetch_mcp_server(
|
||||
request: Request,
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
@@ -1067,8 +1068,30 @@ if MCP_AVAILABLE:
|
||||
"Database not connected. Connect a database to your proxy"
|
||||
)
|
||||
|
||||
# check to see if server exists for all users
|
||||
# check to see if server exists (DB first, then registry for config-based servers)
|
||||
mcp_server = await get_mcp_server(prisma_client, server_id)
|
||||
from_db = mcp_server is not None
|
||||
|
||||
if mcp_server is None:
|
||||
# Fallback: check registry (config-based servers) - list endpoint uses get_registry()
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
||||
client_ip = IPAddressUtils.get_mcp_client_ip(request)
|
||||
registry_server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
|
||||
if registry_server is not None and not global_mcp_server_manager._is_server_accessible_from_ip(
|
||||
registry_server, client_ip
|
||||
):
|
||||
registry_server = None
|
||||
if registry_server is None:
|
||||
# Try lookup by server_name or alias (client may use display name in URL)
|
||||
registry_server = global_mcp_server_manager.get_mcp_server_by_name(
|
||||
server_id, client_ip=client_ip
|
||||
)
|
||||
if registry_server is not None:
|
||||
mcp_server = global_mcp_server_manager._build_mcp_server_table(
|
||||
registry_server
|
||||
)
|
||||
|
||||
if mcp_server is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
@@ -1084,10 +1107,17 @@ if MCP_AVAILABLE:
|
||||
if not is_admin_view:
|
||||
# Perform authz check BEFORE any health check (avoid side-effects for
|
||||
# unauthorized callers).
|
||||
mcp_server_records = await get_all_mcp_servers_for_user(
|
||||
prisma_client, user_api_key_dict
|
||||
)
|
||||
exists = does_mcp_server_exist(mcp_server_records, server_id)
|
||||
if from_db:
|
||||
mcp_server_records = await get_all_mcp_servers_for_user(
|
||||
prisma_client, user_api_key_dict
|
||||
)
|
||||
exists = does_mcp_server_exist(mcp_server_records, server_id)
|
||||
else:
|
||||
# Registry/config server: use same access logic as list endpoint
|
||||
allowed_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers(
|
||||
user_api_key_dict
|
||||
)
|
||||
exists = mcp_server.server_id in allowed_server_ids
|
||||
|
||||
if not exists:
|
||||
raise HTTPException(
|
||||
@@ -1101,7 +1131,8 @@ if MCP_AVAILABLE:
|
||||
)
|
||||
|
||||
# At this point caller is authorized to view the server.
|
||||
await global_mcp_server_manager.add_server(mcp_server)
|
||||
if from_db:
|
||||
await global_mcp_server_manager.add_server(mcp_server)
|
||||
|
||||
# Perform health check on the server using server manager
|
||||
try:
|
||||
|
||||
@@ -76,6 +76,15 @@ def generate_mock_mcp_server_config_record(
|
||||
)
|
||||
|
||||
|
||||
def _make_mock_request(ip: str = "127.0.0.1"):
|
||||
"""Create a mock Request for fetch_mcp_server tests (IP used for access control)."""
|
||||
req = MagicMock()
|
||||
req.client = MagicMock()
|
||||
req.client.host = ip
|
||||
req.headers = {}
|
||||
return req
|
||||
|
||||
|
||||
def generate_mock_user_api_key_auth(
|
||||
user_role: LitellmUserRoles = LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id: str = "test_user_id",
|
||||
@@ -735,7 +744,9 @@ class TestListMCPServers:
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
server_id="server-1", user_api_key_dict=mock_user_auth
|
||||
request=_make_mock_request(),
|
||||
server_id="server-1",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.server_id == "server-1"
|
||||
@@ -788,7 +799,9 @@ class TestListMCPServers:
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
server_id="server-2", user_api_key_dict=mock_user_auth
|
||||
request=_make_mock_request(),
|
||||
server_id="server-2",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.server_id == "server-2"
|
||||
@@ -796,6 +809,290 @@ class TestListMCPServers:
|
||||
assert not hasattr(result, "credentials")
|
||||
assert result.status == "healthy"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_from_registry_config_based(self):
|
||||
"""
|
||||
Test that fetch_mcp_server finds config-based servers when not in DB.
|
||||
Config servers appear in list via get_registry() but were 404 on fetch.
|
||||
"""
|
||||
config_server = generate_mock_mcp_server_config_record(
|
||||
server_id="serper_custom_dev",
|
||||
name="Serper MCP",
|
||||
url="https://serper.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
|
||||
mock_health_result = generate_mock_mcp_server_db_record(
|
||||
server_id="serper_custom_dev", alias="Serper MCP"
|
||||
)
|
||||
mock_health_result.status = "healthy"
|
||||
mock_health_result.last_health_check = datetime.now()
|
||||
mock_health_result.health_check_error = None
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(
|
||||
side_effect=lambda sid: config_server if sid == "serper_custom_dev" else None
|
||||
)
|
||||
mock_manager.get_mcp_server_by_name = MagicMock(return_value=None)
|
||||
mock_manager._build_mcp_server_table = MagicMock(
|
||||
return_value=generate_mock_mcp_server_db_record(
|
||||
server_id="serper_custom_dev",
|
||||
alias="Serper MCP",
|
||||
url="https://serper.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
)
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["serper_custom_dev"]
|
||||
)
|
||||
mock_manager.health_check_server = AsyncMock(return_value=mock_health_result)
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
request=_make_mock_request(),
|
||||
server_id="serper_custom_dev",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.server_id == "serper_custom_dev"
|
||||
assert result.status == "healthy"
|
||||
mock_manager.get_mcp_server_by_id.assert_called_with("serper_custom_dev")
|
||||
mock_manager._build_mcp_server_table.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_from_registry_by_name_passes_client_ip(self):
|
||||
"""
|
||||
When lookup by server_id fails, fallback to get_mcp_server_by_name.
|
||||
Verify client_ip is passed for IP-based access control (security).
|
||||
"""
|
||||
config_server = generate_mock_mcp_server_config_record(
|
||||
server_id="serper_custom_dev",
|
||||
name="Serper MCP",
|
||||
url="https://serper.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(return_value=None)
|
||||
mock_manager.get_mcp_server_by_name = MagicMock(return_value=config_server)
|
||||
mock_manager._build_mcp_server_table = MagicMock(
|
||||
return_value=generate_mock_mcp_server_db_record(
|
||||
server_id="serper_custom_dev",
|
||||
alias="Serper MCP",
|
||||
url="https://serper.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
)
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["serper_custom_dev"]
|
||||
)
|
||||
mock_manager.health_check_server = AsyncMock(
|
||||
return_value=generate_mock_mcp_server_db_record(
|
||||
server_id="serper_custom_dev", alias="Serper MCP"
|
||||
)
|
||||
)
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
request=_make_mock_request(ip="192.168.1.100"),
|
||||
server_id="Serper MCP",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.server_id == "serper_custom_dev"
|
||||
mock_manager.get_mcp_server_by_id.assert_called_with("Serper MCP")
|
||||
mock_manager.get_mcp_server_by_name.assert_called_once_with(
|
||||
"Serper MCP", client_ip="192.168.1.100"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_from_registry_non_admin_denied(self):
|
||||
"""
|
||||
Non-admin user: config server NOT in allowed_server_ids -> 403.
|
||||
"""
|
||||
config_server = generate_mock_mcp_server_config_record(
|
||||
server_id="restricted_server",
|
||||
name="Restricted MCP",
|
||||
url="https://restricted.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(
|
||||
side_effect=lambda sid: config_server if sid == "restricted_server" else None
|
||||
)
|
||||
mock_manager.get_mcp_server_by_name = MagicMock(return_value=None)
|
||||
mock_manager._build_mcp_server_table = MagicMock(
|
||||
return_value=generate_mock_mcp_server_db_record(
|
||||
server_id="restricted_server",
|
||||
alias="Restricted MCP",
|
||||
url="https://restricted.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
)
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["other_server"] # restricted_server NOT in list
|
||||
)
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await fetch_mcp_server(
|
||||
request=_make_mock_request(),
|
||||
server_id="restricted_server",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
mock_manager.get_allowed_mcp_servers.assert_called_once_with(mock_user_auth)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_mcp_server_from_registry_non_admin_granted(self):
|
||||
"""
|
||||
Non-admin user: config server IS in allowed_server_ids -> 200.
|
||||
"""
|
||||
config_server = generate_mock_mcp_server_config_record(
|
||||
server_id="allowed_config_server",
|
||||
name="Allowed MCP",
|
||||
url="https://allowed.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
|
||||
mock_health_result = generate_mock_mcp_server_db_record(
|
||||
server_id="allowed_config_server", alias="Allowed MCP"
|
||||
)
|
||||
mock_health_result.status = "healthy"
|
||||
mock_health_result.last_health_check = datetime.now()
|
||||
mock_health_result.health_check_error = None
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(
|
||||
side_effect=lambda sid: config_server if sid == "allowed_config_server" else None
|
||||
)
|
||||
mock_manager.get_mcp_server_by_name = MagicMock(return_value=None)
|
||||
mock_manager._build_mcp_server_table = MagicMock(
|
||||
return_value=generate_mock_mcp_server_db_record(
|
||||
server_id="allowed_config_server",
|
||||
alias="Allowed MCP",
|
||||
url="https://allowed.example.com/mcp",
|
||||
transport="http",
|
||||
)
|
||||
)
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(
|
||||
return_value=["allowed_config_server"]
|
||||
)
|
||||
mock_manager.health_check_server = AsyncMock(return_value=mock_health_result)
|
||||
|
||||
mock_user_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
fetch_mcp_server,
|
||||
)
|
||||
|
||||
result = await fetch_mcp_server(
|
||||
request=_make_mock_request(),
|
||||
server_id="allowed_config_server",
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
assert result.server_id == "allowed_config_server"
|
||||
assert result.status == "healthy"
|
||||
mock_manager.get_allowed_mcp_servers.assert_called_once_with(mock_user_auth)
|
||||
|
||||
|
||||
class TestTeamScopedMCPServerAccess:
|
||||
"""Tests for cross-team information disclosure and restricted key bypass fixes."""
|
||||
|
||||
Reference in New Issue
Block a user