deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/workflows/sql.py
2026-09-07 18:24:55 +08:00

386 lines
17 KiB
Python

"""SQLAlchemy-backed workflow definition / version store."""
from __future__ import annotations
import json
import uuid
from datetime import UTC, datetime
from typing import Any
from sqlalchemy import func, select, update
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from deerflow.persistence.workflows.base import (
WorkflowDraftConflictError,
WorkflowNotFoundError,
WorkflowStore,
)
from deerflow.persistence.workflows.model import WorkflowDefinitionRow, WorkflowVersionRow
# Concurrent publishers CAS-retry against the workflow's own counter; a handful
# of attempts is far more than any realistic contention window needs.
_PUBLISH_ATTEMPTS = 6
def _iso(value: datetime | None) -> str | None:
return value.isoformat() if isinstance(value, datetime) else None
def _parse_graph(raw: str | dict[str, Any] | None) -> dict[str, Any]:
if raw is None:
return {}
if isinstance(raw, dict):
return raw
text = str(raw).strip()
if not text:
return {}
try:
data = json.loads(text)
except json.JSONDecodeError:
return {}
return data if isinstance(data, dict) else {}
def _dump_graph(graph: dict[str, Any]) -> str:
return json.dumps(graph, ensure_ascii=False, separators=(",", ":"), sort_keys=True)
def _definition_dict(row: WorkflowDefinitionRow, *, include_draft: bool) -> dict[str, Any]:
data = {
"id": row.id,
"name": row.name,
"description": row.description or "",
"owner_id": row.owner_id,
"status": row.status,
"draft_revision": row.draft_revision,
"created_at": _iso(row.created_at),
"updated_at": _iso(row.updated_at),
}
if include_draft:
data["draft_graph"] = _parse_graph(row.draft_graph_json)
# Canvas document as stored (opaque JSON text; "{}" when never saved).
data["draft_canvas_schema"] = row.draft_canvas_schema_json or "{}"
return data
def _version_dict(row: WorkflowVersionRow, *, include_graph: bool = True) -> dict[str, Any]:
data = {
"id": row.id,
"workflow_id": row.workflow_id,
"version_number": row.version_number,
"graph_hash": row.graph_hash,
"input_schema": row.input_schema or {},
"output_schema": row.output_schema or {},
"published_by": row.published_by,
"published_at": _iso(row.published_at),
"change_note": row.change_note or "",
}
if include_graph:
data["graph"] = _parse_graph(row.graph_json)
return data
class WorkflowRepository(WorkflowStore):
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
async def list_definitions(
self,
*,
owner_id: str | None = None,
status: str | None = None,
include_archived: bool = False,
) -> list[dict[str, Any]]:
stmt = select(WorkflowDefinitionRow)
if owner_id:
stmt = stmt.where(WorkflowDefinitionRow.owner_id == owner_id)
if status:
stmt = stmt.where(WorkflowDefinitionRow.status == status)
elif not include_archived:
stmt = stmt.where(WorkflowDefinitionRow.status != "archived")
stmt = stmt.order_by(WorkflowDefinitionRow.updated_at.desc())
async with self._sf() as session:
result = await session.execute(stmt)
return [_definition_dict(row, include_draft=False) for row in result.scalars()]
async def get_definition(self, workflow_id: str, *, include_draft: bool = True) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(WorkflowDefinitionRow, workflow_id)
return _definition_dict(row, include_draft=include_draft) if row is not None else None
async def create_definition(self, data: dict[str, Any]) -> dict[str, Any]:
workflow_id = data.get("id") or f"wf_{uuid.uuid4().hex[:12]}"
graph = data.get("draft_graph") or data.get("graph") or {
"schemaVersion": "1.0",
"id": workflow_id,
"name": data.get("name") or "未命名工作流",
"nodes": [{"id": "start_1", "type": "start", "name": "开始", "config": {}}],
"edges": [],
"inputSchema": {"type": "object", "properties": {}},
"outputSchema": {"type": "object", "properties": {}},
"settings": {},
}
if isinstance(graph, dict) and not graph.get("id"):
graph = {**graph, "id": workflow_id}
if isinstance(graph, dict) and data.get("name") and not graph.get("name"):
graph = {**graph, "name": data["name"]}
row = WorkflowDefinitionRow(
id=workflow_id,
name=data.get("name") or graph.get("name") or "未命名工作流",
description=data.get("description") or "",
owner_id=data["owner_id"],
status=data.get("status") or "draft",
draft_revision=0,
draft_graph_json=_dump_graph(graph),
draft_canvas_schema_json=str(data.get("draft_canvas_schema") or "{}"),
)
async with self._sf() as session:
session.add(row)
await session.flush()
out = _definition_dict(row, include_draft=True)
await session.commit()
return out
async def update_definition(self, workflow_id: str, data: dict[str, Any]) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(WorkflowDefinitionRow, workflow_id)
if row is None:
return None
if "name" in data and data["name"] is not None:
row.name = data["name"]
if "description" in data and data["description"] is not None:
row.description = data["description"]
if "status" in data and data["status"] is not None:
row.status = data["status"]
row.updated_at = datetime.now(UTC)
await session.flush()
out = _definition_dict(row, include_draft=True)
await session.commit()
return out
async def archive_definition(self, workflow_id: str) -> dict[str, Any] | None:
return await self.update_definition(workflow_id, {"status": "archived"})
async def save_draft(
self,
workflow_id: str,
*,
expected_revision: int,
graph: dict[str, Any],
updated_by: str | None = None,
) -> dict[str, Any]:
del updated_by # audit reserved for later
async with self._sf() as session:
row = await session.get(WorkflowDefinitionRow, workflow_id)
if row is None:
raise WorkflowNotFoundError(workflow_id)
if row.draft_revision != expected_revision:
raise WorkflowDraftConflictError(workflow_id, row.draft_revision)
payload = dict(graph)
payload.setdefault("id", workflow_id)
if row.name and not payload.get("name"):
payload["name"] = row.name
# Conditional UPDATE on the revision itself: two sessions that both
# read the same revision cannot both win — the loser's WHERE clause
# matches zero rows and surfaces a deterministic conflict instead of
# silently overwriting the other editor's draft.
values: dict[str, Any] = {
"draft_graph_json": _dump_graph(payload),
"draft_revision": WorkflowDefinitionRow.draft_revision + 1,
"updated_at": datetime.now(UTC),
}
if payload.get("name"):
values["name"] = str(payload["name"])[:255]
result = await session.execute(
update(WorkflowDefinitionRow)
.where(
WorkflowDefinitionRow.id == workflow_id,
WorkflowDefinitionRow.draft_revision == expected_revision,
)
.values(**values)
)
if result.rowcount == 0:
await session.rollback()
current = await session.get(WorkflowDefinitionRow, workflow_id)
if current is None:
raise WorkflowNotFoundError(workflow_id)
raise WorkflowDraftConflictError(workflow_id, current.draft_revision)
await session.commit()
fresh = await session.get(WorkflowDefinitionRow, workflow_id)
assert fresh is not None
return _definition_dict(fresh, include_draft=True)
async def save_studio_draft(
self,
workflow_id: str,
*,
expected_revision: int,
graph: dict[str, Any],
canvas_schema: str = "{}",
updated_by: str | None = None,
) -> dict[str, Any]:
del updated_by # audit reserved for later
async with self._sf() as session:
row = await session.get(WorkflowDefinitionRow, workflow_id)
if row is None:
raise WorkflowNotFoundError(workflow_id)
if row.draft_revision != expected_revision:
raise WorkflowDraftConflictError(workflow_id, row.draft_revision)
payload = dict(graph)
payload.setdefault("id", workflow_id)
if row.name and not payload.get("name"):
payload["name"] = row.name
# One conditional UPDATE writes graph + canvas + revision together:
# the two representations are always from the same save, and a
# concurrent editor loses deterministically instead of interleaving
# a new canvas onto someone else's graph.
values: dict[str, Any] = {
"draft_graph_json": _dump_graph(payload),
"draft_canvas_schema_json": canvas_schema or "{}",
"draft_revision": WorkflowDefinitionRow.draft_revision + 1,
"updated_at": datetime.now(UTC),
}
if payload.get("name"):
values["name"] = str(payload["name"])[:255]
result = await session.execute(
update(WorkflowDefinitionRow)
.where(
WorkflowDefinitionRow.id == workflow_id,
WorkflowDefinitionRow.draft_revision == expected_revision,
)
.values(**values)
)
if result.rowcount == 0:
await session.rollback()
current = await session.get(WorkflowDefinitionRow, workflow_id)
if current is None:
raise WorkflowNotFoundError(workflow_id)
raise WorkflowDraftConflictError(workflow_id, current.draft_revision)
await session.commit()
fresh = await session.get(WorkflowDefinitionRow, workflow_id)
assert fresh is not None
return _definition_dict(fresh, include_draft=True)
async def publish_version(
self,
workflow_id: str,
*,
graph: dict[str, Any],
graph_hash: str,
published_by: str | None = None,
change_note: str = "",
canvas_schema_hash: str | None = None,
) -> dict[str, Any]:
"""Allocate the version number from the workflow's own counter.
``MAX(version_number) + 1`` races between concurrent publishes; here the
number is claimed by CAS-bumping ``next_version_number`` inside the same
transaction that inserts the version row, and the unique
``(workflow_id, version_number)`` constraint is the final safety net.
Concurrent publishers get adjacent, distinct numbers.
"""
last_error: Exception | None = None
for _ in range(_PUBLISH_ATTEMPTS):
try:
async with self._sf() as session:
definition = await session.get(WorkflowDefinitionRow, workflow_id)
if definition is None:
raise WorkflowNotFoundError(workflow_id)
counter_seen = int(definition.next_version_number or 1)
# Legacy rows got the column auto-added with server_default 1
# while versions already exist — never allocate below the max.
latest = (
await session.execute(
select(func.max(WorkflowVersionRow.version_number)).where(
WorkflowVersionRow.workflow_id == workflow_id
)
)
).scalar() or 0
number = max(counter_seen, int(latest) + 1)
claimed = await session.execute(
update(WorkflowDefinitionRow)
.where(
WorkflowDefinitionRow.id == workflow_id,
WorkflowDefinitionRow.next_version_number == counter_seen,
)
.values(next_version_number=number + 1)
)
if claimed.rowcount == 0:
# Another publish landed first — release nothing (the
# counter bump never happened) and retry with fresh state.
await session.rollback()
last_error = None
continue
version = WorkflowVersionRow(
id=f"wv_{uuid.uuid4().hex[:12]}",
workflow_id=workflow_id,
version_number=number,
graph_json=_dump_graph(graph),
graph_hash=graph_hash,
input_schema=graph.get("inputSchema") or graph.get("input_schema") or {},
output_schema=graph.get("outputSchema") or graph.get("output_schema") or {},
published_by=published_by,
change_note=change_note or "",
)
definition.status = "active"
definition.updated_at = datetime.now(UTC)
if canvas_schema_hash:
definition.published_canvas_schema_hash = canvas_schema_hash[:64]
session.add(version)
await session.flush()
out = _version_dict(version, include_graph=True)
await session.commit()
return out
except IntegrityError as exc: # pragma: no cover - CAS makes this near-impossible
last_error = exc
continue
if last_error is not None:
raise last_error
raise RuntimeError(f"failed to allocate a workflow version number for {workflow_id}")
async def list_versions(self, workflow_id: str) -> list[dict[str, Any]]:
stmt = (
select(WorkflowVersionRow)
.where(WorkflowVersionRow.workflow_id == workflow_id)
.order_by(WorkflowVersionRow.version_number.desc())
)
async with self._sf() as session:
result = await session.execute(stmt)
return [_version_dict(row, include_graph=False) for row in result.scalars()]
async def get_version(self, version_id: str) -> dict[str, Any] | None:
async with self._sf() as session:
row = await session.get(WorkflowVersionRow, version_id)
return _version_dict(row, include_graph=True) if row is not None else None
async def list_published_summaries(self, *, exclude_workflow_id: str | None = None) -> list[dict[str, Any]]:
stmt = (
select(WorkflowVersionRow, WorkflowDefinitionRow)
.join(WorkflowDefinitionRow, WorkflowDefinitionRow.id == WorkflowVersionRow.workflow_id)
.where(WorkflowDefinitionRow.status != "archived")
.order_by(WorkflowVersionRow.published_at.desc())
)
if exclude_workflow_id:
stmt = stmt.where(WorkflowVersionRow.workflow_id != exclude_workflow_id)
async with self._sf() as session:
result = await session.execute(stmt)
out: list[dict[str, Any]] = []
seen: set[str] = set()
for version, definition in result.all():
# Latest version per workflow only.
if version.workflow_id in seen:
continue
seen.add(version.workflow_id)
out.append(
{
"workflow_id": definition.id,
"name": definition.name,
"version_id": version.id,
"version_number": version.version_number,
"graph_hash": version.graph_hash,
"published_at": _iso(version.published_at),
}
)
return out