Files
vrsub/src/wov_app/db.py
T
cat-shark 711867e79f feat: 文件夹批量处理引擎(后端)
- BatchWorker 单线程轮询 batch_jobs 表,处理 source=batch 的运行,
  与主调度器互不抢占(next_queued_run 排除 batch 来源)
- 直接读取用户所选文件夹下的视频逐个执行流水线,不上传到工作目录;
  中间态与产物落在视频旁同名文件夹,batch.done.json 完成标记去重
- 支持暂停/继续、失败容错(单视频失败不阻塞后续)、删除任务只清库
- 孤儿清理跳过 source=batch 运行,防止误删用户视频文件夹
- workflow_runs 新增 source 列(upload/batch),旧库自动迁移
2026-08-23 16:25:09 +08:00

612 lines
26 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""SQLite 数据访问层。
所有持久化逻辑集中在本模块,业务代码只依赖 Database 提供的方法。后续切换
PostgreSQL 时只需替换本层实现,不修改调度器与路由的业务逻辑。
"""
from __future__ import annotations
import json
import sqlite3
from contextlib import contextmanager
from pathlib import Path
from typing import Any, Iterator
class Database:
"""SQLite 数据库封装:负责建表以及工作流/任务/产物的 CRUD。"""
def __init__(self, path: Path) -> None:
"""打开数据库并确保父目录存在、表结构已初始化。"""
self.path = path
self.path.parent.mkdir(parents=True, exist_ok=True)
self._init_schema()
@contextmanager
def _connect(self) -> Iterator[sqlite3.Connection]:
"""提供带事务提交的数据库连接上下文。"""
conn = sqlite3.connect(self.path)
# 按列名读取结果,返回 dict 更直观。
conn.row_factory = sqlite3.Row
# 开启外键约束,保证子表记录引用有效。
conn.execute("PRAGMA foreign_keys = ON")
try:
yield conn
conn.commit()
finally:
conn.close()
def _init_schema(self) -> None:
"""创建全部业务表;已存在的表保持不变。"""
with self._connect() as conn:
conn.executescript(
"""
-- 工作流表:只保存概要信息,完整定义存版本表。
CREATE TABLE IF NOT EXISTS workflows (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
description TEXT NOT NULL DEFAULT '',
published INTEGER NOT NULL DEFAULT 0,
latest_version INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
);
-- 工作流版本表:每个版本保存一份 DAG 定义 JSON。
CREATE TABLE IF NOT EXISTS workflow_versions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
workflow_id TEXT NOT NULL,
version INTEGER NOT NULL,
definition_json TEXT NOT NULL,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
UNIQUE(workflow_id, version),
FOREIGN KEY(workflow_id) REFERENCES workflows(id)
);
-- 工作流运行表:记录任务从排队到完成/失败的状态机。
CREATE TABLE IF NOT EXISTS workflow_runs (
id TEXT PRIMARY KEY,
workflow_id TEXT NOT NULL,
workflow_version INTEGER NOT NULL,
status TEXT NOT NULL,
current_node_id TEXT,
progress REAL NOT NULL DEFAULT 0,
error TEXT,
input_uri TEXT,
param_overrides TEXT,
source TEXT NOT NULL DEFAULT 'upload',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL,
FOREIGN KEY(workflow_id) REFERENCES workflows(id)
);
-- 产物表:记录每个任务各节点的输出 URI,按名称唯一。
CREATE TABLE IF NOT EXISTS artifacts (
id INTEGER PRIMARY KEY AUTOINCREMENT,
run_id TEXT NOT NULL,
node_id TEXT NOT NULL,
name TEXT NOT NULL,
uri TEXT NOT NULL,
mime_type TEXT NOT NULL DEFAULT '',
size INTEGER NOT NULL DEFAULT 0,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
UNIQUE(run_id, name),
FOREIGN KEY(run_id) REFERENCES workflow_runs(id)
);
-- 批量处理任务表:一次"文件夹批量处理"对应一条记录,记录目标
-- 文件夹、所选工作流与整体状态。批量引擎与 Web 页面共用。
CREATE TABLE IF NOT EXISTS batch_jobs (
id TEXT PRIMARY KEY,
folder_path TEXT NOT NULL,
workflow_id TEXT NOT NULL,
recursive INTEGER NOT NULL DEFAULT 1,
status TEXT NOT NULL,
progress REAL NOT NULL DEFAULT 0,
total INTEGER NOT NULL DEFAULT 0,
done INTEGER NOT NULL DEFAULT 0,
failed INTEGER NOT NULL DEFAULT 0,
current_video TEXT,
error TEXT,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
-- 批量视频明细表:一次批量任务处理的每个视频一条记录,保存其
-- 对应的工作流 run(断点续跑复用 workflow_runs 的产物状态)。
CREATE TABLE IF NOT EXISTS batch_videos (
id TEXT PRIMARY KEY,
job_id TEXT NOT NULL,
video_path TEXT NOT NULL,
work_dir TEXT NOT NULL,
run_id TEXT,
status TEXT NOT NULL,
error TEXT,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL,
FOREIGN KEY(job_id) REFERENCES batch_jobs(id)
);
"""
)
# 旧库迁移:workflow_runs 补充 param_overrides 列(前端框选覆盖)。
columns = [
row["name"]
for row in conn.execute("PRAGMA table_info(workflow_runs)").fetchall()
]
if "param_overrides" not in columns:
conn.execute("ALTER TABLE workflow_runs ADD COLUMN param_overrides TEXT")
# 旧库迁移:workflow_runs 补充 source 列(upload=网页上传 / batch=批量处理),
# 批量引擎与主调度器据此隔离任务,避免互相抢占。
if "source" not in columns:
conn.execute("ALTER TABLE workflow_runs ADD COLUMN source TEXT NOT NULL DEFAULT 'upload'")
def upsert_workflow(self, workflow: dict[str, Any]) -> None:
"""插入或更新工作流概要信息。"""
with self._connect() as conn:
conn.execute(
"""
INSERT INTO workflows (id, name, description, published, latest_version)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
name = excluded.name,
description = excluded.description,
published = excluded.published,
latest_version = excluded.latest_version
""",
(
workflow["id"],
workflow["name"],
workflow.get("description", ""),
int(workflow.get("published", 0)),
int(workflow.get("latest_version", 0)),
),
)
def get_workflow(self, workflow_id: str) -> dict[str, Any] | None:
"""按 ID 读取工作流概要。"""
with self._connect() as conn:
row = conn.execute("SELECT * FROM workflows WHERE id = ?", (workflow_id,)).fetchone()
return dict(row) if row else None
def list_workflows(self) -> list[dict[str, Any]]:
"""按创建时间倒序返回全部工作流。"""
with self._connect() as conn:
rows = conn.execute("SELECT * FROM workflows ORDER BY created_at DESC").fetchall()
return [dict(row) for row in rows]
def delete_workflow(self, workflow_id: str) -> None:
"""级联删除工作流相关的产物、任务、版本和概要记录。"""
with self._connect() as conn:
# 外键没有级联删除配置,手动按依赖顺序清理。
conn.execute("DELETE FROM artifacts WHERE run_id IN (SELECT id FROM workflow_runs WHERE workflow_id = ?)", (workflow_id,))
conn.execute("DELETE FROM workflow_runs WHERE workflow_id = ?", (workflow_id,))
conn.execute("DELETE FROM workflow_versions WHERE workflow_id = ?", (workflow_id,))
conn.execute("DELETE FROM workflows WHERE id = ?", (workflow_id,))
def create_workflow_version(self, workflow_id: str, version: int, definition: dict[str, Any]) -> None:
"""为工作流新增一个版本,definition 以 JSON 保存。"""
with self._connect() as conn:
conn.execute(
"""
INSERT INTO workflow_versions (workflow_id, version, definition_json)
VALUES (?, ?, ?)
""",
(workflow_id, version, json.dumps(definition, ensure_ascii=False)),
)
def get_latest_workflow_version(self, workflow_id: str) -> dict[str, Any] | None:
"""返回工作流最新版本,并把 definition_json 反序列化为 definition。"""
with self._connect() as conn:
row = conn.execute(
"""
SELECT * FROM workflow_versions
WHERE workflow_id = ?
ORDER BY version DESC
LIMIT 1
""",
(workflow_id,),
).fetchone()
if row is None:
return None
result = dict(row)
# 对外统一暴露 definition 字典,隐藏 JSON 存储细节。
result["definition"] = json.loads(result.pop("definition_json"))
return result
def get_workflow_version(self, workflow_id: str, version: int) -> dict[str, Any] | None:
"""按版本号读取指定工作流版本。"""
with self._connect() as conn:
row = conn.execute(
"""
SELECT * FROM workflow_versions
WHERE workflow_id = ? AND version = ?
""",
(workflow_id, version),
).fetchone()
if row is None:
return None
result = dict(row)
result["definition"] = json.loads(result.pop("definition_json"))
return result
def list_workflow_versions(self, workflow_id: str) -> list[dict[str, Any]]:
"""按版本倒序返回工作流全部版本。"""
with self._connect() as conn:
rows = conn.execute(
"""
SELECT * FROM workflow_versions
WHERE workflow_id = ?
ORDER BY version DESC
""",
(workflow_id,),
).fetchall()
versions = []
for row in rows:
item = dict(row)
item["definition"] = json.loads(item.pop("definition_json"))
versions.append(item)
return versions
def create_run(self, run: dict[str, Any]) -> None:
"""创建一条排队中的工作流运行记录。
source 标识任务来源:upload(网页上传,默认)由主调度器执行;
batch(文件夹批量处理)由批量引擎执行,input_uri 直接指向本地视频。
"""
with self._connect() as conn:
conn.execute(
"""
INSERT INTO workflow_runs (
id, workflow_id, workflow_version, status, current_node_id,
progress, error, input_uri, param_overrides, source,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
run["id"],
run["workflow_id"],
run["workflow_version"],
run["status"],
run.get("current_node_id"),
float(run.get("progress", 0)),
run.get("error"),
run.get("input_uri"),
json.dumps(run["param_overrides"], ensure_ascii=False)
if run.get("param_overrides")
else None,
run.get("source", "upload"),
run["created_at"],
run["updated_at"],
),
)
def get_run(self, run_id: str) -> dict[str, Any] | None:
"""按 ID 读取任务运行记录。"""
with self._connect() as conn:
row = conn.execute("SELECT * FROM workflow_runs WHERE id = ?", (run_id,)).fetchone()
return self._parse_overrides(row) if row else None
@staticmethod
def _parse_overrides(row) -> dict:
"""把查询行中的 param_overrides JSON 字符串解析为字典。"""
result = dict(row)
raw = result.get("param_overrides")
result["param_overrides"] = json.loads(raw) if raw else None
return result
def list_runs(self, limit: int = 20) -> list[dict[str, Any]]:
"""按创建时间倒序返回最近的运行记录。"""
with self._connect() as conn:
rows = conn.execute(
"SELECT * FROM workflow_runs ORDER BY created_at DESC LIMIT ?",
(limit,),
).fetchall()
return [self._parse_overrides(row) for row in rows]
def list_run_ids(self) -> list[str]:
"""返回全部任务 ID,供孤儿数据清理对照磁盘目录使用。"""
with self._connect() as conn:
rows = conn.execute("SELECT id FROM workflow_runs").fetchall()
return [row["id"] for row in rows]
def update_run(self, run_id: str, **fields: Any) -> None:
"""更新运行状态字段,同时刷新 updated_at;未知字段会被忽略。"""
# 只允许更新状态机相关字段,防止任意列被改写。
allowed = {
"status",
"current_node_id",
"progress",
"error",
}
updates = {key: value for key, value in fields.items() if key in allowed}
if not updates:
return
updates["updated_at"] = fields.get("updated_at")
# 动态拼接 SET 子句,键来自白名单,不存在 SQL 注入风险。
assignments = ", ".join(f"{key} = ?" for key in updates)
values = list(updates.values()) + [run_id]
with self._connect() as conn:
conn.execute(f"UPDATE workflow_runs SET {assignments} WHERE id = ?", values)
def reset_run(self, run_id: str, updated_at: str) -> None:
"""把失败任务重置为 QUEUED,并清空进度与旧产物,供重试使用。"""
with self._connect() as conn:
# 清空错误和进度,恢复到首次排队时的状态。
conn.execute(
"""
UPDATE workflow_runs
SET status = 'QUEUED', current_node_id = NULL, progress = 0,
error = NULL, updated_at = ?
WHERE id = ?
""",
(updated_at, run_id),
)
# 删除旧产物,避免重试后残留过期下载链接。
conn.execute("DELETE FROM artifacts WHERE run_id = ?", (run_id,))
def next_queued_run(self) -> dict[str, Any] | None:
"""按创建时间返回最早一条排队(QUEUED)任务。
只取 QUEUEDPAUSED 任务必须由用户显式 resume(转回 QUEUED)后调度器
才重新执行。修复回归——此前把 PAUSED 也当可执行任务拾起,execute_run
会先置 RUNNING 再检查暂停,导致"点击暂停反而开始任务"。
同时排除 source=batch 的批量运行:批量任务由批量引擎使用视频旁的
同名文件夹作为 storage 执行,主调度器拾起会用错存储目录。
"""
with self._connect() as conn:
row = conn.execute(
"""
SELECT * FROM workflow_runs
WHERE status = 'QUEUED' AND source != 'batch'
ORDER BY created_at ASC
LIMIT 1
"""
).fetchone()
return self._parse_overrides(row) if row else None
def recover_interrupted_runs(self, updated_at: str) -> int:
"""重启恢复:把遗留 RUNNING 任务恢复为 QUEUED,返回恢复数量。
进程被杀/重启时 RUNNING 任务没有机会收尾:保持 RUNNING 会永久孤儿
next_queued_run 不拾起)。恢复为 QUEUED 后调度器会从产物表
restore_run_outputs)断点续跑,不重复已完成节点;用户主动暂停的
PAUSED 任务保持不变,等待显式 resume。
"""
with self._connect() as conn:
cur = conn.execute(
"UPDATE workflow_runs SET status = 'QUEUED', updated_at = ? WHERE status = 'RUNNING'",
(updated_at,),
)
return cur.rowcount
def pause_run(self, run_id: str, updated_at: str) -> None:
"""暂停任务:置为 PAUSED;调度器会在节点边界检查并停止推进。"""
with self._connect() as conn:
conn.execute(
"UPDATE workflow_runs SET status = 'PAUSED', updated_at = ? WHERE id = ?",
(updated_at, run_id),
)
def resume_run(self, run_id: str, updated_at: str) -> None:
"""继续任务:PAUSED 恢复为 QUEUED,等待调度器从断点续跑。"""
with self._connect() as conn:
conn.execute(
"UPDATE workflow_runs SET status = 'QUEUED', updated_at = ? WHERE id = ?",
(updated_at, run_id),
)
def restore_run_outputs(self, run_id: str) -> dict[str, dict[str, str]]:
"""从已登记的产物重建各节点输出,供暂停后断点续跑使用。
返回 {节点ID: {输出名: URI}};已完成节点的产物可直接作为后续节点的输入。
"""
outputs: dict[str, dict[str, str]] = {}
for artifact in self.list_artifacts(run_id):
name = artifact["name"]
# 产物名形如 "节点ID.输出名"(如 a.data_uri),还原为 {输出名: URI}。
prefix = artifact["node_id"] + "."
if name.startswith(prefix):
name = name[len(prefix):]
outputs.setdefault(artifact["node_id"], {})[name] = artifact["uri"]
return outputs
def create_artifact(self, artifact: dict[str, Any]) -> None:
"""记录任务产物;同 run 与 name 冲突时覆盖。"""
with self._connect() as conn:
conn.execute(
"""
INSERT OR REPLACE INTO artifacts (
run_id, node_id, name, uri, mime_type, size
)
VALUES (?, ?, ?, ?, ?, ?)
""",
(
artifact["run_id"],
artifact["node_id"],
artifact["name"],
artifact["uri"],
artifact.get("mime_type", ""),
int(artifact.get("size", 0)),
),
)
def list_artifacts(self, run_id: str) -> list[dict[str, Any]]:
"""按创建时间返回任务的全部产物。"""
with self._connect() as conn:
rows = conn.execute(
"SELECT * FROM artifacts WHERE run_id = ? ORDER BY created_at",
(run_id,),
).fetchall()
return [dict(row) for row in rows]
def get_artifact(self, run_id: str, name: str) -> dict[str, Any] | None:
"""按任务与产物名读取单个产物记录。"""
with self._connect() as conn:
row = conn.execute(
"SELECT * FROM artifacts WHERE run_id = ? AND name = ?",
(run_id, name),
).fetchone()
return dict(row) if row else None
def delete_run_artifacts(self, run_id: str) -> None:
"""删除任务的全部产物记录。"""
with self._connect() as conn:
conn.execute("DELETE FROM artifacts WHERE run_id = ?", (run_id,))
def delete_run(self, run_id: str) -> None:
"""删除任务记录本身及其产物记录。
产物表外键引用任务表,必须先删产物再删任务,否则违反外键约束。
"""
with self._connect() as conn:
conn.execute("DELETE FROM artifacts WHERE run_id = ?", (run_id,))
conn.execute("DELETE FROM workflow_runs WHERE id = ?", (run_id,))
# ------------------------------------------------------------------
# 批量处理任务(batch_jobs / batch_videos)数据访问。
# 批量引擎与批量管理页共用这些方法,规则与 workflow_runs 一致。
# ------------------------------------------------------------------
def create_batch_job(self, job: dict[str, Any]) -> None:
"""插入一条批量处理任务记录。"""
with self._connect() as conn:
conn.execute(
"""
INSERT INTO batch_jobs (
id, folder_path, workflow_id, recursive, status, progress,
total, done, failed, current_video, error, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
job["id"],
job["folder_path"],
job["workflow_id"],
int(job.get("recursive", 1)),
job["status"],
float(job.get("progress", 0)),
int(job.get("total", 0)),
int(job.get("done", 0)),
int(job.get("failed", 0)),
job.get("current_video"),
job.get("error"),
job["created_at"],
job["updated_at"],
),
)
def get_batch_job(self, job_id: str) -> dict[str, Any] | None:
"""按 ID 读取批量任务记录。"""
with self._connect() as conn:
row = conn.execute("SELECT * FROM batch_jobs WHERE id = ?", (job_id,)).fetchone()
return dict(row) if row else None
def list_batch_jobs(self, limit: int = 50) -> list[dict[str, Any]]:
"""按创建时间倒序返回最近的批量任务。"""
with self._connect() as conn:
rows = conn.execute(
"SELECT * FROM batch_jobs ORDER BY created_at DESC LIMIT ?",
(limit,),
).fetchall()
return [dict(row) for row in rows]
def list_batch_job_ids(self) -> list[str]:
"""返回全部批量任务 ID,供孤儿清理区分批量运行使用。"""
with self._connect() as conn:
rows = conn.execute("SELECT id FROM batch_jobs").fetchall()
return [row["id"] for row in rows]
def next_queued_batch_job(self) -> dict[str, Any] | None:
"""按创建时间返回最早一条排队(QUEUED)的批量任务。
批量引擎单线程顺序处理,同一时刻只执行一个批量任务。
"""
with self._connect() as conn:
row = conn.execute(
"""
SELECT * FROM batch_jobs
WHERE status = 'QUEUED'
ORDER BY created_at ASC
LIMIT 1
"""
).fetchone()
return dict(row) if row else None
def update_batch_job(self, job_id: str, **fields: Any) -> None:
"""更新批量任务字段,同时刷新 updated_at;未知字段会被忽略。"""
allowed = {
"status",
"progress",
"total",
"done",
"failed",
"current_video",
"error",
}
updates = {key: value for key, value in fields.items() if key in allowed}
if not updates:
return
updates["updated_at"] = fields.get("updated_at")
assignments = ", ".join(f"{key} = ?" for key in updates)
values = list(updates.values()) + [job_id]
with self._connect() as conn:
conn.execute(f"UPDATE batch_jobs SET {assignments} WHERE id = ?", values)
def delete_batch_job(self, job_id: str) -> None:
"""删除批量任务记录及其全部视频明细(不含 workflow_runs)。"""
with self._connect() as conn:
conn.execute("DELETE FROM batch_videos WHERE job_id = ?", (job_id,))
conn.execute("DELETE FROM batch_jobs WHERE id = ?", (job_id,))
def create_batch_video(self, item: dict[str, Any]) -> None:
"""插入一条批量视频明细记录。"""
with self._connect() as conn:
conn.execute(
"""
INSERT INTO batch_videos (
id, job_id, video_path, work_dir, run_id, status, error,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
item["id"],
item["job_id"],
item["video_path"],
item["work_dir"],
item.get("run_id"),
item["status"],
item.get("error"),
item["created_at"],
item["updated_at"],
),
)
def get_batch_video(self, video_id: str) -> dict[str, Any] | None:
"""按 ID 读取批量视频明细。"""
with self._connect() as conn:
row = conn.execute(
"SELECT * FROM batch_videos WHERE id = ?", (video_id,)
).fetchone()
return dict(row) if row else None
def list_batch_videos(self, job_id: str) -> list[dict[str, Any]]:
"""按创建时间返回一次批量任务的全部视频明细。"""
with self._connect() as conn:
rows = conn.execute(
"SELECT * FROM batch_videos WHERE job_id = ? ORDER BY created_at",
(job_id,),
).fetchall()
return [dict(row) for row in rows]
def update_batch_video(self, video_id: str, **fields: Any) -> None:
"""更新批量视频字段,同时刷新 updated_at;未知字段会被忽略。"""
allowed = {"run_id", "status", "error"}
updates = {key: value for key, value in fields.items() if key in allowed}
if not updates:
return
updates["updated_at"] = fields.get("updated_at")
assignments = ", ".join(f"{key} = ?" for key in updates)
values = list(updates.values()) + [video_id]
with self._connect() as conn:
conn.execute(f"UPDATE batch_videos SET {assignments} WHERE id = ?", values)