From ff91bf55622a8c75362bfe29af606ac458493d6f Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 15:01:37 +0000 Subject: [PATCH 1/9] Support *args in plan args --- src/blueapi/worker/__init__.py | 3 +- src/blueapi/worker/task.py | 87 +++++++++++++++------ tests/unit_tests/service/test_rest_api.py | 17 ++-- tests/unit_tests/worker/test_task_worker.py | 54 +++++++------ 4 files changed, 106 insertions(+), 55 deletions(-) diff --git a/src/blueapi/worker/__init__.py b/src/blueapi/worker/__init__.py index 85ae49b453..5002da8a8c 100644 --- a/src/blueapi/worker/__init__.py +++ b/src/blueapi/worker/__init__.py @@ -1,11 +1,12 @@ from .event import ProgressEvent, StatusView, TaskStatus, WorkerEvent, WorkerState -from .task import Task +from .task import Task, TaskParams from .task_worker import TaskWorker, TrackableTask from .worker_errors import WorkerAlreadyStartedError, WorkerBusyError __all__ = [ "TaskWorker", "Task", + "TaskParams", "WorkerEvent", "WorkerState", "StatusView", diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index e2efade0d7..648f442947 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -1,8 +1,9 @@ import logging from collections.abc import Mapping +from inspect import Parameter, signature from typing import Any -from pydantic import BaseModel, Field, TypeAdapter +from pydantic import Field, TypeAdapter from blueapi.core import BlueskyContext from blueapi.utils import BlueapiBaseModel @@ -10,55 +11,93 @@ LOGGER = logging.getLogger(__name__) +class TaskParams(BlueapiBaseModel): + args: tuple[Any, ...] = () + kwargs: Mapping[str, Any] = Field(default_factory=dict) + + class Task(BlueapiBaseModel): """ Task that will run a plan """ name: str = Field(description="Name of plan to run") - params: Mapping[str, Any] = Field( - description="Values for parameters to plan, if any", default_factory=dict + params: TaskParams = Field( + description="Values for parameters to plan, if any", default_factory=TaskParams ) metadata: dict[str, Any] = Field( description="Any metadata to apply to all runs within this task", default_factory=dict, ) - def prepare_params(self, ctx: BlueskyContext) -> Mapping[str, Any]: - model = _lookup_params(ctx, self) - # Re-create dict manually to avoid nesting in model_dump output - return {field: getattr(model, field) for field in model.__pydantic_fields__} + def prepare_params( + self, ctx: BlueskyContext + ) -> tuple[tuple[Any, ...], Mapping[str, Any]]: + return _lookup_params(ctx, self) def do_task(self, ctx: BlueskyContext) -> None: LOGGER.info( f"Asked to run plan {self.name} with {self.params} and " f"metadata {self.metadata} for all runs" ) - - func = ctx.plan_functions[self.name] - prepared_params = self.prepare_params(ctx) + plan = ctx.plan_functions[self.name] + prepared_args, prepared_kwargs = self.prepare_params(ctx) ctx.run_engine.md.update(self.metadata) - result = ctx.run_engine(func(**prepared_params)) + result = ctx.run_engine(plan(*prepared_args, **prepared_kwargs)) if isinstance(result, tuple): # pragma: no cover # this is never true if the run_engine is configured correctly return None return result.plan_result -def _lookup_params(ctx: BlueskyContext, task: Task) -> BaseModel: +def _lookup_params( + ctx: BlueskyContext, task: Task +) -> tuple[tuple[Any, ...], Mapping[str, Any]]: + """ + Validate and prepare the arguments for a plan. """ - Checks plan parameters against context + plan = ctx.plans[task.name] + func = ctx.plan_functions[task.name] - Args: - ctx: Context holding plans and devices - plan: Plan object including schema - params: Parameter values to be validated against schema + sig = signature(func) + bound = sig.bind(*task.params.args, **task.params.kwargs) - Returns: - Mapping[str, Any]: _description_ - """ + # Only validate explicitly supplied arguments. This allows Pydantic's + # default/default_factory to provide injected defaults. + adapter = TypeAdapter(plan.model) + validated = adapter.validate_python(bound.arguments) - plan = ctx.plans[task.name] - model = plan.model - adapter = TypeAdapter(model) - return adapter.validate_python(task.params) + args: list[Any] = [] + kwargs: dict[str, Any] = {} + + for name, parameter in sig.parameters.items(): + supplied = name in bound.arguments + + if supplied: + value = getattr(validated, name) + + match parameter.kind: + case Parameter.POSITIONAL_ONLY: + args.append(value) + + case Parameter.POSITIONAL_OR_KEYWORD: + if name in task.params.kwargs: + kwargs[name] = value + else: + args.append(value) + + case Parameter.VAR_POSITIONAL: + args.extend(value) + + case Parameter.KEYWORD_ONLY: + kwargs[name] = value + + case Parameter.VAR_KEYWORD: + kwargs.update(value) + + else: + # Let the generated Pydantic model provide the default. + value = getattr(validated, name) + kwargs[name] = value + + return tuple(args), kwargs diff --git a/tests/unit_tests/service/test_rest_api.py b/tests/unit_tests/service/test_rest_api.py index a10d5ab3cf..9b6bedf1ba 100644 --- a/tests/unit_tests/service/test_rest_api.py +++ b/tests/unit_tests/service/test_rest_api.py @@ -47,7 +47,7 @@ ) from blueapi.service.runner import WorkerDispatcher from blueapi.worker.event import TaskStatus, WorkerEvent, WorkerState -from blueapi.worker.task import Task +from blueapi.worker.task import Task, TaskParams from blueapi.worker.task_worker import TrackableTask @@ -395,7 +395,10 @@ def test_put_plan_fails_if_not_idle(mock_runner: Mock, client: TestClient) -> No def test_get_tasks(mock_runner: Mock, client: TestClient) -> None: tasks = [ - TrackableTask(task_id="0", task=Task(name="sleep", params={"time": 0.0})), + TrackableTask( + task_id="0", + task=Task(name="sleep", params=TaskParams(kwargs={"time": 0.0})), + ), TrackableTask( task_id="1", task=Task(name="first_task"), @@ -418,7 +421,7 @@ def test_get_tasks(mock_runner: Mock, client: TestClient) -> None: "request_id": None, "task": { "name": "sleep", - "params": {"time": 0.0}, + "params": {"args": [], "kwargs": {"time": 0.0}}, "metadata": {}, }, "outcome": None, @@ -431,7 +434,7 @@ def test_get_tasks(mock_runner: Mock, client: TestClient) -> None: "request_id": None, "task": { "name": "first_task", - "params": {}, + "params": {"args": [], "kwargs": {}}, "metadata": {}, }, "outcome": None, @@ -463,7 +466,7 @@ def test_get_tasks_by_status(mock_runner: Mock, client: TestClient) -> None: "request_id": None, "task": { "name": "third_task", - "params": {}, + "params": {"args": [], "kwargs": {}}, "metadata": {}, }, "outcome": None, @@ -627,7 +630,7 @@ def test_get_task(mock_runner: Mock, client: TestClient): "request_id": None, "task": { "name": "third_task", - "params": {}, + "params": {"args": [], "kwargs": {}}, "metadata": { "foo": "bar", }, @@ -674,7 +677,7 @@ def test_get_all_tasks(mock_runner: Mock, client: TestClient): "task_id": task_id, "task": { "name": "third_task", - "params": {}, + "params": {"args": [], "kwargs": {}}, "metadata": {}, }, "is_complete": False, diff --git a/tests/unit_tests/worker/test_task_worker.py b/tests/unit_tests/worker/test_task_worker.py index 5e6553d2fe..fe1383ed86 100644 --- a/tests/unit_tests/worker/test_task_worker.py +++ b/tests/unit_tests/worker/test_task_worker.py @@ -27,6 +27,7 @@ from blueapi.utils.base_model import BlueapiBaseModel from blueapi.worker import ( Task, + TaskParams, TaskStatus, TaskWorker, TrackableTask, @@ -37,16 +38,16 @@ ) from blueapi.worker.event import TaskResult, TaskStatusEnum -_SIMPLE_TASK = Task(name="sleep", params={"time": 0.0}) -_LONG_TASK = Task(name="sleep", params={"time": 1.0}) +_SIMPLE_TASK = Task(name="sleep", params=TaskParams(kwargs={"time": 0.0})) +_LONG_TASK = Task(name="sleep", params=TaskParams(kwargs={"time": 1.0})) _INDEFINITE_TASK = Task( name="set_absolute", - params={"movable": "fake_device", "value": 4.0}, + params=TaskParams(kwargs={"movable": "fake_device", "value": 4.0}), ) -_FAILING_TASK = Task(name="failing_plan", params={}) +_FAILING_TASK = Task(name="failing_plan", params=TaskParams()) _TASK_WITH_METADATA = Task( name="sleep", - params={"time": 0.0}, + params=TaskParams(kwargs={"time": 0.0}), metadata={ "foo": "bar", "baz": 0, @@ -522,7 +523,9 @@ def assert_running_count_plan_produces_ordered_worker_and_data_events( task: Task | None = None, timeout: float = 5.0, ) -> None: - default_task = Task(name="count", params={"detectors": ["motor"], "num": 1}) + default_task = Task( + name="count", params=TaskParams(kwargs={"detectors": ["motor"], "num": 1}) + ) task = task or default_task event_streams: list[EventStream[Any, int]] = [ @@ -633,7 +636,8 @@ def test_get_tasks(worker: TaskWorker, status, expected_task_ids): "task1": TrackableTask( task_id="task1", task=Task( - name="set_absolute", params={"movable": "fake_device", "value": 4.0} + name="set_absolute", + params=TaskParams(kwargs={"movable": "fake_device", "value": 4.0}), ), is_complete=False, is_pending=False, @@ -641,7 +645,8 @@ def test_get_tasks(worker: TaskWorker, status, expected_task_ids): "task2": TrackableTask( task_id="task2", task=Task( - name="set_absolute", params={"movable": "fake_device", "value": 4.0} + name="set_absolute", + params=TaskParams(kwargs={"movable": "fake_device", "value": 4.0}), ), is_complete=False, is_pending=True, @@ -651,7 +656,8 @@ def test_get_tasks(worker: TaskWorker, status, expected_task_ids): "task3": TrackableTask( task_id="task3", task=Task( - name="set_absolute", params={"movable": "fake_device", "value": 4.0} + name="set_absolute", + params=TaskParams(kwargs={"movable": "fake_device", "value": 4.0}), ), is_complete=True, is_pending=False, @@ -737,8 +743,8 @@ def injected_device_plan( yield from () context.register_plan(injected_device_plan) - params = Task(name="injected_device_plan").prepare_params(context) - assert params["dev"] == fake_device + args, kwargs = Task(name="injected_device_plan").prepare_params(context) + assert kwargs["dev"] == fake_device def test_injected_devices_plan_model( @@ -853,9 +859,9 @@ def test_injected_composite_devices_are_found( context: BlueskyContext, ): context.register_plan(injected_device_plan) - params = Task(name="injected_device_plan").prepare_params(context) - assert params["composite"].fake_device == fake_device - assert params["composite"].second_fake_device == second_fake_device + args, kwargs = Task(name="injected_device_plan").prepare_params(context) + assert kwargs["composite"].fake_device == fake_device + assert kwargs["composite"].second_fake_device == second_fake_device def test_injected_composite_devices_plan_model( @@ -874,9 +880,9 @@ def test_injected_composite_with_pydantic_dataclass( second_fake_device: FakeDevice, ): context.register_plan(injected_dataclass_device_plan) - params = Task(name="injected_dataclass_device_plan").prepare_params(context) - assert params["composite"].fake_device == fake_device - assert params["composite"].second_fake_device == second_fake_device + args, kwargs = Task(name="injected_dataclass_device_plan").prepare_params(context) + assert kwargs["composite"].fake_device == fake_device + assert kwargs["composite"].second_fake_device == second_fake_device def test_injected_composite_with_standard_dataclass( @@ -885,11 +891,11 @@ def test_injected_composite_with_standard_dataclass( second_fake_device: FakeDevice, ): context.register_plan(injected_standard_dataclass_device_plan) - params = Task(name="injected_standard_dataclass_device_plan").prepare_params( + args, kwargs = Task(name="injected_standard_dataclass_device_plan").prepare_params( context ) - assert params["composite"].fake_device == fake_device - assert params["composite"].second_fake_device == second_fake_device + assert kwargs["composite"].fake_device == fake_device + assert kwargs["composite"].second_fake_device == second_fake_device def test_plan_module_with_composite_devices_can_be_loaded_before_device_module( @@ -900,9 +906,11 @@ def test_plan_module_with_composite_devices_can_be_loaded_before_device_module( context_without_devices.register_plan(injected_device_plan) context_without_devices.register_device(fake_device) context_without_devices.register_device(second_fake_device) - params = Task(name="injected_device_plan").prepare_params(context_without_devices) - assert params["composite"].fake_device == fake_device - assert params["composite"].second_fake_device == second_fake_device + args, kwargs = Task(name="injected_device_plan").prepare_params( + context_without_devices + ) + assert kwargs["composite"].fake_device == fake_device + assert kwargs["composite"].second_fake_device == second_fake_device @pytest.mark.parametrize( From fc09b961448b6a80eabb7cbc5a91be67e1048a59 Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 15:20:32 +0000 Subject: [PATCH 2/9] Add tests to cover *args case --- src/blueapi/worker/task.py | 10 +-- tests/unit_tests/worker/test_task_worker.py | 70 +++++++++++++++++++++ 2 files changed, 75 insertions(+), 5 deletions(-) diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index 648f442947..5c39508194 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -1,5 +1,5 @@ import logging -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from inspect import Parameter, signature from typing import Any @@ -12,7 +12,7 @@ class TaskParams(BlueapiBaseModel): - args: tuple[Any, ...] = () + args: Sequence[Any] = [] kwargs: Mapping[str, Any] = Field(default_factory=dict) @@ -32,7 +32,7 @@ class Task(BlueapiBaseModel): def prepare_params( self, ctx: BlueskyContext - ) -> tuple[tuple[Any, ...], Mapping[str, Any]]: + ) -> tuple[list[Any], Mapping[str, Any]]: return _lookup_params(ctx, self) def do_task(self, ctx: BlueskyContext) -> None: @@ -52,7 +52,7 @@ def do_task(self, ctx: BlueskyContext) -> None: def _lookup_params( ctx: BlueskyContext, task: Task -) -> tuple[tuple[Any, ...], Mapping[str, Any]]: +) -> tuple[list[Any], Mapping[str, Any]]: """ Validate and prepare the arguments for a plan. """ @@ -100,4 +100,4 @@ def _lookup_params( value = getattr(validated, name) kwargs[name] = value - return tuple(args), kwargs + return args, kwargs diff --git a/tests/unit_tests/worker/test_task_worker.py b/tests/unit_tests/worker/test_task_worker.py index fe1383ed86..6df6a3b839 100644 --- a/tests/unit_tests/worker/test_task_worker.py +++ b/tests/unit_tests/worker/test_task_worker.py @@ -970,3 +970,73 @@ def test_worker_event_task_id(): def test_worker_event_no_task_id(): event = WorkerEvent(state=WorkerState.IDLE, task_status=None) assert event.task_id is None + + +def test_task_worker_passes_positional_args( + context: BlueskyContext, +) -> None: + + def positional_plan(value: int) -> MsgGenerator: + yield from () + + context.register_plan(positional_plan) + + task = Task(name="positional_plan", params=TaskParams(args=[42])) + args, kwargs = task.prepare_params(context) + + assert args == [42] + assert kwargs == {} + + +def test_task_worker_passes_multiple_positional_args( + context: BlueskyContext, +) -> None: + + def positional_plan(first: int, second: int) -> MsgGenerator: + yield from () + + context.register_plan(positional_plan) + task = Task(name="positional_plan", params=TaskParams(args=[1, 2])) + args, kwargs = task.prepare_params(context) + + assert args == [1, 2] + assert kwargs == {} + + +def test_task_worker_passes_positional_and_keyword_args( + context: BlueskyContext, +) -> None: + + def mixed_plan( + first: int, + second: int, + *, + third: int, + ) -> MsgGenerator: + yield from () + + context.register_plan(mixed_plan) + task = Task( + name="mixed_plan", + params=TaskParams(args=[1, 2], kwargs={"third": 3}), + ) + args, kwargs = task.prepare_params(context) + assert args == [1, 2, 3] + assert kwargs == {} + + +def test_task_worker_resolves_positional_device( + context: BlueskyContext, + fake_device: FakeDevice, +) -> None: + def positional_device_plan(device: FakeDevice) -> MsgGenerator: + yield from () + + context.register_plan(positional_device_plan) + + task = Task( + name="positional_device_plan", params=TaskParams(args=[fake_device.name]) + ) + args, kwargs = task.prepare_params(context) + assert args == [fake_device] + assert kwargs == {} From d6388843b6573839626ace0bc8d9b64fbb90269d Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 15:27:15 +0000 Subject: [PATCH 3/9] Simplify logic --- src/blueapi/worker/task.py | 108 +++++++++++++++++++------------------ 1 file changed, 55 insertions(+), 53 deletions(-) diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index 5c39508194..7e9662cd30 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -33,71 +33,73 @@ class Task(BlueapiBaseModel): def prepare_params( self, ctx: BlueskyContext ) -> tuple[list[Any], Mapping[str, Any]]: - return _lookup_params(ctx, self) + """ + Checks plan parameters against context - def do_task(self, ctx: BlueskyContext) -> None: - LOGGER.info( - f"Asked to run plan {self.name} with {self.params} and " - f"metadata {self.metadata} for all runs" - ) - plan = ctx.plan_functions[self.name] - prepared_args, prepared_kwargs = self.prepare_params(ctx) - ctx.run_engine.md.update(self.metadata) - result = ctx.run_engine(plan(*prepared_args, **prepared_kwargs)) - if isinstance(result, tuple): # pragma: no cover - # this is never true if the run_engine is configured correctly - return None - return result.plan_result + Args: + ctx: Context holding plans and devices + plan: Plan object including schema + params: Parameter values to be validated against schema + Returns: + Mapping[str, Any]: _description_ + """ + plan = ctx.plans[self.name] + func = ctx.plan_functions[self.name] -def _lookup_params( - ctx: BlueskyContext, task: Task -) -> tuple[list[Any], Mapping[str, Any]]: - """ - Validate and prepare the arguments for a plan. - """ - plan = ctx.plans[task.name] - func = ctx.plan_functions[task.name] + sig = signature(func) + bound = sig.bind(*self.params.args, **self.params.kwargs) + + # Only validate explicitly supplied arguments. This allows Pydantic's + # default/default_factory to provide injected defaults. + adapter = TypeAdapter(plan.model) + validated = adapter.validate_python(bound.arguments) - sig = signature(func) - bound = sig.bind(*task.params.args, **task.params.kwargs) + args: list[Any] = [] + kwargs: dict[str, Any] = {} - # Only validate explicitly supplied arguments. This allows Pydantic's - # default/default_factory to provide injected defaults. - adapter = TypeAdapter(plan.model) - validated = adapter.validate_python(bound.arguments) + for name, parameter in sig.parameters.items(): + supplied = name in bound.arguments - args: list[Any] = [] - kwargs: dict[str, Any] = {} + if supplied: + value = getattr(validated, name) - for name, parameter in sig.parameters.items(): - supplied = name in bound.arguments + match parameter.kind: + case Parameter.POSITIONAL_ONLY: + args.append(value) - if supplied: - value = getattr(validated, name) + case Parameter.POSITIONAL_OR_KEYWORD: + if name in self.params.kwargs: + kwargs[name] = value + else: + args.append(value) - match parameter.kind: - case Parameter.POSITIONAL_ONLY: - args.append(value) + case Parameter.VAR_POSITIONAL: + args.extend(value) - case Parameter.POSITIONAL_OR_KEYWORD: - if name in task.params.kwargs: + case Parameter.KEYWORD_ONLY: kwargs[name] = value - else: - args.append(value) - case Parameter.VAR_POSITIONAL: - args.extend(value) + case Parameter.VAR_KEYWORD: + kwargs.update(value) - case Parameter.KEYWORD_ONLY: - kwargs[name] = value + else: + # Let the generated Pydantic model provide the default. + value = getattr(validated, name) + kwargs[name] = value - case Parameter.VAR_KEYWORD: - kwargs.update(value) + return args, kwargs - else: - # Let the generated Pydantic model provide the default. - value = getattr(validated, name) - kwargs[name] = value - - return args, kwargs + def do_task(self, ctx: BlueskyContext) -> None: + LOGGER.info( + f"Asked to run plan {self.name} with {self.params} and " + f"metadata {self.metadata} for all runs" + ) + plan = ctx.plan_functions[self.name] + prepared_args, prepared_kwargs = self.prepare_params(ctx) + ctx.run_engine.md.update(self.metadata) + result = ctx.run_engine(plan(*prepared_args, **prepared_kwargs)) + if isinstance(result, tuple): # pragma: no cover + # this is never true if the run_engine is configured correctly + return None + return result.plan_result From d3d45e107f42274860b2a7e73e0957bd7bfc5ed3 Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 15:27:19 +0000 Subject: [PATCH 4/9] Fix test --- tests/unit_tests/worker/test_task_worker.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/unit_tests/worker/test_task_worker.py b/tests/unit_tests/worker/test_task_worker.py index 6df6a3b839..b4f2460af5 100644 --- a/tests/unit_tests/worker/test_task_worker.py +++ b/tests/unit_tests/worker/test_task_worker.py @@ -1021,8 +1021,8 @@ def mixed_plan( params=TaskParams(args=[1, 2], kwargs={"third": 3}), ) args, kwargs = task.prepare_params(context) - assert args == [1, 2, 3] - assert kwargs == {} + assert args == [1, 2] + assert kwargs == {"third": 3} def test_task_worker_resolves_positional_device( From 44ce2fa5ccefa6feb5074d76c5c4ef424ca58c17 Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 15:29:31 +0000 Subject: [PATCH 5/9] Use default_factory for args --- src/blueapi/worker/task.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index 7e9662cd30..4c64698930 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -12,7 +12,7 @@ class TaskParams(BlueapiBaseModel): - args: Sequence[Any] = [] + args: Sequence[Any] = Field(default_factory=lambda: []) kwargs: Mapping[str, Any] = Field(default_factory=dict) From 32cc60e2e9274952bef8d8717b1a3af79eaca507 Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 16:24:29 +0000 Subject: [PATCH 6/9] Propergate change --- src/blueapi/cli/cli.py | 6 +- src/blueapi/client/client.py | 32 +++++------ src/blueapi/service/main.py | 4 +- src/blueapi/service/model.py | 8 +-- tests/system_tests/test_blueapi_system.py | 69 ++++++++++++----------- tests/unit_tests/cli/test_cli.py | 11 ++-- tests/unit_tests/client/test_client.py | 46 ++++++++++++--- tests/unit_tests/client/test_rest.py | 12 ++-- tests/unit_tests/service/test_rest_api.py | 12 ++-- 9 files changed, 123 insertions(+), 77 deletions(-) diff --git a/src/blueapi/cli/cli.py b/src/blueapi/cli/cli.py index c3f9bde13f..a3a9e56821 100644 --- a/src/blueapi/cli/cli.py +++ b/src/blueapi/cli/cli.py @@ -39,7 +39,7 @@ from blueapi.log import set_up_logging from blueapi.service.authentication import SessionCacheManager, SessionManager from blueapi.service.model import DeviceResponse, PlanResponse, SourceInfo, TaskRequest -from blueapi.worker import ProgressEvent, WorkerEvent +from blueapi.worker import ProgressEvent, TaskParams, WorkerEvent from blueapi.worker.event import TaskError, TaskResult from .scratch import setup_scratch @@ -374,7 +374,9 @@ def run_plan( client = cast(BlueapiClient, obj["client"]) task = TaskRequest( - name=name, params=parameters, instrument_session=instrument_session + name=name, + params=TaskParams(kwargs=parameters), + instrument_session=instrument_session, ) try: diff --git a/src/blueapi/client/client.py b/src/blueapi/client/client.py index 66a79d8dd6..3d77a93daa 100644 --- a/src/blueapi/client/client.py +++ b/src/blueapi/client/client.py @@ -31,6 +31,7 @@ PlanResponse, PythonEnvironmentResponse, SourceInfo, + TaskParams, TaskRequest, TaskResponse, TasksListResponse, @@ -180,7 +181,7 @@ def properties(self) -> dict[str, Any]: def required(self) -> list[str]: return self.model.parameter_schema.get("required", []) - def _build_args(self, *args, **kwargs): + def _build_args(self, *args, **kwargs) -> TaskParams: log.info( "Building args for %s, using %s and %s", "[" + ",".join(self.properties) + "]", @@ -188,27 +189,26 @@ def _build_args(self, *args, **kwargs): kwargs, ) - if len(args) > len(self.properties): + properties = list(self.properties) + + if len(args) > len(properties): raise TypeError(f"{self.name} got too many arguments") - if extra := {k for k in kwargs if k not in self.properties}: + + if extra := {k for k in kwargs if k not in properties}: raise TypeError(f"{self.name} got unexpected arguments: {extra}") - params = {} - # Initially fill parameters using positional args assuming the order - # from the parameter_schema - for req, arg in zip(self.properties, args, strict=False): - params[req] = arg + positional_names = properties[: len(args)] - # Then append any values given via kwargs - for key, value in kwargs.items(): - # If we've already assumed a positional arg was this value, bail out - if key in params: - raise TypeError(f"{self.name} got multiple values for {key}") - params[key] = value + if duplicate := set(positional_names) & kwargs.keys(): + name = next(iter(duplicate)) + raise TypeError(f"{self.name} got multiple values for {name}") - if missing := {k for k in self.required if k not in params}: + supplied = set(positional_names) | kwargs.keys() + + if missing := set(self.required) - supplied: raise TypeError(f"Missing argument(s) for {missing}") - return params + + return TaskParams(args=args, kwargs=kwargs) def __repr__(self) -> str: required = set(self.required) diff --git a/src/blueapi/service/main.py b/src/blueapi/service/main.py index 0d8fbf4774..81b7fd1cc7 100644 --- a/src/blueapi/service/main.py +++ b/src/blueapi/service/main.py @@ -55,7 +55,7 @@ Unauthorized, Update, ) -from blueapi.worker import TrackableTask, WorkerState +from blueapi.worker import TaskParams, TrackableTask, WorkerState from blueapi.worker.event import ProgressEvent, TaskStatusEnum, WorkerEvent from blueapi.worker.worker_errors import WorkerBusyError @@ -320,7 +320,7 @@ def get_device_by_name( example_task_request = TaskRequest( name="count", - params={"detectors": ["x"]}, + params=TaskParams(kwargs={"detectors": ["x"]}), instrument_session="cm12345-1", ) diff --git a/src/blueapi/service/model.py b/src/blueapi/service/model.py index dbaa1d9651..7d7aaabc1d 100644 --- a/src/blueapi/service/model.py +++ b/src/blueapi/service/model.py @@ -1,5 +1,5 @@ import uuid -from collections.abc import Iterable, Mapping +from collections.abc import Iterable from enum import StrEnum from typing import Annotated, Any @@ -11,7 +11,7 @@ from blueapi.core import BLUESKY_PROTOCOLS, Device, Plan from blueapi.core.context import generic_bounds from blueapi.utils import BlueapiBaseModel -from blueapi.worker import WorkerState +from blueapi.worker import TaskParams, WorkerState from blueapi.worker.task_worker import TaskWorker, TrackableTask _UNKNOWN_NAME = "UNKNOWN" @@ -64,8 +64,8 @@ class TaskRequest(BlueapiBaseModel): """ name: str = Field(description="Name of plan to run") - params: Mapping[str, Any] = Field( - description="Values for parameters to plan, if any", default_factory=dict + params: TaskParams = Field( + description="Values for parameters to plan, if any", default_factory=TaskParams ) instrument_session: str = Field( description="Instrument session associated with this task", diff --git a/tests/system_tests/test_blueapi_system.py b/tests/system_tests/test_blueapi_system.py index 8c4dc19cf9..55237850e7 100644 --- a/tests/system_tests/test_blueapi_system.py +++ b/tests/system_tests/test_blueapi_system.py @@ -32,6 +32,7 @@ from blueapi.service.model import ( DeviceResponse, PlanResponse, + TaskParams, TaskRequest, TaskResponse, WorkerTask, @@ -122,7 +123,7 @@ def task_factory( ) -> TaskRequest: return TaskRequest( name="sleep", - params={"time": time}, + params=TaskParams(kwargs={"time": time}), instrument_session=instrument_session if instrument_session else VALID_INSTRUMENT_SESSION[user], @@ -365,7 +366,7 @@ def test_create_task_validation_error(rest_client: BlueapiRestClient): rest_client.create_task( TaskRequest( name="Not-exists", - params={"Not-exists": 0.0}, + params=TaskParams(kwargs={"Not-exists": 0.0}), instrument_session="Not-exists", ) ) @@ -542,7 +543,7 @@ def test_delete_current_environment(client: BlueapiClient): ( TaskRequest( name="count", - params={"detectors": ["det"], "num": 3}, + params=TaskParams(kwargs={"detectors": ["det"], "num": 3}), instrument_session=VALID_INSTRUMENT_SESSION[User.alice], ), User.alice, @@ -550,27 +551,29 @@ def test_delete_current_environment(client: BlueapiClient): ( TaskRequest( name="spec_scan", - params={ - "detectors": ["det"], - "spec": { - "outer": { - "axis": "stage.x", - "start": 0.0, - "stop": 0.4, - "num": 2, - "type": "Linspace", - }, - "inner": { - "axis": "stage.theta", - "start": 5.0, - "stop": 5.3, - "num": 3, - "type": "Linspace", + params=TaskParams( + kwargs={ + "detectors": ["det"], + "spec": { + "outer": { + "axis": "stage.x", + "start": 0.0, + "stop": 0.4, + "num": 2, + "type": "Linspace", + }, + "inner": { + "axis": "stage.theta", + "start": 5.0, + "stop": 5.3, + "num": 3, + "type": "Linspace", + }, + "gap": True, + "type": "Product", }, - "gap": True, - "type": "Product", - }, - }, + } + ), instrument_session=VALID_INSTRUMENT_SESSION[User.bob], ), User.bob, @@ -638,10 +641,12 @@ def on_event(event: AnyEvent) -> None: [ TaskRequest( name="set_absolute", - params={ - "movable": "stage.x", - "value": 1.0, - }, + params=TaskParams( + kwargs={ + "movable": "stage.x", + "value": 1.0, + } + ), instrument_session=VALID_INSTRUMENT_SESSION[User.alice], ), ], @@ -663,11 +668,11 @@ def test_task_submission_after_invalid_task(client_with_stomp: BlueapiClient): res = client_with_stomp.run_task( TaskRequest( name="count", - params={ - "detectors": [ - "det", - ], - }, + params=TaskParams( + kwargs={ + "detectors": ["det"], + } + ), instrument_session=VALID_INSTRUMENT_SESSION[User.alice], ) ) diff --git a/tests/unit_tests/cli/test_cli.py b/tests/unit_tests/cli/test_cli.py index 28cb4fbc92..7131bae224 100644 --- a/tests/unit_tests/cli/test_cli.py +++ b/tests/unit_tests/cli/test_cli.py @@ -50,6 +50,7 @@ TaskRequest, TaskResponse, ) +from blueapi.worker import TaskParams from blueapi.worker.event import ( ProgressEvent, TaskError, @@ -210,7 +211,7 @@ def test_invalid_config_via_env(runner: CliRunner): def test_submit_plan(runner: CliRunner): body_data = { "name": "sleep", - "params": {"time": 5}, + "params": {"args": [], "kwargs": {"time": 5}}, "instrument_session": "cm12345-1", } @@ -270,7 +271,7 @@ def test_run_plan(stomp_client: StompClient, runner: CliRunner): matchers.json_params_matcher( { "name": "sleep", - "params": {"time": 3}, + "params": {"args": [], "kwargs": {"time": 3}}, "instrument_session": "cm12345-1", } ) @@ -386,7 +387,7 @@ def test_run_plan_feedback( ) bc.add_callback.assert_called_once() bc.run_task.assert_called_once_with( - TaskRequest(name="name", params={}, instrument_session="cm12345-1"), + TaskRequest(name="name", params=TaskParams(), instrument_session="cm12345-1"), ) assert res.exit_code == 0 assert res.stdout == message @@ -400,7 +401,7 @@ def test_run_plan_background_without_stomp(runner: CliRunner): matchers.json_params_matcher( { "name": "sleep", - "params": {"time": 3}, + "params": {"args": [], "kwargs": {"time": 3}}, "instrument_session": "cm12345-1", } ) @@ -503,7 +504,7 @@ def test_can_pass_an_instrument_session_with_an_environment_variable( mock_create_task.assert_called_once_with( TaskRequest( name="sleep", - params={"time": 5.0}, + params=TaskParams(kwargs={"time": 5.0}), instrument_session="cm12345-1", ) ) diff --git a/tests/unit_tests/client/test_client.py b/tests/unit_tests/client/test_client.py index d5b0493adc..13502270cc 100644 --- a/tests/unit_tests/client/test_client.py +++ b/tests/unit_tests/client/test_client.py @@ -42,7 +42,14 @@ TasksListResponse, WorkerTask, ) -from blueapi.worker import ProgressEvent, Task, TrackableTask, WorkerEvent, WorkerState +from blueapi.worker import ( + ProgressEvent, + Task, + TaskParams, + TrackableTask, + WorkerEvent, + WorkerState, +) from blueapi.worker.event import TaskError, TaskResult, TaskStatus PLANS = PlanResponse( @@ -72,7 +79,7 @@ ] ) DEVICE = DeviceModel(name="foo", protocols=[]) -TASK = TrackableTask(task_id="foo", task=Task(name="bar", params={})) +TASK = TrackableTask(task_id="foo", task=Task(name="bar", params=TaskParams())) TASKS = TasksListResponse(tasks=[TASK]) ACTIVE_TASK = WorkerTask(task_id="bar") ENVIRONMENT_ID = uuid.uuid4() @@ -900,11 +907,36 @@ def test_plan_empty_fallback_help_text(client): @pytest.mark.parametrize( "args,kwargs,params", [ - p((1,), {}, {"one": 1}, id="required_as_positional"), - p((), {"one": 7}, {"one": 7}, id="required_as_keyword"), - p((1,), {"two": 23}, {"one": 1, "two": 23}, id="all_as_mixed_args_kwargs"), - p((1, 2), {}, {"one": 1, "two": 2}, id="all_as_positional"), - p((), {"one": 21, "two": 42}, {"one": 21, "two": 42}, id="all_as_keyword"), + p( + (1,), + {}, + TaskParams(args=(1,), kwargs={}), + id="required_as_positional", + ), + p( + (), + {"one": 7}, + TaskParams(args=(), kwargs={"one": 7}), + id="required_as_keyword", + ), + p( + (1,), + {"two": 23}, + TaskParams(args=(1,), kwargs={"two": 23}), + id="all_as_mixed_args_kwargs", + ), + p( + (1, 2), + {}, + TaskParams(args=(1, 2), kwargs={}), + id="all_as_positional", + ), + p( + (), + {"one": 21, "two": 42}, + TaskParams(args=(), kwargs={"one": 21, "two": 42}), + id="all_as_keyword", + ), ], ) def test_plan_param_mapping(args, kwargs, params): diff --git a/tests/unit_tests/client/test_rest.py b/tests/unit_tests/client/test_rest.py index 6ecfbaa765..7aa2870886 100644 --- a/tests/unit_tests/client/test_rest.py +++ b/tests/unit_tests/client/test_rest.py @@ -40,11 +40,11 @@ WorkerTask, ) from blueapi.worker.event import WorkerState -from blueapi.worker.task import Task +from blueapi.worker.task import Task, TaskParams from blueapi.worker.task_worker import TrackableTask TASK_REQUEST = TaskRequest( - name="foo", params={"one": "two"}, instrument_session="cm12345-1" + name="foo", params=TaskParams(kwargs={"one": "two"}), instrument_session="cm12345-1" ) @@ -94,7 +94,9 @@ def test_create_task_serialization(): request = TaskRequest( name="demo", instrument_session="cm12345-1", - params={"devices": [DeviceRef(name="foo", cache=Mock(), model=Mock())]}, + params=TaskParams( + kwargs={"devices": [DeviceRef(name="foo", cache=Mock(), model=Mock())]} + ), ) BlueapiRestClient.create_task(rest, request) @@ -107,7 +109,7 @@ def test_create_task_serialization(): data={ "name": "demo", "instrument_session": "cm12345-1", - "params": {"devices": ["foo"]}, + "params": {"args": [], "kwargs": {"devices": ["foo"]}}, }, ) @@ -120,7 +122,7 @@ class CustomType: request = TaskRequest( name="demo", instrument_session="cm12345-1", - params={"devices": [CustomType()]}, + params=TaskParams(kwargs={"devices": [CustomType()]}), ) with pytest.raises(PydanticSerializationError, match="not serializable"): diff --git a/tests/unit_tests/service/test_rest_api.py b/tests/unit_tests/service/test_rest_api.py index 9b6bedf1ba..712b246f5c 100644 --- a/tests/unit_tests/service/test_rest_api.py +++ b/tests/unit_tests/service/test_rest_api.py @@ -172,7 +172,7 @@ def test_rest_config_with_cors( ): task = TaskRequest( name="my-plan", - params={"id": "x"}, + params=TaskParams(kwargs={"id": "x"}), instrument_session=FAKE_INSTRUMENT_SESSION, ) task_id = "f8424be3-203c-494e-b22f-219933b4fa67" @@ -287,7 +287,7 @@ def test_get_non_existent_device_by_name(mock_runner: Mock, client: TestClient) def test_create_task(mock_runner: Mock, client: TestClient) -> None: task = TaskRequest( name="count", - params={"detectors": ["x"]}, + params=TaskParams(kwargs={"detectors": ["x"]}), instrument_session=FAKE_INSTRUMENT_SESSION, ) task_id = str(uuid.uuid4()) @@ -306,7 +306,11 @@ def test_submit_task_requires_permission( mock_opa_client: Mock, access_token: str, ): - task = TaskRequest(name="sleep", params={"time": 2}, instrument_session="cm12345-2") + task = TaskRequest( + name="sleep", + params=TaskParams(kwargs={"time": 2}), + instrument_session="cm12345-2", + ) client_with_opa.headers["Authorization"] = f"Bearer {access_token}" mock_opa_client.can_submit_task.side_effect = HTTPException(status_code=403) mock_runner.run.side_effect = RuntimeError("Task should not be submitted") @@ -323,7 +327,7 @@ def test_create_task_inserts_auth_metadata( ) -> None: task = TaskRequest( name="count", - params={"detectors": ["x"]}, + params=TaskParams(kwargs={"detectors": ["x"]}), instrument_session=FAKE_INSTRUMENT_SESSION, ) client_with_auth.follow_redirects = False From 42bad034038810251b17ca6c8838bfd25ed902bd Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 16:30:03 +0000 Subject: [PATCH 7/9] Fix lint --- tests/unit_tests/cli/test_cli.py | 2 +- tests/unit_tests/service/test_authorization.py | 3 ++- tests/unit_tests/service/test_interface.py | 6 +++--- tests/unit_tests/service/test_protocol.py | 5 ++++- 4 files changed, 10 insertions(+), 6 deletions(-) diff --git a/tests/unit_tests/cli/test_cli.py b/tests/unit_tests/cli/test_cli.py index 7131bae224..d491d45d7d 100644 --- a/tests/unit_tests/cli/test_cli.py +++ b/tests/unit_tests/cli/test_cli.py @@ -1490,6 +1490,6 @@ def test_run_ws_runs_blocking_plan(mock_client: Mock, runner: CliRunner): bc.add_callback.assert_called_once() bc.run_task.assert_not_called() bc.run_blocking.assert_called_once_with( - TaskRequest(name="name", params={}, instrument_session="cm12345-1"), + TaskRequest(name="name", params=TaskParams(), instrument_session="cm12345-1"), ) assert res.exit_code == 0 diff --git a/tests/unit_tests/service/test_authorization.py b/tests/unit_tests/service/test_authorization.py index a2e602f211..2be684c6bd 100644 --- a/tests/unit_tests/service/test_authorization.py +++ b/tests/unit_tests/service/test_authorization.py @@ -14,6 +14,7 @@ validate_tiled_config, ) from blueapi.service.model import TaskRequest +from blueapi.worker import TaskParams # Reusable client patch decorator patch_client_session = patch( @@ -194,7 +195,7 @@ async def test_user_client_can_submit_task(result, context: AbstractContextManag with context: await user_client.can_submit_task( - TaskRequest(name="foo", params={}, instrument_session="cm12345-1") + TaskRequest(name="foo", params=TaskParams(), instrument_session="cm12345-1") ) opa.require_submit_task.assert_called_once_with("cm12345-1", "foo_bar") diff --git a/tests/unit_tests/service/test_interface.py b/tests/unit_tests/service/test_interface.py index 892c5ad2ea..94617381c0 100644 --- a/tests/unit_tests/service/test_interface.py +++ b/tests/unit_tests/service/test_interface.py @@ -46,7 +46,7 @@ WorkerEvent, WorkerState, ) -from blueapi.worker.task import Task +from blueapi.worker.task import Task, TaskParams from blueapi.worker.task_worker import TrackableTask FAKE_INSTRUMENT_SESSION = "cm12345-1" @@ -385,7 +385,7 @@ def test_get_task_by_id( request_id=ANY, task=Task( name="my_plan", - params={}, + params=TaskParams(), metadata=expected_metadata, ), is_complete=False, @@ -417,7 +417,7 @@ def test_submit_task_inserts_metadata(context_mock: MagicMock): request_id=ANY, task=Task( name="my_plan", - params={}, + params=TaskParams(), metadata=expected_metadata, ), is_complete=False, diff --git a/tests/unit_tests/service/test_protocol.py b/tests/unit_tests/service/test_protocol.py index 9178162e18..08a0e39130 100644 --- a/tests/unit_tests/service/test_protocol.py +++ b/tests/unit_tests/service/test_protocol.py @@ -15,6 +15,7 @@ Resume, Submit, ) +from blueapi.worker import TaskParams @pytest.mark.parametrize( @@ -29,7 +30,9 @@ } }""", Submit( - task=TaskRequest(name="foo", params={}, instrument_session="cm12345-1") + task=TaskRequest( + name="foo", params=TaskParams(), instrument_session="cm12345-1" + ) ), ), ('{"kind": "pause"}', Pause()), From 29b05e4173b444e293d16c9404853f1b0aea7ec3 Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 16:35:54 +0000 Subject: [PATCH 8/9] fix SUBMIT_REQUEST --- tests/unit_tests/service/test_rest_api.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit_tests/service/test_rest_api.py b/tests/unit_tests/service/test_rest_api.py index 712b246f5c..9da3a754cb 100644 --- a/tests/unit_tests/service/test_rest_api.py +++ b/tests/unit_tests/service/test_rest_api.py @@ -62,7 +62,7 @@ class MockCountModel(BaseModel): ... "kind": "submit", "task": { "name": "foo", - "params": {"one": "two"}, + "params": {"args": [], "kwargs": {"one": "two"}}, "instrument_session": "cm12345-1", }, } From adfa659bac4a62b1441286d5a0eacb3623fb8bb1 Mon Sep 17 00:00:00 2001 From: Oli Wenman Date: Wed, 9 Sep 2026 16:51:51 +0000 Subject: [PATCH 9/9] Correct doc string --- src/blueapi/worker/task.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index 4c64698930..a7291ce684 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -34,15 +34,13 @@ def prepare_params( self, ctx: BlueskyContext ) -> tuple[list[Any], Mapping[str, Any]]: """ - Checks plan parameters against context + Checks the configured plan parameters against context Args: ctx: Context holding plans and devices - plan: Plan object including schema - params: Parameter values to be validated against schema Returns: - Mapping[str, Any]: _description_ + tuple[list[Any], Mapping[str, Any]]: The prepared parameters for the plan. """ plan = ctx.plans[self.name] func = ctx.plan_functions[self.name]