# -*- coding: utf-8 -*- """任务管理器:队列 + 状态机 + SSE 广播""" import os, json, threading, time, uuid, shutil from dataclasses import dataclass, field, asdict from concurrent.futures import ThreadPoolExecutor from .config import OUTPUT_DIR, MAX_CONCURRENT_TASKS, LLM_CONCURRENCY, WORK_DIRNAME from .pipeline.transcribe import transcribe, extract_and_denoise, VIDEO_EXTS from .pipeline.fix import fix from .pipeline.translate import translate_direct, translate_via_en from .pipeline.srt import write_srts # 去重记录文件:放在被处理视频的所在文件夹下 HISTORY_FILENAME = ".subtitle-history.json" def _history_path(video_path): """视频/文件夹的 history 文件路径(跟视频在同一文件夹)""" if os.path.isdir(video_path): folder = video_path else: folder = os.path.dirname(video_path) return os.path.join(folder, HISTORY_FILENAME) def _load_history(video_path): """从视频所在文件夹加载 history(按 视频|路线 记录)""" hp = _history_path(video_path) if os.path.exists(hp): try: with open(hp, encoding='utf-8') as f: return json.load(f) except Exception: return {} return {} def _save_history(video_path, history): """保存 history 到视频所在文件夹""" hp = _history_path(video_path) try: with open(hp, 'w', encoding='utf-8') as f: json.dump(history, f, ensure_ascii=False, indent=1) except Exception: pass @dataclass class Task: id: str source: str # 文件或文件夹路径 route: str # direct / via_en / transcribe_only language: str = "ja" status: str = "pending" # pending/transcribing/fixing/translating/srt/done/error progress: int = 0 message: str = "" created: float = field(default_factory=time.time) finished: float = None error: str = "" files: list = field(default_factory=list) # 生成的 SRT 文件 transcript_path: str = "" video_errors: list = field(default_factory=list) # 单视频失败记录 [{video, error}] processed_count: int = 0 # 成功处理视频数 failed_count: int = 0 # 失败视频数 total_videos: int = 0 # 总视频数(文件夹) started_at: float = None # 任务开始时间(总体耗时) subtasks: dict = field(default_factory=dict) # {video: {status, progress, message, started_at, elapsed}} 子任务状态 def to_dict(self): d = asdict(self) # 动态耗时:运行中 = now - started_at, 完成 = finished - created if self.finished: d["elapsed"] = round(self.finished - self.created, 1) elif self.started_at: d["elapsed"] = round(time.time() - self.started_at, 1) else: d["elapsed"] = 0 return d class TaskManager: def __init__(self): self.tasks = {} # id -> Task self.queue = [] # 等待队列(串行执行) self.lock = threading.Lock() self.worker = None # 当前执行线程 self.subscribers = [] # SSE 客户端队列 # 并发控制 self.transcribe_lock = threading.BoundedSemaphore(MAX_CONCURRENT_TASKS) # 转录并发(GPU 模型共享,默认 2) self.llm_sem = threading.BoundedSemaphore(LLM_CONCURRENCY) # LLM 并发信号量 self.llm_concurrency = LLM_CONCURRENCY self.transcribe_concurrency = MAX_CONCURRENT_TASKS def set_llm_concurrency(self, n): """动态调整 LLM 并发数(重建信号量)""" n = max(1, min(int(n), 20)) self.llm_concurrency = n self.llm_sem = threading.BoundedSemaphore(n) return n def set_transcribe_concurrency(self, n): """动态调整转录并发数(1-4,重建信号量)""" n = max(1, min(int(n), 4)) self.transcribe_concurrency = n self.transcribe_lock = threading.BoundedSemaphore(n) return n # ---- SSE ---- def subscribe(self, q): self.subscribers.append(q) def unsubscribe(self, q): if q in self.subscribers: self.subscribers.remove(q) def _broadcast(self, task): ev = json.dumps({"type": "task_update", "task": task.to_dict()}, ensure_ascii=False) for q in list(self.subscribers): try: q.put(ev) except Exception: pass # ---- 任务管理 ---- def create_task(self, source, route, language="ja"): # 防重复:同 source+route 已有 pending/running 任务则拒绝(避免重复处理) with self.lock: for t in self.tasks.values(): if t.source == source and t.route == route and t.status in ("pending", "transcribing", "fixing", "translating", "srt"): return None, [], f"该内容正在处理中(任务 {t.id[:8]}),请等待完成后再提交" # 去重检查:文件夹 → 找出未处理/处理过的视频清单 skipped = [] if os.path.isdir(source): history = _load_history(source) for f in sorted(os.listdir(source)): if f.lower().endswith(VIDEO_EXTS): key = os.path.join(source, f) hkey = f"{key}|{route}" if hkey in history: skipped.append(key) if skipped and len(skipped) == sum(1 for f in os.listdir(source) if f.lower().endswith(VIDEO_EXTS)): # 全部已处理过,拒绝创建 return None, skipped, "" elif os.path.isfile(source): # 单文件也检查:同路线已处理过则拒绝 history = _load_history(source) hkey = f"{source}|{route}" if hkey in history: return None, [source], "" task = Task(id=uuid.uuid4().hex[:12], source=source, route=route, language=language) with self.lock: self.tasks[task.id] = task self.queue.append(task.id) self._broadcast(task) self._maybe_start() return task, skipped, "" def _maybe_start(self): with self.lock: if self.worker and self.worker.is_alive(): return running = [t for t in self.tasks.values() if t.status in ("transcribing", "fixing", "translating", "srt")] if len(running) >= MAX_CONCURRENT_TASKS: return if not self.queue: return task_id = self.queue.pop(0) task = self.tasks[task_id] task.status = "transcribing" task.message = "开始处理" self.worker = threading.Thread(target=self._run_task, args=(task,), daemon=True) self.worker.start() self._broadcast(task) def _run_task(self, task): try: self._process(task) task.status = "done" task.progress = 100 if task.failed_count > 0: task.message = f"部分完成: 成功 {task.processed_count}, 失败 {task.failed_count}" else: task.message = f"完成 ({task.processed_count} 个视频)" task.finished = time.time() except Exception as ex: task.status = "error" task.error = str(ex) task.message = f"失败: {ex}" task.finished = time.time() self._broadcast(task) # 处理下一个 self._maybe_start() def _update(self, task, msg, pct=None): task.message = msg if pct is not None: task.progress = int(pct) self._broadcast(task) def _process(self, task): # 收集视频文件 videos = [] skipped_in_run = [] if os.path.isfile(task.source): if task.source.lower().endswith(VIDEO_EXTS): videos = [task.source] elif os.path.isdir(task.source): # 过滤已处理(同路线)的视频 history = _load_history(task.source) for f in sorted(os.listdir(task.source)): if f.lower().endswith(VIDEO_EXTS): key = os.path.join(task.source, f) if f"{key}|{task.route}" in history: skipped_in_run.append(key) else: videos.append(key) if skipped_in_run: task.message = f"跳过 {len(skipped_in_run)} 个已处理视频(同路线去重)" self._update(task, task.message, 1) if not videos: raise ValueError(f"所选内容均已处理过,无需重复处理: {task.source}") out_dir = os.path.join(OUTPUT_DIR, task.id) os.makedirs(out_dir, exist_ok=True) total = len(videos) task.total_videos = total task.started_at = time.time() # 初始化子任务状态 for v in videos: task.subtasks[v] = {"status": "pending", "progress": 0, "message": "等待中", "started_at": None, "elapsed": 0} task.message = f"开始处理 {total} 个视频(转录 {MAX_CONCURRENT_TASKS} + LLM {self.llm_concurrency} 并发)" self._update(task, task.message, 2) # 流水线并发:每个视频一个 worker 线程 # 转录内部拿 GPU 锁(串行),修正/翻译拿 LLM 信号量(限并发) done_count = 0 done_lock = threading.Lock() def _update_overall(): """更新总体进度:已完成/失败/总数""" done_total = task.processed_count + task.failed_count overall_pct = int(done_total / total * 100) if total else 100 # 聚合消息 running = [st for st in task.subtasks.values() if st["status"] in ("transcribing", "fixing", "translating", "srt")] running_names = [v for v, st in task.subtasks.items() if st["status"] in ("transcribing", "fixing", "translating", "srt")] msg = f"总体 {done_total}/{total} 完成" if running_names: msg += f" | 处理中: {', '.join(os.path.basename(n) for n in running_names[:3])}" if task.failed_count: msg += f" | 失败 {task.failed_count}" task.message = msg self._update(task, msg, overall_pct) def _sub_update(video, status, progress, message): """更新单个视频子任务状态(含耗时)""" cur = task.subtasks.get(video, {}) started = cur.get("started_at") if started is None: started = time.time() elapsed = (time.time() - started) if started else 0 task.subtasks[video] = { "status": status, "progress": progress, "message": message, "started_at": started, "elapsed": round(elapsed, 1), } self._broadcast(task) def _video_worker(v_idx, video): nonlocal done_count try: _sub_update(video, "transcribing", 5, "转录中") self._process_one(task, video, out_dir, v_idx, total, _sub_update) with done_lock: task.processed_count += 1 done_count += 1 _sub_update(video, "done", 100, "完成") except Exception as ex: with done_lock: task.failed_count += 1 done_count += 1 task.video_errors.append({"video": video, "error": str(ex)[:300]}) # 记录完整 traceback 到日志(排查用) try: import traceback as _tb tb = _tb.format_exc() with open(os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "_error.log"), 'a', encoding='utf-8') as _ef: _ef.write(f"\n=== {os.path.basename(video)} ===\n{tb}\n") except Exception: pass _sub_update(video, "error", 0, f"失败: {str(ex)[:80]}") finally: _update_overall() # 并发执行(worker 数 = 转录锁 + LLM 并发,转录阶段大部分在等锁) pool_size = max(MAX_CONCURRENT_TASKS, self.llm_concurrency) with ThreadPoolExecutor(max_workers=pool_size) as pool: futures = [pool.submit(_video_worker, i, v) for i, v in enumerate(videos)] # 等待全部完成(含失败) for f in futures: f.result() if task.failed_count > 0: task.message = f"部分完成: 成功 {task.processed_count}, 失败 {task.failed_count}" else: task.message = f"完成 ({task.processed_count} 个视频)" self._update(task, task.message, 100) def _process_one(self, task, video, out_dir, v_idx=0, total=1, sub_update=None): """处理单个视频(转录→修正→翻译→SRT)——支持断点续跑 中间产物在 视频所在文件夹/.subtitle-work/(固定目录,不按视频分) 续跑:有 transcript 无 fixed → 从修正开始;有 fixed → 从翻译开始 转录拿 GPU 锁(串行),修正/翻译拿 LLM 信号量(限并发) """ def sub(status, progress, message): if sub_update: sub_update(video, status, progress, message) base = os.path.splitext(os.path.basename(video))[0] fname = os.path.basename(video) # 中间产物目录:视频所在文件夹 .subtitle-work/ work_dir = os.path.join(os.path.dirname(video), WORK_DIRNAME) os.makedirs(work_dir, exist_ok=True) transcript_json = os.path.join(work_dir, f"{base}_transcript.json") fixed_json = os.path.join(work_dir, f"{base}_fixed.json") translated_json = os.path.join(work_dir, f"{base}_translated.json") wav_path = os.path.join(work_dir, f"{base}_audio.wav") # ===== 阶段 1: 转录(可续跑)===== if os.path.exists(transcript_json): sub("transcribing", 100, f"转录已存在(续跑){fname}") else: # 提取音频 + 降噪(CPU,可并发——在 GPU 锁外做) sub("transcribing", 2, f"提取音频+降噪 {fname}") extract_and_denoise(video, wav_path, denoise=True, progress_cb=lambda msg, pct: sub("transcribing", 2 + int(pct * 0.06), f"{msg} {fname}")) # 转录(GPU 锁) with self.transcribe_lock: sub("transcribing", 10, f"转录中 {fname}") def _tcb(msg, pct, _fname=fname): if msg.startswith("转录中"): sub("transcribing", 10 + int(pct * 0.3), f"{_fname}: {msg.split('转录中')[1].strip()}") else: sub("transcribing", 10 + int(pct * 0.3), f"{msg} {_fname}") _result, _warnings = transcribe(video, transcript_json, language=task.language, progress_cb=_tcb, wav_path=wav_path) if _warnings: task.video_errors.append({"video": video, "error": ";".join(_warnings), "warn": True}) task.transcript_path = transcript_json if task.route == "transcribe_only": task.files.append(transcript_json) self._save_history_record(video, task, [transcript_json]) return with open(transcript_json, encoding='utf-8') as f: items = json.load(f) # ===== 阶段 2+3: 修正 + 翻译(LLM 信号量,可续跑)===== with self.llm_sem: # 修正(可续跑) if os.path.exists(fixed_json): sub("fixing", 100, f"修正已存在(续跑){fname}") with open(fixed_json, encoding='utf-8') as f: items = json.load(f) else: sub("fixing", 45, f"LLM 修正 {fname}") items = fix(items, progress_cb=lambda msg, pct: sub("fixing", 45 + int(pct * 0.15), f"{msg} [{fname}]")) with open(fixed_json, 'w', encoding='utf-8') as f: json.dump(items, f, ensure_ascii=False, indent=1) # 翻译(可续跑:从 translated_json 加载已有结果,只翻缺失段) existing_zh = {} existing_en = {} if os.path.exists(translated_json): try: with open(translated_json, encoding='utf-8') as f: prev = json.load(f) # 提取已有翻译 for i, it in enumerate(prev): if it.get('zh'): existing_zh[i] = it['zh'] if it.get('en'): existing_en[i] = it['en'] sub("translating", 30, f"翻译已有 {len(existing_zh)}/{len(items)} 段(续跑){fname}") except Exception: existing_zh, existing_en = {}, {} sub("translating", 40, f"翻译中 {fname}") if task.route == "direct": items = translate_direct(items, progress_cb=lambda msg, pct: sub("translating", 40 + int(pct * 0.5), f"{msg} [{fname}]"), existing_zh=existing_zh) elif task.route == "via_en": items = translate_via_en(items, progress_cb=lambda msg, pct: sub("translating", 40 + int(pct * 0.5), f"{msg} [{fname}]"), existing_en=existing_en, existing_zh=existing_zh) with open(translated_json, 'w', encoding='utf-8') as f: json.dump(items, f, ensure_ascii=False, indent=1) # 4. SRT(写视频同目录,PotPlayer 需要;副本到 output) sub("srt", 92, f"生成 SRT {fname}") srt_files = write_srts(items, video, task.route, out_dir) task.files.extend(srt_files) sub("srt", 98, f"SRT 完成 {fname}") # 5. 记录 history self._save_history_record(video, task, srt_files) def _save_history_record(self, video, task, srt_files): """记录 history(精确到视频,存在视频所在文件夹)""" try: history = _load_history(video) hkey = f"{video}|{task.route}" # 收集警告(transcribe 的 warnings 记录在 video_errors 里 warn=true) warnings = [] for ve in task.video_errors: if ve.get("video") == video and ve.get("warn"): warnings.append(ve["error"]) history[hkey] = { "video": video, "route": task.route, "files": list(srt_files), "warnings": warnings, "time": time.time(), } _save_history(video, history) except Exception: pass # ---- API ---- def list_tasks(self): return [t.to_dict() for t in sorted(self.tasks.values(), key=lambda t: -t.created)] def get_task(self, tid): t = self.tasks.get(tid) return t.to_dict() if t else None def retry_task(self, tid): t = self.tasks.get(tid) if not t or t.status not in ("error", "done"): return None # 复用 ID 重新入队 t.status = "pending" t.error = "" t.message = "重试" t.progress = 0 with self.lock: self.queue.append(t.id) self._broadcast(t) self._maybe_start() return t.to_dict() def cancel_task(self, tid): """从队列移除(无法中断运行中的,只移出等待队列)""" t = self.tasks.get(tid) if not t: return None if t.status == "pending": with self.lock: if tid in self.queue: self.queue.remove(tid) t.status = "cancelled" t.message = "已取消" self._broadcast(t) return t.to_dict() manager = TaskManager()