mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 15:19:10 +02:00
Fix/cimd flow issue (#4206)
This commit is contained in:
parent
4d5627ba5d
commit
0668eb7b80
4 changed files with 135 additions and 7 deletions
|
|
@ -3,6 +3,9 @@
|
|||
This module provides enhanced versions of MCP SDK authentication middleware
|
||||
that return more helpful error messages for developers troubleshooting
|
||||
authentication issues.
|
||||
|
||||
Implements RFC 6750 §3.1 compliance by distinguishing between missing
|
||||
authentication (no error attribute) and invalid authentication (with error).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -12,7 +15,7 @@ import json
|
|||
from mcp.server.auth.middleware.bearer_auth import (
|
||||
RequireAuthMiddleware as SDKRequireAuthMiddleware,
|
||||
)
|
||||
from starlette.types import Send
|
||||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
|
||||
|
|
@ -25,8 +28,88 @@ class RequireAuthMiddleware(SDKRequireAuthMiddleware):
|
|||
Extends the SDK's RequireAuthMiddleware to provide more actionable
|
||||
error messages when authentication fails. This helps developers
|
||||
understand what went wrong and how to fix it.
|
||||
|
||||
Also implements RFC 6750 §3.1 compliance by distinguishing between
|
||||
missing authentication (initial discovery) and invalid authentication
|
||||
(token validation failure).
|
||||
"""
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
scope: Scope,
|
||||
receive: Receive,
|
||||
send: Send,
|
||||
) -> None:
|
||||
"""Process ASGI scope, distinguishing missing vs invalid auth.
|
||||
|
||||
Per RFC 6750 §3.1:
|
||||
- Missing auth (no Authorization header) → 401 without error attribute
|
||||
- Invalid auth (Authorization header present) → 401 with error attribute
|
||||
|
||||
This ensures OAuth flow initialization works correctly in MCP clients
|
||||
during initial discovery phase.
|
||||
|
||||
Args:
|
||||
scope: ASGI scope
|
||||
receive: ASGI receive callable
|
||||
send: ASGI send callable
|
||||
"""
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
# Check if Authorization header is present
|
||||
headers = scope.get("headers", [])
|
||||
has_auth_header = any(
|
||||
header[0].lower() == b"authorization" for header in headers
|
||||
)
|
||||
|
||||
if not has_auth_header:
|
||||
# Per RFC 6750 §3.1: missing auth should not include error attribute
|
||||
await self._send_missing_auth(send)
|
||||
return
|
||||
|
||||
# Authorization header is present - use parent's validation logic
|
||||
# This will check token validity and call _send_auth_error if invalid
|
||||
await super().__call__(scope, receive, send)
|
||||
|
||||
async def _send_missing_auth(self, send: Send) -> None:
|
||||
"""Send 401 response for missing authentication (RFC 6750 §3.1 compliant).
|
||||
|
||||
When a request lacks any authentication information, per RFC 6750 §3.1:
|
||||
"If the request lacks any authentication information, the error
|
||||
attribute SHOULD NOT be included."
|
||||
|
||||
This allows MCP clients to properly initiate OAuth flow during
|
||||
initial discovery phase.
|
||||
|
||||
Args:
|
||||
send: ASGI send callable
|
||||
"""
|
||||
www_auth_parts = []
|
||||
if self.resource_metadata_url:
|
||||
www_auth_parts.append(f'resource_metadata="{self.resource_metadata_url}"')
|
||||
|
||||
www_authenticate = (
|
||||
("Bearer " + ", ".join(www_auth_parts)) if www_auth_parts else "Bearer"
|
||||
)
|
||||
|
||||
await send(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"status": 401,
|
||||
"headers": [
|
||||
(b"content-length", b"0"),
|
||||
(b"www-authenticate", www_authenticate.encode()),
|
||||
],
|
||||
}
|
||||
)
|
||||
await send({"type": "http.response.body", "body": b""})
|
||||
|
||||
logger.debug(
|
||||
"Missing auth: sent 401 without error attribute (RFC 6750 §3.1 compliant)"
|
||||
)
|
||||
|
||||
async def _send_auth_error(
|
||||
self, send: Send, status_code: int, error: str, description: str
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -2075,6 +2075,12 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
|
|||
revocation_options,
|
||||
)
|
||||
metadata.client_id_metadata_document_supported = True
|
||||
existing = metadata.token_endpoint_auth_methods_supported or []
|
||||
metadata.token_endpoint_auth_methods_supported = [
|
||||
*existing,
|
||||
"private_key_jwt",
|
||||
"none",
|
||||
]
|
||||
handler = MetadataHandler(metadata)
|
||||
methods = route.methods or ["GET", "OPTIONS"]
|
||||
|
||||
|
|
|
|||
|
|
@ -207,7 +207,7 @@ class TestOAuthProxyInitialization:
|
|||
assert proxy._redirect_path == "/auth/callback"
|
||||
|
||||
async def test_metadata_advertises_cimd_support(self, jwt_verifier):
|
||||
"""OAuth metadata should advertise CIMD support when enabled."""
|
||||
"""OAuth metadata should advertise CIMD and public-client auth support."""
|
||||
proxy = OAuthProxy(
|
||||
upstream_authorization_endpoint="https://auth.example.com/authorize",
|
||||
upstream_token_endpoint="https://auth.example.com/token",
|
||||
|
|
@ -231,6 +231,12 @@ class TestOAuthProxyInitialization:
|
|||
assert response.status_code == 200
|
||||
metadata = response.json()
|
||||
assert metadata.get("client_id_metadata_document_supported") is True
|
||||
assert set(metadata.get("token_endpoint_auth_methods_supported")) == {
|
||||
"client_secret_post",
|
||||
"client_secret_basic",
|
||||
"private_key_jwt",
|
||||
"none",
|
||||
}
|
||||
|
||||
|
||||
class TestOptionalClientSecret:
|
||||
|
|
|
|||
|
|
@ -202,8 +202,8 @@ class TestEnhancedRequireAuthMiddleware:
|
|||
base_url="https://test.com",
|
||||
)
|
||||
|
||||
def test_invalid_token_enhanced_error_message(self, jwt_verifier):
|
||||
"""Test that invalid_token errors have enhanced error messages."""
|
||||
def test_missing_auth_no_error_attribute(self, jwt_verifier):
|
||||
"""Test that missing auth returns 401 without error attribute (RFC 6750 §3.1)."""
|
||||
from fastmcp.server.http import create_streamable_http_app
|
||||
|
||||
server = FastMCP("Test Server")
|
||||
|
|
@ -225,6 +225,36 @@ class TestEnhancedRequireAuthMiddleware:
|
|||
assert response.status_code == 401
|
||||
assert "www-authenticate" in response.headers
|
||||
|
||||
# Per RFC 6750 §3.1: no error attribute when auth is missing
|
||||
www_auth = response.headers["www-authenticate"]
|
||||
assert "error=" not in www_auth
|
||||
assert response.content == b""
|
||||
|
||||
def test_invalid_token_enhanced_error_message(self, jwt_verifier):
|
||||
"""Test that invalid_token errors have enhanced error messages."""
|
||||
from fastmcp.server.http import create_streamable_http_app
|
||||
|
||||
server = FastMCP("Test Server")
|
||||
|
||||
@server.tool
|
||||
def test_tool() -> str:
|
||||
return "test"
|
||||
|
||||
app = create_streamable_http_app(
|
||||
server=server,
|
||||
streamable_http_path="/mcp",
|
||||
auth=jwt_verifier,
|
||||
)
|
||||
|
||||
with TestClient(app) as client:
|
||||
# Request WITH an invalid Authorization header
|
||||
response = client.post(
|
||||
"/mcp", headers={"Authorization": "Bearer invalid-token"}
|
||||
)
|
||||
|
||||
assert response.status_code == 401
|
||||
assert "www-authenticate" in response.headers
|
||||
|
||||
# Check enhanced error message
|
||||
data = response.json()
|
||||
assert data["error"] == "invalid_token"
|
||||
|
|
@ -233,7 +263,7 @@ class TestEnhancedRequireAuthMiddleware:
|
|||
assert "automatically re-register" in data["error_description"]
|
||||
|
||||
def test_invalid_token_www_authenticate_header_format(self, jwt_verifier):
|
||||
"""Test that WWW-Authenticate header format matches SDK."""
|
||||
"""Test that invalid token WWW-Authenticate header includes error attribute."""
|
||||
from fastmcp.server.http import create_streamable_http_app
|
||||
|
||||
server = FastMCP("Test Server")
|
||||
|
|
@ -244,12 +274,15 @@ class TestEnhancedRequireAuthMiddleware:
|
|||
)
|
||||
|
||||
with TestClient(app) as client:
|
||||
response = client.post("/mcp")
|
||||
# Request WITH an invalid token
|
||||
response = client.post(
|
||||
"/mcp", headers={"Authorization": "Bearer invalid-token"}
|
||||
)
|
||||
|
||||
assert response.status_code == 401
|
||||
www_auth = response.headers["www-authenticate"]
|
||||
|
||||
# Should follow Bearer challenge format
|
||||
# Should follow Bearer challenge format with error
|
||||
assert www_auth.startswith("Bearer ")
|
||||
assert 'error="invalid_token"' in www_auth
|
||||
assert "error_description=" in www_auth
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue