diff --git a/docker-compose.yml b/docker-compose.yml index b08a995..c9478b7 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -75,4 +75,9 @@ services: - IAM_BASE_URL=${IAM_BASE_URL:-https://iam.digiwincloud.com.cn} - IAM_APP_TOKEN=${IAM_APP_TOKEN:-eyJ0eXAiOiJKV1QiLCJhbGciOiJIUzI1NiJ9.eyJpZCI6ImRhdGEtYnVzaW5lc3MtZGVtbyIsInNpZCI6MH0.Spo64LstbWxjYNefVFAbEbgfjzZoQGNcqKSGuYUOCRk} - IAM_CACHE_TTL=${IAM_CACHE_TTL:-30} + # Redis 缓存配置 + - REDIS_HOST=${REDIS_HOST:-127.0.0.1} + - REDIS_PORT=${REDIS_PORT:-6379} + - REDIS_DB=${REDIS_DB:-0} + - REDIS_PASSWORD=${REDIS_PASSWORD:-digiwin} restart: unless-stopped diff --git a/mcp-auth/backend/app/core/config.py b/mcp-auth/backend/app/core/config.py index aa0867f..322479d 100644 --- a/mcp-auth/backend/app/core/config.py +++ b/mcp-auth/backend/app/core/config.py @@ -23,5 +23,11 @@ class Settings: # 前端静态文件目录(Docker 构建后注入) STATIC_DIR: str = os.getenv("STATIC_DIR", "../frontend/dist") + # Redis 配置 + REDIS_HOST: str = os.getenv("REDIS_HOST", "127.0.0.1") + REDIS_PORT: int = int(os.getenv("REDIS_PORT", "6379")) + REDIS_DB: int = int(os.getenv("REDIS_DB", "0")) + REDIS_PASSWORD: str = os.getenv("REDIS_PASSWORD", "digiwin") + settings = Settings() diff --git a/mcp-auth/backend/app/core/redis.py b/mcp-auth/backend/app/core/redis.py new file mode 100644 index 0000000..62c2220 --- /dev/null +++ b/mcp-auth/backend/app/core/redis.py @@ -0,0 +1,64 @@ +"""Redis 连接池 + 缓存读写封装 + +用于 verify-token 的服务/Token 缓存,减少数据库访问。 +缓存 Key 约定: + mcp:svc:{api_key_hash} → 服务信息(TTL 120s) + mcp:tok:{token_hash} → Token 信息(TTL 30s) + mcp:tok:miss:{token_hash} → 无效标记,防穿透(TTL 30s) +""" + +import json + +import redis.asyncio as aioredis + +from .config import settings + +_pool: aioredis.Redis | None = None + + +async def get_redis() -> aioredis.Redis: + """获取 Redis 连接(进程级单例)。""" + global _pool + if _pool is None: + _pool = aioredis.Redis( + host=settings.REDIS_HOST, + port=settings.REDIS_PORT, + db=settings.REDIS_DB, + password=settings.REDIS_PASSWORD, + decode_responses=True, + ) + return _pool + + +async def close_redis() -> None: + """关闭 Redis 连接(进程退出时调用)。""" + global _pool + if _pool is not None: + await _pool.aclose() + _pool = None + + +async def cache_get(key: str) -> dict | None: + """读取 JSON 缓存,返回 dict 或 None。""" + r = await get_redis() + raw = await r.get(key) + if raw is None: + return None + try: + return json.loads(raw) + except Exception: + return None + + +async def cache_set(key: str, value: dict, ttl: int) -> None: + """写入 JSON 缓存,带 TTL(秒)。""" + r = await get_redis() + await r.setex(key, ttl, json.dumps(value)) + + +async def cache_delete(*keys: str) -> None: + """删除缓存 key。""" + if not keys: + return + r = await get_redis() + await r.delete(*keys) diff --git a/mcp-auth/backend/app/main.py b/mcp-auth/backend/app/main.py index df25b48..f534c16 100644 --- a/mcp-auth/backend/app/main.py +++ b/mcp-auth/backend/app/main.py @@ -15,6 +15,7 @@ from fastapi.staticfiles import StaticFiles from app.core.config import settings from app.core.db import close_pool, get_pool from app.core.iam import close_iam_client +from app.core.redis import close_redis from app.routers import auth, services, stats, tokens, verify @@ -24,6 +25,7 @@ async def lifespan(app: FastAPI): yield await close_pool() await close_iam_client() + await close_redis() app = FastAPI( diff --git a/mcp-auth/backend/app/routers/services.py b/mcp-auth/backend/app/routers/services.py index 24a9be4..3ef4d70 100644 --- a/mcp-auth/backend/app/routers/services.py +++ b/mcp-auth/backend/app/routers/services.py @@ -8,6 +8,7 @@ from pydantic import BaseModel from ..core.db import get_pool from ..core.deps import current_admin +from ..core.redis import cache_delete router = APIRouter(prefix="/api/services", tags=["services"]) @@ -85,7 +86,7 @@ async def revoke_service( ): pool = await get_pool() row = await pool.fetchrow( - "SELECT service_id, status FROM mcp_service WHERE service_id = $1", service_id + "SELECT service_id, status, api_key_hash FROM mcp_service WHERE service_id = $1", service_id ) if row is None: raise HTTPException(404, "服务不存在") @@ -96,6 +97,7 @@ async def revoke_service( "UPDATE mcp_service SET status = 'revoked', revoked_at = NOW() WHERE service_id = $1", service_id, ) + await cache_delete(f"mcp:svc:{row['api_key_hash']}") return {"success": True, "service_id": service_id, "status": "revoked"} @@ -107,7 +109,7 @@ async def enable_service( """启用服务:将已吊销的服务恢复为 active。""" pool = await get_pool() row = await pool.fetchrow( - "SELECT service_id, status FROM mcp_service WHERE service_id = $1", service_id + "SELECT service_id, status, api_key_hash FROM mcp_service WHERE service_id = $1", service_id ) if row is None: raise HTTPException(404, "服务不存在") @@ -118,6 +120,7 @@ async def enable_service( "UPDATE mcp_service SET status = 'active', revoked_at = NULL WHERE service_id = $1", service_id, ) + await cache_delete(f"mcp:svc:{row['api_key_hash']}") return {"success": True, "service_id": service_id, "status": "active"} @@ -128,7 +131,7 @@ async def delete_service( ): pool = await get_pool() row = await pool.fetchrow( - "SELECT service_id, status FROM mcp_service WHERE service_id = $1", service_id + "SELECT service_id, status, api_key_hash FROM mcp_service WHERE service_id = $1", service_id ) if row is None: raise HTTPException(404, "服务不存在") @@ -136,4 +139,5 @@ async def delete_service( raise HTTPException(400, "仅允许删除已吊销的服务,请先调用吊销端点") await pool.execute("DELETE FROM mcp_service WHERE service_id = $1", service_id) + await cache_delete(f"mcp:svc:{row['api_key_hash']}") return {"success": True, "service_id": service_id} diff --git a/mcp-auth/backend/app/routers/tokens.py b/mcp-auth/backend/app/routers/tokens.py index 51b860b..09c094b 100644 --- a/mcp-auth/backend/app/routers/tokens.py +++ b/mcp-auth/backend/app/routers/tokens.py @@ -10,6 +10,7 @@ from pydantic import BaseModel from ..core.db import get_pool from ..core.deps import current_admin +from ..core.redis import cache_delete router = APIRouter(prefix="/api/tokens", tags=["tokens"]) @@ -163,7 +164,7 @@ async def revoke_token( """吊销 token:软删除,status 改为 revoked。吊销后才可删除。""" pool = await get_pool() row = await pool.fetchrow( - "SELECT token_id, status FROM mcp_token WHERE token_id = $1", token_id + "SELECT token_id, status, token_hash FROM mcp_token WHERE token_id = $1", token_id ) if row is None: raise HTTPException(404, "token 不存在") @@ -178,6 +179,7 @@ async def revoke_token( "INSERT INTO mcp_token_log (token_id, event, detail) VALUES ($1, 'revoked', $2)", token_id, json.dumps({"reason": req.reason, "by": admin.get("username")}), ) + await cache_delete(f"mcp:tok:{row['token_hash']}", f"mcp:tok:miss:{row['token_hash']}") return {"success": True, "token_id": token_id, "status": "revoked"} @@ -189,7 +191,7 @@ async def enable_token( """启用 token:将已吊销的 token 恢复为 active。""" pool = await get_pool() row = await pool.fetchrow( - "SELECT token_id, status FROM mcp_token WHERE token_id = $1", token_id + "SELECT token_id, status, token_hash FROM mcp_token WHERE token_id = $1", token_id ) if row is None: raise HTTPException(404, "token 不存在") @@ -204,6 +206,7 @@ async def enable_token( "INSERT INTO mcp_token_log (token_id, event, detail) VALUES ($1, 'enabled', $2)", token_id, json.dumps({"by": admin.get("username")}), ) + await cache_delete(f"mcp:tok:{row['token_hash']}", f"mcp:tok:miss:{row['token_hash']}") return {"success": True, "token_id": token_id, "status": "active"} @@ -215,7 +218,7 @@ async def delete_token( """删除 token:物理删除,仅允许删除已吊销的 token。""" pool = await get_pool() row = await pool.fetchrow( - "SELECT token_id, status FROM mcp_token WHERE token_id = $1", token_id + "SELECT token_id, status, token_hash FROM mcp_token WHERE token_id = $1", token_id ) if row is None: raise HTTPException(404, "token 不存在") @@ -224,6 +227,7 @@ async def delete_token( await pool.execute("DELETE FROM mcp_token_log WHERE token_id = $1", token_id) await pool.execute("DELETE FROM mcp_token WHERE token_id = $1", token_id) + await cache_delete(f"mcp:tok:{row['token_hash']}", f"mcp:tok:miss:{row['token_hash']}") return {"success": True, "token_id": token_id} diff --git a/mcp-auth/backend/app/routers/verify.py b/mcp-auth/backend/app/routers/verify.py index f1f3569..ecfbbec 100644 --- a/mcp-auth/backend/app/routers/verify.py +++ b/mcp-auth/backend/app/routers/verify.py @@ -2,6 +2,11 @@ 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 @@ -12,9 +17,14 @@ 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 @@ -29,11 +39,24 @@ class VerifyResp(BaseModel): async def _resolve_service(x_api_key: str | None = Header(None, alias="X-API-Key")) -> str: """校验 per-service API Key,返回 service_name;不匹配则 401。 - 用 sha256 比对 mcp_service 表,避免明文存储。同时更新 last_used_at。 + 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", @@ -42,7 +65,14 @@ async def _resolve_service(x_api_key: str | None = Header(None, alias="X-API-Key if row is None or row["status"] != "active": raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid or revoked api key") - # 异步更新 last_used_at,不阻塞响应 + # 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"] @@ -53,8 +83,39 @@ async def verify_token(req: VerifyReq, service: str = Depends(_resolve_service)) 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 != "both" and 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 @@ -62,10 +123,21 @@ async def verify_token(req: VerifyReq, service: str = Depends(_resolve_service)) 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) @@ -74,7 +146,7 @@ async def verify_token(req: VerifyReq, service: str = Depends(_resolve_service)) if row["expires_at"] is not None and row["expires_at"].timestamp() < datetime.now(timezone.utc).timestamp(): return VerifyResp(valid=False) - # 服务范围校验:both 放行所有;否则要求精确匹配(由 API Key 识别的 service) + # 服务范围校验:both 放行所有;否则要求精确匹配 scope = row["service_scope"] if scope != "both" and scope != service: return VerifyResp(valid=False) diff --git a/mcp-auth/backend/requirements.txt b/mcp-auth/backend/requirements.txt index 2c62db1..a97bc06 100644 --- a/mcp-auth/backend/requirements.txt +++ b/mcp-auth/backend/requirements.txt @@ -5,3 +5,4 @@ bcrypt>=4.2.0 pyjwt>=2.9.0 pydantic>=2.9.0 httpx>=0.27.0 +redis>=5.0.0