diff --git a/docs/reference/openapi.yaml b/docs/reference/openapi.yaml index e0cd8282df..9182125167 100644 --- a/docs/reference/openapi.yaml +++ b/docs/reference/openapi.yaml @@ -365,7 +365,7 @@ components: title: Errors type: array is_complete: - default: false + readOnly: true title: Is Complete type: boolean is_pending: @@ -389,6 +389,7 @@ components: required: - task_id - task + - is_complete title: TrackableTask type: object ValidationError: @@ -449,7 +450,7 @@ info: name: Apache 2.0 url: https://www.apache.org/licenses/LICENSE-2.0.html title: BlueAPI Control - version: 1.5.0 + version: 1.5.1 openapi: 3.1.0 paths: /api/v1/devices: diff --git a/src/blueapi/config.py b/src/blueapi/config.py index a181f4c344..b93704c57b 100644 --- a/src/blueapi/config.py +++ b/src/blueapi/config.py @@ -324,7 +324,7 @@ class ApplicationConfig(BlueapiBaseModel): """ #: API version to publish in OpenAPI schema - REST_API_VERSION: ClassVar[str] = "1.5.0" + REST_API_VERSION: ClassVar[str] = "1.5.1" LICENSE_INFO: ClassVar[dict[str, str]] = { "name": "Apache 2.0", diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index 9ce373c769..632b719658 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -2,6 +2,7 @@ from collections.abc import Mapping from typing import Any +from bluesky.run_engine import RunEngineResult from pydantic import BaseModel, Field, TypeAdapter from blueapi.core import BlueskyContext @@ -29,7 +30,7 @@ def prepare_params(self, ctx: BlueskyContext) -> Mapping[str, Any]: # Re-create dict manually to avoid nesting in model_dump output return {field: getattr(model, field) for field in model.__pydantic_fields__} - def do_task(self, ctx: BlueskyContext) -> None: + def do_task(self, ctx: BlueskyContext) -> Any: LOGGER.info( f"Asked to run plan {self.name} with {self.params} and " f"metadata {self.metadata} for all runs" @@ -39,9 +40,12 @@ def do_task(self, ctx: BlueskyContext) -> None: prepared_params = self.prepare_params(ctx) ctx.run_engine.md.update(self.metadata) result = ctx.run_engine(func(**prepared_params)) - if isinstance(result, tuple): # pragma: no cover - # this is never true if the run_engine is configured correctly - return None + if not isinstance(result, RunEngineResult): + # this is unreachable unless something has misconfigured it. + raise RuntimeError( + "RunEngine did not return a RunEngineResult - is " + "call_returns_result set on this RunEngine instance?" + ) return result.plan_result diff --git a/src/blueapi/worker/task_worker.py b/src/blueapi/worker/task_worker.py index 4ada3c2a54..6fe6e64684 100644 --- a/src/blueapi/worker/task_worker.py +++ b/src/blueapi/worker/task_worker.py @@ -5,9 +5,10 @@ from dataclasses import dataclass from functools import partial from queue import Full, Queue -from threading import Event, RLock +from threading import Condition, Event, RLock from typing import Any, TypeVar +from bluesky import RunEngineInterrupted from bluesky._vendor.super_state_machine.errors import TransitionError from bluesky.protocols import Status from observability_utils.tracing import ( @@ -19,7 +20,7 @@ from opentelemetry.baggage import get_baggage from opentelemetry.context import Context, get_current from opentelemetry.trace import SpanKind -from pydantic import Field +from pydantic import Field, computed_field, model_validator from pydantic.json_schema import SkipJsonSchema from blueapi.core import ( @@ -68,11 +69,23 @@ class TrackableTask(BlueapiBaseModel): task_id: str task: Task request_id: str | SkipJsonSchema[None] = None - is_complete: bool = False is_pending: bool = True errors: list[str] = Field(default_factory=list) outcome: TaskResult | TaskError | None = None + @computed_field # type: ignore[prop-decorator] + @property + def is_complete(self) -> bool: + return self.outcome is not None + + @model_validator(mode="before") + @classmethod + def remove_complete(cls, values: Any) -> Any: + # fix pydantic falling over itself with computed fields + if isinstance(values, dict): + values.pop("is_complete", None) + return values + def set_result(self, result: Any): self.outcome = TaskResult.from_result(result) @@ -100,16 +113,14 @@ class TaskWorker: _errors: list[str] _warnings: list[str] - # The queue is actually a channel between 2 threads - # most programming languages have a separate abstraction for this - # but Python reuses Queue - # So it's not used as a standard queue, - # but as a box in which to put the "current" task and nothing else - # So the calling thread can only ever submit one plan at a time. + # A channel to the worker thread, not really a queue - maxsize=1 makes it + # a single-item box, so only one plan can be in flight at a time. _task_channel: Queue # type: ignore _current: TrackableTask | None + _pending_cancel: "CancelSignal | None" _status_lock: RLock + _state_change: Condition _status_snapshot: dict[str, StatusView] _completed_statuses: set[str] _worker_events: EventPublisher[WorkerEvent] @@ -139,10 +150,12 @@ def __init__( self._warnings = [] self._task_channel = Queue(maxsize=1) self._current = None + self._pending_cancel = None self._worker_events = EventPublisher() self._progress_events = EventPublisher() self._data_events = EventPublisher() self._status_lock = RLock() + self._state_change = Condition() self._status_snapshot = {} self._completed_statuses = set() self._started = Event() @@ -180,19 +193,57 @@ def cancel_active_task( Returns: The task_id of the active task """ - if self._current is None: + + current = self._current + if current is None: # Persuades type checker that self._current is not None # We only allow this method to be called if a Plan is active raise TransitionError("Attempted to cancel while no active Task") - if failure: - default_reason = "Task failed for unknown reason" - self._ctx.run_engine.abort(reason or default_reason) - add_span_attributes({"Task aborted": reason or default_reason}) + + if self._current_task_otel_context is None: + self._current_task_otel_context = get_current() + + reason = reason or ("Task aborted" if failure else "Task stopped") + if self._state != WorkerState.RUNNING: + with self._state_change: + self._task_channel.put(AbortSignal(failure, reason)) + if not self._state_change.wait_for( + lambda: current.task_id in self._completed_tasks, + timeout=self._start_stop_timeout, + ): + raise TimeoutError( + "Worker did not finish cancelling within " + f"{self._start_stop_timeout} seconds" + ) else: - self._ctx.run_engine.stop() - default_reason = "Cancellation successful: Task stopped without error" - add_span_attributes({"Task stopped": reason or default_reason}) - return self._current.task_id + # The worker thread's own do_task() may be finishing concurrently + # on another thread - don't clobber whatever it already recorded. + # abort()/stop() must be called without holding _state_change: + # _on_state_change() needs that same lock, and abort()/stop() + # don't return until its state-change callbacks have run. + try: + if failure: + with self._status_lock: + if current.outcome is None: + current.outcome = TaskError(type="Abort", message=reason) + self._ctx.run_engine.abort(reason=reason) + else: + with self._status_lock: + if current.outcome is None: + current.outcome = TaskResult.from_result(None) + self._ctx.run_engine.stop() + except TransitionError: + pass + with self._state_change: + if not self._state_change.wait_for( + lambda: current.task_id in self._completed_tasks, + timeout=self._start_stop_timeout, + ): + raise TimeoutError( + "Worker did not finish cancelling within " + f"{self._start_stop_timeout} seconds" + ) + return current.task_id @start_as_current_span(TRACER, "task_id") def get_task_by_id(self, task_id: str) -> TrackableTask | None: @@ -205,6 +256,9 @@ def get_task_by_id(self, task_id: str) -> TrackableTask | None: Optional[TrackableTask[T]]: The task matching the ID, None if the task ID is unknown to the worker. """ + current = self._current + if current is not None and current.task_id == task_id: + return current return self._pending_tasks.get(task_id, None) or self._completed_tasks[task_id] @start_as_current_span(TRACER) @@ -218,11 +272,8 @@ def get_tasks(self, status: TaskStatusEnum | None = None) -> list[TrackableTask] list[TrackableTask]: A list of tasks that match the given status. """ if status == TaskStatusEnum.RUNNING: - return [ - task - for task in self._pending_tasks.values() - if not task.is_pending and not task.is_complete - ] + current = self._current + return [current] if current is not None else [] elif status == TaskStatusEnum.PENDING: return [task for task in self._pending_tasks.values() if task.is_pending] elif status == TaskStatusEnum.COMPLETE: @@ -275,7 +326,6 @@ def submit_task(self, task: Task) -> str: task_id: str = str(uuid.uuid4()) add_span_attributes({"TaskId": task_id}) request_id = get_baggage("correlation_id") - # If request id is not a string, we do not pass it into a TrackableTask if not isinstance(request_id, str): LOGGER.warning(f"Invalid correlation id detected: {request_id}") request_id = None @@ -315,8 +365,9 @@ def mark_task_as_started(event: WorkerEvent, _: str | None) -> None: self._current_task_otel_context = get_current() """ Cache the current trace context as the one for this task id """ self._task_channel.put_nowait(trackable_task) - task_started.wait(timeout=5.0) - if not task_started.is_set(): + if task_started.wait(timeout=5.0): + self._pending_tasks.pop(trackable_task.task_id) + else: raise TimeoutError("Failed to start plan within timeout") except Full as f: LOGGER.error("Cannot submit task while another is running") @@ -343,7 +394,6 @@ def stop(self) -> None: """ LOGGER.info("Attempting to stop worker") - # If the worker has not yet started there is nothing to do. if self._started.is_set(): self._task_channel.put(KillSignal()) else: @@ -410,7 +460,25 @@ def resume(self): Command the worker to resume """ LOGGER.info("Requesting to resume the worker") - self._ctx.run_engine.resume() + self._current_task_otel_context = get_current() + current = self._current + + def resumed() -> bool: + if self._state == WorkerState.PAUSED: + return False + if self._state == WorkerState.IDLE: + return current is None or current.task_id in self._completed_tasks + return True + + with self._state_change: + self._task_channel.put(ResumeSignal()) + if not self._state_change.wait_for( + resumed, + timeout=self._start_stop_timeout, + ): + raise TimeoutError( + f"Worker did not resume within {self._start_stop_timeout} seconds" + ) @start_as_current_span(TRACER) def _cycle_with_error_handling(self) -> None: @@ -436,11 +504,22 @@ def process_task(): LOGGER.info( "Task ran successfully - returned: %s", result, extra=meta ) - self._current.set_result(result) + with self._status_lock: + # cancel_active_task() may have set this concurrently. + if self._current.outcome is None: + self._current.set_result(result) + except RunEngineInterrupted: + # Raised by both a pause (outcome still None) and an + # abort (outcome already TaskError) - only the latter + # is a failure. + if isinstance(self._current.outcome, TaskError): + self._report_error(Exception(self._current.outcome.message)) except Exception as e: LOGGER.error("Task failed", extra=meta) - self._current.set_exception(e) - self._report_error(e) + with self._status_lock: + if self._current.outcome is None: + self._current.set_exception(e) + raise with plan_tag_filter_context(next_task.task.name, LOGGER): if self._current_task_otel_context is not None: @@ -458,10 +537,45 @@ def process_task(): else: process_task() + elif isinstance(next_task, ResumeSignal): + pending_cancel = self._pending_cancel + if pending_cancel is not None: + # A cancel queued after this resume takes priority, so + # this resume never runs. + self._apply_cancel(pending_cancel) + elif self._state == WorkerState.PAUSED: + if self._current is not None: + try: + result = self._ctx.run_engine.resume() + self._current.set_result(result) + except RunEngineInterrupted: + # Plan paused again immediately - not a failure, + # leave the outcome unset so it stays resumable. + LOGGER.debug("RunEngine resume interrupted; ignoring") + else: + LOGGER.warning( + "Received resume signal but no active task, ignoring" + ) + else: + LOGGER.warning( + "Received resume signal but RunEngine is not paused, ignoring" + ) + + elif isinstance(next_task, (CancelSignal, AbortSignal)): + self._apply_cancel(self._pending_cancel or next_task) + elif isinstance(next_task, KillSignal): # If we receive a kill signal we begin to shut the worker down. # Note that the kill signal is explicitly not a type of task as we don't # want it to be part of the worker's public API + self._pending_cancel = None + if self._current is not None and self._state == WorkerState.PAUSED: + self._apply_cancel( + CancelSignal( + failure=True, + reason="Worker is stopping while the task was paused", + ) + ) self._stopping.set() add_span_attributes({"server shutting down": "true"}) else: @@ -469,17 +583,25 @@ def process_task(): except Exception as err: self._report_error(err) finally: - if self._current_task_otel_context is not None: + if self._current_task_otel_context is not None and self._state not in [ + WorkerState.PANICKED, + WorkerState.PAUSED, + ]: self._current_task_otel_context = None if self._current is not None: - self._current.is_complete = True - self._pending_tasks.pop(self._current.task_id) - self._completed_tasks[self._current.task_id] = self._current - self._report_status() - self._errors.clear() - self._warnings.clear() - self._completed_statuses.clear() + # Not done yet while paused - it may still be resumed or cancelled. + if self._state != WorkerState.PAUSED: + self._completed_tasks[self._current.task_id] = self._current + if self._state != WorkerState.PAUSED: + finished_task = self._current + self._current = None + with self._state_change: + self._state_change.notify_all() + self._report_status(finished_task) + self._errors.clear() + self._warnings.clear() + self._completed_statuses.clear() @property def worker_events(self) -> EventStream[WorkerEvent, int]: @@ -520,8 +642,10 @@ def _on_state_change( else: old_state = WorkerState.UNKNOWN LOGGER.debug(f"Notifying state change {old_state} -> {new_state}") - self._state = new_state - self._report_status() + with self._state_change: + self._state = new_state + self._state_change.notify_all() + self._report_status(self._current) def _report_error(self, err: Exception) -> None: LOGGER.error(err, exc_info=True) @@ -529,26 +653,66 @@ def _report_error(self, err: Exception) -> None: self._current.errors.append(str(err)) self._errors.append(str(err)) - @start_as_current_span(TRACER) - def _report_status( - self, + def _finalize_cancel_outcome( + self, current: TrackableTask, signal: "CancelSignal | AbortSignal" ) -> None: + """ + Set current's outcome for a cancellation, unless something else (e.g. + the worker thread finishing do_task() on its own) already got there first. + """ + with self._status_lock: + if current.outcome is None: + if signal.failure: + current.set_exception( + Exception(signal.reason or "Task failed for unknown reason") + ) + else: + current.set_result(None) + + def _apply_cancel(self, signal: "CancelSignal | AbortSignal") -> None: + self._pending_cancel = None + current = self._current + + if signal.failure: + reason = signal.reason or "Task failed for unknown reason" + if current is not None: + self._ctx.run_engine.abort(reason) + self._finalize_cancel_outcome(current, signal) + add_span_attributes({"Task aborted": reason}) + LOGGER.error("Task failed: %s", reason) + if current is not None: + self._report_error(Exception(reason)) + else: + if current is not None: + self._ctx.run_engine.stop() + self._finalize_cancel_outcome(current, signal) + add_span_attributes( + { + "Task stopped": signal.reason + or "Cancellation successful: Task stopped without error" + } + ) + + @start_as_current_span(TRACER) + def _report_status(self, current: TrackableTask | None) -> None: task_status: TaskStatus | None errors = self._errors warnings = self._warnings - if self._current is not None: + if current is not None: + task_complete = current.task_id in self._completed_tasks task_status = TaskStatus( - task_id=self._current.task_id, - task_complete=self._current.is_complete, - task_failed=bool(self._current.errors), - result=self._current.outcome, + task_id=current.task_id, + task_complete=task_complete, + task_failed=bool(current.errors) + or isinstance(current.outcome, TaskError), + result=current.outcome, ) - correlation_id = self._current.task_id + correlation_id = current.task_id add_span_attributes( { - "task_id": self._current.task_id, - "task_complete": self._current.is_complete, - "task_failed": self._current.errors, + "task_id": current.task_id, + "task_complete": task_complete, + "task_failed": current.errors, } ) else: @@ -598,7 +762,7 @@ def _on_document(self, name: str, document: Mapping[str, Any]) -> None: ) else: - raise KeyError( + raise RuntimeError( "Trying to emit a document despite the fact that the RunEngine is idle" ) @@ -683,6 +847,31 @@ class KillSignal: ... +@dataclass +class ResumeSignal: + """ + Object put in the worker's task queue to tell it to resume if paused. + """ + + pass + + +@dataclass +class AbortSignal: + failure: bool + reason: str + + +@dataclass +class CancelSignal: + """ + Object put in the worker's task queue to tell it to cancel the current task. + """ + + failure: bool + reason: str | None + + def run_worker_in_own_thread( worker: TaskWorker, executor: ThreadPoolExecutor | None = None ) -> Future: diff --git a/tests/system_tests/test_blueapi_system.py b/tests/system_tests/test_blueapi_system.py index a7e4c86413..970bd909e2 100644 --- a/tests/system_tests/test_blueapi_system.py +++ b/tests/system_tests/test_blueapi_system.py @@ -409,11 +409,14 @@ def test_delete_non_existent_task(rest_client: BlueapiRestClient): rest_client.clear_task("Not-exists") -def test_put_worker_task(rest_client: BlueapiRestClient, small_task: TaskRequest): - created_task = rest_client.create_task(small_task) +def test_put_worker_task(rest_client: BlueapiRestClient, long_task: TaskRequest): + # long_task, since a near-instant task could complete (clearing the + # active task) before get_active_task() below is called. + created_task = rest_client.create_task(long_task) rest_client.update_worker_task(WorkerTask(task_id=created_task.task_id)) active_task = rest_client.get_active_task() assert active_task.task_id == created_task.task_id + rest_client.cancel_current_task(WorkerState.ABORTING) rest_client.clear_task(created_task.task_id) @@ -806,7 +809,9 @@ def test_any_user_can_retrieve_active_task( task_id = ( client_factory[AdminUser.admin] .create_and_start_task( - task_factory(AdminUser.admin, VALID_INSTRUMENT_SESSION[AdminUser.admin]) + task_factory( + AdminUser.admin, VALID_INSTRUMENT_SESSION[AdminUser.admin], time=1 + ) ) .task_id ) @@ -814,6 +819,8 @@ def test_any_user_can_retrieve_active_task( for user in User: assert client_factory[user].get_active_task().task_id == task_id + client_factory[AdminUser.admin].abort() + def test_non_admin_can_only_start_own_tasks( client_factory: dict[ValidUser, BlueapiClient], diff --git a/tests/unit_tests/service/test_interface.py b/tests/unit_tests/service/test_interface.py index ac756fe851..f9050851cd 100644 --- a/tests/unit_tests/service/test_interface.py +++ b/tests/unit_tests/service/test_interface.py @@ -386,7 +386,6 @@ def test_get_task_by_id( params={}, metadata=expected_metadata, ), - is_complete=False, is_pending=True, errors=[], ) @@ -416,7 +415,6 @@ def test_submit_task_inserts_metadata(context_mock: MagicMock): params={}, metadata=metadata, ), - is_complete=False, is_pending=True, errors=[], ) diff --git a/tests/unit_tests/service/test_rest_api.py b/tests/unit_tests/service/test_rest_api.py index dedd9d7ff3..fe73240601 100644 --- a/tests/unit_tests/service/test_rest_api.py +++ b/tests/unit_tests/service/test_rest_api.py @@ -43,7 +43,7 @@ WorkerTask, ) from blueapi.service.runner import WorkerDispatcher -from blueapi.worker.event import WorkerState +from blueapi.worker.event import TaskResult, WorkerState from blueapi.worker.task import Task from blueapi.worker.task_worker import TrackableTask @@ -371,7 +371,7 @@ def test_put_plan_fails_if_not_idle(mock_runner: Mock, client: TestClient) -> No # Set to non idle mock_runner.run.return_value = TrackableTask( - task=Task(name="none"), task_id=task_id_current, is_complete=False + task=Task(name="none"), task_id=task_id_current ) resp = client.put("/worker/task", json={"task_id": task_id_new}) @@ -386,7 +386,6 @@ def test_get_tasks(mock_runner: Mock, client: TestClient) -> None: TrackableTask( task_id="1", task=Task(name="first_task"), - is_complete=False, is_pending=True, ), ] @@ -433,7 +432,7 @@ def test_get_tasks_by_status(mock_runner: Mock, client: TestClient) -> None: TrackableTask( task_id="3", task=Task(name="third_task"), - is_complete=True, + outcome=TaskResult.from_result(42), is_pending=False, ), ] @@ -453,7 +452,7 @@ def test_get_tasks_by_status(mock_runner: Mock, client: TestClient) -> None: "params": {}, "metadata": {}, }, - "outcome": None, + "outcome": {"outcome": "success", "type": "int", "result": 42}, "task_id": "3", } ] @@ -534,7 +533,7 @@ def test_set_active_task_active_task_complete( mock_runner.run.return_value = TrackableTask( task_id="1", task=Task(name="a_completed_task"), - is_complete=True, + outcome=TaskResult.from_result(42), is_pending=False, ) @@ -553,7 +552,6 @@ def test_set_active_task_worker_already_running( mock_runner.run.return_value = TrackableTask( task_id="1", task=Task(name="a_running_task"), - is_complete=False, is_pending=False, ) diff --git a/tests/unit_tests/worker/test_task_worker.py b/tests/unit_tests/worker/test_task_worker.py index 588c37d8e4..fa34057755 100644 --- a/tests/unit_tests/worker/test_task_worker.py +++ b/tests/unit_tests/worker/test_task_worker.py @@ -1,16 +1,19 @@ import dataclasses import itertools import threading +import time from collections.abc import Callable, Iterable -from concurrent.futures import Future +from concurrent.futures import Future, ThreadPoolExecutor from pathlib import Path from queue import Full from typing import Any, TypeVar from unittest.mock import ANY, MagicMock, Mock, patch +import bluesky.plan_stubs as plan_stubs import pydantic import pytest from bluesky.protocols import Movable, Readable, Status +from bluesky.run_engine import RunEngineResult from bluesky.utils import MsgGenerator from dodal.common import inject from dodal.common.types import UpdatingPathProvider @@ -35,7 +38,8 @@ WorkerEvent, WorkerState, ) -from blueapi.worker.event import TaskResult, TaskStatusEnum +from blueapi.worker.event import TaskError, TaskResult, TaskStatusEnum +from blueapi.worker.task_worker import CancelSignal _SIMPLE_TASK = Task(name="sleep", params={"time": 0.0}) _LONG_TASK = Task(name="sleep", params={"time": 1.0}) @@ -44,6 +48,10 @@ params={"movable": "fake_device", "value": 4.0}, ) _FAILING_TASK = Task(name="failing_plan", params={}) +_PAUSING_TASK = Task(name="pausing_plan", params={}) +_PAUSING_TASK_SLOW_CLEANUP = Task(name="pausing_plan_with_slow_cleanup", params={}) +_TWICE_PAUSING_TASK = Task(name="twice_pausing_plan", params={}) +_ABORT_CLEANUP_DELAY = 1.0 _TASK_WITH_METADATA = Task( name="sleep", params={"time": 0.0}, @@ -79,6 +87,22 @@ def failing_plan() -> MsgGenerator: raise KeyError("I failed") +def pausing_plan() -> MsgGenerator: + yield from plan_stubs.pause() + + +def pausing_plan_with_slow_cleanup() -> MsgGenerator: + try: + yield from plan_stubs.pause() + finally: + time.sleep(_ABORT_CLEANUP_DELAY) + + +def twice_pausing_plan() -> MsgGenerator: + yield from plan_stubs.pause() + yield from plan_stubs.pause() + + @dataclasses.dataclass class ComplexReturn: foo: int @@ -344,9 +368,9 @@ def test_plan_failure_recorded_in_active_task(worker: TaskWorker) -> None: assert events[-1].task_status.task_failed assert events[-1].errors == ["'I failed'"] - active_task = worker.get_active_task() - assert active_task is not None - assert active_task.errors == ["'I failed'"] + completed_task = worker._completed_tasks.get(task_id) + assert completed_task is not None + assert completed_task.errors == ["'I failed'"] def test_task_not_run_twice(worker: TaskWorker) -> None: @@ -423,6 +447,14 @@ def raise_full(item): def test_metadata_passed_to_context(context: BlueskyContext): context.run_engine = Mock() context.run_engine.md = {} + context.run_engine.return_value = RunEngineResult( + run_start_uids=(), + plan_result=None, + exit_status="success", + interrupted=False, + reason="", + exception=None, + ) _TASK_WITH_METADATA.do_task(context) for metadata in [("foo", "bar"), ("baz", 0)]: assert metadata in context.run_engine.md.items() @@ -455,6 +487,19 @@ def begin_task_and_wait_until_complete( return events.result(timeout=timeout) +def begin_task_and_wait_until_paused( + worker: TaskWorker, + task_id: str, + timeout: float = 5.0, +) -> None: + paused_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.PAUSED, + ) + worker.begin_task(task_id) + paused_future.result(timeout=timeout) + + # # Event stream helpers # @@ -629,21 +674,18 @@ def callback(unused: Future[list[Any]], stream=stream, sub=sub): ], ) def test_get_tasks(worker: TaskWorker, status, expected_task_ids): + # A task that has started is tracked via _current, not _pending_tasks. + worker._current = TrackableTask( + task_id="task1", + task=Task(name="set_absolute", params={"movable": "fake_device", "value": 4.0}), + is_pending=False, + ) worker._pending_tasks = { - "task1": TrackableTask( - task_id="task1", - task=Task( - name="set_absolute", params={"movable": "fake_device", "value": 4.0} - ), - is_complete=False, - is_pending=False, - ), "task2": TrackableTask( task_id="task2", task=Task( name="set_absolute", params={"movable": "fake_device", "value": 4.0} ), - is_complete=False, is_pending=True, ), } @@ -653,7 +695,7 @@ def test_get_tasks(worker: TaskWorker, status, expected_task_ids): task=Task( name="set_absolute", params={"movable": "fake_device", "value": 4.0} ), - is_complete=True, + outcome=TaskResult.from_result(None), is_pending=False, ), } @@ -668,7 +710,10 @@ def test_submitting_completed_task_fails(worker: TaskWorker): with pytest.raises(ValueError): worker._submit_trackable_task( TrackableTask( - task_id="task1", task=_SIMPLE_TASK, is_complete=True, is_pending=False + task_id="task1", + task=_SIMPLE_TASK, + outcome=TaskResult.from_result(None), + is_pending=False, ) ) @@ -803,10 +848,6 @@ def test_cycle_without_otel_context(mock_logger: Mock, inert_worker: TaskWorker) inert_worker._cycle() assert inert_worker._current_task_otel_context is None - # Bad way to tell that this branch has been run, but I can't think of a better way - # Have to set these values to match output - task.is_complete = False - task.is_pending = True mock_logger.info.assert_called_with( "Task ran successfully - returned: %s", None, extra={"task_id": "0"} ) @@ -922,3 +963,324 @@ def test_task_result_serialization(plan_result, task_result, type_name): res = TaskResult.from_result(plan_result) assert res.result == task_result assert res.type == type_name + + +def test_pause_does_not_publish_error_event(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + events_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.PAUSED, + ) + worker.begin_task(task_id) + events = events_future.result(timeout=5.0) + + assert all(e.errors == [] for e in events) + + +def test_paused_task_tracked_via_current_not_pending(worker: TaskWorker) -> None: + # A task leaves _pending_tasks as soon as it starts - self._current is + # the sole source of truth for it from then on, paused or not. + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + assert task_id not in worker._pending_tasks + assert task_id not in worker._completed_tasks + current = worker._current + assert current is not None + assert current.task_id == task_id + assert worker.get_task_by_id(task_id) is current + + +def test_worker_state_is_paused(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + assert worker.state == WorkerState.PAUSED + + +def test_resume_after_pause_completes_task(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + complete_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.is_complete(), + ) + worker.resume() + events = complete_future.result(timeout=5.0) + + assert events[-1].task_status is not None + assert events[-1].task_status.task_complete + assert isinstance(events[-1].task_status.result, TaskResult) + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + + +def test_cancel_active_task_abort(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + cancel_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: ( + e.state == WorkerState.IDLE + and e.task_status is not None + and e.task_status.task_complete + ), + ) + worker.cancel_active_task(failure=True) + events = cancel_future.result(timeout=5.0) + + assert events[-1].errors == ["Task aborted"] + assert events[-1].task_status is not None + assert events[-1].task_status.task_failed + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + + +def test_cancel_active_task_graceful(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + cancel_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: ( + e.state == WorkerState.IDLE + and e.task_status is not None + and e.task_status.task_complete + ), + ) + worker.cancel_active_task(failure=False) + events = cancel_future.result(timeout=5.0) + + assert events[-1].errors == [] + assert events[-1].task_status is not None + assert events[-1].task_status.task_complete + assert isinstance(events[-1].task_status.result, TaskResult) + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + + +def test_cancel_running_task_records_failure(worker: TaskWorker) -> None: + # _LONG_TASK uses asyncio.sleep so the RE event loop is free to process the + # abort signal — unlike FakeDevice which blocks the loop with a sync wait. + task_id = worker.submit_task(_LONG_TASK) + + running_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.RUNNING, + ) + worker.begin_task(task_id) + running_future.result(timeout=5.0) + + cancel_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.is_complete(), + ) + worker.cancel_active_task(failure=True, reason="mid-run abort") + events = cancel_future.result(timeout=5.0) + + assert events[-1].task_status is not None + assert events[-1].task_status.task_failed + assert isinstance(events[-1].task_status.result, TaskError) + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + + +def test_cancel_running_task_gracefully(worker: TaskWorker) -> None: + # cancel_active_task(failure=False) on a running (not paused) task must + # not be reported as a failure. + task_id = worker.submit_task(_LONG_TASK) + + running_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.RUNNING, + ) + worker.begin_task(task_id) + running_future.result(timeout=5.0) + + cancel_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.is_complete(), + ) + worker.cancel_active_task(failure=False) + events = cancel_future.result(timeout=5.0) + + assert events[-1].errors == [] + assert events[-1].task_status is not None + assert not events[-1].task_status.task_failed + assert isinstance(events[-1].task_status.result, TaskResult) + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + + +def test_cancel_active_task_does_not_overwrite_existing_outcome( + worker: TaskWorker, +) -> None: + task_id = worker.submit_task(_LONG_TASK) + + running_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.RUNNING, + ) + worker.begin_task(task_id) + running_future.result(timeout=5.0) + + current = worker._current + assert current is not None + with worker._status_lock: + current.set_exception(Exception("worker thread got there first")) + + worker.cancel_active_task(failure=True, reason="caller's reason") + + assert isinstance(current.outcome, TaskError) + assert current.outcome.message == "worker thread got there first" + + +def test_cancel_active_task_blocks_caller_until_paused_cleanup_finishes( + worker: TaskWorker, +) -> None: + worker._ctx.register_plan(pausing_plan_with_slow_cleanup) + task_id = worker.submit_task(_PAUSING_TASK_SLOW_CLEANUP) + + begin_task_and_wait_until_paused(worker, task_id) + + start = time.monotonic() + worker.cancel_active_task(failure=True) + elapsed = time.monotonic() - start + + assert elapsed >= _ABORT_CLEANUP_DELAY, ( + f"cancel_active_task(failure=True) returned after only {elapsed:.2f}s, " + f"before the RunEngine's {_ABORT_CLEANUP_DELAY}s abort cleanup could " + "have finished." + ) + assert task_id in worker._completed_tasks + assert worker._completed_tasks[task_id].outcome is not None + assert isinstance(worker._completed_tasks[task_id].outcome, TaskError) + + +def test_concurrent_resume_and_cancel_do_not_corrupt_state( + worker: TaskWorker, +) -> None: + worker._ctx.register_plan(pausing_plan_with_slow_cleanup) + task_id = worker.submit_task(_PAUSING_TASK_SLOW_CLEANUP) + + begin_task_and_wait_until_paused(worker, task_id) + + with ThreadPoolExecutor(1) as executor: + resume_future = executor.submit(worker.resume) + time.sleep(_ABORT_CLEANUP_DELAY / 4) + worker.cancel_active_task(failure=True, reason="changed my mind") + resume_future.result(timeout=5.0) + + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + completed = worker._completed_tasks[task_id] + assert completed.outcome is not None + + +def test_second_cancel_while_first_still_queued_supersedes_it( + worker: TaskWorker, +) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + complete_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.is_complete(), + ) + first_signal = CancelSignal(failure=False, reason="first, graceful") + worker._pending_cancel = first_signal + worker._task_channel.put_nowait(first_signal) + second_signal = CancelSignal(failure=True, reason="second, urgent") + worker._pending_cancel = second_signal + events = complete_future.result(timeout=5.0) + + assert events[-1].task_status is not None + assert events[-1].task_status.task_failed + assert isinstance(events[-1].task_status.result, TaskError) + assert events[-1].errors == ["second, urgent"] + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + + +def test_stop_while_paused_completes_task(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + + begin_task_and_wait_until_paused(worker, task_id) + + worker.stop() + + assert task_id in worker._completed_tasks + assert task_id not in worker._pending_tasks + assert worker._completed_tasks[task_id].is_complete + assert worker._completed_tasks[task_id].outcome is not None + + +def test_resume_when_not_paused_does_nothing(worker: TaskWorker) -> None: + task_id = worker.submit_task(_SIMPLE_TASK) + begin_task_and_wait_until_complete(worker, task_id) + + complete_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.IDLE, + ) + worker.resume() + events = complete_future.result(timeout=5.0) + + assert all(e.errors == [] for e in events) + + +def test_resume_re_pauses_when_plan_pauses_again(worker: TaskWorker) -> None: + # If the plan pauses again immediately on resume, that's not a failure + # and must not leave the task half-finished. + worker._ctx.register_plan(twice_pausing_plan) + task_id = worker.submit_task(_TWICE_PAUSING_TASK) + begin_task_and_wait_until_paused(worker, task_id) + + paused_again_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.state == WorkerState.PAUSED, + ) + worker.resume() + paused_again_future.result(timeout=5.0) + + assert worker.state == WorkerState.PAUSED + current = worker._current + assert current is not None + assert current.outcome is None + assert task_id not in worker._pending_tasks + assert task_id not in worker._completed_tasks + + +def test_can_run_task_after_resume(worker: TaskWorker) -> None: + worker._ctx.register_plan(pausing_plan) + task_id = worker.submit_task(_PAUSING_TASK) + begin_task_and_wait_until_paused(worker, task_id) + + complete_future: Future[list[WorkerEvent]] = take_events( + worker.worker_events, + lambda e: e.is_complete(), + ) + worker.resume() + complete_future.result(timeout=5.0) + + task_id_2 = worker.submit_task(_SIMPLE_TASK) + events = begin_task_and_wait_until_complete(worker, task_id_2) + assert events[-1].task_status is not None + assert events[-1].task_status.task_complete