nex_docus/backend/app/services/project_export_service.py

203 lines
6.7 KiB
Python

"""
项目导出业务服务
"""
import asyncio
import threading
import time
import uuid
import zipfile
from pathlib import Path
from fastapi import HTTPException
from app.services.storage import storage_service
class ProjectExportService:
"""管理项目导出任务及其文件打包流程。"""
def __init__(self, ttl_seconds: int = 3600):
self.ttl_seconds = ttl_seconds
self._tasks: dict[str, dict] = {}
self._tasks_lock = threading.Lock()
@staticmethod
def _build_zip_filename(project_name: str) -> str:
safe_project_name = project_name.replace("/", "_").replace("\\", "_")
return f"{safe_project_name}.zip"
@staticmethod
def _serialize_task(task: dict) -> dict:
total_files = task.get("total_files", 0) or 0
processed_files = task.get("processed_files", 0) or 0
progress = task.get("progress")
if progress is None:
progress = int(processed_files * 100 / total_files) if total_files else 0
return {
"task_id": task["task_id"],
"project_id": task["project_id"],
"status": task["status"],
"message": task.get("message", ""),
"progress": progress,
"processed_files": processed_files,
"total_files": total_files,
"file_count": total_files,
"zip_filename": task["zip_filename"],
"error": task.get("error"),
"created_at": task.get("created_at"),
"completed_at": task.get("completed_at"),
}
def _remove_task(self, task_id: str) -> None:
with self._tasks_lock:
task = self._tasks.pop(task_id, None)
if not task:
return
file_path = task.get("file_path")
if file_path:
try:
Path(file_path).unlink(missing_ok=True)
except Exception:
pass
def cleanup_expired_tasks(self) -> None:
now = time.time()
expired_task_ids = []
with self._tasks_lock:
for task_id, task in self._tasks.items():
created_at = task.get("created_at", now)
completed_at = task.get("completed_at")
if completed_at and now - completed_at > self.ttl_seconds:
expired_task_ids.append(task_id)
elif task.get("status") == "failed" and now - created_at > self.ttl_seconds:
expired_task_ids.append(task_id)
for task_id in expired_task_ids:
self._remove_task(task_id)
def _update_task(self, task_id: str, **fields) -> None:
with self._tasks_lock:
task = self._tasks.get(task_id)
if task:
task.update(fields)
def get_task_or_404(self, task_id: str) -> dict:
self.cleanup_expired_tasks()
with self._tasks_lock:
task = self._tasks.get(task_id)
if not task:
raise HTTPException(status_code=404, detail="导出任务不存在或已过期")
return task
def get_owned_task_or_404(self, project_id: int, task_id: str, user_id: int) -> dict:
task = self.get_task_or_404(task_id)
if task["project_id"] != project_id or task["user_id"] != user_id:
raise HTTPException(status_code=404, detail="导出任务不存在")
return task
def cleanup_task(self, task_id: str) -> None:
self._remove_task(task_id)
def serialize_task(self, task: dict) -> dict:
return self._serialize_task(task)
def _run_export_task(self, task_id: str, source_dir: Path) -> None:
try:
self._update_task(
task_id,
status="scanning",
progress=0,
message="正在统计导出文件...",
processed_files=0,
total_files=0,
)
files = [file_path for file_path in source_dir.rglob("*") if file_path.is_file()]
total_files = len(files)
storage_service.temp_root.mkdir(parents=True, exist_ok=True)
zip_path = storage_service.temp_root / f"{task_id}.zip"
self._update_task(
task_id,
status="zipping",
progress=0 if total_files else 100,
message="正在打包项目文件...",
total_files=total_files,
file_path=str(zip_path),
)
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zip_file:
for index, file_path in enumerate(files, start=1):
arcname = file_path.relative_to(source_dir)
zip_file.write(file_path, arcname)
progress = int(index * 100 / total_files) if total_files else 100
self._update_task(
task_id,
processed_files=index,
progress=progress,
)
self._update_task(
task_id,
status="completed",
progress=100,
message="项目导出已完成",
completed_at=time.time(),
)
except Exception as exc:
self._update_task(
task_id,
status="failed",
message="项目导出失败",
error=str(exc),
completed_at=time.time(),
)
task = self.get_task_or_404(task_id)
file_path = task.get("file_path")
if file_path:
try:
Path(file_path).unlink(missing_ok=True)
except Exception:
pass
self._update_task(task_id, file_path=None)
async def start_export(self, project_id: int, user_id: int, project_name: str, source_dir: Path) -> dict:
if not source_dir.exists() or not source_dir.is_dir():
raise HTTPException(status_code=404, detail="项目目录不存在")
self.cleanup_expired_tasks()
task_id = uuid.uuid4().hex
task = {
"task_id": task_id,
"project_id": project_id,
"user_id": user_id,
"status": "pending",
"message": "导出任务已创建",
"progress": 0,
"processed_files": 0,
"total_files": 0,
"zip_filename": self._build_zip_filename(project_name),
"file_path": None,
"error": None,
"created_at": time.time(),
"completed_at": None,
}
with self._tasks_lock:
self._tasks[task_id] = task
asyncio.create_task(asyncio.to_thread(self._run_export_task, task_id, source_dir))
return self._serialize_task(task)
project_export_service = ProjectExportService()