"""内部 token 校验 API — 供 MCP 服务调用(不查库直连,走 HTTP)。 MCP 服务(ERP/CRM)通过此端点校验 Bearer Token,后端统一查 mcp_auth.mcp_token 表。 用 per-service API Key(X-API-Key)鉴权,从 key 识别调用方服务,无需 MCP 服务自报 service。 Redis 缓存层: mcp:svc:{api_key_hash} → 服务信息(TTL 120s) mcp:tok:{token_hash} → Token 信息(TTL 30s) mcp:tok:miss:{token_hash} → 无效标记,防穿透(TTL 30s) """ import asyncio import hashlib from datetime import datetime, timezone from fastapi import APIRouter, Depends, Header, HTTPException, status from pydantic import BaseModel from ..core.db import get_pool from ..core.redis import cache_get, cache_set, cache_delete router = APIRouter(prefix="/api/auth", tags=["verify"]) # 缓存 TTL _SVC_TTL = 120 # 服务缓存 120s(服务极少变化) _TOK_TTL = 30 # Token 缓存 30s(与 MCP 服务侧对齐) class VerifyReq(BaseModel): token: str # 明文 Bearer Token class VerifyResp(BaseModel): valid: bool client_id: str | None = None service_scope: str | None = None async def _resolve_service(x_api_key: str | None = Header(None, alias="X-API-Key")) -> str: """校验 per-service API Key,返回 service_name;不匹配则 401。 Redis 缓存 mcp:svc:{api_key_hash},命中则跳过 DB。 """ if not x_api_key: raise HTTPException(status.HTTP_401_UNAUTHORIZED, "missing X-API-Key") api_key_hash = hashlib.sha256(x_api_key.encode()).hexdigest() cache_key = f"mcp:svc:{api_key_hash}" # 1. 查缓存 cached = await cache_get(cache_key) if cached is not None: if cached.get("status") != "active": raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid or revoked api key") # 异步更新 last_used_at pool = await get_pool() asyncio.create_task(_touch_service(pool, cached["service_id"])) return cached["service_name"] # 2. 查 DB pool = await get_pool() row = await pool.fetchrow( "SELECT service_id, service_name, status FROM mcp_service WHERE api_key_hash = $1", api_key_hash, ) if row is None or row["status"] != "active": raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid or revoked api key") # 3. 写缓存 await cache_set(cache_key, { "service_id": row["service_id"], "service_name": row["service_name"], "status": row["status"], }, _SVC_TTL) # 异步更新 last_used_at asyncio.create_task(_touch_service(pool, row["service_id"])) return row["service_name"] @router.post("/verify-token", response_model=VerifyResp) async def verify_token(req: VerifyReq, service: str = Depends(_resolve_service)): """校验 Bearer Token:sha256 比对 + 状态/过期/服务范围检查 + 更新 last_used。 service 由 X-API-Key 自动识别,MCP 服务无需在 body 中传 service。 返回 valid=true 时附带 client_id 和 service_scope;校验失败返回 valid=false(非 401)。 Redis 缓存 mcp:tok:{token_hash},命中则跳过 DB。 无效结果也缓存(mcp:tok:miss:),防穿透。 """ token_hash = hashlib.sha256(req.token.encode()).hexdigest() cache_key = f"mcp:tok:{token_hash}" miss_key = f"mcp:tok:miss:{token_hash}" # 1. 查无效标记(防穿透) miss_cached = await cache_get(miss_key) if miss_cached is not None: return VerifyResp(valid=False) # 2. 查 Token 缓存 cached = await cache_get(cache_key) if cached is not None: # 校验状态 if cached.get("status") != "active": return VerifyResp(valid=False) # 校验过期 if cached.get("expires_at") and cached["expires_at"] > 0: if cached["expires_at"] < datetime.now(timezone.utc).timestamp(): return VerifyResp(valid=False) # 校验服务范围 scope = cached.get("service_scope") if scope != service: return VerifyResp(valid=False) # 异步更新 last_used pool = await get_pool() asyncio.create_task(_touch_token(pool, cached["token_id"], service)) return VerifyResp(valid=True, client_id=cached.get("client_id"), service_scope=scope) # 3. 查 DB pool = await get_pool() row = await pool.fetchrow( """SELECT token_id, client_id, status, expires_at, service_scope FROM mcp_token WHERE token_hash = $1""", token_hash, ) # 不存在 → 缓存无效标记 if row is None: await cache_set(miss_key, {"valid": False}, _TOK_TTL) return VerifyResp(valid=False) # 写缓存(expires_at 转为时间戳,便于序列化) cache_val = { "token_id": row["token_id"], "client_id": row["client_id"], "status": row["status"], "expires_at": row["expires_at"].timestamp() if row["expires_at"] else None, "service_scope": row["service_scope"], } await cache_set(cache_key, cache_val, _TOK_TTL) # 已吊销 if row["status"] != "active": return VerifyResp(valid=False) # 已过期 if row["expires_at"] is not None and row["expires_at"].timestamp() < datetime.now(timezone.utc).timestamp(): return VerifyResp(valid=False) # 服务范围校验:要求精确匹配 scope = row["service_scope"] if scope != service: return VerifyResp(valid=False) # 异步更新 last_used_at / last_used_svc asyncio.create_task(_touch_token(pool, row["token_id"], service)) return VerifyResp(valid=True, client_id=row["client_id"], service_scope=scope) async def _touch_service(pool, service_id: int) -> None: try: await pool.execute( "UPDATE mcp_service SET last_used_at = now() WHERE service_id = $1", service_id ) except Exception: pass async def _touch_token(pool, token_id: int, service: str) -> None: try: await pool.execute( "UPDATE mcp_token SET last_used_at = now(), last_used_svc = $2 WHERE token_id = $1", token_id, service, ) except Exception: pass