|
| 1 | +"""Authentication module for Late MCP HTTP server.""" |
| 2 | + |
| 3 | +import os |
| 4 | +import secrets |
| 5 | + |
| 6 | +from starlette.requests import Request |
| 7 | +from starlette.responses import JSONResponse |
| 8 | + |
| 9 | + |
| 10 | +def get_server_api_key() -> str: |
| 11 | + """ |
| 12 | + Get API key from environment variable. |
| 13 | +
|
| 14 | + Returns: |
| 15 | + The MCP server API key. |
| 16 | +
|
| 17 | + Raises: |
| 18 | + ValueError: If MCP_SERVER_API_KEY is not set. |
| 19 | + """ |
| 20 | + api_key = os.getenv("MCP_SERVER_API_KEY") |
| 21 | + if not api_key: |
| 22 | + raise ValueError( |
| 23 | + "MCP_SERVER_API_KEY environment variable not set. " |
| 24 | + "Please set it to secure your MCP server." |
| 25 | + ) |
| 26 | + return api_key |
| 27 | + |
| 28 | + |
| 29 | +def extract_api_key(request: Request) -> str | None: |
| 30 | + """ |
| 31 | + Extract API key from request (header or query param). |
| 32 | +
|
| 33 | + Checks in order: |
| 34 | + 1. Authorization header (Bearer token) |
| 35 | + 2. X-API-Key header |
| 36 | + 3. api_key query parameter |
| 37 | +
|
| 38 | + Args: |
| 39 | + request: The incoming Starlette request. |
| 40 | +
|
| 41 | + Returns: |
| 42 | + The extracted API key, or None if not found. |
| 43 | + """ |
| 44 | + # Try Authorization header first: "Bearer <key>" |
| 45 | + auth_header = request.headers.get("Authorization") |
| 46 | + if auth_header and auth_header.startswith("Bearer "): |
| 47 | + return auth_header[7:] # Remove "Bearer " prefix |
| 48 | + |
| 49 | + # Try X-API-Key header |
| 50 | + api_key_header = request.headers.get("X-API-Key") |
| 51 | + if api_key_header: |
| 52 | + return api_key_header |
| 53 | + |
| 54 | + # Try query parameter as fallback |
| 55 | + return request.query_params.get("api_key") |
| 56 | + |
| 57 | + |
| 58 | +def verify_api_key(request: Request) -> bool: |
| 59 | + """ |
| 60 | + Verify API key from request matches server key. |
| 61 | +
|
| 62 | + Uses secrets.compare_digest for timing-attack resistance. |
| 63 | +
|
| 64 | + Args: |
| 65 | + request: The incoming Starlette request. |
| 66 | +
|
| 67 | + Returns: |
| 68 | + True if API key is valid, False otherwise. |
| 69 | + """ |
| 70 | + try: |
| 71 | + expected_key = get_server_api_key() |
| 72 | + provided_key = extract_api_key(request) |
| 73 | + |
| 74 | + if not provided_key: |
| 75 | + return False |
| 76 | + |
| 77 | + # Use secrets.compare_digest for timing-attack resistance |
| 78 | + return secrets.compare_digest(expected_key, provided_key) |
| 79 | + except Exception: |
| 80 | + # If any error occurs (e.g., env var not set), deny access |
| 81 | + return False |
| 82 | + |
| 83 | + |
| 84 | +async def require_api_key(request: Request, call_next): |
| 85 | + """ |
| 86 | + Middleware to require API key on all requests except health check. |
| 87 | +
|
| 88 | + Args: |
| 89 | + request: The incoming Starlette request. |
| 90 | + call_next: The next middleware or route handler. |
| 91 | +
|
| 92 | + Returns: |
| 93 | + The response from the next handler, or 401 if unauthorized. |
| 94 | + """ |
| 95 | + # Allow health check without authentication |
| 96 | + if request.url.path == "/health": |
| 97 | + return await call_next(request) |
| 98 | + |
| 99 | + # Verify API key for all other requests |
| 100 | + if not verify_api_key(request): |
| 101 | + return JSONResponse( |
| 102 | + {"error": "Invalid or missing API key"}, status_code=401 |
| 103 | + ) |
| 104 | + |
| 105 | + return await call_next(request) |
0 commit comments