Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
162 changes: 162 additions & 0 deletions backend/app/core/workflow_runs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
"""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,
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():
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:
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)}))
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()
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)
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()
76 changes: 48 additions & 28 deletions backend/app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -182,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...")
Expand Down Expand Up @@ -421,8 +423,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"}

Expand Down Expand Up @@ -1201,8 +1204,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"}

Expand Down Expand Up @@ -1268,15 +1272,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}")
Expand All @@ -1290,26 +1297,37 @@ async def websocket_run_workflow(websocket: WebSocket) -> None:
return

await websocket.accept()
run_id = None
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), "run_id": run_id})
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,
Expand Down Expand Up @@ -1387,6 +1405,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)
Expand Down
5 changes: 5 additions & 0 deletions backend/app/schemas/task.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
8 changes: 8 additions & 0 deletions backend/app/storage/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
18 changes: 18 additions & 0 deletions backend/app/storage/task_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
8 changes: 8 additions & 0 deletions backend/tests/test_task_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading
Loading