deerflow-code/offline-backend-20260512/backend/app/gateway/workflow_dispatcher.py
2026-09-07 18:24:55 +08:00

152 lines
5.3 KiB
Python

"""Lease-based dispatcher for workflow runs.
A single background loop scans for claimable runs (``queued``, or ``running``
with an expired lease) and hands each successfully claimed run to the executor.
Because the claim is an atomic conditional UPDATE, several workers can run this
loop concurrently and a run that dies with its worker is picked up by whoever
notices the expired lease first — that is the restart-recovery story.
``nudge()`` is called by the start/resume routes so a fresh run does not wait
for the next poll tick.
"""
from __future__ import annotations
import asyncio
import logging
import os
import socket
import uuid
from datetime import UTC, datetime, timedelta
from typing import Any
from app.gateway.workflow_executor import LeaseParams, WorkflowRunExecutor
from deerflow.config.workflow_config import WorkflowConfig
from deerflow.persistence.workflow_events.base import WorkflowEventStore
from deerflow.persistence.workflow_runs.base import WorkflowRunStore
from deerflow.workflows.runtime import PersistingWorkflowEventSink
logger = logging.getLogger(__name__)
_POLL_INTERVAL_SECONDS = 2.0
_CLAIM_BATCH = 5
def _worker_id() -> str:
return f"{socket.gethostname()[:24]}-{os.getpid()}-{uuid.uuid4().hex[:6]}"
class WorkflowRunDispatcher:
def __init__(
self,
store: WorkflowRunStore,
executor: WorkflowRunExecutor,
*,
event_store: WorkflowEventStore | None = None,
live_publisher: Any = None,
config: WorkflowConfig | None = None,
) -> None:
self._store = store
self._executor = executor
self._events = event_store
self._publish = live_publisher
self._config = config or WorkflowConfig()
self._worker_id = _worker_id()
self._wake = asyncio.Event()
self._task: asyncio.Task | None = None
self._stopping = False
@property
def worker_id(self) -> str:
return self._worker_id
def start(self) -> None:
if self._task is not None and not self._task.done():
return
self._stopping = False
self._task = asyncio.create_task(self._loop(), name="workflow-run-dispatcher")
logger.info("workflow dispatcher started (worker=%s)", self._worker_id)
async def stop(self) -> None:
self._stopping = True
self._wake.set()
if self._task is not None:
self._task.cancel()
try:
await self._task
except (asyncio.CancelledError, Exception): # noqa: BLE001
pass
self._task = None
await self._executor.drain()
def nudge(self) -> None:
self._wake.set()
async def _loop(self) -> None:
await self._safe_dispatch()
while not self._stopping:
try:
await asyncio.wait_for(self._wake.wait(), timeout=_POLL_INTERVAL_SECONDS)
except TimeoutError:
pass
self._wake.clear()
if self._stopping:
return
await self._safe_dispatch()
async def _safe_dispatch(self) -> None:
try:
await self.dispatch_once()
except asyncio.CancelledError:
raise
except Exception: # noqa: BLE001 - the loop must survive a bad round
logger.exception("workflow dispatch round failed")
async def dispatch_once(self) -> int:
now = datetime.now(UTC)
await self._reap_cancel_requested()
candidates = await self._store.list_claimable(now=now, limit=_CLAIM_BATCH)
if not candidates:
return 0
ttl = max(30, self._config.lease_ttl_seconds)
started = 0
for run_id in candidates:
row = await self._store.claim_run(
run_id,
lease_owner=self._worker_id,
lease_until=now + timedelta(seconds=ttl),
)
if row is None:
continue
if self._executor.start(LeaseParams(run_id=run_id, lease_owner=self._worker_id, attempt=int(row.get("attempt") or 1))):
started += 1
return started
async def _reap_cancel_requested(self) -> None:
"""Finalize cancels whose worker is gone (no live lease to notice them)."""
rows = await self._store.list_runs(status="cancel_requested", limit=20)
now = datetime.now(UTC)
for row in rows:
lease_until = row.get("lease_until")
if lease_until:
try:
if datetime.fromisoformat(str(lease_until)) > now:
continue # its worker is alive and will handle the cancel
except ValueError:
pass
finalized = await self._store.finalize_cancel(row["id"])
if finalized is not None and self._events is not None:
# Nobody owns this run, so the dispatcher owes the stream its
# terminal frame — otherwise open SSE connections hang.
sink = PersistingWorkflowEventSink(
self._events,
run_id=row["id"],
workflow_id=row["workflow_id"],
version_id=row["workflow_version_id"],
publish=self._publish,
)
await sink.emit("run.cancelled", data={"reapedBy": self._worker_id})
__all__ = ["WorkflowRunDispatcher"]