"""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, 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) ); """ ) # 旧库迁移: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") 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: """创建一条排队中的工作流运行记录。""" 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, 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["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 再检查暂停,导致"点击暂停反而开始任务"。 """ with self._connect() as conn: row = conn.execute( """ SELECT * FROM workflow_runs WHERE status = 'QUEUED' 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,))