"""Worker: dequeue tasks and run executions. See specs/05-protocols.md sections 1-8."""
from __future__ import annotations
import asyncio
import importlib
import inspect
from collections.abc import Callable
from contextlib import suppress
from dataclasses import dataclass
from datetime import timedelta
from typing import Any
from flowli.codec import digest, unstructure
from flowli.domain import (
ROOT_FID,
Actor,
Cancelled,
ClaimedTask,
Condition,
Eid,
Entry,
Execution,
Failed,
FrameKind,
FrameRef,
Lease,
LeaseLost,
MemoTable,
Message,
NondeterminismError,
Ports,
Provenance,
Site,
Task,
TaskKind,
Timer,
WorkflowNotRegistered,
parse_eid,
)
from flowli.log import bound, get_logger
from .context import Context, Raised, Returned, Suspended, Wait
from .engine import Engine, UnknownExecution
log = get_logger("flowli.worker")
@dataclass(frozen=True, slots=True)
class Done:
pass
@dataclass(frozen=True, slots=True)
class Retry:
delay: timedelta
TaskResult = Done | Retry
class _RenewLoop:
"""Renews a lease every ttl/3 seconds. On LeaseLost, calls on_lost once."""
def __init__(self, lease: Lease, ttl: float, on_lost: Callable[[], None]) -> None:
self._lease = lease
self._interval = ttl / 3
self._on_lost = on_lost
self.lost = False
self._task: asyncio.Task[None] | None = None # pragma: no mutate
async def __aenter__(self) -> _RenewLoop:
self._task = asyncio.create_task(self._run())
return self
async def __aexit__(self, *exc: object) -> None:
if self._task is not None:
self._task.cancel()
with suppress(asyncio.CancelledError):
await self._task
async def _run(self) -> None:
while True:
await asyncio.sleep(self._interval)
try:
await self._lease.renew()
except LeaseLost:
self.lost = True
self._on_lost()
return
[docs]
class Worker:
def __init__(self, engine: Engine, queues: list[str], worker_id: str) -> None:
self.engine = engine
self.queues = list(queues)
self.worker_id = worker_id
self.actor = Actor.worker(worker_id)
self._current: asyncio.Task[Any] | None = None # pragma: no mutate
@property
def ports(self) -> Ports:
return self.engine.ports
# --- procedure 1: worker loop -----------------------------------------------------
[docs]
async def run_once(self) -> bool:
"""Dequeue at most one task and process it. Return True when a task was processed."""
with bound(worker_id=self.worker_id):
claimed = await self._dequeue()
if claimed is None:
return False
await self._process(claimed)
return True
[docs]
async def run_forever(self, stop: asyncio.Event | None = None) -> None:
"""Loop until `stop`. One failed iteration is logged, not fatal: the bucket arbitrates,
so a lost race or a storage hiccup is retried on the next pass."""
stop = stop or asyncio.Event()
while not stop.is_set():
try:
busy = await self.run_once()
except asyncio.CancelledError:
raise
except Exception:
log.exception("worker_iteration_failed", worker_id=self.worker_id)
busy = False # pragma: no mutate
if not busy:
with suppress(asyncio.TimeoutError):
await asyncio.wait_for(stop.wait(), timeout=self.engine.config.poll_interval)
async def _dequeue(self) -> ClaimedTask | None:
for queue in self.queues:
claimed = await self.ports.queue.dequeue(
queue, self.worker_id, self.engine.config.task_ttl
)
if claimed is not None:
return claimed
return None
async def _process(self, claimed: ClaimedTask) -> None:
"""Run the task under its lease. A lost lease means: stop, do not ack."""
task = claimed.task
result: TaskResult
with bound(task_id=task.task_id, queue=task.queue, task_kind=task.kind.value):
log.debug("task_dequeued", reason=task.reason, eid=str(task.target.eid))
async with _RenewLoop(
claimed.lease, self.engine.config.task_ttl, self._cancel_current
) as renew:
self._current = asyncio.ensure_future(self._dispatch(task))
try:
result = await self._current
except asyncio.CancelledError:
if renew.lost:
log.warning("task_lease_lost", eid=str(task.target.eid))
return
raise
except LeaseLost as exc:
log.warning("task_lease_lost", eid=str(task.target.eid), detail=str(exc))
return
finally:
self._current = None # pragma: no mutate
match result:
case Done():
await self.ports.queue.ack(claimed)
log.debug("task_acked")
case Retry(delay):
await self.ports.queue.nack(claimed, delay)
log.info("task_nacked", delay_s=delay.total_seconds())
def _cancel_current(self) -> None:
if self._current is not None and not self._current.done():
self._current.cancel()
async def _dispatch(self, task: Task) -> TaskResult:
match task.kind:
case TaskKind.START | TaskKind.RESUME:
return await self.run_execution(task)
case TaskKind.RUN_STEP:
return await self.run_detached_step(task)
case TaskKind.DELEGATE:
log.warning("delegate_task_on_worker_queue")
return Retry(timedelta(seconds=self.engine.config.nack_delay))
raise AssertionError(task.kind) # pragma: no mutate
# --- procedure 2: run an execution ---------------------------------------------------
[docs]
async def run_execution(self, task: Task) -> TaskResult:
eid = task.target.eid
engine = self.engine
lease = await self.ports.ownership.acquire(eid, self.worker_id, engine.config.exec_ttl)
if lease is None:
log.info("execution_owned_elsewhere", eid=str(eid))
return Done()
site = engine.site.with_epoch(lease.epoch)
with bound(eid=str(eid), epoch=lease.epoch):
return await self._run_owned(task, lease, site)
async def _run_owned(self, task: Task, lease: Lease, site: Site) -> TaskResult:
eid = task.target.eid
engine = self.engine
try:
try:
execution = await engine.execution(eid)
except UnknownExecution:
log.info("execution_unknown", hint="archived?")
await lease.release()
return Done()
try:
wf = engine.workflow_ref(execution)
except WorkflowNotRegistered:
log.warning(
"workflow_not_registered",
workflow=execution.workflow,
version=execution.version,
)
await lease.release()
return Retry(timedelta(seconds=engine.config.nack_delay))
memo = await engine.memo(eid)
prov = engine.provenance(
self.actor, workflow=wf.name, version=wf.version, site=site, frame_name="worker"
)
if memo.is_terminal:
await lease.release()
return Done()
if _cancel_requested(lease.state):
await self._cancel(execution, memo, lease, prov)
return Done()
if task.kind is TaskKind.START and memo.tail == 0:
await self._record(eid, Entry.execution_started(prov, execution.args), memo, lease)
log.info("execution_started", workflow=wf.name, version=wf.version)
else:
await self._record(
eid, Entry.execution_resumed(prov, lease.epoch, task.reason), memo, lease
)
log.info(
"execution_resumed", workflow=wf.name, version=wf.version, reason=task.reason
)
ctx = Context(
eid=eid,
workflow=wf,
memo=memo,
ports=self.ports,
actor=self.actor,
site=site,
clock=engine.clock,
code_ref=engine.config.code_ref,
starter=engine,
workflow_of=engine.registry.of,
cancel_requested=lambda: _cancel_requested(lease.state),
)
args = execution.args or {}
async with _RenewLoop(lease, engine.config.exec_ttl, self._cancel_current) as renew:
try:
call_args = args.get("args", []) # pragma: no mutate
call_kwargs = args.get("kwargs", {}) # pragma: no mutate
outcome = await ctx.run(*call_args, **call_kwargs)
except asyncio.CancelledError:
if renew.lost:
raise LeaseLost(f"execution {eid} lease lost during run") from None
raise
match outcome:
case Returned(value):
await self._record(eid, Entry.execution_completed(prov, value), memo, lease)
await engine.notify_parent(
execution, {"status": "completed", "value": value}, prov
)
await lease.release({"status": "completed"})
log.info("execution_completed")
case Raised(error) if isinstance(error, Cancelled):
await self._cancel(execution, memo, lease, prov)
case Raised(error) if isinstance(error, NondeterminismError):
await self._record(
eid, Entry.execution_suspended(prov, [Condition.operator()]), memo, lease
)
await lease.release({"status": "suspended", "blocked": "nondeterminism"})
log.error("execution_blocked", reason="nondeterminism", fid=error.fid)
case Raised(error):
failed = Failed.from_exception(error, retryable=False) # pragma: no mutate
await self._record(eid, Entry.execution_failed(prov, failed), memo, lease)
await engine.notify_parent(
execution, {"status": "failed", "error": str(error)}, prov
)
await lease.release({"status": "failed"})
log.warning(
"execution_failed", error_type=failed.error_type, message=failed.message
)
case Suspended(waits):
await self._suspend(execution, memo, lease, prov, waits)
return Done()
except LeaseLost:
log.warning("execution_lease_lost", outcome="discarded")
raise
async def _record(
self, eid: Eid, entry: Entry, memo: MemoTable, lease: Lease | None = None
) -> None:
"""Append an execution.* entry to the journal, fold it, and announce it.
With a lease, first make one guarded write on it: a fenced worker gets LeaseLost
here instead of publishing a lifecycle entry it no longer owns. Frame entries are
not checked: the memo rule makes duplicates harmless (02-journal.md section 3).
"""
from flowli.domain import Sequenced
if lease is not None:
await lease.refresh_state()
seq = await self.ports.journal.append(eid, entry)
memo.apply(Sequenced(seq, entry))
await self.ports.control.announce(
Entry(entry.type, ROOT_FID, {"eid": str(eid), **entry.payload}, entry.provenance)
)
async def _cancel(
self, execution: Execution, memo: MemoTable, lease: Lease, prov: Provenance
) -> None:
by = (lease.state or {}).get("cancel_requested") or unstructure(prov.actor)
await self._record(execution.eid, Entry.execution_cancelled(prov, by), memo, lease)
await self.engine.notify_parent(
execution, {"status": "cancelled", "error": "cancelled"}, prov
)
await lease.release({"status": "cancelled"})
log.info("execution_cancelled", by=by)
# --- procedure 3: suspend ---------------------------------------------------------------
async def _suspend(
self,
execution: Execution,
memo: MemoTable,
lease: Lease,
prov: Provenance,
waits: tuple[Wait, ...],
) -> None:
eid = execution.eid
ports = self.ports
channels: list[str] = []
children: list[str] = []
for w in waits:
ref = FrameRef(eid, w.fid)
kind, name = Condition.parse(w.on)
if kind == Condition.CHANNEL and name is not None:
channels.append(name)
await ports.channel.register_wait(name, ref)
elif kind == Condition.CHILD and name is not None:
children.append(name)
if w.deadline is not None:
await ports.timers.schedule(Timer(w.deadline, ref))
await self._record(eid, Entry.execution_suspended(prov, [w.on for w in waits]), memo, lease)
await lease.release({"status": "suspended"})
log.info("execution_suspended", on=[w.on for w in waits])
# check again after the release: close the race with senders
for channel in channels:
messages = await ports.channel.read(channel, after=memo.last_consumed(channel))
if messages:
seq = messages[0].seq
await self.engine.enqueue_resume(
eid, f"message:{channel}:{seq}", f"message:{channel}", prov
)
for child_eid in children:
with suppress(Exception):
if (await self.engine.status(parse_eid(child_eid))).is_terminal:
await self.engine.enqueue_resume(
eid, f"child:{child_eid}", f"child:{child_eid}", prov
)
# --- procedure 7: detached step ------------------------------------------------------------
[docs]
async def run_detached_step(self, task: Task) -> TaskResult:
"""payload: {"fn": "module:qualname", "args": [...], "kwargs": {...}, "name": str}."""
ref = task.target
payload = task.payload or {}
fn = _resolve(payload["fn"])
name = payload.get("name", fn.__name__)
execution = await self.engine.execution(ref.eid)
memo = await self.engine.memo(ref.eid)
if ref.fid in memo.memos:
return Done()
attempt = memo.next_attempt(ref.fid)
prov = self.engine.provenance(
self.actor,
workflow=execution.workflow,
version=execution.version,
frame_kind=FrameKind.STEP.value,
frame_name=name,
attempt=attempt,
)
args_digest = digest({"args": payload.get("args", []), "kwargs": payload.get("kwargs", {})})
await self.ports.journal.append(
ref.eid,
Entry.frame_started(prov, ref.fid, FrameKind.STEP.value, name, args_digest, attempt),
)
body: dict[str, Any]
try:
result = fn(*payload.get("args", []), **payload.get("kwargs", {}))
if inspect.isawaitable(result):
result = await result
except Exception as exc:
failed = Failed.from_exception(exc)
await self.ports.journal.append(
ref.eid, Entry.frame_failed(prov, ref.fid, attempt, failed)
)
body = {"status": "failed", "error": unstructure(failed)}
else:
await self.ports.journal.append(
ref.eid, Entry.frame_completed(prov, ref.fid, attempt, result)
)
body = {"status": "completed", "value": result}
await self.ports.channel.send(Message(ref.step_channel, 0, body, prov)) # pragma: no mutate
await self.engine.enqueue_resume(ref.eid, f"step:{ref.fid}", f"step:{ref.fid}", prov)
log.info("detached_step_done", eid=str(ref.eid), fid=ref.fid, status=body["status"])
return Done()
def _cancel_requested(state: Any) -> bool:
return isinstance(state, dict) and bool(state.get("cancel_requested"))
def _resolve(path: str) -> Callable[..., Any]:
module, _, qualname = path.partition(":") # pragma: no mutate
obj: Any = importlib.import_module(module)
for part in qualname.split("."):
obj = getattr(obj, part)
return obj # type: ignore[no-any-return]