177 lines
6.1 KiB
Python
177 lines
6.1 KiB
Python
"""内部 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
|