添加服务注册及api鉴权

This commit is contained in:
2026-08-31 13:06:37 +08:00
parent cbff2283e2
commit 6fe61ff094
15 changed files with 764 additions and 218 deletions
+1
View File
@@ -1,3 +1,4 @@
mcp>=1.0.0
pydantic>=2.0.0
asyncpg>=0.30.0
httpx>=0.27.0
+63 -90
View File
@@ -1,73 +1,60 @@
"""动态 Bearer Token 鉴权(Authorization)。
"""动态 Bearer Token 鉴权(通过 mcp-auth 后端 API 校验,不直连鉴权库)。
token 校验查 mcp_auth.mcp_token 表(sha256 hash 比对),进程内 LRU 缓存 30s。
token 表由独立管理后台(mcp-auth-admin)维护,MCP 服务只读。
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_DB_HOST 鉴权库主机
MCP_AUTH_DB_PORT 鉴权库端口,默认 5432
MCP_AUTH_DB_USER 鉴权库用户
MCP_AUTH_DB_PASSWORD 鉴权库密码
MCP_AUTH_DB_NAME 鉴权库名,默认 mcp_auth
MCP_AUTH_API_URL mcp-auth 后端地址,如 http://localhost:8000
MCP_AUTH_API_KEY 本服务的 API Key(由 mcp-auth 后端签发,per-service 唯一)
MCP_AUTH_CACHE_TTL 缓存秒数,默认 30
MCP_AUTH_SERVICE 当前服务标识(erp/crm),用于 service_scope 校验和 last_used_svc
"""
import asyncio
import hashlib
import os
import time
import asyncpg
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
# 鉴权库连接池(独立于业务库,进程级单例)
_auth_pool: asyncpg.Pool | None = None
_CACHE_TTL = int(os.getenv("MCP_AUTH_CACHE_TTL", "30"))
async def get_auth_pool() -> asyncpg.Pool:
"""获取鉴权库连接池(首次调用时惰性创建)。"""
global _auth_pool
if _auth_pool is None:
_auth_pool = await asyncpg.create_pool(
host=os.getenv("MCP_AUTH_DB_HOST", "127.0.0.1"),
port=int(os.getenv("MCP_AUTH_DB_PORT", "5432")),
user=os.getenv("MCP_AUTH_DB_USER", "postgres"),
password=os.getenv("MCP_AUTH_DB_PASSWORD", "digiwin"),
database=os.getenv("MCP_AUTH_DB_NAME", "mcp_auth"),
min_size=1,
max_size=5,
)
return _auth_pool
# httpx 异步客户端(进程级单例,复用连接池)
_http_client: httpx.AsyncClient | None = None
async def close_auth_pool() -> None:
"""关闭鉴权库连接池(进程退出时调用)。"""
global _auth_pool
if _auth_pool is not None:
await _auth_pool.close()
_auth_pool = 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
class DbTokenVerifier:
"""查库校验 Bearer Token,命中返回 AccessToken,未命中/已吊销/已过期/服务范围不匹配返回 None(401)。
async def close_auth_client() -> None:
"""关闭 httpx 客户端(进程退出时调用)。"""
global _http_client
if _http_client is not None:
await _http_client.aclose()
_http_client = None
进程内 LRU 缓存(token_hash -> (row|None, fetched_at)),TTL 由 MCP_AUTH_CACHE_TTL 控制。
吊销后最长 TTL 秒内仍可能命中旧状态;需要即时生效可调用 invalidate(token)。
class ApiTokenVerifier:
"""通过 mcp-auth 后端 API 校验 Bearer Token。
进程内 LRU 缓存(token_hash -> (valid, client_id, fetched_at)),TTL 由 MCP_AUTH_CACHE_TTL 控制。
"""
def __init__(self, service: str):
self._service = service # 'erp' / 'crm'
self._cache: dict[str, tuple[asyncpg.Record | None, float]] = {}
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 清全部,否则清单个。供管理后台通知后调用(可选)。"""
"""清缓存:token=None 清全部,否则清单个。"""
if token is None:
self._cache.clear()
else:
@@ -77,66 +64,52 @@ class DbTokenVerifier:
token_hash = hashlib.sha256(token.encode()).hexdigest()
# 1. 查缓存
cached = self._cache.get(token_hash)
now = time.time()
if cached is not None and (now - cached[1]) < _CACHE_TTL:
row = cached[0]
else:
# 2. 查库
pool = await get_auth_pool()
row = await pool.fetchrow(
"""SELECT client_id, status, expires_at, service_scope
FROM mcp_token WHERE token_hash = $1""",
token_hash,
)
self._cache[token_hash] = (row, now)
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)
# 3. 校验
if row is None:
return None
if row["status"] != "active":
return None
if row["expires_at"] is not None and row["expires_at"].timestamp() < now:
return None
# service_scope='both' 对所有服务放行;否则要求精确匹配当前服务
if row["service_scope"] != "both" and row["service_scope"] != self._service:
return None
# 4. 异步更新 last_used_*,不阻塞响应;失败忽略
asyncio.create_task(self._touch(token_hash))
return AccessToken(
token=token,
client_id=row["client_id"],
scopes=[],
expires_at=None,
)
async def _touch(self, token_hash: str) -> None:
"""异步更新 last_used_at/last_used_svc。"""
# 2. 调用后端校验 API(后端从 X-API-Key 识别 service,无需传 service)
try:
pool = await get_auth_pool()
await pool.execute(
"""UPDATE mcp_token
SET last_used_at = now(), last_used_svc = $2
WHERE token_hash = $1""",
token_hash, self._service,
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:
pass # 审计字段更新失败不影响鉴权
# 后端不可达,缓存短时间避免雪崩
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, DbTokenVerifier]:
def get_auth(port: int) -> tuple[AuthSettings, ApiTokenVerifier]:
"""构建 MCPServer 的 (auth, token_verifier) 参数。
连接池在首次 verify_token 时惰性创建,此处不连库。
service 标识取自 MCP_AUTH_SERVICE 环境变量,默认按 port 推断(8001→erp / 8002→crm)。
service 由后端从 API Key 自动识别,MCP 服务无需配置 MCP_AUTH_SERVICE。
"""
url = os.getenv("MCP_PUBLIC_URL") or f"http://localhost:{port}"
service = os.getenv("MCP_AUTH_SERVICE") or ("erp" if port == 8001 else "crm" if port == 8002 else "unknown")
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),
DbTokenVerifier(service=service),
ApiTokenVerifier(api_url=api_url, api_key=api_key),
)