120 lines
4.4 KiB
Python
120 lines
4.4 KiB
Python
"""动态 Bearer Token 鉴权(通过 mcp-auth 后端 API 校验,不直连鉴权库)。
|
||
|
||
token 校验由 MCP 服务调用 mcp-auth 后端 `POST /api/auth/verify-token` 完成,
|
||
后端统一查 mcp_auth.mcp_token 表。后端通过 per-service API Key 识别调用方服务,
|
||
MCP 服务无需自报 service。进程内 LRU 缓存 30s,减少 API 调用。
|
||
|
||
环境变量:
|
||
MCP_PUBLIC_URL 服务对外地址(OAuth 资源元数据),默认 http://localhost:{port}
|
||
MCP_AUTH_API_URL mcp-auth 后端地址,如 http://localhost:8000
|
||
MCP_AUTH_API_KEY 本服务的 API Key(由 mcp-auth 后端签发,per-service 唯一)
|
||
MCP_AUTH_CACHE_TTL 缓存秒数,默认 30
|
||
"""
|
||
|
||
import hashlib
|
||
import os
|
||
import time
|
||
|
||
import httpx
|
||
from mcp.server.auth.middleware.auth_context import get_access_token
|
||
from mcp.server.auth.provider import AccessToken
|
||
from mcp.server.auth.settings import AuthSettings
|
||
|
||
|
||
_CACHE_TTL = int(os.getenv("MCP_AUTH_CACHE_TTL", "30"))
|
||
|
||
# httpx 异步客户端(进程级单例,复用连接池)
|
||
_http_client: httpx.AsyncClient | None = None
|
||
|
||
|
||
def _get_http_client() -> httpx.AsyncClient:
|
||
global _http_client
|
||
if _http_client is None:
|
||
_http_client = httpx.AsyncClient(timeout=5.0)
|
||
return _http_client
|
||
|
||
|
||
async def close_auth_client() -> None:
|
||
"""关闭 httpx 客户端(进程退出时调用)。"""
|
||
global _http_client
|
||
if _http_client is not None:
|
||
await _http_client.aclose()
|
||
_http_client = None
|
||
|
||
|
||
class ApiTokenVerifier:
|
||
"""通过 mcp-auth 后端 API 校验 Bearer Token。
|
||
|
||
进程内 LRU 缓存(token_hash -> (valid, client_id, fetched_at)),TTL 由 MCP_AUTH_CACHE_TTL 控制。
|
||
"""
|
||
|
||
def __init__(self, api_url: str, api_key: str):
|
||
self._api_url = api_url.rstrip("/")
|
||
self._api_key = api_key
|
||
self._cache: dict[str, tuple[bool, str | None, float]] = {}
|
||
|
||
def invalidate(self, token: str | None = None) -> None:
|
||
"""清缓存:token=None 清全部,否则清单个。"""
|
||
if token is None:
|
||
self._cache.clear()
|
||
else:
|
||
self._cache.pop(hashlib.sha256(token.encode()).hexdigest(), None)
|
||
|
||
async def verify_token(self, token: str) -> AccessToken | None:
|
||
token_hash = hashlib.sha256(token.encode()).hexdigest()
|
||
|
||
# 1. 查缓存
|
||
now = time.time()
|
||
cached = self._cache.get(token_hash)
|
||
if cached is not None and (now - cached[2]) < _CACHE_TTL:
|
||
valid, client_id, _ = cached
|
||
if not valid:
|
||
return None
|
||
return AccessToken(token=token, client_id=client_id, scopes=[], expires_at=None)
|
||
|
||
# 2. 调用后端校验 API(后端从 X-API-Key 识别 service,无需传 service)
|
||
try:
|
||
client = _get_http_client()
|
||
resp = await client.post(
|
||
f"{self._api_url}/api/auth/verify-token",
|
||
headers={"X-API-Key": self._api_key},
|
||
json={"token": token},
|
||
)
|
||
except Exception:
|
||
# 后端不可达,缓存短时间避免雪崩
|
||
self._cache[token_hash] = (False, None, now)
|
||
return None
|
||
|
||
if resp.status_code == 200:
|
||
data = resp.json()
|
||
valid = data.get("valid", False)
|
||
client_id = data.get("client_id")
|
||
self._cache[token_hash] = (valid, client_id, now)
|
||
if valid:
|
||
return AccessToken(token=token, client_id=client_id, scopes=[], expires_at=None)
|
||
return None
|
||
else:
|
||
# 非 200(含 401 API Key 错误),缓存避免雪崩
|
||
self._cache[token_hash] = (False, None, now)
|
||
return None
|
||
|
||
|
||
def get_auth(port: int) -> tuple[AuthSettings, ApiTokenVerifier]:
|
||
"""构建 MCPServer 的 (auth, token_verifier) 参数。
|
||
|
||
service 由后端从 API Key 自动识别,MCP 服务无需配置 MCP_AUTH_SERVICE。
|
||
"""
|
||
url = os.getenv("MCP_PUBLIC_URL") or f"http://localhost:{port}"
|
||
api_url = os.getenv("MCP_AUTH_API_URL", "http://localhost:8000")
|
||
api_key = os.getenv("MCP_AUTH_API_KEY", "")
|
||
return (
|
||
AuthSettings(issuer_url=url, resource_server_url=url),
|
||
ApiTokenVerifier(api_url=api_url, api_key=api_key),
|
||
)
|
||
|
||
|
||
def get_caller() -> str:
|
||
"""工具内获取当前调用方标识(未认证时返回 anonymous)。"""
|
||
access_token = get_access_token()
|
||
return access_token.client_id if access_token else "anonymous"
|