Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions src/blueapi/cli/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
32 changes: 16 additions & 16 deletions src/blueapi/client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
PlanResponse,
PythonEnvironmentResponse,
SourceInfo,
TaskParams,
TaskRequest,
TaskResponse,
TasksListResponse,
Expand Down Expand Up @@ -180,35 +181,34 @@ 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) + "]",
args,
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)
Expand Down
4 changes: 2 additions & 2 deletions src/blueapi/service/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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",
)

Expand Down
8 changes: 4 additions & 4 deletions src/blueapi/service/model.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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"
Expand Down Expand Up @@ -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",
Expand Down
3 changes: 2 additions & 1 deletion src/blueapi/worker/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
Expand Down
101 changes: 70 additions & 31 deletions src/blueapi/worker/task.py
Original file line number Diff line number Diff line change
@@ -1,64 +1,103 @@
import logging
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
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

LOGGER = logging.getLogger(__name__)


class TaskParams(BlueapiBaseModel):
args: Sequence[Any] = Field(default_factory=lambda: [])
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[list[Any], Mapping[str, Any]]:
"""
Checks the configured plan parameters against context

Args:
ctx: Context holding plans and devices

Returns:
tuple[list[Any], Mapping[str, Any]]: The prepared parameters for the plan.
"""
plan = ctx.plans[self.name]
func = ctx.plan_functions[self.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)

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 self.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 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"
)

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:
"""
Checks 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_
"""

plan = ctx.plans[task.name]
model = plan.model
adapter = TypeAdapter(model)
return adapter.validate_python(task.params)
69 changes: 37 additions & 32 deletions tests/system_tests/test_blueapi_system.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from blueapi.service.model import (
DeviceResponse,
PlanResponse,
TaskParams,
TaskRequest,
TaskResponse,
WorkerTask,
Expand Down Expand Up @@ -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],
Expand Down Expand Up @@ -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",
)
)
Expand Down Expand Up @@ -542,35 +543,37 @@ 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,
),
(
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,
Expand Down Expand Up @@ -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],
),
],
Expand All @@ -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],
)
)
Expand Down
Loading
Loading