mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(mcp): return 401+WWW-Authenticate (not 500) for unauthenticated MCP bootstrap
This commit is contained in:
parent
144279eb57
commit
57e8c6f59d
2 changed files with 287 additions and 0 deletions
|
|
@ -2499,6 +2499,55 @@ if MCP_AVAILABLE:
|
|||
verbose_logger.exception(f"Error executing local tool {name}: {str(e)}")
|
||||
return [TextContent(text=f"Error: {str(e)}", type="text")]
|
||||
|
||||
def _maybe_raise_oauth_bootstrap_challenge(scope: Scope, path: str) -> None:
|
||||
"""
|
||||
For an unauthenticated cold-start request to an OAuth2-configured MCP
|
||||
server, raise 401 + WWW-Authenticate so the client can discover OAuth
|
||||
metadata via /.well-known and start PKCE.
|
||||
|
||||
Skips when:
|
||||
- The path doesn't resolve to a named MCP server.
|
||||
- No resolved server has auth_type == oauth2.
|
||||
- The request carries any Authorization header (let the existing auth
|
||||
+ 401 logic downstream handle those).
|
||||
|
||||
Without this pre-check, a request with no Authorization header reaches
|
||||
strict API-key validation in extract_mcp_auth_context, which raises a
|
||||
ProxyException that the catch-all coerces to 500 — so the client never
|
||||
discovers the OAuth metadata and can't bootstrap PKCE.
|
||||
"""
|
||||
for header_name, _ in scope.get("headers", []) or []:
|
||||
if isinstance(header_name, bytes):
|
||||
if header_name.lower() == b"authorization":
|
||||
return
|
||||
elif (
|
||||
isinstance(header_name, str) and header_name.lower() == "authorization"
|
||||
):
|
||||
return
|
||||
|
||||
mcp_servers = _get_mcp_servers_in_path(path)
|
||||
if not mcp_servers:
|
||||
return
|
||||
|
||||
request = StarletteRequest(scope)
|
||||
client_ip = IPAddressUtils.get_mcp_client_ip(request)
|
||||
|
||||
for server_name in mcp_servers:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_name(
|
||||
server_name, client_ip=client_ip
|
||||
)
|
||||
if server and server.auth_type == MCPAuth.oauth2:
|
||||
base_url = get_request_base_url(request)
|
||||
authorization_uri = (
|
||||
f"Bearer authorization_uri="
|
||||
f"{base_url}/.well-known/oauth-authorization-server/{server_name}"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Unauthorized",
|
||||
headers={"www-authenticate": authorization_uri},
|
||||
)
|
||||
|
||||
def _get_mcp_servers_in_path(path: str) -> Optional[List[str]]:
|
||||
"""
|
||||
Get the MCP servers from the path
|
||||
|
|
@ -2730,6 +2779,15 @@ if MCP_AVAILABLE:
|
|||
"""Handle MCP requests through StreamableHTTP."""
|
||||
try:
|
||||
path = scope.get("path", "")
|
||||
|
||||
# Pre-emptive OAuth challenge for the cold-start bootstrap case.
|
||||
# If the path resolves to an OAuth2-configured server AND the
|
||||
# request carries no Authorization header, fire 401 +
|
||||
# WWW-Authenticate immediately. Otherwise the empty key would
|
||||
# raise ProxyException downstream and the catch-all would coerce
|
||||
# that to 500, so the client never discovers OAuth metadata.
|
||||
_maybe_raise_oauth_bootstrap_challenge(scope, path)
|
||||
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
|
|
@ -2868,6 +2926,7 @@ if MCP_AVAILABLE:
|
|||
"""Handle MCP requests through SSE."""
|
||||
try:
|
||||
path = scope.get("path", "")
|
||||
_maybe_raise_oauth_bootstrap_challenge(scope, path)
|
||||
(
|
||||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -591,3 +591,231 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
|
||||
assert mock_get_stored_token.await_count == 1
|
||||
assert mock_handle_request.await_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_bootstrap_returns_401_without_mocking_extract_mcp_auth_context():
|
||||
"""
|
||||
Regression: an unauthenticated POST to /mcp/{server} where the server is
|
||||
OAuth2-configured must return 401 + WWW-Authenticate, not 500.
|
||||
|
||||
This test deliberately does NOT mock extract_mcp_auth_context. It registers
|
||||
a real OAuth2 server in the manager and lets the request flow through the
|
||||
pre-check helper. Without the pre-check, the empty Authorization header
|
||||
would reach strict API-key validation and raise a ProxyException that the
|
||||
catch-all coerces to 500.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
public_server = MCPServer(
|
||||
server_id="bootstrap_test_server",
|
||||
name="bootstrap_test_server",
|
||||
server_name="bootstrap_test_server",
|
||||
alias="bootstrap_test_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://idp.example/authorize",
|
||||
token_url="https://idp.example/token",
|
||||
)
|
||||
global_mcp_server_manager.registry[public_server.server_id] = public_server
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bootstrap_test_server",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.IPAddressUtils.get_mcp_client_ip",
|
||||
return_value=None,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 401
|
||||
assert "www-authenticate" in exc.headers
|
||||
assert (
|
||||
"/.well-known/oauth-authorization-server/bootstrap_test_server"
|
||||
in exc.headers["www-authenticate"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_bootstrap_skips_when_authorization_header_present():
|
||||
"""
|
||||
Pre-check must NOT fire when the client sends any Authorization header —
|
||||
let the existing auth flow + 401 logic decide.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
public_server = MCPServer(
|
||||
server_id="bootstrap_test_server",
|
||||
name="bootstrap_test_server",
|
||||
server_name="bootstrap_test_server",
|
||||
alias="bootstrap_test_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="real-client-id",
|
||||
client_secret=None,
|
||||
authorization_url="https://idp.example/authorize",
|
||||
token_url="https://idp.example/token",
|
||||
)
|
||||
global_mcp_server_manager.registry[public_server.server_id] = public_server
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bootstrap_test_server",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"authorization", b"Bearer sk-some-litellm-key"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
|
||||
sentinel_extract = AsyncMock(
|
||||
side_effect=RuntimeError("downstream_reached_as_expected")
|
||||
)
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
sentinel_extract,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.IPAddressUtils.get_mcp_client_ip",
|
||||
return_value=None,
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert sentinel_extract.await_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_bootstrap_skips_when_path_does_not_resolve_to_named_server():
|
||||
"""
|
||||
Pre-check no-op when path is /mcp (root) without a server name — flows
|
||||
into the existing extract_mcp_auth_context path.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
|
||||
sentinel_extract = AsyncMock(
|
||||
side_effect=RuntimeError("downstream_reached_as_expected")
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
sentinel_extract,
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
|
||||
assert sentinel_extract.await_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_bootstrap_skips_for_non_oauth_server():
|
||||
"""
|
||||
Pre-check no-op for path-resolved server whose auth_type != oauth2.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
api_key_server = MCPServer(
|
||||
server_id="api_key_server",
|
||||
name="api_key_server",
|
||||
server_name="api_key_server",
|
||||
alias="api_key_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.api_key,
|
||||
)
|
||||
global_mcp_server_manager.registry[api_key_server.server_id] = api_key_server
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/api_key_server",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
|
||||
sentinel_extract = AsyncMock(
|
||||
side_effect=RuntimeError("downstream_reached_as_expected")
|
||||
)
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
sentinel_extract,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.IPAddressUtils.get_mcp_client_ip",
|
||||
return_value=None,
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
assert sentinel_extract.await_count == 1
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue