Spaces:
Sleeping
Sleeping
| """Task execution for claims returned by messaging tree aggregates.""" | |
| from __future__ import annotations | |
| import asyncio | |
| import contextlib | |
| from collections.abc import Awaitable, Callable | |
| from dataclasses import dataclass | |
| from loguru import logger | |
| from ..safe_diagnostics import format_exception_for_log | |
| from .runtime import MessageTree | |
| from .transitions import CancellationReason, NodeClaim, QueueEntry | |
| NodeProcessor = Callable[[NodeClaim], Awaitable[None]] | |
| QueueUpdateCallback = Callable[[tuple[QueueEntry, ...]], Awaitable[None]] | |
| NodeStartedCallback = Callable[[NodeClaim], Awaitable[None]] | |
| ClaimFailureCallback = Callable[[NodeClaim], Awaitable[None]] | |
| ClaimFinishedCallback = Callable[[MessageTree, NodeClaim], Awaitable[None]] | |
| class _TaskSlot: | |
| tree: MessageTree | |
| claim: NodeClaim | |
| task: asyncio.Task[None] | None = None | |
| runner_started: bool = False | |
| transitioned: bool = False | |
| recovery_task: asyncio.Task[None] | None = None | |
| cancellation_requested: bool = False | |
| cancellation_reason: CancellationReason | None = None | |
| class CancelledTask: | |
| """Task handle plus whether the node runner owns cancellation UI.""" | |
| task: asyncio.Task[None] | |
| runner_started: bool | |
| class TreeQueueProcessor: | |
| """Own asyncio tasks while MessageTree owns scheduling state.""" | |
| def __init__( | |
| self, | |
| node_processor: NodeProcessor, | |
| *, | |
| claim_failure_callback: ClaimFailureCallback, | |
| claim_finished_callback: ClaimFinishedCallback, | |
| queue_update_callback: QueueUpdateCallback | None = None, | |
| node_started_callback: NodeStartedCallback | None = None, | |
| log_messaging_error_details: bool = False, | |
| ) -> None: | |
| self._node_processor = node_processor | |
| self._claim_failure_callback = claim_failure_callback | |
| self._claim_finished_callback = claim_finished_callback | |
| self._queue_update_callback = queue_update_callback | |
| self._node_started_callback = node_started_callback | |
| self._log_messaging_error_details = log_messaging_error_details | |
| self._tasks: dict[str, _TaskSlot] = {} | |
| self._completion_failures: list[Exception] = [] | |
| self._idle = asyncio.Event() | |
| self._idle.set() | |
| def _key(claim: NodeClaim) -> str: | |
| return claim.claim_id | |
| def launch( | |
| self, | |
| tree: MessageTree, | |
| claim: NodeClaim, | |
| *, | |
| announce_started: bool = False, | |
| queue: tuple[QueueEntry, ...] = (), | |
| ) -> None: | |
| """Attach a task synchronously before another coroutine can cancel it.""" | |
| key = self._key(claim) | |
| if key in self._tasks: | |
| raise RuntimeError(f"Claim {key} already has a task") | |
| slot = _TaskSlot(tree=tree, claim=claim) | |
| self._tasks[key] = slot | |
| self._idle.clear() | |
| ownership_ready = asyncio.Event() | |
| claim_runner = self._run_claim( | |
| slot, | |
| ownership_ready=ownership_ready, | |
| announce_started=announce_started, | |
| queue=queue, | |
| ) | |
| try: | |
| task = asyncio.create_task( | |
| claim_runner, | |
| name=(f"messaging-claim-{claim.identity.root_id}-{claim.claim_id[:8]}"), | |
| ) | |
| except BaseException: | |
| claim_runner.close() | |
| if self._tasks.get(key) is slot: | |
| self._tasks.pop(key) | |
| if not self._tasks: | |
| self._idle.set() | |
| raise | |
| slot.task = task | |
| task.add_done_callback(lambda _task, claim_key=key: self._task_done(claim_key)) | |
| ownership_ready.set() | |
| def _task_done(self, key: str) -> None: | |
| """Recover a claim if its task was cancelled before entering its body.""" | |
| slot = self._tasks.get(key) | |
| if slot is None or slot.transitioned or slot.recovery_task is not None: | |
| return | |
| slot.recovery_task = asyncio.create_task( | |
| self._recover_unentered_task(slot), | |
| name=f"messaging-claim-recovery-{key[:8]}", | |
| ) | |
| async def _recover_unentered_task(self, slot: _TaskSlot) -> None: | |
| task = slot.task | |
| if task is not None: | |
| with contextlib.suppress(asyncio.CancelledError, Exception): | |
| await task | |
| if not slot.transitioned: | |
| await self._finish_and_continue(slot) | |
| async def _notify_queue_updated(self, queue: tuple[QueueEntry, ...]) -> None: | |
| if self._queue_update_callback is None: | |
| return | |
| try: | |
| await self._queue_update_callback(queue) | |
| except Exception as exc: | |
| logger.warning( | |
| "Queue update callback failed: {}", | |
| format_exception_for_log( | |
| exc, | |
| log_full_message=self._log_messaging_error_details, | |
| ), | |
| ) | |
| async def notify_queue_updated(self, queue: tuple[QueueEntry, ...]) -> None: | |
| """Publish a transition-owned queue snapshot.""" | |
| await self._notify_queue_updated(queue) | |
| async def _notify_node_started(self, claim: NodeClaim) -> None: | |
| if self._node_started_callback is None: | |
| return | |
| try: | |
| await self._node_started_callback(claim) | |
| except Exception as exc: | |
| logger.warning( | |
| "Node started callback failed: {}", | |
| format_exception_for_log( | |
| exc, | |
| log_full_message=self._log_messaging_error_details, | |
| ), | |
| ) | |
| async def _run_claim( | |
| self, | |
| slot: _TaskSlot, | |
| *, | |
| ownership_ready: asyncio.Event, | |
| announce_started: bool, | |
| queue: tuple[QueueEntry, ...], | |
| ) -> None: | |
| await ownership_ready.wait() | |
| claim = slot.claim | |
| try: | |
| if announce_started: | |
| await self._notify_node_started(claim) | |
| await self._notify_queue_updated(queue) | |
| if slot.cancellation_requested: | |
| if slot.cancellation_reason is None: | |
| raise asyncio.CancelledError | |
| raise asyncio.CancelledError(slot.cancellation_reason) | |
| slot.runner_started = True | |
| await self._node_processor(claim) | |
| except asyncio.CancelledError: | |
| logger.info("Task for node {} was cancelled", claim.node.node_id) | |
| raise | |
| except Exception as exc: | |
| logger.error( | |
| "Error processing node {}: {}", | |
| claim.node.node_id, | |
| format_exception_for_log( | |
| exc, | |
| log_full_message=self._log_messaging_error_details, | |
| ), | |
| ) | |
| await self._claim_failure_callback(claim) | |
| finally: | |
| if not slot.transitioned: | |
| await self._finish_and_continue(slot) | |
| async def _finish_and_continue(self, slot: _TaskSlot) -> None: | |
| current = asyncio.current_task() | |
| if current is not None: | |
| while current.cancelling(): | |
| current.uncancel() | |
| try: | |
| while True: | |
| try: | |
| await self._claim_finished_callback(slot.tree, slot.claim) | |
| slot.transitioned = True | |
| break | |
| except asyncio.CancelledError: | |
| if current is not None: | |
| while current.cancelling(): | |
| current.uncancel() | |
| continue | |
| except Exception as exc: | |
| self._completion_failures.append(exc) | |
| logger.error( | |
| "Claim completion callback failed for node {}: {}", | |
| slot.claim.node.node_id, | |
| format_exception_for_log( | |
| exc, | |
| log_full_message=self._log_messaging_error_details, | |
| ), | |
| ) | |
| finally: | |
| key = self._key(slot.claim) | |
| if self._tasks.get(key) is slot: | |
| self._tasks.pop(key) | |
| if not self._tasks: | |
| self._idle.set() | |
| def cancel( | |
| self, | |
| claim: NodeClaim, | |
| reason: CancellationReason | None, | |
| ) -> CancelledTask | None: | |
| """Cancel exactly the task bound to one aggregate claim.""" | |
| slot = self._tasks.get(self._key(claim)) | |
| if slot is None: | |
| return None | |
| slot.cancellation_requested = True | |
| slot.cancellation_reason = reason | |
| task = slot.task | |
| if task is None or task.done(): | |
| return None | |
| if reason is None: | |
| task.cancel() | |
| else: | |
| task.cancel(reason) | |
| if slot.runner_started: | |
| return CancelledTask(task=task, runner_started=True) | |
| if slot.recovery_task is None: | |
| slot.recovery_task = asyncio.create_task( | |
| self._recover_unentered_task(slot), | |
| name=f"messaging-claim-recovery-{claim.claim_id[:8]}", | |
| ) | |
| return CancelledTask(task=slot.recovery_task, runner_started=False) | |
| def task_count(self) -> int: | |
| """Return the number of attached claims for observability.""" | |
| return len(self._tasks) | |
| async def wait_idle(self) -> None: | |
| """Wait for every task and hand completion failures to the caller once.""" | |
| await self._idle.wait() | |
| if not self._completion_failures: | |
| return | |
| failures = self._completion_failures | |
| self._completion_failures = [] | |
| if len(failures) == 1: | |
| raise failures[0] | |
| raise ExceptionGroup("Messaging claim completion failures", failures) | |
| __all__ = ["CancelledTask", "TreeQueueProcessor"] | |