resolved greptile comemnt

This commit is contained in:
shivam 2026-03-03 15:36:40 -08:00
parent a4041d728d
commit c96fcb9ab5
2 changed files with 100 additions and 6 deletions

View file

@ -673,6 +673,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),
):
@ -695,11 +696,14 @@ if MCP_AVAILABLE:
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 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=None
server_id, client_ip=client_ip
)
if registry_server is not None:
mcp_server = global_mcp_server_manager._build_mcp_server_table(

View file

@ -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"
@ -861,7 +874,9 @@ class TestListMCPServers:
)
result = await fetch_mcp_server(
server_id="serper_custom_dev", user_api_key_dict=mock_user_auth
request=_make_mock_request(),
server_id="serper_custom_dev",
user_api_key_dict=mock_user_auth,
)
assert result.server_id == "serper_custom_dev"
@ -869,6 +884,77 @@ class TestListMCPServers:
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):
"""
@ -926,7 +1012,9 @@ class TestListMCPServers:
with pytest.raises(HTTPException) as exc_info:
await fetch_mcp_server(
server_id="restricted_server", user_api_key_dict=mock_user_auth
request=_make_mock_request(),
server_id="restricted_server",
user_api_key_dict=mock_user_auth,
)
assert exc_info.value.status_code == 403
@ -996,7 +1084,9 @@ class TestListMCPServers:
)
result = await fetch_mcp_server(
server_id="allowed_config_server", user_api_key_dict=mock_user_auth
request=_make_mock_request(),
server_id="allowed_config_server",
user_api_key_dict=mock_user_auth,
)
assert result.server_id == "allowed_config_server"