x
This commit is contained in:
+25
-20
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user