feat: prepare ARR for controlled public deployment
This commit is contained in:
93
arr_mcp/auth.py
Normal file
93
arr_mcp/auth.py
Normal file
@@ -0,0 +1,93 @@
|
||||
"""Transport-level bearer protection for the ARR MCP ASGI application."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Tuple
|
||||
|
||||
|
||||
MCP_BEARER_ENV = "ARR_MCP_BEARER_TOKEN"
|
||||
ASGIApp = Callable[
|
||||
[Dict[str, Any], Callable[[], Awaitable[Dict[str, Any]]], Callable[[Dict[str, Any]], Awaitable[None]]],
|
||||
Awaitable[None],
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BearerAuthConfig:
|
||||
token: str = field(repr=False)
|
||||
|
||||
@classmethod
|
||||
def from_environment(cls) -> "BearerAuthConfig":
|
||||
value = os.environ.get(MCP_BEARER_ENV, "")
|
||||
try:
|
||||
encoded = value.encode("ascii")
|
||||
except UnicodeEncodeError:
|
||||
encoded = b""
|
||||
if (
|
||||
not 32 <= len(encoded) <= 256
|
||||
or any(byte <= 32 or byte == 127 for byte in encoded)
|
||||
):
|
||||
raise ValueError(f"{MCP_BEARER_ENV} is unavailable or invalid")
|
||||
return cls(value)
|
||||
|
||||
|
||||
class BearerAuthASGI:
|
||||
"""Pure ASGI middleware; never logs or exposes the configured token."""
|
||||
|
||||
def __init__(self, application: ASGIApp, config: BearerAuthConfig) -> None:
|
||||
self._application = application
|
||||
self._expected_sha256 = hashlib.sha256(
|
||||
("Bearer " + config.token).encode("ascii")
|
||||
).digest()
|
||||
|
||||
@staticmethod
|
||||
def _authorization_values(scope: Dict[str, Any]) -> List[bytes]:
|
||||
headers: List[Tuple[bytes, bytes]] = scope.get("headers", [])
|
||||
return [
|
||||
value
|
||||
for name, value in headers
|
||||
if name.lower() == b"authorization"
|
||||
]
|
||||
|
||||
def _authorized(self, scope: Dict[str, Any]) -> bool:
|
||||
values = self._authorization_values(scope)
|
||||
if len(values) != 1:
|
||||
return False
|
||||
return hmac.compare_digest(
|
||||
hashlib.sha256(values[0]).digest(),
|
||||
self._expected_sha256,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _reject(
|
||||
send: Callable[[Dict[str, Any]], Awaitable[None]],
|
||||
) -> None:
|
||||
body = b'{"error":"unauthorized"}'
|
||||
await send(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"status": 401,
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"content-length", str(len(body)).encode("ascii")),
|
||||
(b"cache-control", b"no-store"),
|
||||
(b"www-authenticate", b"Bearer"),
|
||||
],
|
||||
}
|
||||
)
|
||||
await send({"type": "http.response.body", "body": body})
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
scope: Dict[str, Any],
|
||||
receive: Callable[[], Awaitable[Dict[str, Any]]],
|
||||
send: Callable[[Dict[str, Any]], Awaitable[None]],
|
||||
) -> None:
|
||||
if scope.get("type") == "http" and not self._authorized(scope):
|
||||
await self._reject(send)
|
||||
return
|
||||
await self._application(scope, receive, send)
|
||||
Reference in New Issue
Block a user