Files
mcp-auth/mcp-for-erp/src/auth.py
T
2026-08-31 13:06:37 +08:00

120 lines
4.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""动态 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"