Files
2026-09-02 17:31:52 +08:00

177 lines
6.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""内部 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