mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(mcp): gate OAuth authorize/token/register/discovery on auth_type=oauth2 (#31736)
Some checks are pending
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Some checks are pending
GitHub Actions Security Analysis / zizmor (push) Waiting to run
* fix(mcp): gate OAuth authorize/token/register/discovery on auth_type=oauth2 A non-oauth2 MCP server (notably auth_type=none, access-group gated) has no client_id and no authorization URL, yet the gateway OAuth endpoints did not check auth_type. authorize() raised "client_id is required" before the auth_type was ever examined, and the .well-known discovery builders always advertised authorization_servers / authorization_endpoint / token_endpoint / registration_endpoint, so spec-compliant MCP clients were pointed at an OAuth flow that can never succeed. Add an auth_type != oauth2 guard to the authorize, token, register, protected-resource and authorization-server paths (covering the internal UI OAuth endpoints too). The discovery guard sits after the OAuth pass-through branch so genuine pass-through servers keep proxying their upstream metadata. oauth2 servers are unaffected. * fix(mcp): accurate non-oauth2 message; 404 unknown discovery names to close enumeration oracle Address review feedback on the auth_type gate. The 400 message no longer claims access is governed by access groups, which is only true for auth_type=none; it now states that the gateway runs the OAuth client_id/authorize/token/register flow only for oauth2 servers and that the server is reached using its configured auth_type, which is accurate for every non-oauth2 type (api_key, oauth2_token_exchange, etc.). The discovery gate previously 404'd a named non-oauth2 server but still returned 200 metadata for an unknown name, which both serves a broken document for a typo and lets an unauthenticated caller enumerate non-OAuth server names by comparing 404 vs 200. A named discovery request now returns 200 only when it resolves to an oauth2 server; unknown (or hidden) and non-oauth2 names return the same 404. Root discovery and pass-through servers are unaffected. * Apply suggestions from code review Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>
This commit is contained in:
parent
f5f8ba93fa
commit
c370503091
4 changed files with 456 additions and 0 deletions
|
|
@ -390,6 +390,46 @@ async def _store_per_user_token_server_side(
|
|||
)
|
||||
|
||||
|
||||
def _raise_if_not_oauth2(mcp_server: MCPServer) -> None:
|
||||
"""Reject a non-oauth2 server from the gateway's OAuth authorize/token/register flow."""
|
||||
if mcp_server.auth_type == MCPAuth.oauth2:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "server_not_oauth2",
|
||||
"message": (
|
||||
f"MCP server '{mcp_server.server_name or mcp_server.name}' does not use OAuth "
|
||||
f"(auth_type={mcp_server.auth_type}). This server does not support the authorization-code "
|
||||
"flow; it has no client_id, authorize, token, or registration endpoint. "
|
||||
"Access is controlled by the server's configured auth_type and access groups"
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _raise_unless_oauth2_discovery_server(
|
||||
mcp_server: Optional[MCPServer],
|
||||
mcp_server_name: Optional[str],
|
||||
description: str,
|
||||
) -> None:
|
||||
"""404 a NAMED discovery request unless it resolves to an oauth2 server.
|
||||
|
||||
A named server that is unknown (or hidden from the caller) and one that exists
|
||||
but is non-oauth2 both return the same 404, so the well-known discovery paths
|
||||
cannot be used to enumerate non-OAuth server names. Root discovery (no name) is
|
||||
unaffected, and pass-through servers are resolved by the caller before this runs.
|
||||
"""
|
||||
if mcp_server_name is None:
|
||||
return
|
||||
if mcp_server is not None and mcp_server.auth_type == MCPAuth.oauth2:
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"MCP server '{mcp_server_name}' is {description}",
|
||||
)
|
||||
|
||||
|
||||
async def authorize_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -457,6 +497,7 @@ async def exchange_token_with_server(
|
|||
refresh_token: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
):
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
if grant_type not in ("authorization_code", "refresh_token"):
|
||||
raise HTTPException(status_code=400, detail="Unsupported grant_type")
|
||||
|
||||
|
|
@ -582,6 +623,7 @@ async def register_client_with_server(
|
|||
token_endpoint_auth_method: Optional[str],
|
||||
fallback_client_id: Optional[str] = None,
|
||||
):
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
request_base_url = get_request_base_url(request)
|
||||
dummy_return = {
|
||||
"client_id": fallback_client_id or mcp_server.server_name,
|
||||
|
|
@ -655,6 +697,7 @@ async def authorize(
|
|||
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")
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
# Use server's stored client_id when caller doesn't supply one.
|
||||
# Raise a clear error instead of passing an empty string — an empty
|
||||
# client_id would silently produce a broken authorization URL.
|
||||
|
|
@ -1063,6 +1106,8 @@ async def _build_oauth_protected_resource_response(
|
|||
detail=(f"Upstream oauth-protected-resource metadata unavailable for MCP server {mcp_server.name!r}"),
|
||||
)
|
||||
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
|
||||
|
||||
return {
|
||||
"authorization_servers": [
|
||||
(f"{request_base_url}/{mcp_server_name}" if mcp_server_name else f"{request_base_url}")
|
||||
|
|
@ -1149,6 +1194,8 @@ def _build_oauth_authorization_server_response(
|
|||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
|
||||
|
||||
return {
|
||||
"issuer": request_base_url, # point to your proxy
|
||||
"authorization_endpoint": authorization_endpoint,
|
||||
|
|
|
|||
|
|
@ -132,6 +132,7 @@ if MCP_AVAILABLE:
|
|||
update_mcp_server,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_raise_if_not_oauth2,
|
||||
authorize_with_server,
|
||||
exchange_token_with_server,
|
||||
get_request_base_url,
|
||||
|
|
@ -1611,6 +1612,7 @@ if MCP_AVAILABLE:
|
|||
scope: Optional[str] = None,
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
# Use the server's stored client_id when the caller doesn't supply one
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
if not resolved_client_id:
|
||||
|
|
@ -1655,6 +1657,7 @@ if MCP_AVAILABLE:
|
|||
scope: Optional[str] = Form(None),
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
if not resolved_client_id:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -3011,3 +3011,320 @@ async def test_token_endpoint_client_secret_basic_without_secret_returns_400():
|
|||
code_verifier="verifier",
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# Non-oauth2 (auth_type=none, access-group gated) servers must not be
|
||||
# driven through the gateway OAuth authorize/token/register/discovery
|
||||
# flow, and must not be advertised as OAuth-protected in discovery docs.
|
||||
# -------------------------------------------------------------------
|
||||
|
||||
|
||||
def _access_group_none_server(server_name="access_group_server"):
|
||||
"""A non-oauth2, access-group gated MCP server: no client_id, no OAuth."""
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
return MCPServer(
|
||||
server_id=server_name,
|
||||
name=server_name,
|
||||
server_name=server_name,
|
||||
alias=server_name,
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
access_groups=["eng"],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_endpoint_rejects_non_oauth2_server():
|
||||
"""authorize() against a none-auth server returns an accurate 'does not use OAuth' 400,
|
||||
not the misleading 'client_id is required' that fired before the auth_type was checked."""
|
||||
try:
|
||||
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()
|
||||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await authorize(
|
||||
request=mock_request,
|
||||
client_id=None,
|
||||
mcp_server_name="access_group_server",
|
||||
redirect_uri="http://127.0.0.1:60108/callback",
|
||||
state="test_state",
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
detail_text = str(exc_info.value.detail)
|
||||
assert "does not use OAuth" in detail_text
|
||||
assert "client_id is required" not in detail_text
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_rejects_non_oauth2_server():
|
||||
"""token_endpoint() against a none-auth server returns 'does not use OAuth' 400 instead
|
||||
of the misleading 'token url is not set'."""
|
||||
try:
|
||||
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()
|
||||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="authorization_code",
|
||||
code="auth-code",
|
||||
redirect_uri="http://localhost/callback",
|
||||
client_id="some-client",
|
||||
mcp_server_name="access_group_server",
|
||||
client_secret=None,
|
||||
code_verifier="verifier",
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
detail_text = str(exc_info.value.detail)
|
||||
assert "does not use OAuth" in detail_text
|
||||
assert "token url is not set" not in detail_text
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_client_rejects_non_oauth2_server():
|
||||
"""register_client() against a named none-auth server returns 'does not use OAuth' 400
|
||||
instead of the misleading 'authorization url is not set'."""
|
||||
try:
|
||||
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()
|
||||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
|
||||
new=AsyncMock(return_value={}),
|
||||
):
|
||||
await register_client(request=mock_request, mcp_server_name="access_group_server")
|
||||
assert exc_info.value.status_code == 400
|
||||
detail_text = str(exc_info.value.detail)
|
||||
assert "does not use OAuth" in detail_text
|
||||
assert "authorization url is not set" not in detail_text
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_protected_resource_404_for_non_oauth2_server():
|
||||
"""Discovery must not advertise a none-auth server as an OAuth-protected resource."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_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()
|
||||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _build_oauth_protected_resource_response(
|
||||
request=mock_request,
|
||||
mcp_server_name="access_group_server",
|
||||
use_standard_pattern=False,
|
||||
)
|
||||
assert exc_info.value.status_code == 404
|
||||
assert "not an OAuth-protected resource" in str(exc_info.value.detail)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_authorization_server_404_for_non_oauth2_server():
|
||||
"""Discovery must not advertise a none-auth server as an OAuth authorization server."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_authorization_server_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()
|
||||
server = _access_group_none_server()
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_build_oauth_authorization_server_response(
|
||||
request=mock_request,
|
||||
mcp_server_name="access_group_server",
|
||||
)
|
||||
assert exc_info.value.status_code == 404
|
||||
assert "not an OAuth authorization server" in str(exc_info.value.detail)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_protected_resource_passthrough_none_auth_not_404():
|
||||
"""Regression guard for the protected-resource auth_type gate placement: a none-auth
|
||||
server that opted into OAuth pass-through must still proxy upstream metadata, it must
|
||||
NOT be 404'd. The gate has to sit after the pass-through branch."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_protected_resource_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
passthrough_server = MCPServer(
|
||||
server_id="passthrough_server",
|
||||
name="passthrough_server",
|
||||
server_name="passthrough_server",
|
||||
alias="passthrough_server",
|
||||
url="https://upstream.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
oauth_passthrough=True,
|
||||
extra_headers=["Authorization"],
|
||||
)
|
||||
global_mcp_server_manager.registry[passthrough_server.server_id] = passthrough_server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.fetch_upstream_oauth_protected_resource",
|
||||
new=AsyncMock(return_value={"authorization_servers": ["https://upstream-idp.example.com"]}),
|
||||
):
|
||||
response = await _build_oauth_protected_resource_response(
|
||||
request=mock_request,
|
||||
mcp_server_name="passthrough_server",
|
||||
use_standard_pattern=False,
|
||||
)
|
||||
assert response["authorization_servers"] == ["https://upstream-idp.example.com"]
|
||||
assert response["resource"].endswith("/passthrough_server/mcp")
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_protected_resource_404_for_unknown_server_name():
|
||||
"""A discovery request for an unknown server name returns the same 404 as a non-oauth2
|
||||
server (not a 200 metadata doc with broken URLs), so the well-known paths cannot be used
|
||||
to enumerate non-OAuth server names."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_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()
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _build_oauth_protected_resource_response(
|
||||
request=mock_request,
|
||||
mcp_server_name="does_not_exist",
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_authorization_server_404_for_unknown_server_name():
|
||||
"""A named authorization-server discovery request for an unknown server returns 404, not a
|
||||
200 metadata document pointing at non-existent /{name}/authorize and /{name}/token."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
_build_oauth_authorization_server_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()
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_build_oauth_authorization_server_response(
|
||||
request=mock_request,
|
||||
mcp_server_name="does_not_exist",
|
||||
)
|
||||
assert exc_info.value.status_code == 404
|
||||
|
|
|
|||
|
|
@ -2068,6 +2068,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
server.auth_type = MCPAuth.oauth2
|
||||
authorize_response = MagicMock()
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
|
|
@ -2110,6 +2111,91 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
scope="scope1",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_authorize_rejects_non_oauth2_server(self):
|
||||
"""mcp_authorize must reject a none-auth server with an accurate 'does not use OAuth'
|
||||
400 before the client_id check, never delegating to authorize_with_server."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
mcp_authorize,
|
||||
)
|
||||
|
||||
server = generate_mock_mcp_server_config_record(server_id="none-server")
|
||||
server.auth_type = MCPAuth.none
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
|
||||
return_value=server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
|
||||
AsyncMock(),
|
||||
) as authorize_mock,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await mcp_authorize(
|
||||
request=MagicMock(),
|
||||
server_id="none-server",
|
||||
user_api_key_dict=admin_auth,
|
||||
client_id=None,
|
||||
redirect_uri="https://example.com/callback",
|
||||
state="state123",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
detail_text = str(exc_info.value.detail)
|
||||
assert "does not use OAuth" in detail_text
|
||||
assert "missing_client_id" not in detail_text
|
||||
authorize_mock.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_token_rejects_non_oauth2_server(self):
|
||||
"""mcp_token must reject a none-auth server with 'does not use OAuth' 400 before the
|
||||
client_id check, never delegating to exchange_token_with_server."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
mcp_token,
|
||||
)
|
||||
|
||||
server = generate_mock_mcp_server_config_record(server_id="none-server")
|
||||
server.auth_type = MCPAuth.none
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
|
||||
return_value=server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
|
||||
AsyncMock(),
|
||||
) as exchange_mock,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await mcp_token(
|
||||
request=MagicMock(),
|
||||
server_id="none-server",
|
||||
user_api_key_dict=admin_auth,
|
||||
grant_type="authorization_code",
|
||||
code="code-123",
|
||||
redirect_uri="https://example.com/callback",
|
||||
client_id=None,
|
||||
client_secret=None,
|
||||
code_verifier="verifier",
|
||||
refresh_token=None,
|
||||
scope=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
detail_text = str(exc_info.value.detail)
|
||||
assert "does not use OAuth" in detail_text
|
||||
assert "missing_client_id" not in detail_text
|
||||
exchange_mock.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_token_proxies_to_exchange_endpoint(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
|
|
@ -2118,6 +2204,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
server.auth_type = MCPAuth.oauth2
|
||||
exchange_response = {"access_token": "token"}
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
|
|
@ -2170,6 +2257,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
server.auth_type = MCPAuth.oauth2
|
||||
exchange_response = {"access_token": "new-token", "refresh_token": "new-rt"}
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
|
|
@ -2222,6 +2310,7 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
server.auth_type = MCPAuth.oauth2
|
||||
register_response = {"client_id": "generated"}
|
||||
request_body = {
|
||||
"client_name": "LiteLLM",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue