This commit is contained in:
2026-08-31 18:01:34 +08:00
parent 3ea59551db
commit bee439cf34
46 changed files with 1884 additions and 1898 deletions
+25 -20
View File
@@ -28,6 +28,7 @@ def create_app(
service_name: str,
version: str = "4.0.0",
mount_html: bool = False,
serve_frontend_static: bool = False,
startup_hooks: Optional[List[Callable[[], Awaitable[None]]]] = None,
register_routers: Optional[Callable[[FastAPI], None]] = None,
) -> FastAPI:
@@ -38,6 +39,7 @@ def create_app(
service_name: 服务名(用于 /health 响应)
version: 版本号
mount_html: 是否挂载 /html 静态目录(moldinsight 需要)
serve_frontend_static: 是否由后端托管 /static 与 SPA fallback(默认关闭,前端独立部署)
startup_hooks: 额外的 startup 钩子列表(在数据库/RustFS/Redis 初始化后执行)
register_routers: 回调函数,用于注册业务路由
"""
@@ -70,7 +72,7 @@ def create_app(
# 跳过静态资源和健康检查的详细日志
path = request.url.path
is_static = path.startswith("/static") or path == "/health"
is_static = (serve_frontend_static and path.startswith("/static")) or path == "/health"
if not is_static:
log_level = "warning" if response.status_code >= 400 else "info"
@@ -92,16 +94,18 @@ def create_app(
# ── 目录准备 ─────────────────────────────────────────────────
Path("uploads").mkdir(exist_ok=True)
Path("static").mkdir(exist_ok=True)
if serve_frontend_static:
Path("static").mkdir(exist_ok=True)
if mount_html:
Path("html_output").mkdir(exist_ok=True)
# ── 静态文件挂载 ─────────────────────────────────────────────
app.mount(
"/static",
StaticFiles(directory=os.path.join(os.getcwd(), "static")),
name="static",
)
if serve_frontend_static:
app.mount(
"/static",
StaticFiles(directory=os.path.join(os.getcwd(), "static")),
name="static",
)
if mount_html:
app.mount(
"/html",
@@ -189,19 +193,20 @@ def create_app(
"database_error": db_error,
}
# ── SPA fallback(排除 /api 前缀,避免吞掉 API 404)────────
@app.get("/{full_path:path}")
async def spa_fallback(full_path: str):
# API 路径不走 SPA fallback,让 FastAPI 正常返回 404 JSON
if full_path.startswith("api/") or full_path.startswith("api"):
raise _api_not_found(full_path)
# 健康检查 / 文档路径也排除
if full_path.startswith("docs") or full_path.startswith("openapi"):
raise _api_not_found(full_path)
static_index = os.path.join(os.getcwd(), "static", "index.html")
if os.path.exists(static_index):
return FileResponse(static_index)
return JSONResponse({"detail": "SPA index not found"}, status_code=404)
# ── SPA fallback(独立前端部署时默认关闭)────────────────────
if serve_frontend_static:
@app.get("/{full_path:path}")
async def spa_fallback(full_path: str):
# API 路径不走 SPA fallback,让 FastAPI 正常返回 404 JSON
if full_path.startswith("api/") or full_path.startswith("api"):
raise _api_not_found(full_path)
# 健康检查 / 文档路径也排除
if full_path.startswith("docs") or full_path.startswith("openapi"):
raise _api_not_found(full_path)
static_index = os.path.join(os.getcwd(), "static", "index.html")
if os.path.exists(static_index):
return FileResponse(static_index)
return JSONResponse({"detail": "SPA index not found"}, status_code=404)
return app
+25 -30
View File
@@ -13,6 +13,7 @@ class Settings:
self.HOST = os.getenv("HOST", "0.0.0.0")
self.PORT = int(os.getenv("PORT", "8000"))
self.DEBUG = os.getenv("DEBUG", "false").lower() == "true"
self.SERVE_FRONTEND_STATIC = os.getenv("SERVE_FRONTEND_STATIC", "false").lower() == "true"
self.UPLOAD_DIR = os.getenv("UPLOAD_DIR", "./uploads")
self.MAX_FILE_SIZE = int(os.getenv("MAX_FILE_SIZE", "104857600"))
@@ -28,32 +29,12 @@ class Settings:
self.RUSTFS_TIMEOUT = int(os.getenv("RUSTFS_TIMEOUT", "30"))
self.RUSTFS_PRESIGNED_URL_EXPIRES = int(os.getenv("RUSTFS_PRESIGNED_URL_EXPIRES", "3600"))
db_host = os.getenv("DB_HOST")
db_port_str = os.getenv("DB_PORT")
db_name = os.getenv("DB_NAME")
db_user = os.getenv("DB_USER")
db_password = os.getenv("DB_PASSWORD")
missing_configs = []
if not db_host:
missing_configs.append("DB_HOST")
if not db_port_str:
missing_configs.append("DB_PORT")
if not db_name:
missing_configs.append("DB_NAME")
if not db_user:
missing_configs.append("DB_USER")
if not db_password:
missing_configs.append("DB_PASSWORD")
if missing_configs:
raise ValueError(f"数据库配置缺失,请在.env文件中设置: {', '.join(missing_configs)}")
self.DB_HOST = db_host
self.DB_PORT = int(db_port_str)
self.DB_NAME = db_name
self.DB_USER = db_user
self.DB_PASSWORD = db_password
# 数据库配置改为惰性校验:允许在无 DB 环境下 import 项目模块(测试/静态分析)
self.DB_HOST = os.getenv("DB_HOST")
self.DB_PORT = int(os.getenv("DB_PORT")) if os.getenv("DB_PORT") else None
self.DB_NAME = os.getenv("DB_NAME")
self.DB_USER = os.getenv("DB_USER")
self.DB_PASSWORD = os.getenv("DB_PASSWORD")
self.SECRET_KEY = os.getenv("SECRET_KEY")
self.ALGORITHM = os.getenv("ALGORITHM", "HS256")
@@ -90,10 +71,24 @@ class Settings:
@property
def DATABASE_URL(self) -> str:
if self.DB_PASSWORD:
safe_password = urllib.parse.quote(self.DB_PASSWORD.encode("utf-8"), safe="")
else:
safe_password = ""
missing_configs = []
if not self.DB_HOST:
missing_configs.append("DB_HOST")
if not self.DB_PORT:
missing_configs.append("DB_PORT")
if not self.DB_NAME:
missing_configs.append("DB_NAME")
if not self.DB_USER:
missing_configs.append("DB_USER")
if self.DB_PASSWORD is None:
missing_configs.append("DB_PASSWORD")
if missing_configs:
raise ValueError(
f"数据库配置缺失,请在.env文件中设置: {', '.join(missing_configs)}"
)
safe_password = urllib.parse.quote((self.DB_PASSWORD or "").encode("utf-8"), safe="")
return f"postgresql+asyncpg://{self.DB_USER}:{safe_password}@{self.DB_HOST}:{self.DB_PORT}/{self.DB_NAME}"
@property
+6 -4
View File
@@ -45,16 +45,18 @@ class DatabaseManager:
Args:
role: 连接角色,"web" 或 "celery",决定连接池大小
"""
if not settings.DATABASE_URL:
logger.warning("未配置数据库连接,跳过数据库初始化")
try:
database_url = settings.DATABASE_URL
except ValueError as e:
logger.warning(f"未配置数据库连接,跳过数据库初始化: {e}")
self.is_connected = False
return
try:
pool_cfg = _get_pool_config(role)
# 创建异步引擎
self.engine = create_async_engine(
settings.DATABASE_URL,
database_url,
echo=settings.DEBUG,
pool_size=pool_cfg["pool_size"],
max_overflow=pool_cfg["max_overflow"],
+121 -35
View File
@@ -1,20 +1,27 @@
# services/redis_task_manager.py
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理"""
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理。
存储格式:Redis Hash(field -> JSON 字符串)。
- update_task 走 HSET 字段级原子更新,消除旧 get->merge->set 三步竞态
(后台处理流程与导出端点并发写同一任务时丢更新);
- 进度 tick 只重写变化字段,不再全量重写整个任务 blob;
- 兼容读旧 string 格式(升级前写入的在途任务),新写入一律 Hash。
"""
import json
import os
from typing import Dict, Any, Optional
from datetime import datetime
import redis.asyncio as aioredis
from shared.config.settings import settings
from shared.utils.logger import get_logger
logger = get_logger(__name__)
class RedisTaskManager:
"""基于 Redis 的任务状态管理"""
"""基于 Redis Hash 的任务状态管理"""
_instance: Optional["RedisTaskManager"] = None
@@ -31,14 +38,14 @@ class RedisTaskManager:
return cls._instance
async def connect(self):
"""连接 Redis"""
"""连接 Redis(配置统一来自 shared.config.settings,不再硬编码主机名)"""
if self._connected and self._redis:
return
host = os.getenv("REDIS_HOST", "szcjw")
port = int(os.getenv("REDIS_PORT", "6379"))
password = os.getenv("REDIS_PASSWORD", "")
db = int(os.getenv("REDIS_DB", "0"))
host = settings.REDIS_HOST
port = settings.REDIS_PORT
password = settings.REDIS_PASSWORD
db = settings.REDIS_DB
try:
self._redis = aioredis.Redis(
@@ -84,6 +91,16 @@ class RedisTaskManager:
def is_connected(self) -> bool:
return self._connected and self._redis is not None
@property
def redis_client(self) -> aioredis.Redis:
"""暴露底层客户端(batch 元数据等非任务结构数据使用)。
未连接时抛出明确错误,而不是让调用方踩 AttributeError。
"""
if not self.is_connected or self._redis is None:
raise RuntimeError("Redis 未连接,无法直接访问 redis_client")
return self._redis
# ---- 内存回退 ----
_fallback_tasks: Dict[str, Dict[str, Any]] = {}
@@ -102,55 +119,128 @@ class RedisTaskManager:
def _fallback_count(self) -> int:
return len(self._fallback_tasks)
# ---- 内部工具 ----
def _key(self, task_id: str) -> str:
return f"{self._prefix}{task_id}"
@staticmethod
def _dump_mapping(data: Dict[str, Any]) -> Dict[str, str]:
"""把任务 dict 序列化为 Hash mapping(field -> JSON 字符串)"""
serializable = RedisTaskManager._make_serializable(data)
return {k: json.dumps(v, ensure_ascii=False) for k, v in serializable.items()}
async def _load_hash(self, key: str) -> Optional[Dict[str, Any]]:
raw = await self._redis.hgetall(key)
if not raw:
return None
result = {}
for field, value in raw.items():
try:
result[field] = json.loads(value)
except (json.JSONDecodeError, TypeError):
result[field] = value
return result
async def _load_any(self, key: str) -> Optional[Dict[str, Any]]:
"""读取任务数据,自动识别 Hash(新)与 string(旧)格式。"""
key_type = await self._redis.type(key)
if key_type == "hash":
return await self._load_hash(key)
if key_type == "string":
legacy = await self._redis.get(key)
if not legacy:
return None
try:
return json.loads(legacy)
except json.JSONDecodeError:
logger.warning(f"任务数据解析失败(旧 string 格式): {key}")
return None
return None
# ---- 公共接口 ----
async def set_task(self, task_id: str, data: Dict[str, Any], ttl: Optional[int] = None):
"""设置任务数据"""
"""整包写入任务数据(Hash,覆盖旧值,含旧 string 格式清理)"""
effective_ttl = ttl or self._ttl
# 确保数据可序列化
serializable = self._make_serializable(data)
mapping = self._dump_mapping(data)
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.setex(key, effective_ttl, json.dumps(serializable, ensure_ascii=False))
key = self._key(task_id)
# DEL 先清掉可能存在的旧 string/Hash,保证覆盖语义
pipe = self._redis.pipeline()
pipe.delete(key)
pipe.hset(key, mapping=mapping)
pipe.expire(key, effective_ttl)
await pipe.execute()
return
except Exception as e:
logger.warning(f"Redis 写入失败,回退到内存: {e}")
self._fallback_set(task_id, serializable)
self._fallback_set(task_id, self._make_serializable(data))
async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
"""获取任务数据"""
"""获取任务数据(Hash / 旧 string 兼容)"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
raw = await self._redis.get(key)
if raw:
return json.loads(raw)
return None
return await self._load_any(self._key(task_id))
except Exception as e:
logger.warning(f"Redis 读取失败,回退到内存: {e}")
return self._fallback_get(task_id)
async def update_task(self, task_id: str, updates: Dict[str, Any]):
"""更新任务的部分字段"""
current = await self.get_task(task_id)
"""字段级原子更新(HSET),无读改写竞态。
兼容旧 string 格式:先迁移为 Hash 再更新。
"""
mapping = self._dump_mapping(updates)
if self.is_connected:
try:
key = self._key(task_id)
key_type = await self._redis.type(key)
if key_type == "none":
logger.warning(f"任务 {task_id} 不存在,无法更新")
return
if key_type == "string":
# 旧格式迁移:string -> Hash
legacy = await self._redis.get(key)
try:
base = json.loads(legacy) if legacy else {}
except json.JSONDecodeError:
base = {}
base.update(mapping)
pipe = self._redis.pipeline()
pipe.delete(key)
pipe.hset(key, mapping=self._dump_mapping(base))
pipe.expire(key, self._ttl)
await pipe.execute()
return
await self._redis.hset(key, mapping=mapping)
await self._redis.expire(key, self._ttl)
return
except Exception as e:
logger.warning(f"Redis 更新失败,回退到内存: {e}")
# 内存回退保持读改写语义(单进程内存无并发竞态)
current = self._fallback_get(task_id)
if current is None:
logger.warning(f"任务 {task_id} 不存在,无法更新")
return
current.update(self._make_serializable(updates))
await self.set_task(task_id, current)
self._fallback_set(task_id, current)
async def delete_task(self, task_id: str):
"""删除任务"""
"""删除任务(DEL 对 Hash/string 均有效)"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.delete(key)
await self._redis.delete(self._key(task_id))
return
except Exception as e:
logger.warning(f"Redis 删除失败,回退到内存: {e}")
@@ -162,16 +252,12 @@ class RedisTaskManager:
if self.is_connected:
try:
pattern = f"{self._prefix}*"
keys = []
async for key in self._redis.scan_iter(match=pattern):
keys.append(key)
result = {}
for key in keys:
async for key in self._redis.scan_iter(match=pattern):
task_id = key.replace(self._prefix, "")
raw = await self._redis.get(key)
if raw:
result[task_id] = json.loads(raw)
task = await self._load_any(key)
if task:
result[task_id] = task
return result
except Exception as e:
logger.warning(f"Redis 扫描失败,回退到内存: {e}")