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..5497024054 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) -> RunEngineResult | tuple[str, ...]: LOGGER.info( f"Asked to run plan {self.name} with {self.params} and " f"metadata {self.metadata} for all runs" @@ -38,11 +39,7 @@ def do_task(self, ctx: BlueskyContext) -> None: func = ctx.plan_functions[self.name] 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 - return result.plan_result + return ctx.run_engine(func(**prepared_params)) def _lookup_params(ctx: BlueskyContext, task: Task) -> BaseModel: diff --git a/src/blueapi/worker/task_worker.py b/src/blueapi/worker/task_worker.py index 4ada3c2a54..046e18c142 100644 --- a/src/blueapi/worker/task_worker.py +++ b/src/blueapi/worker/task_worker.py @@ -5,11 +5,13 @@ 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._vendor.super_state_machine.errors import TransitionError from bluesky.protocols import Status +from bluesky.run_engine import RunEngineResult +from bluesky.utils import RunEngineInterrupted from observability_utils.tracing import ( add_span_attributes, get_tracer, @@ -19,7 +21,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 +70,22 @@ 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 + @property + def is_complete(self) -> bool: + return self.outcome is not None + + @model_validator(mode="before") + @classmethod + def remove_complete(cls, values: dict[str, Any]) -> dict[str, Any]: + # fix pydantic falling over itself with computed fields + values.pop("is_complete", None) + return values + def set_result(self, result: Any): self.outcome = TaskResult.from_result(result) @@ -110,6 +123,7 @@ class TaskWorker: _task_channel: Queue # type: ignore _current: TrackableTask | None _status_lock: RLock + _state_change: Condition _status_snapshot: dict[str, StatusView] _completed_statuses: set[str] _worker_events: EventPublisher[WorkerEvent] @@ -143,6 +157,7 @@ def __init__( 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,22 +195,32 @@ 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)) + self._state_change.wait_for(lambda: self._state == WorkerState.IDLE) 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 + if failure: + current.outcome = TaskError(type="Abort", message=reason) + self._ctx.run_engine.abort(reason=reason) + else: + current.outcome = TaskResult.from_result(None) + self._ctx.run_engine.stop() + return current.task_id @start_as_current_span(TRACER, "task_id") - def get_task_by_id(self, task_id: str) -> TrackableTask | None: + def get_task_by_id(self, task_id: str) -> TrackableTask: """ Returns a task matching the task ID supplied, if the worker knows of it. @@ -205,6 +230,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 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 +246,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 else [] elif status == TaskStatusEnum.PENDING: return [task for task in self._pending_tasks.values() if task.is_pending] elif status == TaskStatusEnum.COMPLETE: @@ -315,8 +340,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") @@ -410,7 +436,10 @@ 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() + with self._state_change: + self._task_channel.put(ResumeSignal()) + self._state_change.wait_for(lambda: self._state != WorkerState.PAUSED) @start_as_current_span(TRACER) def _cycle_with_error_handling(self) -> None: @@ -423,20 +452,50 @@ def _cycle_with_error_handling(self) -> None: def _cycle(self) -> None: try: LOGGER.info("Awaiting task") - next_task: TrackableTask | KillSignal = self._task_channel.get() + next_task: TrackableTask | KillSignal | ResumeSignal | AbortSignal = ( + self._task_channel.get() + ) + if isinstance(next_task, ResumeSignal): + if self._current is None or self._current.is_complete: + LOGGER.debug( + "Ignoring resume signal - nothing to resume: current=%s", + self._current, + ) + return + # If there is an incomplete task, treat it as the next task to run + next_task = self._current + if isinstance(next_task, TrackableTask): def process_task(): LOGGER.info(f"Got new task: {next_task}") self._current = next_task - self._current.is_pending = False meta = {"task_id": self._current.task_id} try: - result = self._current.task.do_task(self._ctx) - LOGGER.info( - "Task ran successfully - returned: %s", result, extra=meta - ) - self._current.set_result(result) + if self._current.is_pending: + self._current.is_pending = False + result = self._current.task.do_task(self._ctx) + else: + LOGGER.debug("Resuming previous task") + result = self._ctx.run_engine.resume() + if isinstance(result, RunEngineResult): + # Should always be a RunEngineResult if the run + # engine is configured correctly + self._current.set_result(result.plan_result) + LOGGER.info( + "Task ran successfully - returned: %s", + result.plan_result, + extra=meta, + ) + else: + LOGGER.warn( + "Task ran successfully but did not return " + "RunEngineResult. Is `call_returns_result` set?" + ) + except RunEngineInterrupted as rei: + LOGGER.info("Task interrupted", extra=meta) + if self._ctx.run_engine.state != "paused": + self._report_error(rei) except Exception as e: LOGGER.error("Task failed", extra=meta) self._current.set_exception(e) @@ -457,6 +516,16 @@ def process_task(): process_task() else: process_task() + elif isinstance(next_task, AbortSignal): + if self._current and not self._current.is_complete: + if next_task.failure: + result = self._ctx.run_engine.abort(reason=next_task.reason) + self._current.outcome = TaskError( + type="Abort", message=next_task.reason + ) + else: + result = self._ctx.run_engine.stop() + self._current.set_result(result.plan_result) elif isinstance(next_task, KillSignal): # If we receive a kill signal we begin to shut the worker down. @@ -472,9 +541,7 @@ def process_task(): if self._current_task_otel_context is not None: 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) + if self._current is not None and self._current.is_complete: self._completed_tasks[self._current.task_id] = self._current self._report_status() self._errors.clear() @@ -514,14 +581,18 @@ def _on_state_change( raw_new_state: RawRunEngineState, raw_old_state: RawRunEngineState | None = None, ) -> None: - new_state = WorkerState.from_bluesky_state(raw_new_state) - if raw_old_state: - old_state = WorkerState.from_bluesky_state(raw_old_state) - else: - old_state = WorkerState.UNKNOWN - LOGGER.debug(f"Notifying state change {old_state} -> {new_state}") - self._state = new_state - self._report_status() + print("Start of _on_state_change") + with self._state_change: + new_state = WorkerState.from_bluesky_state(raw_new_state) + if raw_old_state: + old_state = WorkerState.from_bluesky_state(raw_old_state) + else: + old_state = WorkerState.UNKNOWN + LOGGER.debug(f"Notifying state change {old_state} -> {new_state}") + self._state = new_state + self._state_change.notify() + self._report_status() + print("End of _on_state_change") def _report_error(self, err: Exception) -> None: LOGGER.error(err, exc_info=True) @@ -683,6 +754,15 @@ class KillSignal: ... +class ResumeSignal: ... + + +@dataclass +class AbortSignal: + failure: bool + reason: str + + def run_worker_in_own_thread( worker: TaskWorker, executor: ThreadPoolExecutor | None = None ) -> Future: diff --git a/tests/unit_tests/service/test_rest_api.py b/tests/unit_tests/service/test_rest_api.py index dedd9d7ff3..16bc4e2970 100644 --- a/tests/unit_tests/service/test_rest_api.py +++ b/tests/unit_tests/service/test_rest_api.py @@ -45,7 +45,7 @@ from blueapi.service.runner import WorkerDispatcher from blueapi.worker.event import WorkerState from blueapi.worker.task import Task -from blueapi.worker.task_worker import TrackableTask +from blueapi.worker.task_worker import TaskResult, TrackableTask class MockCountModel(BaseModel): ... @@ -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..a3811ac5a6 100644 --- a/tests/unit_tests/worker/test_task_worker.py +++ b/tests/unit_tests/worker/test_task_worker.py @@ -53,6 +53,8 @@ }, ) +_SUCCESS = TaskResult(type="int", result=42) + class FakeDevice(Movable[float]): event: threading.Event @@ -629,21 +631,19 @@ def callback(unused: Future[list[Any]], stream=stream, sub=sub): ], ) def test_get_tasks(worker: TaskWorker, status, expected_task_ids): + worker._current = TrackableTask( + task_id="task1", + task=Task(name="set_absolute", params={"movable": "fake_device", "value": 4.0}), + outcome=None, + 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, + outcome=_SUCCESS, is_pending=True, ), } @@ -653,7 +653,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=_SUCCESS, is_pending=False, ), } @@ -668,7 +668,7 @@ 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, is_pending=False, outcome=_SUCCESS ) ) @@ -805,8 +805,6 @@ def test_cycle_without_otel_context(mock_logger: Mock, inert_worker: TaskWorker) 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"} )