- BatchWorker 单线程轮询 batch_jobs 表,处理 source=batch 的运行, 与主调度器互不抢占(next_queued_run 排除 batch 来源) - 直接读取用户所选文件夹下的视频逐个执行流水线,不上传到工作目录; 中间态与产物落在视频旁同名文件夹,batch.done.json 完成标记去重 - 支持暂停/继续、失败容错(单视频失败不阻塞后续)、删除任务只清库 - 孤儿清理跳过 source=batch 运行,防止误删用户视频文件夹 - workflow_runs 新增 source 列(upload/batch),旧库自动迁移
612 lines
26 KiB
Python
612 lines
26 KiB
Python
"""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)任务。
|
||
|
||
只取 QUEUED:PAUSED 任务必须由用户显式 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)
|