import logging import time from fastapi import Request, Response from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint from starlette.responses import JSONResponse logger = logging.getLogger("videogen") # 网站关闭时仍然放行的路径前缀 _WHITELIST_PREFIXES = ( "/api/auth/", # 登录/注册/短信/验证码 "/api/admin/", # 后台管理(确保管理员能登录并重新开启网站) "/api/captcha/", # 图形验证码 "/internal/", # 健康检查/文档 "/uploads/", # 静态文件 "/api/generation-records/callback", # 火山引擎回调 "/api/payments/wechat/callback", # 微信回调 "/api/payments/alipay/callback", # 支付宝回调 ) # 完全匹配的白名单路径 _WHITELIST_EXACT = ( "/api/auth/site-info", "/internal/health", ) # site_enabled 缓存有效期(秒) _CACHE_TTL = 5.0 class SiteStatusMiddleware(BaseHTTPMiddleware): """网站访问开关中间件。 检查 SystemConfig 中 site_enabled 的值: - 配置不存在或值不为 "false" → 放行 - 值为 "false" 且不在白名单 → 检查开发者令牌,不匹配则返回 503 { code: "SITE_CLOSED" } """ def __init__(self, app): super().__init__(app) self._cache_enabled = True # 缓存的站点状态(True=开启) self._cache_at: float = 0.0 # 缓存时间戳 async def _is_site_enabled(self) -> bool: """查询 site_enabled 配置,带短缓存避免每个请求查 DB。""" now = time.time() if now - self._cache_at < _CACHE_TTL: return self._cache_enabled try: from app.models.base import async_session from app.models.system_config import SystemConfig from sqlalchemy import select async with async_session() as db: result = await db.execute( select(SystemConfig.value).where(SystemConfig.key == "site_enabled").limit(1) ) value = result.scalar_one_or_none() self._cache_enabled = value != "false" self._cache_at = now except Exception: # 查询异常时默认放行,避免数据库故障导致全站不可用 logger.exception("Failed to read site_enabled config, defaulting to enabled") self._cache_enabled = True self._cache_at = now return self._cache_enabled @staticmethod def _is_whitelisted(path: str) -> bool: """检查路径是否在白名单中。""" if path in _WHITELIST_EXACT: return True return any(path.startswith(prefix) for prefix in _WHITELIST_PREFIXES) @staticmethod async def _check_dev_token(token: str) -> bool: """校验开发者访问令牌是否匹配配置。""" try: from app.models.base import async_session from app.models.system_config import SystemConfig from sqlalchemy import select async with async_session() as db: result = await db.execute( select(SystemConfig.value) .where(SystemConfig.key == "site_dev_access_token") .limit(1) ) expected = result.scalar_one_or_none() return expected is not None and token == expected except Exception: logger.exception("Failed to verify dev access token") return False async def dispatch( self, request: Request, call_next: RequestResponseEndpoint ) -> Response: path = request.url.path # 白名单直接放行 if self._is_whitelisted(path): return await call_next(request) # 网站开启时放行 if await self._is_site_enabled(): return await call_next(request) # 网站已关闭 — 校验开发者令牌 token = request.query_params.get("dev_access") or request.headers.get("x-dev-access") if token and await self._check_dev_token(token): return await call_next(request) # 拦截请求,返回维护状态 return JSONResponse( status_code=503, content={ "detail": { "code": "SITE_CLOSED", "message": "系统正在升级维护,请稍后再试", } }, )