新增mcp_auth项目
This commit is contained in:
+116
-33
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user