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:
devin-ai-integration[bot] 2026-09-03 15:04:45 -07:00 committed by GitHub
parent 3244a034ac
commit f2f65a6e8b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 209 additions and 31 deletions

View file

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

View file

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

View file

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

View file

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