添加服务注册及api鉴权
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
mcp>=1.0.0
|
||||
pydantic>=2.0.0
|
||||
asyncpg>=0.30.0
|
||||
httpx>=0.27.0
|
||||
|
||||
+63
-90
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user