优化
This commit is contained in:
@@ -3,17 +3,22 @@ moldinsight/api/batch_router.py — 批量分析端点
|
||||
|
||||
- POST /api/batch-upload 批量上传多文件,返回 batch_id + 各 task_id
|
||||
- GET /api/batch/{batch_id} 聚合查询批量任务进度
|
||||
|
||||
批次 2(D7):批量元数据以 PG 为单一事实源——ProcessingTask.batch_id
|
||||
列聚合查询,替代此前的 Redis key + 进程内存降级存储。
|
||||
"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import List, Dict, Any, Optional
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from fastapi import APIRouter, UploadFile, File, Form, HTTPException, Depends
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from shared.database.database import get_db_session
|
||||
from shared.services.auth_service import get_current_active_user
|
||||
from shared.models.database import User
|
||||
from shared.models.database import User, ProcessingTask, STPFile
|
||||
from shared.models.schemas import ProcessingStatus, create_task_info
|
||||
from shared.services.redis_task_manager import redis_task_manager
|
||||
from shared.utils.file_handler import FileHandler
|
||||
@@ -27,17 +32,6 @@ router = APIRouter()
|
||||
|
||||
file_handler = FileHandler()
|
||||
|
||||
# ─── 批量元数据 Redis key 约定 ──────────────────────────────────────
|
||||
_BATCH_KEY_PREFIX = "batch:"
|
||||
_BATCH_TTL = 86400 # 24h
|
||||
|
||||
# Redis 不可用时的进程内降级存储(同进程内可查,跨进程/重启不可见)
|
||||
_batch_meta_memory: Dict[str, dict] = {}
|
||||
|
||||
|
||||
def _batch_redis_key(batch_id: str) -> str:
|
||||
return f"{_BATCH_KEY_PREFIX}{batch_id}"
|
||||
|
||||
|
||||
@router.post("/batch-upload")
|
||||
async def batch_upload(
|
||||
@@ -91,8 +85,11 @@ async def batch_upload(
|
||||
)
|
||||
|
||||
await storage_service.create_processing_task(
|
||||
db_session, task_id, stp_file.id, parameters=process_params,
|
||||
db_session, task_id, stp_file.id,
|
||||
parameters=process_params, batch_id=batch_id,
|
||||
)
|
||||
# D9:STPFile + ProcessingTask 原子提交,分派前置事务收口
|
||||
await db_session.commit()
|
||||
|
||||
task_info = create_task_info(
|
||||
task_id=task_id,
|
||||
@@ -108,7 +105,7 @@ async def batch_upload(
|
||||
await redis_task_manager.set_task(task_id, task_info)
|
||||
|
||||
# 调度处理
|
||||
dispatch_processing(task_id, str(file_path), stp_file.id, process_params)
|
||||
dispatch_processing(task_id, stp_file.id, process_params)
|
||||
|
||||
tasks.append({
|
||||
"filename": file.filename,
|
||||
@@ -129,17 +126,6 @@ async def batch_upload(
|
||||
"error": str(exc),
|
||||
})
|
||||
|
||||
# 将 batch 元数据写入 Redis;Redis 不可用时降级到进程内存储(任务状态本身有内存回退)
|
||||
batch_meta = {
|
||||
"batch_id": batch_id,
|
||||
"user_id": current_user.id,
|
||||
"created_at": str(datetime.now()),
|
||||
"task_ids": [t["task_id"] for t in tasks if t.get("task_id")],
|
||||
"total": len(tasks),
|
||||
"params": process_params,
|
||||
}
|
||||
_save_batch_meta(batch_id, batch_meta)
|
||||
|
||||
return {
|
||||
"batch_id": batch_id,
|
||||
"total": len(tasks),
|
||||
@@ -148,67 +134,41 @@ async def batch_upload(
|
||||
}
|
||||
|
||||
|
||||
def _save_batch_meta(batch_id: str, batch_meta: dict):
|
||||
"""批量元数据持久化:优先 Redis(跨进程、带 TTL),降级进程内 dict。"""
|
||||
import json as _json
|
||||
|
||||
if redis_task_manager.is_connected:
|
||||
try:
|
||||
redis_task_manager.redis_client.set(
|
||||
_batch_redis_key(batch_id),
|
||||
_json.dumps(batch_meta),
|
||||
ex=_BATCH_TTL,
|
||||
)
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.warning(f"[BATCH] batch 元数据写 Redis 失败,降级内存: {exc}")
|
||||
_batch_meta_memory[batch_id] = batch_meta
|
||||
|
||||
|
||||
async def _load_batch_meta(batch_id: str) -> Optional[dict]:
|
||||
"""读取批量元数据,Redis 优先,内存兜底;不存在返回 None。"""
|
||||
import json as _json
|
||||
|
||||
if redis_task_manager.is_connected:
|
||||
try:
|
||||
raw = await redis_task_manager.redis_client.get(_batch_redis_key(batch_id))
|
||||
if raw:
|
||||
return _json.loads(raw)
|
||||
except Exception as exc:
|
||||
logger.warning(f"[BATCH] batch 元数据读 Redis 失败: {exc}")
|
||||
return _batch_meta_memory.get(batch_id)
|
||||
|
||||
|
||||
@router.get("/batch/{batch_id}")
|
||||
async def get_batch_status(
|
||||
batch_id: str,
|
||||
db_session: AsyncSession = Depends(get_db_session),
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
):
|
||||
"""聚合查询批量任务进度"""
|
||||
batch_meta = await _load_batch_meta(batch_id)
|
||||
if not batch_meta:
|
||||
"""聚合查询批量任务进度(D7:以 PG 为单一事实源,按 batch_id 聚合;Redis 仅热缓存)"""
|
||||
rows = (await db_session.execute(
|
||||
select(ProcessingTask, STPFile)
|
||||
.join(STPFile, ProcessingTask.stp_file_id == STPFile.id)
|
||||
.where(ProcessingTask.batch_id == batch_id)
|
||||
.options(joinedload(STPFile.html_file))
|
||||
.order_by(ProcessingTask.id)
|
||||
)).unique().all()
|
||||
|
||||
if not rows:
|
||||
raise HTTPException(404, "批量任务不存在或已过期")
|
||||
|
||||
# 权限检查
|
||||
if batch_meta.get("user_id") and batch_meta["user_id"] != current_user.id:
|
||||
# 归属校验:同批任务属于同一上传用户,任一不匹配即拒绝(无主不等于公共)
|
||||
if any(getattr(stp, "user_id", None) != current_user.id for _, stp in rows):
|
||||
raise HTTPException(403, "无权访问该批量任务")
|
||||
|
||||
task_ids = batch_meta.get("task_ids", [])
|
||||
task_statuses = []
|
||||
completed = 0
|
||||
failed = 0
|
||||
processing = 0
|
||||
earliest_created = None
|
||||
|
||||
for tid in task_ids:
|
||||
task_data = await redis_task_manager.get_task(tid)
|
||||
if not task_data:
|
||||
task_statuses.append({"task_id": tid, "status": "unknown"})
|
||||
continue
|
||||
status = task_data.get("status", "unknown")
|
||||
progress = task_data.get("progress", 0)
|
||||
filename = task_data.get("filename", "")
|
||||
error = task_data.get("error", "")
|
||||
html_file = task_data.get("html_file", "")
|
||||
for task, stp in rows:
|
||||
status = task.status or "unknown"
|
||||
|
||||
if earliest_created is None or (
|
||||
task.created_time and task.created_time < earliest_created
|
||||
):
|
||||
earliest_created = task.created_time
|
||||
|
||||
if status == ProcessingStatus.COMPLETED:
|
||||
completed += 1
|
||||
@@ -217,19 +177,24 @@ async def get_batch_status(
|
||||
else:
|
||||
processing += 1
|
||||
|
||||
html_file = ""
|
||||
if stp.html_file and stp.html_file.filename:
|
||||
html_file = f"/html/{stp.html_file.filename}"
|
||||
|
||||
task_statuses.append({
|
||||
"task_id": tid,
|
||||
"task_id": task.task_id,
|
||||
"status": status,
|
||||
"progress": progress,
|
||||
"filename": filename,
|
||||
"error": error,
|
||||
"progress": task.progress or 0,
|
||||
"current_step": task.current_step,
|
||||
"filename": stp.original_filename or "",
|
||||
"error": task.error_message or "",
|
||||
"html_file": html_file,
|
||||
})
|
||||
|
||||
total = len(task_ids)
|
||||
total = len(rows)
|
||||
return {
|
||||
"batch_id": batch_id,
|
||||
"created_at": batch_meta.get("created_at"),
|
||||
"created_at": earliest_created.isoformat() if earliest_created else None,
|
||||
"total": total,
|
||||
"completed": completed,
|
||||
"failed": failed,
|
||||
|
||||
Reference in New Issue
Block a user