152 lines
5.3 KiB
Python
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"]
|