386 lines
17 KiB
Python
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
|