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
141 changes: 84 additions & 57 deletions src/blueapi/client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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(
Expand All @@ -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:
Expand Down Expand Up @@ -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"
2 changes: 2 additions & 0 deletions src/blueapi/core/bluesky_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
38 changes: 34 additions & 4 deletions src/blueapi/core/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand All @@ -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 _:
Expand Down
44 changes: 39 additions & 5 deletions src/blueapi/service/model.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand All @@ -102,13 +129,20 @@ 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":
return cls(
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()
},
)


Expand Down
9 changes: 8 additions & 1 deletion src/blueapi/worker/task.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading