Files
CapacityReport/app/api/routers/tasks.py
T

225 lines
7.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from datetime import datetime
from pathlib import Path
from threading import Thread
from typing import Any
from fastapi import APIRouter, Body, HTTPException
from app import state
from app.config import AppConfig
from app.processor import DataProcessor, ProcessLogger
from app.services.license import LicenseError, LicenseInfo, check_processing_allowed
router = APIRouter(tags=["tasks"])
@router.get("/api/task/status")
async def get_global_task_status():
if state.global_task_lock["locked"]:
task_id = state.global_task_lock["task_id"]
if task_id and _task_finished(task_id):
state.reset_task_lock()
return {"has_active": False}
return {
"has_active": True,
"task_id": task_id,
"stage": state.global_task_lock["stage"],
"started_at": state.global_task_lock["started_at"],
"logs": [],
}
active_tasks = {
task_id: task
for task_id, task in state.processing_tasks.items()
if task.get("status") == "processing"
}
if active_tasks:
task_id = next(iter(active_tasks))
return {
"has_active": True,
"task_id": task_id,
"stage": active_tasks[task_id].get("stage", "processing"),
"logs": active_tasks[task_id].get("logs", []),
}
return {"has_active": False}
@router.post("/api/task/lock")
async def lock_task(task_id: str = Body(..., embed=True)):
if state.global_task_lock["locked"]:
raise HTTPException(status_code=409, detail="已有任务在运行")
state.global_task_lock.update(
{
"locked": True,
"task_id": task_id,
"stage": "uploading",
"started_at": datetime.now().isoformat(),
}
)
return {"success": True, "message": "任务已锁定"}
@router.post("/api/task/unlock")
async def unlock_task(task_id: str | None = Body(None, embed=True)):
if task_id and state.global_task_lock["task_id"] != task_id:
raise HTTPException(status_code=403, detail="无权解锁此任务")
state.reset_task_lock()
return {"success": True, "message": "任务已解锁"}
@router.get("/api/process/active")
async def get_active_task():
return await get_global_task_status()
@router.post("/api/process/start")
async def start_processing(task_id: str = Body(..., embed=True)):
record = state.history_manager.get(task_id)
if not record:
raise HTTPException(status_code=404, detail="任务不存在")
if record.status == "processing":
raise HTTPException(status_code=400, detail="任务正在处理中")
work_dir = Path(record.work_dir)
if not work_dir.exists():
raise HTTPException(status_code=400, detail="工作目录不存在")
logs: list[str] = []
current_stage = "processing"
def log_callback(message: str) -> None:
logs.append(message)
_set_task_stage(task_id, current_stage, logs)
def stage_callback(stage: str) -> None:
nonlocal current_stage
current_stage = stage
_set_task_stage(task_id, current_stage, logs)
logger = ProcessLogger(
log_file=work_dir / "log.txt",
callback=log_callback,
stage_callback=stage_callback,
)
app_config = state.current_config()
state.history_manager.update(task_id, status="processing")
state.processing_tasks[task_id] = {"logs": [], "status": "processing", "stage": current_stage}
state.global_task_lock.update(
{
"locked": True,
"task_id": task_id,
"stage": "processing",
"started_at": datetime.now().isoformat(),
}
)
thread = Thread(target=_run_processing, args=(task_id, work_dir, logger, app_config), daemon=True)
thread.start()
return {"success": True, "message": "处理任务已启动", "task_id": task_id}
@router.post("/api/process/status")
async def get_processing_status(task_id: str = Body(..., embed=True)):
if task_id in state.processing_tasks:
task_info = state.processing_tasks[task_id]
logs = task_info.get("logs") or state.history_manager.get_logs(task_id)
return {
"task_id": task_id,
"status": task_info["status"],
"stage": task_info.get("stage"),
"logs": logs,
"error": task_info.get("error"),
"error_detail": task_info.get("error_detail"),
}
record = state.history_manager.get(task_id)
if not record:
raise HTTPException(status_code=404, detail="任务不存在")
return {
"task_id": task_id,
"status": record.status,
"stage": record.status,
"logs": state.history_manager.get_logs(task_id),
"elapsed_time": record.elapsed_time,
"error": record.error,
}
def _task_finished(task_id: str) -> bool:
record = state.history_manager.get(task_id)
if record and record.status in {"completed", "failed"}:
return True
task_info: dict[str, Any] | None = state.processing_tasks.get(task_id)
return bool(task_info and task_info.get("status") in {"completed", "failed"})
def _set_task_stage(task_id: str, stage: str, logs: list[str], status: str = "processing") -> None:
state.processing_tasks[task_id] = {
"logs": logs.copy(),
"status": status,
"stage": stage,
}
if state.global_task_lock["task_id"] == task_id:
state.global_task_lock["stage"] = stage
def _run_processing(task_id: str, work_dir: Path, logger: ProcessLogger, app_config: AppConfig) -> None:
try:
logger.set_stage("license")
_log_license_check(logger, check_processing_allowed(work_dir))
processor = DataProcessor(app_config, work_dir, logger)
result = processor.process()
status = "completed" if result.get("success") else "failed"
error = result.get("error")
state.history_manager.update(
task_id,
status=status,
elapsed_time=result.get("elapsed_time", 0),
error=error,
result_tables=["4G_结果表", "5G_结果表"],
)
state.processing_tasks[task_id] = {
"logs": state.history_manager.get_logs(task_id),
"status": status,
"stage": status,
"error": error,
}
except Exception as exc:
error_detail = exc.to_detail() if isinstance(exc, LicenseError) else None
logger.error(str(exc))
state.history_manager.update(task_id, status="failed", error=str(exc))
state.processing_tasks[task_id] = {
"logs": state.history_manager.get_logs(task_id),
"status": "failed",
"stage": "failed",
"error": str(exc),
"error_detail": error_detail,
}
finally:
try:
state.apply_history_retention()
except Exception as exc:
print(f"自动清理处理历史失败: {exc}")
state.reset_task_lock()
def _log_license_check(logger: ProcessLogger, info: LicenseInfo) -> None:
if info.current_date:
logger.info(
f"授权校验通过,数据日期: {info.current_date.isoformat()},"
f"到期日期: {info.expires_on.isoformat()}"
)
elif info.zip_count:
logger.warning("未从 ZIP 文件名识别到日期,已跳过授权日期比对")
else:
logger.warning("未找到 ZIP 文件,已跳过授权日期比对")