"""转写管线编排:提取音频 → ASR → 断句 → 翻译 → 写 SRT。 任务状态机: queued → extracting → transcribing → segmenting → translating → done 任一步失败 → failed 模型不共驻:ASR 与翻译分阶段加载,翻译时先卸载 Whisper 释放显存跑大 batch。 管线在后台线程跑(每个任务一个线程),通过 DB 更新状态与进度。 """ from __future__ import annotations import logging import threading import traceback from datetime import datetime, timezone from pathlib import Path from ..config import get_settings from ..database import get_session_local from ..models.task import Task from . import ffmpeg_service, asr_service, segmenter, translate_service, srt_writer from .types import Subtitle logger = logging.getLogger("audio2text.pipeline") # 进度锚点(各阶段在 0-100 中的占比) P_EXTRACT = 5.0 P_TRANSCRIBE_START = 5.0 P_TRANSCRIBE_END = 55.0 P_SEGMENT_START = 55.0 P_SEGMENT_END = 60.0 P_TRANSLATE_START = 60.0 P_TRANSLATE_END = 98.0 P_DONE = 100.0 def enqueue_task(task_id: int) -> None: """把任务交给后台线程处理(非阻塞,供 upload_service.complete 调用)。""" t = threading.Thread(target=_run_task, args=(task_id,), daemon=True) t.start() logger.info("任务 %d 已入队(后台线程 %s)。", task_id, t.name) def _run_task(task_id: int) -> None: """后台执行完整管线。所有异常都被捕获并写入 task.error。""" db = get_session_local()() try: task = db.get(Task, task_id) if task is None: logger.error("任务 %d 不存在。", task_id) return _pipeline(db, task) except Exception as exc: logger.exception("任务 %d 失败:%s", task_id, exc) _mark_failed(db, task_id, str(exc)) finally: db.close() def _pipeline(db, task: Task) -> None: s = get_settings() src = s.upload_dir() / task.source_path # ---------- 1. 提取音频 ---------- _set_status(db, task, "extracting", P_EXTRACT) wav = s.work_dir() / f"task_{task.id}.wav" ffmpeg_service.extract_audio(src, wav) # 按配置决定是否删原始视频(提取成功后) if s.processing.delete_original_after_extract and src.is_file(): try: src.unlink() logger.info("已删除原始视频 %s(delete_original_after_extract=true)。", src.name) except OSError as exc: logger.warning("删除原始视频失败 %s: %s", src, exc) # ---------- 2. 语音识别 ---------- _set_status(db, task, "transcribing", P_TRANSCRIBE_START) segments = asr_service.transcribe(wav) _set_status(db, task, "transcribing", P_TRANSCRIBE_END, note=f"识别出 {len(segments)} 段") # ---------- 3. 断句 + 时间戳重算 ---------- _set_status(db, task, "segmenting", P_SEGMENT_START) subs = segmenter.resegment(segments) _set_status(db, task, "segmenting", P_SEGMENT_END, note=f"重组为 {len(subs)} 条字幕") # ---------- 4. 翻译 ---------- _set_status(db, task, "translating", P_TRANSLATE_START) # 翻译阶段:model_manager 会自动卸载 ASR、加载翻译器(独占显存) zh_texts = translate_service.translate(subs) _set_status(db, task, "translating", P_TRANSLATE_END, note=f"翻译 {len(zh_texts)} 条") # ---------- 5. 写 SRT ---------- out_dir = s.output_dir() stem = Path(task.filename).stem en_path = out_dir / f"task_{task.id}/{stem}.en.srt" zh_path = out_dir / f"task_{task.id}/{stem}.zh.srt" bi_path = out_dir / f"task_{task.id}/{stem}.srt" srt_writer.write_srt(subs, en_path) # 中文 SRT(用译文 + 同时间戳) zh_subs = [Subtitle(text=zh, start=sub.start, end=sub.end) for zh, sub in zip(zh_texts, subs)] srt_writer.write_srt(zh_subs, zh_path) srt_writer.write_bilingual_srt(subs, zh_texts, bi_path) # 记录相对路径 task.en_srt_path = str(en_path.relative_to(out_dir)) task.zh_srt_path = str(zh_path.relative_to(out_dir)) task.bilingual_srt_path = str(bi_path.relative_to(out_dir)) task.status = "done" task.progress = P_DONE task.updated_at = datetime.now(timezone.utc) db.commit() # 清理中间音频 if not s.processing.keep_audio and wav.is_file(): try: wav.unlink() except OSError: pass logger.info("任务 %d 完成:%s", task.id, bi_path.name) # ---------------- DB 状态更新 ---------------- def _set_status(db, task: Task, status: str, progress: float, note: str = "") -> None: task.status = status task.progress = progress task.updated_at = datetime.now(timezone.utc) db.commit() if note: logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note) else: logger.info("任务 %d [%s %.0f%%]", task.id, status, progress) def _mark_failed(db, task_id: int, error: str) -> None: try: task = db.get(Task, task_id) if task is None: return task.status = "failed" task.error = error[:2000] task.updated_at = datetime.now(timezone.utc) db.commit() except Exception: # pragma: no cover logger.error("写入失败状态时又失败:\n%s", traceback.format_exc())