"""工作流调度器。 轮询 SQLite 中的排队任务,按工作流 DAG 的拓扑顺序依次调用进程内节点 处理器,并把节点产物登记为任务产物。单体版使用单线程顺序执行,节点在 同一进程内直接调用,不再经过子进程与 HTTP 协议。 """ from __future__ import annotations import threading import time from datetime import datetime, timezone from pathlib import Path from wov_sdk.models import InvokeRequest, WorkflowDefinition from wov_app import registry from wov_app.config import SCHEDULER_INTERVAL_SECONDS from wov_app.db import Database from wov_app.logging import get_logger from wov_app.storage import atomic_copy # 调度器运行日志:节点进度、暂停/续跑等状态变化。 logger = get_logger("scheduler") def _now_iso() -> str: """返回当前 UTC 时间的 ISO 格式字符串。""" return datetime.now(timezone.utc).isoformat() def topological_sort(definition: WorkflowDefinition) -> list[str]: """对工作流 DAG 做拓扑排序,返回可执行的节点 ID 顺序。""" nodes = {node.id: node for node in definition.nodes} # 统计每个节点的入度,并记录依赖关系。 indegree = {node_id: 0 for node_id in nodes} dependents: dict[str, list[str]] = {node_id: [] for node_id in nodes} for edge in definition.edges: # 边引用了不存在的节点时直接报错。 if edge.from_node not in nodes or edge.to_node not in nodes: raise ValueError(f"unknown edge: {edge.from_node} -> {edge.to_node}") indegree[edge.to_node] += 1 dependents[edge.from_node].append(edge.to_node) # Kahn 算法:从入度为 0 的节点开始逐层取出。 queue = [node_id for node_id, degree in indegree.items() if degree == 0] ordered: list[str] = [] while queue: current = queue.pop(0) ordered.append(current) for dependent in dependents[current]: indegree[dependent] -= 1 if indegree[dependent] == 0: queue.append(dependent) # 排序结果数量不足说明存在环,无法确定执行顺序。 if len(ordered) != len(nodes): raise ValueError("workflow contains a cycle") return ordered class WorkflowScheduler: """后台任务调度器:单线程轮询并执行排队中的工作流运行。""" def __init__( self, db: Database, storage_dir: Path, interval_seconds: float | None = None, ) -> None: """保存依赖并初始化轮询线程控制字段。""" self.db = db self.storage_dir = storage_dir self.interval_seconds = interval_seconds or SCHEDULER_INTERVAL_SECONDS self._thread: threading.Thread | None = None self._stopping = False def start(self) -> None: """启动调度线程;重复调用无副作用。""" if self._thread is not None: return self._stopping = False self._thread = threading.Thread( target=self._loop, name="wov-workflow-scheduler", daemon=True, ) self._thread.start() def stop(self) -> None: """请求停止并等待轮询线程退出。""" self._stopping = True if self._thread is not None: self._thread.join(timeout=5) self._thread = None def _loop(self) -> None: """轮询循环:有排队任务就立即执行,否则休眠一个间隔。""" while not self._stopping: try: run = self.db.next_queued_run() if run is not None: self.execute_run(run["id"]) else: time.sleep(self.interval_seconds) except Exception: # noqa: BLE001 # 单次轮询异常不能杀死调度线程:否则任务会永远停在 QUEUED # 无人拾起。记录后跳过本轮,下一轮继续。 logger.exception("调度器轮询异常,跳过本轮") time.sleep(self.interval_seconds) def _resolve_ref( self, ref: str, run_input_uri: str | None, outputs_by_node: dict[str, dict[str, str]], ) -> str | None: """解析输入引用:input.xxx 取任务入口,node.key 取前序节点产物。""" # 入口引用以 input. 为前缀。 if ref.startswith("input."): return run_input_uri # 其余引用必须形如 "节点ID.输出名"。 node_id, separator, key = ref.partition(".") if not separator: return None return outputs_by_node.get(node_id, {}).get(key) def execute_run(self, run_id: str, stop_after: str | None = None) -> None: """执行单个任务:加载 DAG、按拓扑顺序调用节点并登记产物。 stop_after 指定"只执行到该节点"(批量分块流水线的阶段执行):该节点完成 后任务保持 RUNNING 不收尾,下一次调用从产物表跳过已完成节点继续后面的 阶段;不传时执行整条 DAG 并收尾(登记 final_outputs、标 COMPLETED)。 """ run = self.db.get_run(run_id) # 任务不存在或不在可执行状态时直接返回,避免重复执行。RUNNING 只来自 # 分阶段执行的上一个阶段(任务保持 RUNNING 等下一阶段)或进程异常中断的 # 残留,续跑时已完成节点由产物表跳过;调度器只拾取 QUEUED 任务、批量 # 引擎单线程推进,不会出现两个驱动方重复执行同一任务。 if run is None or run["status"] not in ("QUEUED", "PAUSED", "RUNNING"): return # 已暂停的任务不自动续跑:直接返回保持 PAUSED,等用户显式 resume # (resume 转回 QUEUED 后才执行);否则暂停会被立刻覆盖成 RUNNING。 if run["status"] == "PAUSED": return # 工作流或版本记录丢失时把任务标记为失败。 workflow = self.db.get_workflow(run["workflow_id"]) if workflow is None: self.db.update_run(run_id, status="FAILED", error="workflow not found", updated_at=_now_iso()) return version = self.db.get_workflow_version(run["workflow_id"], run["workflow_version"]) if version is None: self.db.update_run(run_id, status="FAILED", error="workflow version not found", updated_at=_now_iso()) return # 解析并校验 DAG,随后计算拓扑执行顺序。 # 任何预检异常(缺字段/边引用不存在/环形依赖)都必须在这里把任务标 # FAILED:否则队首记录会一直停在 QUEUED 堵塞后续任务。 try: definition = WorkflowDefinition.from_dict(version["definition"]) definition.validate() ordered = topological_sort(definition) # 阶段节点必须存在于 DAG:写错会让任务永远停在 RUNNING 无人推进。 if stop_after is not None and stop_after not in ordered: raise ValueError(f"stop_after node not in workflow: {stop_after}") except Exception as exc: # noqa: BLE001 logger.exception("任务 %s 工作流定义无效,标记失败: %s", run_id, exc) self.db.update_run( run_id, status="FAILED", error=str(exc), updated_at=_now_iso(), ) return # 从已登记产物重建已完成节点的输出,支持暂停后断点续跑。 outputs_by_node = self.db.restore_run_outputs(run_id) run_started = time.monotonic() # 清除可能残留的暂停信号(重启/异常中断后),避免本次执行误触发节点内暂停。 (self.storage_dir / "runs" / run_id / "paused.flag").unlink(missing_ok=True) self.db.update_run(run_id, status="RUNNING", progress=0, updated_at=_now_iso()) try: for index, node_id in enumerate(ordered): # 暂停检查:用户暂停后调度器在节点边界停下,保持 PAUSED 等待续跑。 current = self.db.get_run(run_id) if current is None or current["status"] == "PAUSED": logger.info("任务 %s 已暂停,停止在节点 %s 之前", run_id, node_id) return # 断点续跑:跳过已产出结果的节点(其产物已作为输入可用)。 if node_id in outputs_by_node: # 阶段边界落在已完成的节点上:直接结束本阶段。 if node_id == stop_after: break continue # 当前节点进度 = 已完成节点数 / 总节点数。 node_spec = next(item for item in definition.nodes if item.id == node_id) self.db.update_run( run_id, current_node_id=node_id, progress=index / len(ordered), updated_at=_now_iso(), ) # 解析节点声明的每个输入引用,缺任一输入即失败。 invoke_inputs: dict[str, str] = {} for input_name, ref in node_spec.inputs.items(): value = self._resolve_ref(ref, run.get("input_uri"), outputs_by_node) if value is None: raise ValueError(f"missing input {input_name} for node {node_id}") invoke_inputs[input_name] = value # 前端框选的 crop 等参数覆盖:按节点 ID 合并进节点参数。 node_params = dict(node_spec.params) overrides = run.get("param_overrides") or {} node_params.update(overrides.get(node_id, {})) # 每个任务的每个节点使用独立产物目录,避免并发冲突。 node_started = time.monotonic() output_dir = ( self.storage_dir / "runs" / run_id / "steps" / node_id ) response = registry.invoke( node_spec.node_type, InvokeRequest( run_id=run_id, node_instance_id="", inputs=invoke_inputs, params=node_params, output_dir=str(output_dir), ), ) # 节点返回非 completed 即视为步骤失败。 if response.status != "completed": raise RuntimeError(response.error or f"node {node_id} failed") # 记录节点输出,供后续节点引用和最终产物映射使用。 outputs_by_node[node_id] = { str(key): str(value) for key, value in response.outputs.items() } for key, uri in outputs_by_node[node_id].items(): # 产物名带节点前缀,例如 asr.srt_uri,避免跨节点重名。 artifact = { "run_id": run_id, "node_id": node_id, "name": f"{node_id}.{key}", "uri": uri, "mime_type": self._mime_type(uri), "size": self._file_size(uri), } self.db.create_artifact(artifact) # 进度日志:节点序号/总数、耗时与任务累计运行时间(数据速度可观测)。 logger.info( "任务 %s 进度 %d/%d 节点: %s 耗时 %.1fs, 运行累计 %.1fs", run_id, index + 1, len(ordered), node_id, time.monotonic() - node_started, time.monotonic() - run_started, ) # 阶段边界:本阶段节点已完成,不再执行后续节点。 if node_id == stop_after: break # 分阶段执行:本阶段节点已全部完成(含本轮跳过的情况),任务保持 # RUNNING 等下一个阶段,不做 final_outputs 与完成标记。 if stop_after is not None: logger.info( "任务 %s 阶段完成: 已执行到节点 %s(分阶段执行,保持 RUNNING 等下一阶段)", run_id, stop_after, ) return # 处理 final_outputs,为用户端提供简洁的下载别名。 for alias, ref in definition.final_outputs.items(): resolved = self._resolve_ref(ref, run.get("input_uri"), outputs_by_node) if resolved is None: raise ValueError(f"missing final output: {alias} ({ref})") # 成品按 上传文件名.标识.时间戳 命名,保留节点原始文件与 URI。 # 同一个 run 重复收尾使用稳定路径,不再因暂停/重启反复改名。 resolved = self._final_artifact_uri(resolved, run, definition, alias, ref) self.db.create_artifact( { "run_id": run_id, "node_id": ref.partition(".")[0], "name": alias, "uri": resolved, "mime_type": self._mime_type(resolved), "size": self._file_size(resolved), } ) # 全部节点成功后标记完成;期间被暂停则保持 PAUSED,等待续跑补做收尾。 if self.db.get_run(run_id)["status"] == "PAUSED": logger.info("任务 %s 节点全部完成但已暂停,保持 PAUSED", run_id) return self.db.update_run( run_id, status="COMPLETED", current_node_id=None, progress=1.0, updated_at=_now_iso(), ) except Exception as exc: # noqa: BLE001 # 节点执行中被暂停(节点内检测到 paused.flag 而中止):保持 PAUSED # 等待用户 resume 从断点续跑,而不是把暂停误报为 FAILED。 current = self.db.get_run(run_id) if current is not None and current["status"] == "PAUSED": logger.info("任务 %s 节点内被暂停,保持 PAUSED: %s", run_id, exc) return # 其余异常:结束任务并记录错误,等待用户重试。 self.db.update_run( run_id, status="FAILED", error=str(exc), updated_at=_now_iso(), ) def _final_artifact_uri( self, resolved: str, run: dict, definition: WorkflowDefinition, alias: str, ref: str, ) -> str: """生成命名成品副本并返回稳定 URI,节点原始产物始终保留。 标识优先取产出节点的 target_language,否则回退别名;时间戳固定为 run 创建时间。每个别名单独目录,防止相同语言/扩展名的输出互相覆盖。 复制使用原子替换;收尾重复执行覆盖相同目标,不生成新的时间戳副本。 """ source = Path(resolved) if not source.is_file(): # 兼容历史记录:原文件已改名,但最终别名仍指向有效路径时复用。 existing = self.db.get_artifact(run["id"], alias) if existing is not None and Path(existing["uri"]).is_file(): return existing["uri"] raise ValueError(f"missing final output file: {alias} ({resolved})") # 基础名来自上传文件名;无上传文件时退回通用名称 subtitle。 base = Path(run["input_uri"]).stem if run.get("input_uri") else "subtitle" # 通过最终输出引用定位产出节点,取其语言参数作为标识。 node_id = ref.partition(".")[0] node = next((item for item in definition.nodes if item.id == node_id), None) tag = (node.params.get("target_language") if node else None) or alias # 编码目录别名避免路径分隔符;保留常规 cn_srt/ass 名称便于排查。 from urllib.parse import quote timestamp = datetime.fromisoformat(run["created_at"]).strftime("%Y%m%d%H%M%S") final_dir = self.storage_dir / "runs" / run["id"] / "finals" / ("output-" + quote(alias, safe="")) filename = f"{base}.{tag}.{timestamp}{source.suffix}" if Path(filename).name != filename or "\\" in filename: raise ValueError(f"invalid final output filename: {alias}") new_path = final_dir / filename atomic_copy(source, new_path) return str(new_path) @staticmethod def _mime_type(uri: str) -> str: """按扩展名推断产物 MIME 类型,未知类型使用通用二进制类型。""" path = Path(uri) suffix = path.suffix.lower() return { ".srt": "application/x-subrip", ".ass": "text/plain", ".wav": "audio/wav", ".mp4": "video/mp4", ".txt": "text/plain", }.get(suffix, "application/octet-stream") @staticmethod def _file_size(uri: str) -> int: """读取产物文件大小;文件缺失或路径非法时按 0 处理。 需同时捕获 ValueError:节点返回的产物 URI 可能来自把二进制文件当 文本读取的输入(含 \x00 等字符),Path.stat() 会抛 "embedded null character in path"。此处只影响产物元数据登记大小,不应让整个任务 失败;真实产物不存在时 _final_artifact_uri 已负责报错。 """ try: return Path(uri).stat().st_size except (OSError, ValueError): return 0