diff --git a/src/blueapi/client/client.py b/src/blueapi/client/client.py index 3d77a93da..8800d424d 100644 --- a/src/blueapi/client/client.py +++ b/src/blueapi/client/client.py @@ -5,15 +5,13 @@ from concurrent.futures import Future from contextlib import suppress from functools import cached_property +from inspect import Parameter from pathlib import Path from typing import Any, Self from bluesky_stomp.messaging import MessageContext, StompClient from bluesky_stomp.models import Broker -from observability_utils.tracing import ( - get_tracer, - start_as_current_span, -) +from observability_utils.tracing import get_tracer, start_as_current_span from blueapi.config import ( ApplicationConfig, @@ -58,13 +56,6 @@ _REPR_MAX_LENGTH = 100 _REPR_MAX_ARGS_INLINE = 3 -_JSON_TYPE_MAP = { - "string": "str", - "integer": "int", - "boolean": "bool", - "number": "float", - "object": "dict", -} class MissingInstrumentSessionError(Exception): @@ -156,6 +147,7 @@ def __init__(self, name, model: PlanModel, client: "BlueapiClient"): self.model = model self._client = client self.__doc__ = model.description + self.parameter_kinds = model.parameter_kinds def __call__(self, *args, **kwargs) -> Any: req = TaskRequest( @@ -182,48 +174,102 @@ def required(self) -> list[str]: return self.model.parameter_schema.get("required", []) def _build_args(self, *args, **kwargs) -> TaskParams: - log.info( - "Building args for %s, using %s and %s", - "[" + ",".join(self.properties) + "]", - args, - kwargs, - ) - - properties = list(self.properties) - - if len(args) > len(properties): - raise TypeError(f"{self.name} got too many arguments") + kinds = self.parameter_kinds - if extra := {k for k in kwargs if k not in properties}: - raise TypeError(f"{self.name} got unexpected arguments: {extra}") + positional_parameters = [ + name + for name, kind in kinds.items() + if kind + in (Parameter.POSITIONAL_ONLY.name, Parameter.POSITIONAL_OR_KEYWORD.name) + ] - positional_names = properties[: len(args)] + var_positional = any( + kind == Parameter.VAR_POSITIONAL.name for kind in kinds.values() + ) + var_keyword = any(kind == Parameter.VAR_KEYWORD.name for kind in kinds.values()) - if duplicate := set(positional_names) & kwargs.keys(): - name = next(iter(duplicate)) - raise TypeError(f"{self.name} got multiple values for {name}") + if not var_positional and len(args) > len(positional_parameters): + raise TypeError(f"{self.name} got too many arguments") - supplied = set(positional_names) | kwargs.keys() + # Check positional/keyword collisions. + for name in positional_parameters[: len(args)]: + if name in kwargs: + raise TypeError(f"{self.name} got multiple values for {name}") + + # Keyword arguments must correspond to a named parameter, unless + # the function has **kwargs. + if not var_keyword: + unexpected = set(kwargs) - set(kinds) + if unexpected: + raise TypeError(f"{self.name} got unexpected arguments: {unexpected}") + + # A keyword-only parameter cannot be supplied positionally. + # Since positional_parameters excludes KEYWORD_ONLY, this is + # automatically handled above. + + supplied = set(kwargs) | set(positional_parameters[: len(args)]) + + required = { + name + for name in self.required + if kinds[name] + not in (Parameter.VAR_POSITIONAL.name, Parameter.VAR_KEYWORD.name) + } - if missing := set(self.required) - supplied: + if missing := required - supplied: raise TypeError(f"Missing argument(s) for {missing}") return TaskParams(args=args, kwargs=kwargs) def __repr__(self) -> str: - required = set(self.required) + kinds = self.model.parameter_kinds def _format_arg(name: str, info: dict[str, Any]) -> str: - typ = _pretty_type(info) - default = info.get("default") + typ = self.model.parameter_types[name] + kind = kinds[name] - if name in required: + if kind == Parameter.VAR_POSITIONAL.name: + return f"*{name}: {typ}" + + if kind == Parameter.VAR_KEYWORD.name: + return f"**{name}: {typ}" + + if name in self.required: return f"{name}: {typ}" - if default := info.get("default"): - return f"{name}: {typ} = {default!r}" - return f"{name}: {typ} | None = None" - args = [_format_arg(name, info) for name, info in self.properties.items()] + if "default" in info: + return f"{name}: {typ} = {info['default']!r}" + + return f"{name}: {typ} = None" + + args: list[str] = [] + names = list(self.properties) + + has_var_positional = Parameter.VAR_POSITIONAL.name in kinds.values() + + for i, name in enumerate(names): + kind = kinds[name] + + # Positional-only parameters need a "/" after the last one. + if kind == Parameter.POSITIONAL_ONLY.name: + args.append(_format_arg(name, self.properties[name])) + + if ( + i + 1 == len(names) + or kinds[names[i + 1]] != Parameter.POSITIONAL_ONLY.name + ): + args.append("/") + + continue + + # Keyword-only parameters need a "*" separator if there is + # no *args parameter to provide the separator. + if kind == Parameter.KEYWORD_ONLY.name and not has_var_positional: + if not args or args[-1] != "*": + args.append("*") + + args.append(_format_arg(name, self.properties[name])) + single_line = f"{self.name}({', '.join(args)})" if len(single_line) <= _REPR_MAX_LENGTH and len(args) <= _REPR_MAX_ARGS_INLINE: @@ -834,22 +880,3 @@ class PlanFailedError(Exception): def __init__(self, typ: str, message: str): super().__init__(message) self._type = typ - - -def _pretty_type(schema: dict[str, Any]) -> str: - if "$ref" in schema: - return schema["$ref"].split("/")[-1] - - if schema.get("type") == "array": - item_schema = schema.get("items", {}) - inner = _pretty_type(item_schema) - return f"list[{inner}]" - - if "anyOf" in schema: - return " | ".join(_pretty_type(s) for s in schema["anyOf"]) - - json_type = schema.get("type") - if isinstance(json_type, str): - return _JSON_TYPE_MAP.get(json_type, json_type.split(".")[-1]) - - return "Any" diff --git a/src/blueapi/core/bluesky_types.py b/src/blueapi/core/bluesky_types.py index 64a2a5289..c3fb4be4f 100644 --- a/src/blueapi/core/bluesky_types.py +++ b/src/blueapi/core/bluesky_types.py @@ -93,6 +93,8 @@ class Plan(BlueapiBaseModel): model: type[BaseModel] = Field( description="Validation model of the parameters for the plan" ) + parameter_kinds: dict[str, str] = {} + parameter_types: dict[str, Any] = {} class DataEvent(BlueapiBaseModel): diff --git a/src/blueapi/core/context.py b/src/blueapi/core/context.py index 76ce74e64..e3d564863 100644 --- a/src/blueapi/core/context.py +++ b/src/blueapi/core/context.py @@ -326,9 +326,21 @@ def my_plan(a: int, b: str): __config__=BlueapiPlanModelConfig, **self._type_spec_for_function(plan), # type: ignore ) + parameter_kinds = { + name: parameter.kind.name + for name, parameter in signature(plan).parameters.items() + } + parameter_types = { + name: parameter.annotation + for name, parameter in signature(plan).parameters.items() + } LOGGER.debug("Registering plan %s from %s", plan.__name__, plan.__module__) self.plans[plan.__name__] = Plan( - name=plan.__name__, model=model, description=plan.__doc__ + name=plan.__name__, + model=model, + description=plan.__doc__, + parameter_kinds=parameter_kinds, + parameter_types=parameter_types, ) self.plan_functions[plan.__name__] = plan return plan @@ -439,19 +451,31 @@ def _type_spec_for_function( Returns: Mapping of {name: (type, default)} to be used by pydantic for deserialising - function arguments + function arguments """ args = signature(func).parameters types = get_type_hints(func) new_args: dict[str, tuple[type, FieldInfo]] = {} + for name, para in args.items(): arg_type = types.get(name, Parameter.empty) + if arg_type is Parameter.empty: raise ValueError( f"Type annotation is required for '{name}' in '{func.__name__}'" ) - no_default = para.default is Parameter.empty + match para.kind: + case Parameter.VAR_POSITIONAL: + arg_type = tuple[arg_type, ...] + case Parameter.VAR_KEYWORD: + arg_type = dict[str, arg_type] + + no_default = para.default is Parameter.empty and para.kind not in ( + Parameter.VAR_POSITIONAL, + Parameter.VAR_KEYWORD, + ) + if ( isclass(arg_type) and (issubclass(arg_type, BaseModel) or is_dataclass(arg_type)) @@ -462,9 +486,15 @@ def _type_spec_for_function( info = FieldInfo(default_factory=default_factory) else: _type = self._convert_type(arg_type, no_default) + match para.default: case Parameter.empty: - info = FieldInfo(default_factory=None) + if para.kind is Parameter.VAR_POSITIONAL: + info = FieldInfo(default_factory=tuple) + elif para.kind is Parameter.VAR_KEYWORD: + info = FieldInfo(default_factory=dict) + else: + info = FieldInfo(default_factory=None) case None: info = FieldInfo(default_factory=lambda: None) case _: diff --git a/src/blueapi/service/model.py b/src/blueapi/service/model.py index 7d7aaabc1..56565285d 100644 --- a/src/blueapi/service/model.py +++ b/src/blueapi/service/model.py @@ -1,7 +1,8 @@ import uuid from collections.abc import Iterable from enum import StrEnum -from typing import Annotated, Any +from types import NoneType, UnionType +from typing import Annotated, Any, Union, get_args, get_origin from bluesky.protocols import HasName from pydantic import Field @@ -88,11 +89,37 @@ class DeviceResponse(BlueapiBaseModel): devices: list[DeviceModel] = Field(description="Devices available to use in plans") -class PlanModel(BlueapiBaseModel): - """ - Representation of a plan - """ +def _pretty_annotation(annotation: Any) -> str: + # Unwrap Annotated[T, ...] + if get_origin(annotation) is Annotated: + annotation = get_args(annotation)[0] + + # None + if annotation is NoneType: + return "None" + + origin = get_origin(annotation) + args = get_args(annotation) + + # PEP 604: T | U + if origin is UnionType: + return " | ".join(_pretty_annotation(arg) for arg in args) + # typing.Union[T, U] + if origin is Union: + return " | ".join(_pretty_annotation(arg) for arg in args) + + # Generic types, e.g. Movable[float], tuple[T, ...], list[T] + if origin is not None: + origin_name = getattr(origin, "__name__", str(origin)) + formatted_args = ", ".join(_pretty_annotation(arg) for arg in args) + return f"{origin_name}[{formatted_args}]" + + # Plain classes + return getattr(annotation, "__name__", str(annotation)) + + +class PlanModel(BlueapiBaseModel): name: str = Field(description="Name of the plan") description: str | SkipJsonSchema[None] = Field( description="Docstring of the plan", default=None @@ -102,6 +129,8 @@ class PlanModel(BlueapiBaseModel): alias="schema", default_factory=dict, ) + parameter_kinds: dict[str, str] = Field(default_factory=dict) + parameter_types: dict[str, str] = Field(default_factory=dict) @classmethod def from_plan(cls, plan: Plan) -> "PlanModel": @@ -109,6 +138,11 @@ def from_plan(cls, plan: Plan) -> "PlanModel": name=plan.name, schema=plan.model.model_json_schema(), description=plan.description, + parameter_kinds=plan.parameter_kinds, + parameter_types={ + name: _pretty_annotation(annotation) + for name, annotation in plan.parameter_types.items() + }, ) diff --git a/src/blueapi/worker/task.py b/src/blueapi/worker/task.py index a7291ce68..816392eb0 100644 --- a/src/blueapi/worker/task.py +++ b/src/blueapi/worker/task.py @@ -82,7 +82,14 @@ def prepare_params( kwargs.update(value) else: - # Let the generated Pydantic model provide the default. + # Defaults need to be materialised for ordinary parameters, + # but variadic parameters are absent when not supplied. + if parameter.kind in ( + Parameter.VAR_POSITIONAL, + Parameter.VAR_KEYWORD, + ): + continue + value = getattr(validated, name) kwargs[name] = value