优化
This commit is contained in:
@@ -23,9 +23,11 @@ class Settings:
|
||||
self.MESH_QUALITY = os.getenv("MESH_QUALITY", "high")
|
||||
self.PARALLEL_PROCESSING = os.getenv("PARALLEL_PROCESSING", "true").lower() == "true"
|
||||
|
||||
self.RUSTFS_ENDPOINT = os.getenv("RUSTFS_ENDPOINT") or os.getenv("MINIO_ENDPOINT") or "http://localhost:8080"
|
||||
self.RUSTFS_ACCESS_KEY = os.getenv("RUSTFS_ACCESS_KEY") or os.getenv("MINIO_ACCESS_KEY") or "your-access-key"
|
||||
self.RUSTFS_SECRET_KEY = os.getenv("RUSTFS_SECRET_KEY") or os.getenv("MINIO_SECRET_KEY") or "your-secret-key"
|
||||
# RUSTFS_* 不给代码兜底默认值(含 MINIO_* 兼容别名):
|
||||
# 缺失时由 rustfs_storage.connect 抛出明确配置错误,而不是拿占位口令连库
|
||||
self.RUSTFS_ENDPOINT = os.getenv("RUSTFS_ENDPOINT") or os.getenv("MINIO_ENDPOINT")
|
||||
self.RUSTFS_ACCESS_KEY = os.getenv("RUSTFS_ACCESS_KEY") or os.getenv("MINIO_ACCESS_KEY")
|
||||
self.RUSTFS_SECRET_KEY = os.getenv("RUSTFS_SECRET_KEY") or os.getenv("MINIO_SECRET_KEY")
|
||||
self.RUSTFS_TIMEOUT = int(os.getenv("RUSTFS_TIMEOUT", "30"))
|
||||
self.RUSTFS_PRESIGNED_URL_EXPIRES = int(os.getenv("RUSTFS_PRESIGNED_URL_EXPIRES", "3600"))
|
||||
|
||||
@@ -36,6 +38,10 @@ class Settings:
|
||||
self.DB_USER = os.getenv("DB_USER")
|
||||
self.DB_PASSWORD = os.getenv("DB_PASSWORD")
|
||||
|
||||
# 启动时是否自动执行 alembic 迁移(D12):多副本同时启动会并发迁移,
|
||||
# 生产多副本应设 false,改由部署流程单点执行 alembic CLI 或本模块 __main__
|
||||
self.AUTO_MIGRATE = os.getenv("AUTO_MIGRATE", "true").lower() == "true"
|
||||
|
||||
self.SECRET_KEY = os.getenv("SECRET_KEY")
|
||||
self.ALGORITHM = os.getenv("ALGORITHM", "HS256")
|
||||
self.ACCESS_TOKEN_EXPIRE_MINUTES = int(os.getenv("ACCESS_TOKEN_EXPIRE_MINUTES", "1440"))
|
||||
|
||||
@@ -118,6 +118,13 @@ async def init_roles(session, perm_map):
|
||||
|
||||
async def create_admin_user(session):
|
||||
"""创建默认管理员"""
|
||||
# compose 不再给 ADMIN_PASSWORD 弱默认(D14):缺失时显式失败,
|
||||
# 而不是静默创建空口令管理员
|
||||
if not settings.ADMIN_PASSWORD:
|
||||
raise RuntimeError(
|
||||
"ADMIN_PASSWORD 未配置:请在 .env 中设置管理员初始密码后重启"
|
||||
)
|
||||
|
||||
result = await session.execute(select(User).where(User.username == settings.ADMIN_USERNAME))
|
||||
existing_admin = result.scalar_one_or_none()
|
||||
|
||||
@@ -150,7 +157,13 @@ async def init_database(keep_connected: bool = True):
|
||||
"""初始化数据库"""
|
||||
try:
|
||||
await db_manager.connect()
|
||||
await _run_alembic_migrations()
|
||||
if settings.AUTO_MIGRATE:
|
||||
await _run_alembic_migrations()
|
||||
else:
|
||||
logger.info(
|
||||
"AUTO_MIGRATE=false:跳过启动期 alembic 迁移,"
|
||||
"schema 由部署流程单点执行(alembic CLI 或 python -m shared.database.init_db)"
|
||||
)
|
||||
|
||||
async with db_manager.session() as session:
|
||||
perm_map = await init_permissions(session)
|
||||
|
||||
@@ -282,6 +282,10 @@ class ProcessingTask(Base):
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
task_id = Column(String(36), unique=True, index=True, nullable=False)
|
||||
stp_file_id = Column(Integer, ForeignKey("stp_files.id"), nullable=False, index=True)
|
||||
|
||||
# 批量上传聚合 ID(批次 2:批量元数据入库——PG 为单一事实源,
|
||||
# 同批任务经此列聚合查询,不再依赖 Redis/进程内存存批量元数据)
|
||||
batch_id = Column(String(36), nullable=True, index=True)
|
||||
|
||||
# 任务类型和状态
|
||||
task_type = Column(String(50), default="stp_parsing") # stp_parsing, geometry_analysis, mold_generation
|
||||
|
||||
@@ -212,10 +212,15 @@ async def create_user(
|
||||
if existing_email.scalar_one_or_none():
|
||||
raise HTTPException(status_code=400, detail="邮箱已存在")
|
||||
|
||||
try:
|
||||
hashed_password = get_password_hash(user_data.password)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
user = User(
|
||||
username=user_data.username,
|
||||
email=user_data.email,
|
||||
hashed_password=get_password_hash(user_data.password),
|
||||
hashed_password=hashed_password,
|
||||
full_name=user_data.full_name,
|
||||
is_active=True
|
||||
)
|
||||
@@ -313,7 +318,10 @@ async def reset_user_password(
|
||||
if not user:
|
||||
raise HTTPException(status_code=404, detail="用户不存在")
|
||||
|
||||
user.hashed_password = get_password_hash(new_password)
|
||||
try:
|
||||
user.hashed_password = get_password_hash(new_password)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
await db_session.commit()
|
||||
|
||||
logger.info(f"管理员 {current_user.username} 重置了用户 {user.username} 的密码")
|
||||
|
||||
@@ -20,13 +20,28 @@ pwd_context = bcrypt
|
||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False)
|
||||
|
||||
|
||||
def _require_secret_key() -> str:
|
||||
"""SECRET_KEY 惰性校验:未配置时给出明确错误,而不是让 jwt.encode/decode 报晦涩 TypeError。"""
|
||||
if not settings.SECRET_KEY:
|
||||
raise RuntimeError("SECRET_KEY 未配置:请在 .env 中设置后重启服务(认证功能不可用)")
|
||||
return settings.SECRET_KEY
|
||||
|
||||
|
||||
def verify_password(plain_password: str, hashed_password: str) -> bool:
|
||||
return pwd_context.checkpw(plain_password.encode('utf-8'), hashed_password.encode('utf-8'))
|
||||
# 比较侧按 bcrypt 语义截断到 72 字节:兼容历史上被截断存储的口令,
|
||||
# 且避免 checkpw 对超长输入直接抛 ValueError(登录会变 500);
|
||||
# 新口令的超长拒绝在 get_password_hash 中完成
|
||||
password_bytes = plain_password.encode('utf-8')[:72]
|
||||
try:
|
||||
return pwd_context.checkpw(password_bytes, hashed_password.encode('utf-8'))
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def get_password_hash(password: str) -> str:
|
||||
# bcrypt 算法上限 72 字节:超长密码必须显式拒绝,静默截断会改变有效密码
|
||||
if len(password.encode('utf-8')) > 72:
|
||||
password = password[:72]
|
||||
raise ValueError("密码长度超过 72 字节限制,请使用更短的密码")
|
||||
return pwd_context.hashpw(password.encode('utf-8'), pwd_context.gensalt()).decode('utf-8')
|
||||
|
||||
|
||||
@@ -37,7 +52,7 @@ def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -
|
||||
else:
|
||||
expire = datetime.utcnow() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
to_encode.update({"exp": expire})
|
||||
encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
|
||||
encoded_jwt = jwt.encode(to_encode, _require_secret_key(), algorithm=settings.ALGORITHM)
|
||||
return encoded_jwt
|
||||
|
||||
|
||||
@@ -49,7 +64,7 @@ async def get_current_user(
|
||||
return None
|
||||
|
||||
try:
|
||||
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
|
||||
payload = jwt.decode(token, _require_secret_key(), algorithms=[settings.ALGORITHM])
|
||||
username: str = payload.get("sub")
|
||||
if username is None:
|
||||
logger.warning(f"[AUTH] Token 中缺少 sub 字段")
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
# services/redis_task_manager.py
|
||||
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理。
|
||||
"""Redis 任务管理器 - 任务状态热缓存(D7:不再有进程内存回退)。
|
||||
|
||||
存储格式:Redis Hash(field -> JSON 字符串)。
|
||||
- update_task 走 HSET 字段级原子更新,消除旧 get->merge->set 三步竞态
|
||||
(后台处理流程与导出端点并发写同一任务时丢更新);
|
||||
- 进度 tick 只重写变化字段,不再全量重写整个任务 blob;
|
||||
- 兼容读旧 string 格式(升级前写入的在途任务),新写入一律 Hash。
|
||||
- 兼容读旧 string 格式(升级前写入的在途任务),新写入一律 Hash;
|
||||
- **PG 是任务状态单一事实源**:Redis 不可用时本管理器不再降级进程内 dict
|
||||
(多副本下各进程内存互相不可见,造成同一任务不同副本读到不同状态),
|
||||
而是 no-op / 返回 None——状态查询路径(TaskQueryService)自然落到 PG。
|
||||
"""
|
||||
|
||||
import json
|
||||
@@ -101,24 +104,6 @@ class RedisTaskManager:
|
||||
raise RuntimeError("Redis 未连接,无法直接访问 redis_client")
|
||||
return self._redis
|
||||
|
||||
# ---- 内存回退 ----
|
||||
_fallback_tasks: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
def _fallback_set(self, task_id: str, data: Dict[str, Any]):
|
||||
self._fallback_tasks[task_id] = data
|
||||
|
||||
def _fallback_get(self, task_id: str) -> Optional[Dict[str, Any]]:
|
||||
return self._fallback_tasks.get(task_id)
|
||||
|
||||
def _fallback_delete(self, task_id: str):
|
||||
self._fallback_tasks.pop(task_id, None)
|
||||
|
||||
def _fallback_all(self) -> Dict[str, Dict[str, Any]]:
|
||||
return dict(self._fallback_tasks)
|
||||
|
||||
def _fallback_count(self) -> int:
|
||||
return len(self._fallback_tasks)
|
||||
|
||||
# ---- 内部工具 ----
|
||||
|
||||
def _key(self, task_id: str) -> str:
|
||||
@@ -161,143 +146,116 @@ class RedisTaskManager:
|
||||
# ---- 公共接口 ----
|
||||
|
||||
async def set_task(self, task_id: str, data: Dict[str, Any], ttl: Optional[int] = None):
|
||||
"""整包写入任务数据(Hash,覆盖旧值,含旧 string 格式清理)"""
|
||||
"""整包写入任务数据(Hash,覆盖旧值,含旧 string 格式清理)。
|
||||
|
||||
Redis 不可用时 no-op:任务状态事实源在 PG,缓存缺失不影响正确性。
|
||||
"""
|
||||
if not self.is_connected:
|
||||
return
|
||||
|
||||
effective_ttl = ttl or self._ttl
|
||||
mapping = self._dump_mapping(data)
|
||||
|
||||
if self.is_connected:
|
||||
try:
|
||||
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, self._make_serializable(data))
|
||||
try:
|
||||
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()
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 写入失败(任务状态以 PG 为准): task={task_id}, {e}")
|
||||
|
||||
async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""获取任务数据(Hash / 旧 string 兼容)"""
|
||||
if self.is_connected:
|
||||
try:
|
||||
return await self._load_any(self._key(task_id))
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 读取失败,回退到内存: {e}")
|
||||
"""获取任务数据(Hash / 旧 string 兼容)。
|
||||
|
||||
return self._fallback_get(task_id)
|
||||
Redis 不可用 / 未命中返回 None,调用方落到 PG 路径。
|
||||
"""
|
||||
if not self.is_connected:
|
||||
return None
|
||||
try:
|
||||
return await self._load_any(self._key(task_id))
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 读取失败(任务状态以 PG 为准): task={task_id}, {e}")
|
||||
return None
|
||||
|
||||
async def update_task(self, task_id: str, updates: Dict[str, Any]):
|
||||
"""字段级原子更新(HSET),无读改写竞态。
|
||||
|
||||
兼容旧 string 格式:先迁移为 Hash 再更新。
|
||||
Redis 不可用时 no-op(状态事实源在 PG)。
|
||||
"""
|
||||
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} 不存在,无法更新")
|
||||
if not self.is_connected:
|
||||
return
|
||||
|
||||
current.update(self._make_serializable(updates))
|
||||
self._fallback_set(task_id, current)
|
||||
mapping = self._dump_mapping(updates)
|
||||
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)
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 更新失败(任务状态以 PG 为准): task={task_id}, {e}")
|
||||
|
||||
async def delete_task(self, task_id: str):
|
||||
"""删除任务(DEL 对 Hash/string 均有效)"""
|
||||
if self.is_connected:
|
||||
try:
|
||||
await self._redis.delete(self._key(task_id))
|
||||
return
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 删除失败,回退到内存: {e}")
|
||||
|
||||
self._fallback_delete(task_id)
|
||||
"""删除任务(DEL 对 Hash/string 均有效);Redis 不可用时 no-op"""
|
||||
if not self.is_connected:
|
||||
return
|
||||
try:
|
||||
await self._redis.delete(self._key(task_id))
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 删除失败: task={task_id}, {e}")
|
||||
|
||||
async def get_all_tasks(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""获取所有任务"""
|
||||
if self.is_connected:
|
||||
try:
|
||||
pattern = f"{self._prefix}*"
|
||||
result = {}
|
||||
async for key in self._redis.scan_iter(match=pattern):
|
||||
task_id = key.replace(self._prefix, "")
|
||||
task = await self._load_any(key)
|
||||
if task:
|
||||
result[task_id] = task
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 扫描失败,回退到内存: {e}")
|
||||
|
||||
return self._fallback_all()
|
||||
"""获取所有任务;Redis 不可用时返回空 dict(调用方需容忍)"""
|
||||
if not self.is_connected:
|
||||
return {}
|
||||
try:
|
||||
pattern = f"{self._prefix}*"
|
||||
result = {}
|
||||
async for key in self._redis.scan_iter(match=pattern):
|
||||
task_id = key.replace(self._prefix, "")
|
||||
task = await self._load_any(key)
|
||||
if task:
|
||||
result[task_id] = task
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 扫描失败: {e}")
|
||||
return {}
|
||||
|
||||
async def get_task_count(self) -> int:
|
||||
"""获取任务总数"""
|
||||
if self.is_connected:
|
||||
try:
|
||||
pattern = f"{self._prefix}*"
|
||||
count = 0
|
||||
async for _ in self._redis.scan_iter(match=pattern):
|
||||
count += 1
|
||||
return count
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 计数失败,回退到内存: {e}")
|
||||
|
||||
return self._fallback_count()
|
||||
|
||||
async def cleanup_old_tasks(self, max_age_seconds: int = 86400 * 7):
|
||||
"""清理过期任务(Redis 由 TTL 自动管理,内存回退需手动清理)"""
|
||||
now = datetime.now()
|
||||
to_delete = []
|
||||
|
||||
for task_id, task in self._fallback_tasks.items():
|
||||
completed_at = task.get("completed_at")
|
||||
if completed_at:
|
||||
try:
|
||||
completed_dt = datetime.fromisoformat(completed_at)
|
||||
if (now - completed_dt).total_seconds() > max_age_seconds:
|
||||
to_delete.append(task_id)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
for task_id in to_delete:
|
||||
del self._fallback_tasks[task_id]
|
||||
|
||||
if to_delete:
|
||||
logger.info(f"清理了 {len(to_delete)} 个过期内存任务")
|
||||
"""获取任务总数;Redis 不可用时返回 0(调用方需容忍)"""
|
||||
if not self.is_connected:
|
||||
return 0
|
||||
try:
|
||||
pattern = f"{self._prefix}*"
|
||||
count = 0
|
||||
async for _ in self._redis.scan_iter(match=pattern):
|
||||
count += 1
|
||||
return count
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis 计数失败: {e}")
|
||||
return 0
|
||||
|
||||
# ---- 工具方法 ----
|
||||
|
||||
|
||||
Reference in New Issue
Block a user