"""分片上传服务:会话管理 + 分片落盘 + 拼接 + 创建转写任务。 存储布局:: //0.part 分片暂存 //1.part ... ///. complete 后的正式视频 complete 创建转写 Task,管线触发由 controller 调用 scheduler.enqueue_task, 本服务不依赖 scheduler(避免循环依赖)。 """ from __future__ import annotations import logging import os import shutil import threading import uuid from datetime import datetime, timedelta, timezone from pathlib import Path from fastapi import HTTPException from sqlalchemy.orm import Session from ..config import get_settings from ..models.task import Task from ..models.upload_session import UploadSession from ..schemas.task import ( ChunkUploadResponse, CompleteResponse, CreateSessionRequest, CreateSessionResponse, SessionStatusResponse, ) logger = logging.getLogger("audio2text.upload") # 按 upload_id 串行化 uploaded_chunks 的读-改-写,避免并发分片写入丢失更新。 # SQLite 无行锁,JSON 列的 append 操作不是原子的,必须进程内加锁。 _chunk_locks: dict[str, threading.Lock] = {} _chunk_locks_guard = threading.Lock() def _get_chunk_lock(upload_id: str) -> threading.Lock: with _chunk_locks_guard: lock = _chunk_locks.get(upload_id) if lock is None: lock = threading.Lock() _chunk_locks[upload_id] = lock return lock class UploadService: def __init__(self, db: Session) -> None: self.db = db s = get_settings() self.upload_root = s.upload_dir() self.work_root = s.work_dir() self.chunk_bytes = s.storage.chunk_bytes self.session_ttl = s.storage.chunk_session_ttl_seconds # ---------------- 会话生命周期 ---------------- def create_session(self, body: CreateSessionRequest) -> CreateSessionResponse: upload_id = uuid.uuid4().hex session = UploadSession( upload_id=upload_id, filename=body.filename, size_bytes=body.size_bytes, chunk_size=body.chunk_size, total_chunks=body.total_chunks, uploaded_chunks=[], status="pending", ) self.db.add(session) self.db.commit() self._session_dir(upload_id).mkdir(parents=True, exist_ok=True) return CreateSessionResponse( upload_id=upload_id, filename=body.filename, size_bytes=body.size_bytes, chunk_size=body.chunk_size, total_chunks=body.total_chunks, ) def get_status(self, upload_id: str) -> SessionStatusResponse: session = self._require_session(upload_id) return SessionStatusResponse( upload_id=session.upload_id, filename=session.filename, size_bytes=session.size_bytes, chunk_size=session.chunk_size, total_chunks=session.total_chunks, uploaded_chunks=list(session.uploaded_chunks or []), completed=(session.status == "completed"), task_id=session.task_id, ) # ---------------- 分片写入 ---------------- def write_chunk(self, upload_id: str, index: int, data: bytes) -> list[int]: session = self._require_session(upload_id) self._validate_index(session, index) session_dir = self._session_dir(upload_id) session_dir.mkdir(parents=True, exist_ok=True) chunk_path = session_dir / f"{index}.part" # 先落盘(无锁,IO 可并行),再持锁更新 DB 计数 try: with chunk_path.open("wb") as out: out.write(data) out.flush() os.fsync(out.fileno()) except Exception: chunk_path.unlink(missing_ok=True) raise # 持锁重读+更新:避免并发分片各自 read old list → append → write,后者覆盖前者 lock = _get_chunk_lock(upload_id) with lock: self.db.refresh(session) # 拿最新 uploaded_chunks,不被其他并发请求的旧值覆盖 uploaded = list(session.uploaded_chunks or []) if index not in uploaded: uploaded.append(index) session.uploaded_chunks = uploaded session.updated_at = datetime.now(timezone.utc) self.db.commit() return sorted(uploaded) # ---------------- 拼接 + 创建任务 ---------------- # 允许的音视频扩展名白名单(防可执行文件落盘到上传目录) _ALLOWED_EXTS = frozenset({ ".mp4", ".mkv", ".avi", ".mov", ".webm", ".flv", ".mp3", ".wav", ".flac", ".aac", ".m4a", ".ogg", ".wma", }) # ---------------- 拼接 + 创建任务 ---------------- def complete(self, upload_id: str) -> CompleteResponse: session = self._require_session(upload_id) # 幂等:已 complete 则返回已建任务 if session.status == "completed" and session.task_id is not None: task = self.db.get(Task, session.task_id) if task is not None: return CompleteResponse( task_id=task.id, filename=task.filename, size_bytes=session.size_bytes, status=task.status, ) uploaded = set(session.uploaded_chunks or []) missing = [i for i in range(session.total_chunks) if i not in uploaded] if missing: raise HTTPException( 409, f"分片不齐全:缺失 {len(missing)} 个,例如 {sorted(missing)[:10]}", ) final_path = self._assemble(session) rel = str(final_path.relative_to(self.upload_root)) # 单事务:建 Task + 更新 session 状态 + 关联 task_id 一次 commit # 避免双 commit 之间崩溃产生孤儿 Task(Task 已建但 session.task_id 为空) task = Task( filename=session.filename, source_path=rel, status="queued", progress=0.0, ) self.db.add(task) session.status = "completed" session.final_path = rel session.task_id = None # 占位,flush 后用 task.id 赋值 session.updated_at = datetime.now(timezone.utc) self.db.flush() # 拿到 task.id(不 commit,仍在事务内) session.task_id = task.id self.db.commit() self.db.refresh(task) # 清理分片暂存(commit 后,即使清理失败也不影响已建任务) self._cleanup_session_dir(upload_id) # 清理进程内锁,避免长期运行后 _chunk_locks 无限增长 with _chunk_locks_guard: _chunk_locks.pop(upload_id, None) logger.info("上传完成 task_id=%s file=%s size=%d", task.id, session.filename, session.size_bytes) return CompleteResponse( task_id=task.id, filename=task.filename, size_bytes=session.size_bytes, status=task.status, ) # ---------------- 过期会话清理 ---------------- def reap_stale_sessions(self) -> int: """清理被放弃的会话:pending 且 updated_at 超过 ttl。 时间比较统一用 aware UTC datetime,避免 naive datetime 的 .timestamp() 按本地时区算导致的偏移。 """ cutoff = datetime.now(timezone.utc) - timedelta(seconds=self.session_ttl) sessions = ( self.db.query(UploadSession) .filter(UploadSession.status == "pending") .all() ) n = 0 for session in sessions: updated = session.updated_at if updated is None: continue # SQLite 存 naive datetime,统一补 UTC 后比较 if updated.tzinfo is None: updated = updated.replace(tzinfo=timezone.utc) if updated < cutoff: self._cleanup_session_dir(session.upload_id) self.db.delete(session) n += 1 logger.info("清理过期分片会话 upload_id=%s file=%s", session.upload_id, session.filename) if n: self.db.commit() return n # ---------------- 内部 ---------------- def _require_session(self, upload_id: str) -> UploadSession: if not upload_id: raise HTTPException(400, "upload_id 不能为空") session = ( self.db.query(UploadSession) .filter(UploadSession.upload_id == upload_id) .first() ) if session is None: raise HTTPException(404, f"会话不存在或已过期:{upload_id}") return session @staticmethod def _validate_index(session: UploadSession, index: int) -> None: if index < 0 or index >= session.total_chunks: raise HTTPException(400, f"分片下标越界:{index} 不在 [0, {session.total_chunks})") def _session_dir(self, upload_id: str) -> Path: return self.work_root / upload_id def _assemble(self, session: UploadSession) -> Path: """按 index 顺序拼接全部分片为正式视频文件。""" # 扩展名取自客户端 filename,但做白名单净化:不在允许列表内则回退 .bin ext = Path(session.filename).suffix.lower() if ext not in self._ALLOWED_EXTS: ext = ".bin" now = datetime.now(timezone.utc) sub = self.upload_root / f"{now:%Y}" / f"{now:%m}" sub.mkdir(parents=True, exist_ok=True) final = sub / f"{uuid.uuid4().hex}{ext}" part_path = final.with_suffix(final.suffix + ".part") session_dir = self._session_dir(session.upload_id) try: with part_path.open("wb") as out: for index in range(session.total_chunks): chunk_path = session_dir / f"{index}.part" if not chunk_path.is_file(): raise HTTPException(409, f"拼接时发现分片缺失:{index}.part") with chunk_path.open("rb") as src: while buf := src.read(self.chunk_bytes): out.write(buf) out.flush() os.fsync(out.fileno()) os.replace(part_path, final) except Exception: part_path.unlink(missing_ok=True) raise return final def _cleanup_session_dir(self, upload_id: str) -> None: session_dir = self._session_dir(upload_id) try: if session_dir.exists(): shutil.rmtree(session_dir, ignore_errors=True) except Exception as exc: # pragma: no cover logger.warning("清理会话目录失败 upload_id=%s: %s", upload_id, exc)