chore(mcp): tighten OAuth root endpoint resolution

This commit is contained in:
user
2026-04-29 18:36:17 -07:00
parent d7431c9db9
commit 8af0544ef0
2 changed files with 285 additions and 16 deletions
@@ -131,6 +131,22 @@ def decode_state_hash(encrypted_state: str) -> dict:
return state_data
def _get_validated_client_redirect_uri(state_data: Dict[str, Any]) -> str:
"""Return a loopback client redirect URI from OAuth state."""
redirect_uri = state_data.get("client_redirect_uri") or state_data.get("base_url")
if not redirect_uri or not isinstance(redirect_uri, str):
raise HTTPException(status_code=400, detail="Invalid redirect URI")
validate_loopback_redirect_uri(redirect_uri)
return redirect_uri
def _append_query_params(url: str, params: Dict[str, str]) -> str:
parsed = urlparse(url)
query_params = parse_qsl(parsed.query, keep_blank_values=True)
query_params.extend(params.items())
return urlunparse(parsed._replace(query=urlencode(query_params)))
def _resolve_oauth2_server_for_root_endpoints(
client_ip: Optional[str] = None,
) -> Optional[MCPServer]:
@@ -568,7 +584,7 @@ async def authorize(
else None
)
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints()
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
raise HTTPException(status_code=404, detail="MCP server not found")
# Use server's stored client_id when caller doesn't supply one.
@@ -630,7 +646,7 @@ async def token_endpoint(
lookup_name, client_ip=client_ip
)
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints()
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
raise HTTPException(status_code=404, detail="MCP server not found")
return await exchange_token_with_server(
@@ -651,7 +667,6 @@ async def token_endpoint(
async def callback(code: str, state: str):
try:
state_data = decode_state_hash(state)
base_url = state_data["base_url"]
original_state = state_data["original_state"]
# Re-validate loopback at the sink. /authorize rejects non-loopback
@@ -659,10 +674,10 @@ async def callback(code: str, state: str):
# minted before that check was added have no expiry and remain
# valid indefinitely. Validating here blocks the open-redirect +
# code-theft primitive even for pre-fix states.
validate_loopback_redirect_uri(base_url)
redirect_uri = _get_validated_client_redirect_uri(state_data)
params = {"code": code, "state": original_state}
complete_returned_url = f"{base_url}?{urlencode(params)}"
complete_returned_url = _append_query_params(redirect_uri, params)
return RedirectResponse(url=complete_returned_url, status_code=302)
except HTTPException:
@@ -719,16 +734,16 @@ def _build_oauth_protected_resource_response(
)
request_base_url = get_request_base_url(request)
client_ip = IPAddressUtils.get_mcp_client_ip(request)
# When no server name provided, try to resolve the single OAuth2 server
if mcp_server_name is None:
resolved = _resolve_oauth2_server_for_root_endpoints()
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
mcp_server_name = resolved.server_name or resolved.name
mcp_server: Optional[MCPServer] = None
if mcp_server_name:
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)
@@ -835,10 +850,11 @@ def _build_oauth_authorization_server_response(
)
request_base_url = get_request_base_url(request)
client_ip = IPAddressUtils.get_mcp_client_ip(request)
# When no server name provided, try to resolve the single OAuth2 server
if mcp_server_name is None:
resolved = _resolve_oauth2_server_for_root_endpoints()
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
mcp_server_name = resolved.server_name or resolved.name
@@ -855,7 +871,6 @@ def _build_oauth_authorization_server_response(
mcp_server: Optional[MCPServer] = None
if mcp_server_name:
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)
@@ -1007,8 +1022,9 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
"client_secret": "dummy",
"redirect_uris": [f"{request_base_url}/callback"],
}
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if not mcp_server_name:
resolved = _resolve_oauth2_server_for_root_endpoints()
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
return await register_client_with_server(
request=request,
@@ -1021,7 +1037,6 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
)
return dummy_return
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)
@@ -230,7 +230,6 @@ async def test_authorize_endpoint_forwards_pkce_parameters():
async def test_token_endpoint_forwards_code_verifier():
"""Test that token endpoint forwards code_verifier for PKCE flow"""
try:
import httpx
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
@@ -632,8 +631,7 @@ async def test_token_endpoint_respects_x_forwarded_proto():
) as mock_get_client:
mock_get_client.return_value = mock_async_client
# Call token endpoint
response = await token_endpoint(
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="test_code",
@@ -933,8 +931,7 @@ async def test_token_endpoint_respects_x_forwarded_host():
) as mock_get_client:
mock_get_client.return_value = mock_async_client
# Call token endpoint
response = await token_endpoint(
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="test_code",
@@ -1240,6 +1237,7 @@ def _create_oauth2_server(
alias="test_oauth",
client_id="test_client_id",
client_secret="test_client_secret",
available_on_public_internet=True,
):
"""Helper to create a mock OAuth2 MCPServer."""
from litellm.proxy._types import MCPTransport
@@ -1258,6 +1256,7 @@ def _create_oauth2_server(
authorization_url="https://provider.com/oauth/authorize",
token_url="https://provider.com/oauth/token",
scopes=["read", "write"],
available_on_public_internet=available_on_public_internet,
)
@@ -1352,6 +1351,47 @@ async def test_authorize_root_fails_with_multiple_oauth2_servers():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_authorize_root_does_not_resolve_private_server_for_external_client():
"""Root /authorize must not auto-select an MCP server hidden from the caller IP."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
authorize,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
):
with pytest.raises(HTTPException) as exc_info:
await authorize(
request=mock_request,
client_id="dummy_client",
mcp_server_name=None,
redirect_uri="http://localhost:62646/callback",
state="test_state",
)
assert exc_info.value.status_code == 404
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_token_root_resolves_single_oauth2_server():
"""When /token is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
@@ -1417,6 +1457,50 @@ async def test_token_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_token_root_does_not_resolve_private_server_for_external_client():
"""Root /token must not exchange codes for a hidden MCP server."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
token_endpoint,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
):
with pytest.raises(HTTPException) as exc_info:
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="test_auth_code",
redirect_uri="http://localhost:62646/callback",
client_id="dummy_client",
mcp_server_name=None,
client_secret=None,
code_verifier="test_verifier",
)
assert exc_info.value.status_code == 404
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_register_root_resolves_single_oauth2_server():
"""When /register is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
@@ -1454,6 +1538,48 @@ async def test_register_root_resolves_single_oauth2_server():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_register_root_does_not_resolve_private_server_for_external_client():
"""Root /register must not reveal or use a hidden MCP server."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={}),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
),
):
result = await register_client(request=mock_request, mcp_server_name=None)
assert result["client_id"] == "dummy_client"
assert result["redirect_uris"] == ["https://llm.example.com/callback"]
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_discovery_root_includes_server_name_prefix():
"""When root discovery is hit and exactly 1 OAuth2 server exists, include server name in URLs."""
@@ -1493,6 +1619,54 @@ async def test_discovery_root_includes_server_name_prefix():
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_discovery_root_does_not_expose_private_server_for_external_client():
"""Root discovery must use caller visibility before adding server-specific metadata."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_build_oauth_authorization_server_response,
_build_oauth_protected_resource_response,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://llm.example.com/"
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
):
authorization_response = _build_oauth_authorization_server_response(
request=mock_request,
mcp_server_name=None,
)
resource_response = _build_oauth_protected_resource_response(
request=mock_request,
mcp_server_name=None,
use_standard_pattern=False,
)
assert "/test_oauth/" not in authorization_response["authorization_endpoint"]
assert "/test_oauth/" not in authorization_response["token_endpoint"]
assert authorization_response["scopes_supported"] == []
assert resource_response["authorization_servers"] == ["https://llm.example.com"]
assert resource_response["scopes_supported"] == []
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_oauth_callback_redirects_with_state():
"""Test OAuth callback endpoint properly decodes state and redirects to client callback URL."""
@@ -1536,6 +1710,44 @@ async def test_oauth_callback_redirects_with_state():
mock_decode.assert_called_once_with("encrypted_state_value")
@pytest.mark.asyncio
async def test_oauth_callback_preserves_client_redirect_uri_query():
"""The callback should append code/state without dropping a client's existing query."""
try:
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"base_url": "http://localhost:3000/ui/mcp/oauth/callback",
"original_state": "test-uuid-state-123",
"code_challenge": "test_challenge",
"code_challenge_method": "S256",
"client_redirect_uri": (
"http://localhost:3000/ui/mcp/oauth/callback?session=abc"
),
}
response = await callback(
code="test_authorization_code_12345",
state="encrypted_state_value",
)
assert response.status_code == 302
parsed_location = urlparse(response.headers["location"])
query_params = parse_qs(parsed_location.query)
assert query_params["session"] == ["abc"]
assert query_params["code"] == ["test_authorization_code_12345"]
assert query_params["state"] == ["test-uuid-state-123"]
@pytest.mark.asyncio
async def test_oauth_callback_handles_invalid_state():
"""Test OAuth callback returns error page when state decryption fails."""
@@ -1948,6 +2160,48 @@ async def test_callback_revalidates_loopback_on_decoded_base_url():
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_callback_revalidates_loopback_on_decoded_client_redirect_uri():
"""If a state contains a full client_redirect_uri, validate that exact sink."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"base_url": "http://localhost:3000/cb",
"original_state": "s",
"code_challenge": None,
"code_challenge_method": None,
"client_redirect_uri": "https://attacker.example.com/cb",
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="stolen_code", state="encrypted_stale_state")
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_callback_rejects_state_missing_redirect_uri():
"""Malformed state without a redirect target should fail with a structured 400."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash"
) as mock_decode:
mock_decode.return_value = {
"original_state": "s",
"code_challenge": None,
"code_challenge_method": None,
}
with pytest.raises(HTTPException) as exc_info:
await callback(code="code", state="encrypted_malformed_state")
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_token_endpoint_sets_no_store_cache_control():
"""RFC 6749 §5.1 / OAuth 2.1 draft-15 §4.1.3: the token response