新增mcp_auth项目

This commit is contained in:
2026-08-31 11:02:07 +08:00
parent efe6fbe770
commit cbff2283e2
33 changed files with 4723 additions and 73 deletions
+116 -33
View File
@@ -1,59 +1,142 @@
"""静态 Bearer Token 鉴权(Authorization)。
"""动态 Bearer Token 鉴权(Authorization)。
token → 客户端 映射来自环境变量 MCP_AUTH_TOKENS,格式(逗号分隔,每对 token:client_id):
MCP_AUTH_TOKENS=token-trae:trae,token-partner:partner-a
token 校验查 mcp_auth.mcp_token 表(sha256 hash 比对),进程内 LRU 缓存 30s。
token 表由独立管理后台(mcp-auth-admin)维护,MCP 服务只读。
每个客户端配置自己的 token 调用,服务端校验失败返回 401;
工具内可通过 get_caller() 获取当前调用方标识。
环境变量:
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_CACHE_TTL 缓存秒数,默认 30
MCP_AUTH_SERVICE 当前服务标识(erp/crm),用于 service_scope 校验和 last_used_svc
"""
import asyncio
import hashlib
import os
import time
import asyncpg
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
class StaticTokenVerifier:
"""静态 token 查表校验:命中返回 AccessToken(client_id 即客户端标识),未命中返回 None(401)。"""
# 鉴权库连接池(独立于业务库,进程级单例)
_auth_pool: asyncpg.Pool | None = None
def __init__(self, tokens: dict[str, str]):
self._tokens = tokens
_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
async def close_auth_pool() -> None:
"""关闭鉴权库连接池(进程退出时调用)。"""
global _auth_pool
if _auth_pool is not None:
await _auth_pool.close()
_auth_pool = None
class DbTokenVerifier:
"""查库校验 Bearer Token,命中返回 AccessToken,未命中/已吊销/已过期/服务范围不匹配返回 None(401)。
进程内 LRU 缓存(token_hash -> (row|None, fetched_at)),TTL 由 MCP_AUTH_CACHE_TTL 控制。
吊销后最长 TTL 秒内仍可能命中旧状态;需要即时生效可调用 invalidate(token)。
"""
def __init__(self, service: str):
self._service = service # 'erp' / 'crm'
self._cache: dict[str, tuple[asyncpg.Record | 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:
client_id = self._tokens.get(token)
if client_id is None:
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)
# 3. 校验
if row is None:
return None
return AccessToken(token=token, client_id=client_id, scopes=[], expires_at=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。"""
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,
)
except Exception:
pass # 审计字段更新失败不影响鉴权
def _parse_tokens(raw: str) -> dict[str, str]:
"""解析 'token1:client1,token2:client2' → {token: client_id}"""
tokens: dict[str, str] = {}
for item in raw.split(","):
item = item.strip()
if not item:
continue
token, _, client_id = item.partition(":")
token, client_id = token.strip(), client_id.strip()
if token and client_id:
tokens[token] = client_id
return tokens
def get_auth(port: int) -> tuple[AuthSettings, StaticTokenVerifier]:
def get_auth(port: int) -> tuple[AuthSettings, DbTokenVerifier]:
"""构建 MCPServer 的 (auth, token_verifier) 参数。
服务对外地址默认 http://localhost:{port},部署时用 MCP_PUBLIC_URL 覆盖
(如 http://192.168.1.119:8002),用于 OAuth 资源元数据发现。
连接池在首次 verify_token 时惰性创建,此处不连库。
service 标识取自 MCP_AUTH_SERVICE 环境变量,默认按 port 推断(8001→erp / 8002→crm)。
"""
url = os.getenv("MCP_PUBLIC_URL") or f"http://localhost:{port}"
tokens = _parse_tokens(os.getenv("MCP_AUTH_TOKENS", ""))
if not tokens:
raise RuntimeError("环境变量 MCP_AUTH_TOKENS 未配置,格式:token1:client1,token2:client2")
service = os.getenv("MCP_AUTH_SERVICE") or ("erp" if port == 8001 else "crm" if port == 8002 else "unknown")
return (
AuthSettings(issuer_url=url, resource_server_url=url),
StaticTokenVerifier(tokens),
DbTokenVerifier(service=service),
)