132 lines
4.4 KiB
Python
132 lines
4.4 KiB
Python
"""Lease-based dispatcher. ``start()`` is only for ``worker_enabled`` workers.
|
|
|
|
Gateway-only processes keep ``worker_enabled=false`` and never call ``start()``.
|
|
Tests drive ``dispatch_once()`` without mounting a background loop.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
import os
|
|
import socket
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
|
|
from app.report_collaboration.execution.executor import LeaseParams, ReportCollaborationExecutor
|
|
from deerflow.config.report_collaboration_config import ReportCollaborationConfig
|
|
from deerflow.persistence.report_collaboration.base import ReportCollaborationStore
|
|
|
|
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 ReportCollaborationDispatcher:
|
|
def __init__(
|
|
self,
|
|
store: ReportCollaborationStore,
|
|
executor: ReportCollaborationExecutor,
|
|
*,
|
|
config: ReportCollaborationConfig | None = None,
|
|
worker_id: str | None = None,
|
|
) -> None:
|
|
self._store = store
|
|
self._executor = executor
|
|
self._config = config or ReportCollaborationConfig()
|
|
self._worker_id = worker_id or _worker_id()
|
|
self._wake = asyncio.Event()
|
|
self._task: asyncio.Task[None] | None = None
|
|
self._stopping = False
|
|
|
|
@property
|
|
def worker_id(self) -> str:
|
|
return self._worker_id
|
|
|
|
def start(self) -> None:
|
|
if not self._config.worker_enabled:
|
|
logger.warning("report collaboration dispatcher.start ignored because worker_enabled is false")
|
|
return
|
|
if self._task is not None and not self._task.done():
|
|
return
|
|
self._stopping = False
|
|
self._task = asyncio.create_task(self._loop(), name="report-collaboration-dispatcher")
|
|
logger.info("report collaboration 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
|
|
logger.exception("report collaboration dispatch round failed")
|
|
|
|
async def dispatch_once(self) -> int:
|
|
await self._reap_cancel_requested()
|
|
now = datetime.now(UTC)
|
|
candidates = await self._store.list_claimable(now=now, limit=_CLAIM_BATCH)
|
|
if not candidates:
|
|
return 0
|
|
ttl = max(30, int(self._config.lease_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:
|
|
for run_id in await self._store.list_cancel_pending():
|
|
record = await self._store.get_run_record(run_id)
|
|
if record is None:
|
|
continue
|
|
owner = record.get("lease_owner")
|
|
if owner and self._executor.owns(run_id):
|
|
continue
|
|
finalized = await self._store.finalize_cancel(run_id)
|
|
if finalized is not None:
|
|
from app.report_collaboration.execution.recovery import cancel_in_flight_nodes
|
|
|
|
await cancel_in_flight_nodes(self._store, run_id)
|
|
|
|
|
|
__all__ = ["ReportCollaborationDispatcher"]
|