mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(mcp): resolve OAuth broker endpoints by server_id with IP access checks (#39432)
* fix(mcp): resolve OAuth broker endpoints by server_id with IP access checks Resolve named OAuth lookups through server IDs while retaining client IP checks\n\nCo-authored-by: KK291860 <krishnakumar.kocherykumaran@sephora.com> Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: retrigger e2e pipeline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
3244a034ac
commit
f2f65a6e8b
4 changed files with 209 additions and 31 deletions
|
|
@ -448,6 +448,17 @@ def _append_query_params(url: str, params: dict[str, str]) -> str:
|
|||
return urlunparse(parsed._replace(query=urlencode(query_params)))
|
||||
|
||||
|
||||
def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCPServer | None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
|
||||
if by_name is not None:
|
||||
return by_name
|
||||
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
|
||||
|
||||
|
||||
def _resolve_oauth2_server_for_root_endpoints(
|
||||
client_ip: str | None = None,
|
||||
) -> MCPServer | None:
|
||||
|
|
@ -1766,10 +1777,6 @@ async def authorize(
|
|||
resource: str | None = None,
|
||||
):
|
||||
# Redirect to real OAuth provider with PKCE support
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
|
||||
if is_proxy_api_resource(request, resource):
|
||||
return await native_client_authorize(
|
||||
|
|
@ -1797,9 +1804,7 @@ async def authorize(
|
|||
|
||||
lookup_name: Final[str | None] = mcp_server_name or client_id
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
mcp_server = (
|
||||
global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) if lookup_name else None
|
||||
)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip) if lookup_name else None
|
||||
if mcp_server is None and mcp_server_name is None:
|
||||
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if mcp_server is None:
|
||||
|
|
@ -1855,10 +1860,6 @@ async def token_endpoint(
|
|||
3. Return the token
|
||||
4. Return a virtual key in this response
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
|
||||
master_key,
|
||||
|
|
@ -1882,7 +1883,7 @@ async def token_endpoint(
|
|||
|
||||
lookup_name: Final = mcp_server_name or client_id
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip)
|
||||
if mcp_server is None and mcp_server_name is None:
|
||||
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if mcp_server is None:
|
||||
|
|
@ -2288,10 +2289,6 @@ async def _build_oauth_protected_resource_response(
|
|||
Returns:
|
||||
OAuth protected resource metadata dict
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
explicitly_named: Final = mcp_server_name is not None
|
||||
|
|
@ -2304,7 +2301,7 @@ async def _build_oauth_protected_resource_response(
|
|||
|
||||
mcp_server: MCPServer | None = None
|
||||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
|
||||
# Build resource URL based on the pattern
|
||||
if mcp_server_name:
|
||||
|
|
@ -2562,10 +2559,6 @@ def _build_oauth_authorization_server_response(
|
|||
registry lookups; unlike :func:`_build_oauth_protected_resource_response`
|
||||
it does not need to await any upstream IO.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
explicitly_named: Final = mcp_server_name is not None
|
||||
|
|
@ -2583,7 +2576,7 @@ def _build_oauth_authorization_server_response(
|
|||
|
||||
mcp_server: MCPServer | None = None
|
||||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
|
||||
|
||||
|
|
@ -2709,10 +2702,6 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
|
|||
@router.post("/{mcp_server_name}/register")
|
||||
@router.post("/register")
|
||||
async def register_client(request: Request, mcp_server_name: str | None = None):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
# Get the correct base URL considering X-Forwarded-* headers
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
|
||||
|
|
@ -2748,7 +2737,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
)
|
||||
return dummy_return
|
||||
|
||||
mcp_server: Final = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
if mcp_server is None:
|
||||
return dummy_return
|
||||
return await register_client_with_server(
|
||||
|
|
|
|||
|
|
@ -6147,13 +6147,13 @@ class MCPServerManager:
|
|||
internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges"))
|
||||
return IPAddressUtils.is_internal_ip(client_ip, internal_networks)
|
||||
|
||||
def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None:
|
||||
"""
|
||||
Get the MCP Server from the server id
|
||||
"""
|
||||
def get_mcp_server_by_id(self, server_id: str, client_ip: str | None = None) -> MCPServer | None:
|
||||
"""Get the MCP Server from the server id."""
|
||||
registry: Final = self.get_registry()
|
||||
for server in registry.values():
|
||||
if server.server_id == server_id:
|
||||
if not self._is_server_accessible_from_ip(server, client_ip):
|
||||
return None
|
||||
return server
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -3508,6 +3508,169 @@ def _create_oauth2_server(
|
|||
)
|
||||
|
||||
|
||||
def _create_id_lookup_oauth2_server():
|
||||
return _create_oauth2_server(
|
||||
server_id="oauth-server-id",
|
||||
name="oauth-server-name",
|
||||
server_name="oauth-server-name",
|
||||
alias="oauth-server-alias",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "encrypt_value_helper", return_value="encrypted-state"), # test-quality-ok: flow seam
|
||||
):
|
||||
response = await discoverable_endpoints.authorize(
|
||||
request=request,
|
||||
client_id=server.client_id,
|
||||
mcp_server_name=server.server_id,
|
||||
redirect_uri="http://localhost:62646/callback",
|
||||
state="test_state",
|
||||
)
|
||||
|
||||
assert response.status_code == 307
|
||||
assert "https://provider.com/oauth/authorize" in response.headers["location"]
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"access_token": "token", "token_type": "Bearer"}
|
||||
response.raise_for_status = MagicMock()
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(return_value=response)
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam
|
||||
):
|
||||
result = await discoverable_endpoints.token_endpoint(
|
||||
request=request,
|
||||
grant_type="authorization_code",
|
||||
code="test_code",
|
||||
redirect_uri="http://localhost:62646/callback",
|
||||
client_id=server.client_id,
|
||||
mcp_server_name=server.server_id,
|
||||
client_secret=server.client_secret,
|
||||
)
|
||||
|
||||
assert json.loads(result.body)["access_token"] == "token"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_client_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server().model_copy(
|
||||
update={"client_id": None, "client_secret": None, "registration_url": "https://provider.com/oauth/register"}
|
||||
)
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"client_id": "registered-client", "client_secret": "registered-secret"}
|
||||
response.raise_for_status = MagicMock()
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(return_value=response)
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam
|
||||
patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam
|
||||
):
|
||||
result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id)
|
||||
|
||||
assert json.loads(result.body)["client_id"] == "registered-client"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protected_resource_metadata_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
):
|
||||
result = await discoverable_endpoints._build_oauth_protected_resource_response(
|
||||
request=request,
|
||||
mcp_server_name=server.server_id,
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
|
||||
assert result["authorization_servers"] == ["https://llm.example.com/mcp"]
|
||||
assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
):
|
||||
result = discoverable_endpoints._build_oauth_authorization_server_response(
|
||||
request=request,
|
||||
mcp_server_name=server.server_id,
|
||||
)
|
||||
|
||||
assert result["scopes_supported"] == server.scopes
|
||||
assert result["issuer"] == f"https://llm.example.com/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_root_resolves_single_oauth2_server():
|
||||
"""When /authorize is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
|
||||
|
|
|
|||
|
|
@ -129,6 +129,32 @@ class TestMCPServerManager:
|
|||
assert added_server.args == ["-m", "server"]
|
||||
assert added_server.env == {"DEBUG": "1", "TEST": "1"}
|
||||
|
||||
def test_get_mcp_server_by_id_allows_internal_or_unspecified_client_ip(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="private-server",
|
||||
name="private-server",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
|
||||
assert manager.get_mcp_server_by_id(server.server_id) is server
|
||||
assert manager.get_mcp_server_by_id(server.server_id, client_ip="10.0.0.1") is server
|
||||
|
||||
def test_get_mcp_server_by_id_rejects_private_server_for_public_ip(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="private-server",
|
||||
name="private-server",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
|
||||
with patch.object(manager, "_get_general_settings", return_value={}):
|
||||
assert manager.get_mcp_server_by_id(server.server_id, client_ip="8.8.8.8") is None
|
||||
|
||||
async def test_create_mcp_client_stdio(self):
|
||||
"""Test creating MCP client for stdio transport"""
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue