取消
This commit is contained in:
@@ -1622,41 +1622,6 @@ async def update_system_config(
|
||||
return config
|
||||
|
||||
|
||||
@router.post("/system-configs/regenerate-dev-token")
|
||||
async def regenerate_dev_token(
|
||||
admin: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""重新生成开发者访问令牌。"""
|
||||
import secrets
|
||||
result = await db.execute(
|
||||
select(SystemConfig).where(SystemConfig.key == "site_dev_access_token").limit(1)
|
||||
)
|
||||
config = result.scalar_one_or_none()
|
||||
new_token = secrets.token_urlsafe(32)
|
||||
if config:
|
||||
config.value = new_token
|
||||
else:
|
||||
config = SystemConfig(
|
||||
id=generate_id(),
|
||||
key="site_dev_access_token",
|
||||
value=new_token,
|
||||
description="开发者访问令牌(网站关闭时用于绕过限制)",
|
||||
)
|
||||
db.add(config)
|
||||
await db.flush()
|
||||
await log_operation(
|
||||
db,
|
||||
admin.id,
|
||||
admin.username,
|
||||
"重新生成开发者访问令牌",
|
||||
"POST",
|
||||
"/admin/system-configs/regenerate-dev-token",
|
||||
)
|
||||
await db.commit()
|
||||
return {"token": new_token}
|
||||
|
||||
|
||||
# ── Operation Logs ──────────────────────────────────────
|
||||
|
||||
@router.get("/operation-logs")
|
||||
|
||||
@@ -345,19 +345,12 @@ async def get_site_info(db: AsyncSession = Depends(get_db)):
|
||||
base_url = settings.BASE_URL.rstrip("/")
|
||||
return f"{base_url}{path}"
|
||||
|
||||
# 读取网站开关状态(不存在时默认开启)
|
||||
site_enabled_result = await db.execute(
|
||||
select(SystemConfig.value).where(SystemConfig.key == "site_enabled").limit(1)
|
||||
)
|
||||
site_enabled_value = site_enabled_result.scalar_one_or_none()
|
||||
|
||||
return {
|
||||
"site_name": info.get("site_name", "VideoGen.AI"),
|
||||
"site_logo": to_full_url(info.get("site_logo")),
|
||||
"user_agreement_privacy_url": to_full_url(info.get("user_agreement_privacy_url")),
|
||||
"site_copyright": info.get("site_copyright", "© 2024 民众智创 版权所有"),
|
||||
"operation_manual": info.get("operation_manual", ""),
|
||||
"site_enabled": site_enabled_value != "false",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -16,7 +16,6 @@ from app.middleware.logging import RequestLoggingMiddleware
|
||||
from app.middleware.anti_crawler import AntiCrawlerMiddleware
|
||||
from app.middleware.rate_limit import RateLimitMiddleware
|
||||
from app.middleware.request_encrypt import RequestEncryptMiddleware
|
||||
from app.middleware.site_status import SiteStatusMiddleware
|
||||
from app.services.log_config import decrypt_data
|
||||
|
||||
logging.basicConfig(level=logging.INFO if settings.DEBUG else logging.WARNING)
|
||||
@@ -30,7 +29,6 @@ async def lifespan(app: FastAPI):
|
||||
os.makedirs(settings.UPLOAD_LOCAL_PATH, exist_ok=True)
|
||||
await init_database()
|
||||
await init_redis()
|
||||
await _ensure_site_configs()
|
||||
# await _seed_data()
|
||||
|
||||
# Start task queue (handles both video and image generation)
|
||||
@@ -515,40 +513,6 @@ async def _seed_data():
|
||||
await db.commit()
|
||||
|
||||
|
||||
async def _ensure_site_configs() -> None:
|
||||
"""确保网站开关相关配置项存在(不存在则自动创建)。"""
|
||||
import secrets
|
||||
from app.models.base import async_session
|
||||
from app.models.system_config import SystemConfig
|
||||
from app.utils.id_gen import generate_id
|
||||
from sqlalchemy import select
|
||||
|
||||
async with async_session() as db:
|
||||
# site_enabled
|
||||
existing = await db.execute(
|
||||
select(SystemConfig).where(SystemConfig.key == "site_enabled").limit(1)
|
||||
)
|
||||
if not existing.scalar_one_or_none():
|
||||
db.add(SystemConfig(
|
||||
id=generate_id(),
|
||||
key="site_enabled",
|
||||
value="true",
|
||||
description="网站访问开关(true=开启,false=关闭)",
|
||||
))
|
||||
# site_dev_access_token
|
||||
existing_token = await db.execute(
|
||||
select(SystemConfig).where(SystemConfig.key == "site_dev_access_token").limit(1)
|
||||
)
|
||||
if not existing_token.scalar_one_or_none():
|
||||
db.add(SystemConfig(
|
||||
id=generate_id(),
|
||||
key="site_dev_access_token",
|
||||
value=secrets.token_urlsafe(32),
|
||||
description="开发者访问令牌(网站关闭时用于绕过限制)",
|
||||
))
|
||||
await db.commit()
|
||||
|
||||
|
||||
def create_app() -> FastAPI:
|
||||
application = FastAPI(
|
||||
title=settings.APP_NAME,
|
||||
@@ -558,8 +522,7 @@ def create_app() -> FastAPI:
|
||||
redoc_url="/internal/api-redoc",
|
||||
)
|
||||
|
||||
# Middleware (outermost first) — SiteStatus 放最外层,早于其他中间件拦截
|
||||
application.add_middleware(SiteStatusMiddleware)
|
||||
# Middleware (outermost first)
|
||||
application.add_middleware(RequestLoggingMiddleware)
|
||||
application.add_middleware(AntiCrawlerMiddleware)
|
||||
application.add_middleware(RateLimitMiddleware)
|
||||
|
||||
@@ -1,125 +0,0 @@
|
||||
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": "系统正在升级维护,请稍后再试",
|
||||
}
|
||||
},
|
||||
)
|
||||
Reference in New Issue
Block a user