From 76ff2eed0161d30a04e189b3af5d4fb532045d3f Mon Sep 17 00:00:00 2001 From: Berry Wahlberg <40695099+BerryUIKI@users.noreply.github.com> Date: Fri, 9 Oct 2026 01:07:42 +0800 Subject: [PATCH 1/4] feat(workflow): persist ordered events for resumable subscriptions --- backend/app/storage/db.py | 8 ++++++++ backend/app/storage/task_store.py | 18 ++++++++++++++++++ backend/tests/test_task_store.py | 8 ++++++++ 3 files changed, 34 insertions(+) diff --git a/backend/app/storage/db.py b/backend/app/storage/db.py index b205905..a3cee71 100644 --- a/backend/app/storage/db.py +++ b/backend/app/storage/db.py @@ -111,6 +111,14 @@ async def _init_schema(self, conn: aiosqlite.Connection) -> None: CREATE INDEX IF NOT EXISTS idx_tasks_run ON tasks(run_id); + CREATE TABLE IF NOT EXISTS workflow_events ( + run_id TEXT NOT NULL, + sequence INTEGER NOT NULL, + event_json TEXT NOT NULL, + PRIMARY KEY (run_id, sequence), + FOREIGN KEY (run_id) REFERENCES runs(id) ON DELETE CASCADE + ); + CREATE TABLE IF NOT EXISTS cache_entries ( node_hash TEXT PRIMARY KEY, output_json TEXT NOT NULL, diff --git a/backend/app/storage/task_store.py b/backend/app/storage/task_store.py index 27f9744..88abfea 100644 --- a/backend/app/storage/task_store.py +++ b/backend/app/storage/task_store.py @@ -97,6 +97,24 @@ async def reconcile_interrupted(self) -> int: await conn.commit() return cursor.rowcount + async def append_event(self, run_id: str, event: dict[str, Any]) -> None: + conn = await self.manager.get_connection() + await conn.execute( + """INSERT INTO workflow_events (run_id, sequence, event_json) + SELECT ?, COALESCE(MAX(sequence), 0) + 1, ? FROM workflow_events WHERE run_id = ?""", + (run_id, json.dumps(redact_submission({**event, "run_id": run_id})), run_id), + ) + await conn.commit() + + async def get_events(self, run_id: str, after_sequence: int = 0) -> list[dict[str, Any]]: + conn = await self.manager.get_connection() + async with conn.execute( + "SELECT sequence, event_json FROM workflow_events WHERE run_id = ? AND sequence > ? ORDER BY sequence LIMIT 500", + (run_id, after_sequence), + ) as cursor: + rows = await cursor.fetchall() + return [{**json.loads(row["event_json"]), "sequence": row["sequence"]} for row in rows] + task_store = TaskStore() execution_task: ContextVar[tuple[TaskStore, TaskRecord] | None] = ContextVar("execution_task", default=None) diff --git a/backend/tests/test_task_store.py b/backend/tests/test_task_store.py index 89b6673..31ba2a9 100644 --- a/backend/tests/test_task_store.py +++ b/backend/tests/test_task_store.py @@ -33,5 +33,13 @@ async def test_durable_history_and_restart_recovery(tmp_path: Path) -> None: assert history["remote"].params["nested"] == {"chroma_key": "green"} assert await store.list_tasks(project_id="another") == [] assert "api_key" not in (await store.get_run("remote")).request + await store.append_event("done", {"type": "GRAPH_STARTED"}) + await store.append_event("done", {"type": "GRAPH_FINISHED", "status": "completed"}) + await manager.close() + events = await store.get_events("done", after_sequence=1) + assert len(events) == 1 + assert events[0]["sequence"] == 2 + assert events[0]["run_id"] == "done" + assert events[0]["status"] == "completed" finally: await manager.close() From 0ba00711707ffd7b8a0268bcab926b819ce7fd33 Mon Sep 17 00:00:00 2001 From: Berry Wahlberg <40695099+BerryUIKI@users.noreply.github.com> Date: Fri, 9 Oct 2026 01:09:55 +0800 Subject: [PATCH 2/4] fix(workflow): decouple idempotent execution from subscriptions --- backend/app/core/workflow_runs.py | 147 ++++++++++++++++++++++++++++ backend/app/main.py | 74 ++++++++------ backend/app/schemas/task.py | 5 + backend/tests/test_workflow_runs.py | 51 ++++++++++ 4 files changed, 249 insertions(+), 28 deletions(-) create mode 100644 backend/app/core/workflow_runs.py create mode 100644 backend/tests/test_workflow_runs.py diff --git a/backend/app/core/workflow_runs.py b/backend/app/core/workflow_runs.py new file mode 100644 index 0000000..6bd6d88 --- /dev/null +++ b/backend/app/core/workflow_runs.py @@ -0,0 +1,147 @@ +"""Idempotent workflow submission and durable subscriptions independent of transport.""" + +import asyncio +import json +import uuid +from collections.abc import AsyncIterator, Awaitable, Callable +from typing import Any + +from app.core.task_registry import task_registry +from app.schemas.task import RunRecord, TaskRecord, WorkflowRunRequest +from app.storage.task_store import TERMINAL_STATUSES, TaskStore, execution_task, redact_submission, task_store + + +class WorkflowEventSink: + """Persist events and node outcomes before exposing them to subscribers.""" + + def __init__(self, request: WorkflowRunRequest, store: TaskStore, changed: asyncio.Event, cancel: asyncio.Event) -> None: + self.request = request + self.store = store + self.changed = changed + self.cancel_event = cancel + self.finished = False + self.tasks = {node.id: TaskRecord(id=f"{request.run_id}:{node.id}", run_id=request.run_id, + node_id=node.id, node_type=node.type, params=node.params) for node in request.graph.nodes} + + async def initialize(self) -> None: + for task in self.tasks.values(): + await self.store.save_task(task) + + async def bind_inputs(self, node_id: str, inputs: dict[str, Any]) -> None: + task = self.tasks[node_id] + task.inputs = inputs + execution_task.set((self.store, task)) + await self.store.save_task(task) + + async def send_text(self, data: str) -> None: + event = json.loads(data) + task = self.tasks.get(event.get("node_id")) + if task: + if event["type"] == "NODE_STATUS": + task.status = {"completed": "succeeded", "error": "failed", "idle": "queued"}.get(event["status"], event["status"]) + elif event["type"] == "NODE_OUTPUT": + task.outputs = event["output"] + elif event["type"] == "NODE_ERROR": + task.status = "failed" + task.error = event["message"] + await self.store.save_task(task) + if event["type"] == "GRAPH_STARTED": + await self.store.finish_run(self.request.run_id, "running") + elif event["type"] == "GRAPH_FINISHED": + outcome = {"completed": "succeeded"}.get(event["status"], event["status"]) + for unfinished in self.tasks.values(): + if unfinished.status not in TERMINAL_STATUSES: + unfinished.status = "cancelled" if outcome == "cancelled" else "interrupted" + await self.store.save_task(unfinished) + await self.store.finish_run(self.request.run_id, outcome) + self.finished = True + await self.store.append_event(self.request.run_id, event) + self.changed.set() + + async def close(self) -> None: + """Execution has no ownership of a subscriber's socket.""" + + +WorkflowExecutor = Callable[[WorkflowRunRequest, WorkflowEventSink], Awaitable[None]] + + +class WorkflowRunService: + def __init__(self, store: TaskStore = task_store, concurrency: int = 4) -> None: + self.store = store + self._admission = asyncio.Lock() + self._slots = asyncio.Semaphore(concurrency) + self._workers: dict[str, asyncio.Task[None]] = {} + self._changes: dict[str, asyncio.Event] = {} + self._cancellations: dict[str, asyncio.Event] = {} + + async def submit(self, request: WorkflowRunRequest, executor: WorkflowExecutor) -> str: + request = request.model_copy(deep=True) + request.run_id = request.run_id or str(uuid.uuid4()) + async with self._admission: + existing = await self.store.get_run(request.run_id) + if existing: + if existing.request != redact_submission(request.model_dump()): + raise ValueError("Run ID already belongs to a different immutable submission") + return request.run_id + await self.store.create_run(RunRecord(id=request.run_id, project_id=request.project_id, + target_node_id=request.target_node_id, request=request.model_dump())) + changed = self._changes.setdefault(request.run_id, asyncio.Event()) + cancel = self._cancellations.setdefault(request.run_id, asyncio.Event()) + task_registry.register_task(request.run_id, "workflow_graph", cancel) + sink = WorkflowEventSink(request, self.store, changed, cancel) + await sink.initialize() + self._workers[request.run_id] = asyncio.create_task(self._execute(request, sink, executor)) + return request.run_id + + async def _execute(self, request: WorkflowRunRequest, sink: WorkflowEventSink, executor: WorkflowExecutor) -> None: + try: + async with self._slots: + await executor(request, sink) + if not sink.finished: + await sink.send_text(json.dumps({"type": "GRAPH_FINISHED", "status": "failed", "execution_time_ms": 0})) + except asyncio.CancelledError: + await self.store.finish_run(request.run_id, "interrupted") + raise + except Exception as error: + await sink.send_text(json.dumps({"type": "ERROR", "message": str(error)})) + await sink.send_text(json.dumps({"type": "GRAPH_FINISHED", "status": "failed", "execution_time_ms": 0})) + finally: + task_registry.unregister_task(request.run_id) + self._workers.pop(request.run_id, None) + self._cancellations.pop(request.run_id, None) + sink.changed.set() + + def cancel(self, run_id: str) -> bool: + event = self._cancellations.get(run_id) + if event is None: + return False + event.set() + return True + + async def subscribe(self, run_id: str, after_sequence: int = 0) -> AsyncIterator[dict[str, Any]]: + if await self.store.get_run(run_id) is None: + raise LookupError("Run not found; reconnect never submits a new run") + changed = self._changes.setdefault(run_id, asyncio.Event()) + while True: + changed.clear() + events = await self.store.get_events(run_id, after_sequence) + for event in events: + after_sequence = event["sequence"] + yield event + if event["type"] == "GRAPH_FINISHED": + return + if len(events) == 500: + continue + run = await self.store.get_run(run_id) + if run.status in TERMINAL_STATUSES: + # A restart may have reconciled the run before a terminal event was recorded. + yield {"type": "GRAPH_FINISHED", "run_id": run_id, "status": "completed" if run.status in {"succeeded", "cached"} else run.status, + "execution_time_ms": 0, "sequence": after_sequence + 1} + return + try: + await asyncio.wait_for(changed.wait(), timeout=30) + except asyncio.TimeoutError: + continue + + +workflow_runs = WorkflowRunService() diff --git a/backend/app/main.py b/backend/app/main.py index 465533c..fd9eeea 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -42,6 +42,7 @@ ) from app.core.session import session_manager, ALLOWED_ORIGINS from app.core.task_registry import task_registry +from app.core.workflow_runs import WorkflowEventSink, workflow_runs from app.core.execution_contract import ExecutionContractError, validate_execution_contract from app.core.media_validator import ( validate_and_inspect_media, @@ -118,7 +119,7 @@ ) from app.schemas.node import NodeDefinition from app.schemas.project import Project, ProjectCreate, ProjectUpdate, AssetRecord -from app.schemas.task import WorkflowRunRequest +from app.schemas.task import WorkflowRunRequest, WorkflowSubscription from app.schemas.task import TaskRecord from app.storage.task_store import task_store from app.schemas.workflow import ExecutionPlan, PlannedNodeStep, WorkflowGraph @@ -421,8 +422,9 @@ async def generation_task_history(project_id: Optional[str] = None) -> List[Task @app.post("/api/v1/workflow/cancel/{run_id}") async def cancel_workflow_run(run_id: str) -> Dict[str, Any]: """Cancel an active DAG workflow run (R14).""" - if run_id in active_cancellations: - active_cancellations[run_id].set() + if workflow_runs.cancel(run_id) or run_id in active_cancellations: + if run_id in active_cancellations: + active_cancellations[run_id].set() return {"run_id": run_id, "success": True, "status": "cancelled"} return {"run_id": run_id, "success": False, "status": "not_found"} @@ -1201,8 +1203,9 @@ async def get_asset_content(asset_id: str) -> FileResponse: @app.post("/api/v1/workflow/cancel/{run_id}") async def cancel_workflow_run(run_id: str) -> dict[str, Any]: - if run_id in active_cancellations: - active_cancellations[run_id].set() + if workflow_runs.cancel(run_id) or run_id in active_cancellations: + if run_id in active_cancellations: + active_cancellations[run_id].set() return {"success": True, "run_id": run_id, "status": "cancel-requested"} return {"success": False, "message": f"Run '{run_id}' not found or already completed"} @@ -1268,15 +1271,18 @@ async def generate_plan(graph: WorkflowGraph, target_node: str | None = None) -> ) +@app.post("/api/v1/workflow/submit") +async def submit_workflow(request: WorkflowRunRequest) -> dict[str, str]: + """Admit an immutable run once; subscribers never resubmit inference.""" + try: + run_id = await workflow_runs.submit(request, _execute_workflow_request) + except ValueError as error: + raise HTTPException(status_code=409, detail=str(error)) + return {"run_id": run_id, "status": "accepted"} + + @app.websocket("/ws/workflow/run") async def websocket_run_workflow(websocket: WebSocket) -> None: - """ - WebSocket endpoint for real-time workflow execution with: - - Targeted single-node execution - - Port-aware semantic caching - - Strict failure boundaries (aborts dependent child nodes) - - Run cancellation support - """ origin = websocket.headers.get("origin") if origin and not session_manager.is_origin_allowed(origin): logger.warning(f"Rejecting WebSocket handshake from unauthorized origin: {origin}") @@ -1291,25 +1297,35 @@ async def websocket_run_workflow(websocket: WebSocket) -> None: await websocket.accept() try: - raw = await websocket.receive_text() - parsed = json.loads(raw) - # Parse either WorkflowRunRequest or raw WorkflowGraph - if "graph" in parsed: - run_request = WorkflowRunRequest.model_validate(parsed) - graph = run_request.graph - target_node = run_request.target_node_id - run_id = run_request.run_id or str(uuid.uuid4()) + parsed = json.loads(await websocket.receive_text()) + after_sequence = 0 + if parsed.get("type") == "SUBSCRIBE": + subscription = WorkflowSubscription.model_validate(parsed) + run_id = subscription.run_id + after_sequence = subscription.after_sequence else: - graph = WorkflowGraph.model_validate(parsed) - target_node = None - run_id = str(uuid.uuid4()) - except Exception as e: - await websocket.send_text(json.dumps({"type": "ERROR", "message": f"Invalid graph payload: {e}"})) - await websocket.close() - return + request = WorkflowRunRequest.model_validate(parsed if "graph" in parsed else {"graph": parsed}) + run_id = await workflow_runs.submit(request, _execute_workflow_request) + async for event in workflow_runs.subscribe(run_id, after_sequence): + await websocket.send_json(event) + except WebSocketDisconnect: + logger.info("Workflow subscriber disconnected; execution continues independently") + except (ValueError, LookupError) as error: + await websocket.send_json({"type": "ERROR", "message": str(error)}) + finally: + try: + await websocket.close() + except RuntimeError: + pass + +async def _execute_workflow_request(request: WorkflowRunRequest, websocket: WorkflowEventSink) -> None: + """Execute into a durable event sink without owning any client connection.""" + graph = request.graph + target_node = request.target_node_id + run_id = request.run_id # Register cancellation token - cancel_event = asyncio.Event() + cancel_event = websocket.cancel_event active_cancellations[run_id] = cancel_event task_registry.register_task( task_id=run_id, @@ -1387,6 +1403,8 @@ async def websocket_run_workflow(websocket: WebSocket) -> None: inputs[edge.target_handle] = val bindings.append((edge.target_handle, compute_content_hash(val), edge.source_handle)) + await websocket.bind_inputs(node.id, inputs) + # Compute port-aware semantic hash if bindings: node_hash = compute_semantic_node_hash(node.type, node.params, bindings) diff --git a/backend/app/schemas/task.py b/backend/app/schemas/task.py index de9ad99..c100280 100644 --- a/backend/app/schemas/task.py +++ b/backend/app/schemas/task.py @@ -50,3 +50,8 @@ class WorkflowRunRequest(BaseModel): target_node_id: Optional[str] = None run_id: Optional[str] = None project_id: Optional[str] = None + + +class WorkflowSubscription(BaseModel): + run_id: str + after_sequence: int = Field(default=0, ge=0) diff --git a/backend/tests/test_workflow_runs.py b/backend/tests/test_workflow_runs.py new file mode 100644 index 0000000..d4673e6 --- /dev/null +++ b/backend/tests/test_workflow_runs.py @@ -0,0 +1,51 @@ +"""A disconnected subscriber must never trigger a second inference call.""" + +import asyncio +import json +from pathlib import Path + +import pytest + +from app.core.workflow_runs import WorkflowRunService +from app.schemas.task import WorkflowRunRequest +from app.storage.db import DatabaseManager +from app.storage.task_store import TaskStore + + +@pytest.mark.asyncio +async def test_reconnect_and_duplicate_submission_execute_once(tmp_path: Path) -> None: + manager = DatabaseManager(tmp_path / "runs.db") + service = WorkflowRunService(TaskStore(manager)) + request = WorkflowRunRequest(run_id="one-run", graph={"nodes": []}) + calls = 0 + release = asyncio.Event() + + async def execute(request, sink) -> None: + nonlocal calls + calls += 1 + await sink.send_text(json.dumps({"type": "GRAPH_STARTED", "total_nodes": 0, "cached_nodes": 0})) + await release.wait() + await sink.send_text(json.dumps({"type": "GRAPH_FINISHED", "status": "completed", "execution_time_ms": 1})) + + try: + await service.submit(request, execute) + first = service.subscribe("one-run") + started = await anext(first) + await first.aclose() + assert await service.submit(request, execute) == "one-run" + release.set() + resumed = [event async for event in service.subscribe("one-run", started["sequence"])] + assert resumed[-1]["status"] == "completed" + assert calls == 1 + with pytest.raises(ValueError, match="different immutable"): + await service.submit(WorkflowRunRequest(run_id="one-run", graph={"nodes": [{"id": "changed", "type": "input.text"}]}), execute) + await asyncio.gather(*service._workers.values()) + await manager.close() + restarted = WorkflowRunService(TaskStore(manager)) + replay = [event async for event in restarted.subscribe("one-run", started["sequence"])] + assert replay[-1]["status"] == "completed" + assert calls == 1 + finally: + release.set() + await asyncio.gather(*service._workers.values()) + await manager.close() From 79248761e25257586c483e90da27935b8c7ad5a5 Mon Sep 17 00:00:00 2001 From: Berry Wahlberg <40695099+BerryUIKI@users.noreply.github.com> Date: Fri, 9 Oct 2026 01:11:36 +0800 Subject: [PATCH 3/4] fix(canvas): reconnect to existing runs and reject obsolete events --- backend/app/main.py | 3 +- frontend/src/stores/useCanvasStore.ts | 70 +++++++++---------- .../src/tests/workflowSubscriptions.test.ts | 63 +++++++++++++++++ 3 files changed, 97 insertions(+), 39 deletions(-) create mode 100644 frontend/src/tests/workflowSubscriptions.test.ts diff --git a/backend/app/main.py b/backend/app/main.py index fd9eeea..be93b02 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -1296,6 +1296,7 @@ async def websocket_run_workflow(websocket: WebSocket) -> None: return await websocket.accept() + run_id = None try: parsed = json.loads(await websocket.receive_text()) after_sequence = 0 @@ -1311,7 +1312,7 @@ async def websocket_run_workflow(websocket: WebSocket) -> None: except WebSocketDisconnect: logger.info("Workflow subscriber disconnected; execution continues independently") except (ValueError, LookupError) as error: - await websocket.send_json({"type": "ERROR", "message": str(error)}) + await websocket.send_json({"type": "ERROR", "message": str(error), "run_id": run_id}) finally: try: await websocket.close() diff --git a/frontend/src/stores/useCanvasStore.ts b/frontend/src/stores/useCanvasStore.ts index 0693eac..d5034cc 100644 --- a/frontend/src/stores/useCanvasStore.ts +++ b/frontend/src/stores/useCanvasStore.ts @@ -13,6 +13,15 @@ import { CustomNodeData, ExecutionStatus, NodeDefinition, NodeOutputValue, NodeP import { CanvasNodeData, ImageCardData } from '../types/creative'; import { GenerationHistoryItem, ProjectCanvasData, ProjectViewport } from '../types/project'; +interface WorkflowStreamEvent { + type: string; + run_id?: string; + sequence?: number; + node_id: string; + status: ExecutionStatus; + output: Record; +} + const isImageCardNode = (n: Node): n is Node => n.type === 'imageCard'; const isWorkflowNode = (n: Node): n is Node => n.type === 'workflowNode' || !n.type; @@ -490,51 +499,33 @@ export const useCanvasStore = create((set, get) => ({ let reconnectAttempts = 0; const maxReconnectAttempts = 3; - let heartbeatTimer: any = null; let isTerminated = false; - - const stopHeartbeat = () => { - if (heartbeatTimer) { - clearInterval(heartbeatTimer); - heartbeatTimer = null; - } - }; + let lastSequence = 0; + let socket: WebSocket | null = null; const cleanupExecution = () => { isTerminated = true; - stopHeartbeat(); - set({ isExecuting: false, currentRunId: null }); + socket?.close(); + if (get().currentRunId === runId) set({ isExecuting: false, currentRunId: null }); }; const connect = () => { - if (isTerminated || !get().isExecuting) return; + if (isTerminated || get().currentRunId !== runId) return; try { const ws = new WebSocket(wsUrl); + socket = ws; ws.onopen = () => { - reconnectAttempts = 0; - ws.send(JSON.stringify(runPayload)); - - // Start client-side keepalive ping to maintain connection through proxies - stopHeartbeat(); - heartbeatTimer = setInterval(() => { - if (ws.readyState === WebSocket.OPEN) { - try { - ws.send(JSON.stringify({ type: 'PING' })); - } catch { - // ignore send error - } - } - }, 15000); + ws.send(JSON.stringify({ type: 'SUBSCRIBE', run_id: runId, after_sequence: lastSequence })); }; ws.onmessage = (event) => { try { - const msg = JSON.parse(event.data); - if (msg.type === 'PONG') { - return; - } + const msg: WorkflowStreamEvent = JSON.parse(event.data); + if (get().currentRunId !== runId || msg.run_id !== runId) return; + if (msg.sequence !== undefined && msg.sequence <= lastSequence) return; + if (msg.sequence !== undefined) lastSequence = msg.sequence; if (msg.type === 'NODE_STATUS') { setNodeStatus(msg.node_id, msg.status); } else if (msg.type === 'NODE_OUTPUT') { @@ -547,28 +538,23 @@ export const useCanvasStore = create((set, get) => ({ if (msg.node_id) { setNodeStatus(msg.node_id, 'error'); } - cleanupExecution(); + if (msg.type === 'ERROR') cleanupExecution(); } } catch { // ignore parsing error } }; - ws.onerror = () => { - stopHeartbeat(); - }; - ws.onclose = () => { - stopHeartbeat(); // If clean termination occurred or run concluded, do not reconnect - if (isTerminated || !get().isExecuting) return; + if (isTerminated || get().currentRunId !== runId) return; // Attempt reconnection if abruptly disconnected if (reconnectAttempts < maxReconnectAttempts) { reconnectAttempts += 1; const delay = Math.min(1000 * Math.pow(2, reconnectAttempts - 1), 4000); setTimeout(() => { - if (get().isExecuting && !isTerminated) { + if (get().currentRunId === runId && !isTerminated) { connect(); } }, delay); @@ -581,6 +567,14 @@ export const useCanvasStore = create((set, get) => ({ } }; - connect(); + try { + const response = await fetch('/api/v1/workflow/submit', { + method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify(runPayload), + }); + if (!response.ok) throw new Error(`Workflow submission failed (${response.status})`); + connect(); + } catch { + cleanupExecution(); + } }, })); diff --git a/frontend/src/tests/workflowSubscriptions.test.ts b/frontend/src/tests/workflowSubscriptions.test.ts new file mode 100644 index 0000000..51e5ca0 --- /dev/null +++ b/frontend/src/tests/workflowSubscriptions.test.ts @@ -0,0 +1,63 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { useCanvasStore } from '../stores/useCanvasStore'; + +class FakeSocket { + static instances: FakeSocket[] = []; + sent: Record[] = []; + onopen: (() => void) | null = null; + onclose: (() => void) | null = null; + onmessage: ((event: MessageEvent) => void) | null = null; + constructor() { FakeSocket.instances.push(this); } + send(value: string) { this.sent.push(JSON.parse(value)); } + close() { this.onclose?.(); } + emit(value: Record) { this.onmessage?.({ data: JSON.stringify(value) } as MessageEvent); } +} + +describe('Workflow subscriptions', () => { + beforeEach(() => { + vi.useFakeTimers(); + FakeSocket.instances = []; + vi.stubGlobal('window', { location: { protocol: 'http:', host: 'localhost:8000' } }); + vi.stubGlobal('WebSocket', FakeSocket); + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ ok: true })); + useCanvasStore.setState({ nodes: [{ id: 'text', type: 'workflowNode', position: { x: 0, y: 0 }, + data: { definition: { type: 'input.text', title: 'Text', category: 'input', description: '', inputs: [], outputs: [], parameters: [] }, params: { value: 'hello' }, status: 'idle' } }], + edges: [], currentRunId: null, isExecuting: false }); + }); + afterEach(() => { vi.useRealTimers(); vi.unstubAllGlobals(); }); + + it('submits once and reconnects with the last observed sequence', async () => { + await useCanvasStore.getState().runWorkflow(); + const runId = useCanvasStore.getState().currentRunId; + const first = FakeSocket.instances[0]; + first.onopen?.(); + expect(first.sent[0]).toEqual({ type: 'SUBSCRIBE', run_id: runId, after_sequence: 0 }); + first.emit({ type: 'NODE_STATUS', run_id: runId, sequence: 3, node_id: 'text', status: 'running' }); + first.close(); + await vi.advanceTimersByTimeAsync(1000); + const resumed = FakeSocket.instances[1]; + resumed.onopen?.(); + expect(resumed.sent[0]).toEqual({ type: 'SUBSCRIBE', run_id: runId, after_sequence: 3 }); + expect(fetch).toHaveBeenCalledTimes(1); + resumed.emit({ type: 'GRAPH_FINISHED', run_id: runId, sequence: 4, status: 'completed' }); + expect(useCanvasStore.getState().isExecuting).toBe(false); + await vi.advanceTimersByTimeAsync(30000); + expect(resumed.sent).toHaveLength(1); // No unmatched application PING protocol. + }); + + it('ignores stale events and cleanup from an obsolete run', async () => { + await useCanvasStore.getState().runWorkflow(); + const oldId = useCanvasStore.getState().currentRunId; + const old = FakeSocket.instances[0]; + await vi.advanceTimersByTimeAsync(1); + await useCanvasStore.getState().runWorkflow(); + const currentId = useCanvasStore.getState().currentRunId; + expect(currentId).not.toBe(oldId); + old.emit({ type: 'NODE_OUTPUT', run_id: oldId, sequence: 1, node_id: 'text', output: { text: 'obsolete' } }); + old.emit({ type: 'GRAPH_FINISHED', run_id: oldId, sequence: 2, status: 'completed' }); + old.close(); + expect(useCanvasStore.getState().currentRunId).toBe(currentId); + expect(useCanvasStore.getState().nodes[0].data.output).toBeUndefined(); + expect(useCanvasStore.getState().isExecuting).toBe(true); + }); +}); From b6097ab9382556a712f93f0eeb9f54c1e37da7dc Mon Sep 17 00:00:00 2001 From: Berry Wahlberg <40695099+BerryUIKI@users.noreply.github.com> Date: Fri, 9 Oct 2026 01:19:14 +0800 Subject: [PATCH 4/4] test(workflow): verify disconnect replay and durable execution outcomes --- backend/app/core/workflow_runs.py | 19 +++++++++-- backend/app/main.py | 1 + backend/tests/test_workflow_runs.py | 50 +++++++++++++++++++++++++++++ 3 files changed, 68 insertions(+), 2 deletions(-) diff --git a/backend/app/core/workflow_runs.py b/backend/app/core/workflow_runs.py index 6bd6d88..5e30ed1 100644 --- a/backend/app/core/workflow_runs.py +++ b/backend/app/core/workflow_runs.py @@ -21,7 +21,8 @@ def __init__(self, request: WorkflowRunRequest, store: TaskStore, changed: async self.cancel_event = cancel self.finished = False self.tasks = {node.id: TaskRecord(id=f"{request.run_id}:{node.id}", run_id=request.run_id, - node_id=node.id, node_type=node.type, params=node.params) for node in request.graph.nodes} + node_id=node.id, node_type=node.type, params=node.params, + metadata={"engine": "cloud" if node.type in {"text.llm", "image.generate"} else "local"}) for node in request.graph.nodes} async def initialize(self) -> None: for task in self.tasks.values(): @@ -100,7 +101,14 @@ async def _execute(self, request: WorkflowRunRequest, sink: WorkflowEventSink, e if not sink.finished: await sink.send_text(json.dumps({"type": "GRAPH_FINISHED", "status": "failed", "execution_time_ms": 0})) except asyncio.CancelledError: - await self.store.finish_run(request.run_id, "interrupted") + uncertain = False + for task in sink.tasks.values(): + if task.status not in TERMINAL_STATUSES: + task.status = "outcome-unknown" if task.metadata.get("engine") == "cloud" else "interrupted" + uncertain |= task.status == "outcome-unknown" + task.error = "Application stopped before completion. Check the engine/provider before retrying." + await self.store.save_task(task) + await self.store.finish_run(request.run_id, "outcome-unknown" if uncertain else "interrupted") raise except Exception as error: await sink.send_text(json.dumps({"type": "ERROR", "message": str(error)})) @@ -110,6 +118,13 @@ async def _execute(self, request: WorkflowRunRequest, sink: WorkflowEventSink, e self._workers.pop(request.run_id, None) self._cancellations.pop(request.run_id, None) sink.changed.set() + self._changes.pop(request.run_id, None) + + async def shutdown(self) -> None: + workers = list(self._workers.values()) + for worker in workers: + worker.cancel() + await asyncio.gather(*workers, return_exceptions=True) def cancel(self, run_id: str) -> bool: event = self._cancellations.get(run_id) diff --git a/backend/app/main.py b/backend/app/main.py index be93b02..2c2c120 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -183,6 +183,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]: yield logger.info("Berry AI Studio API shutting down...") + await workflow_runs.shutdown() try: if llama_server_supervisor.is_running(): logger.info("Stopping embedded llama-server process...") diff --git a/backend/tests/test_workflow_runs.py b/backend/tests/test_workflow_runs.py index d4673e6..4d981db 100644 --- a/backend/tests/test_workflow_runs.py +++ b/backend/tests/test_workflow_runs.py @@ -2,14 +2,18 @@ import asyncio import json +import threading from pathlib import Path import pytest +from fastapi.testclient import TestClient +from unittest.mock import patch from app.core.workflow_runs import WorkflowRunService from app.schemas.task import WorkflowRunRequest from app.storage.db import DatabaseManager from app.storage.task_store import TaskStore +from app.schemas.events import NodeOutputEvent, NodeStatusEvent @pytest.mark.asyncio @@ -49,3 +53,49 @@ async def execute(request, sink) -> None: release.set() await asyncio.gather(*service._workers.values()) await manager.close() + + +def test_websocket_disconnect_resumes_single_inference(tmp_path: Path) -> None: + from app.main import app + from app.core.cache import CacheStore + + manager = DatabaseManager(tmp_path / "socket.db") + store = TaskStore(manager) + service = WorkflowRunService(store) + release = threading.Event() + calls = 0 + + async def controlled_runner(node_id: str, params: dict): + nonlocal calls + calls += 1 + yield NodeStatusEvent(node_id=node_id, status="running") + await asyncio.to_thread(release.wait, 3) + yield NodeOutputEvent(node_id=node_id, output={"text": "one output"}) + yield NodeStatusEvent(node_id=node_id, status="completed") + + with patch("app.main.workflow_runs", service), patch("app.main.task_store", store), patch( + "app.main.run_input_text_node", controlled_runner + ), patch("app.main.cache_store", CacheStore(manager=manager)), patch("app.main.llama_server_supervisor.is_installed", return_value=False): + with TestClient(app) as client: + submitted = client.post("/api/v1/workflow/submit", json={"run_id": "socket-run", "graph": {"nodes": [{"id": "text", "type": "input.text"}]}}) + assert submitted.status_code == 200 + with client.websocket_connect("/ws/workflow/run") as first: + first.send_json({"type": "SUBSCRIBE", "run_id": "socket-run"}) + started = first.receive_json() + running = first.receive_json() + assert running["status"] == "running" + release.set() + with client.websocket_connect("/ws/workflow/run") as resumed: + resumed.send_json({"type": "SUBSCRIBE", "run_id": "socket-run", "after_sequence": running["sequence"]}) + received = [] + while not received or received[-1]["type"] != "GRAPH_FINISHED": + received.append(resumed.receive_json()) + assert calls == 1 + assert any(event["type"] == "NODE_OUTPUT" for event in received) + assert received[-1]["status"] == "completed" + assert started["run_id"] == "socket-run" + history = client.get("/api/v1/tasks/history").json() + assert history[0]["outputs"] == {"text": "one output"} + assert history[0]["status"] == "succeeded" + assert client.post("/api/v1/workflow/submit", json={"run_id": "socket-run", "graph": {"nodes": []}}).status_code == 409 + asyncio.run(manager.close())