Skip to content
Merged
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
31 changes: 15 additions & 16 deletions src/packages/harp-device/src/harp/device/schema/_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,14 +104,14 @@ class ConverterContext:
name: str # yml field key ("__value__" for a whole-register value)
interface_type: Optional[str] # the DSL interfaceType (None = raw/native)
mask: Optional[int] # bit mask, when the value is bit-packed
length: int # element count this value spans (0 = unset -> scalar)
length: int # register elements this value spans, at least one
element: np.dtype # base element dtype of the register, from PayloadType
element_size: int # base element byte size of the register

@property
def span(self) -> int:
"""Byte span of the value (element count * element size)."""
return max(1, self.length) * self.element_size
"""Byte span of the value, its element count times the element size."""
return self.length * self.element_size

@property
def member_dtype(self) -> np.dtype:
Expand Down Expand Up @@ -343,7 +343,7 @@ def _extension(self, symbol: str, ctx: ConverterContext) -> Converter[Any]:
def _default(self, member: PayloadMember, type_name: str, ctx: ConverterContext) -> Any:
"""The typed default value of the field, or ``_NO_DEFAULT`` when it has none."""
_default_value = member.defaultValue if member.defaultValue is not None else member.minValue
if _default_value is None or (member.length or 0) > 1:
if _default_value is None or ctx.length > 1:
return _NO_DEFAULT
value = float(_default_value.root)
group_mask = self._find_mask(type_name)
Expand All @@ -367,17 +367,17 @@ def _build_field(self, key: str, member: PayloadMember, reg: Register) -> Any:
# ``key`` stays the verbatim yml name: it feeds ``ConverterContext.name``, and
# a custom converter symbol is derived from the pre-rename key ("Data" ->
# "DataConverter"). The renamed attribute name is applied by the caller.
elem_np = _ELEMENT[reg.type]
elem_size = np.dtype(elem_np).itemsize
elem = np.dtype(_ELEMENT[reg.type])
elem_size = elem.itemsize
offset = member.offset or 0
it = member.interfaceType.root if member.interfaceType else None
type_name = it or (member.maskType.root if member.maskType else "")
ctx = ConverterContext(
name=key,
interface_type=it,
mask=member.mask,
length=member.length or 0,
element=np.dtype(elem_np),
length=member.length or 1,
element=elem,
element_size=elem_size,
)
default = self._default(member, type_name, ctx)
Expand Down Expand Up @@ -416,16 +416,16 @@ def _build_payload(self, name: str, reg: Register) -> type:

def _new_payload(self, class_name: str, owner: str, reg: Register) -> type:
elem_np = _ELEMENT[reg.type]
elem_size = np.dtype(elem_np).itemsize
length = reg.length or 1
elem = np.dtype(elem_np)
elem_size = elem.itemsize

if reg.payloadSpec is not None:
renamed = self._rename("field", owner, reg.payloadSpec, field_name, reserved=True)
namespace = {
renamed[key]: self._build_field(key, member, reg)
for key, member in reg.payloadSpec.items()
}
kwds = {"length": length} if length > 1 else {}
kwds = {"length": reg.length}
return _new_class(class_name, (StructPayload[elem_np],), namespace, kwds)

# anonymous single-value payload
Expand All @@ -451,8 +451,8 @@ def _new_payload(self, class_name: str, owner: str, reg: Register) -> type:
name="__value__",
interface_type=it,
mask=None,
length=length,
element=np.dtype(elem_np),
length=reg.length or 1,
element=elem,
element_size=elem_size,
)
descriptor = Field(self._resolve_converter(ctx))
Expand All @@ -464,7 +464,6 @@ def _class_name(self, name: str, reg: Register) -> str:
return f"_{name}" if reg.visibility is Visibility.private else name

def _build_register(self, name: str, class_name: str, reg: Register) -> type[RegisterBase[Any]]:
length = reg.length or 1
it = reg.interfaceType.root if reg.interfaceType else None

# A plain scalar/array register needs no payload wrapper: its whole value is a
Expand All @@ -475,8 +474,8 @@ def _build_register(self, name: str, class_name: str, reg: Register) -> type[Reg
and reg.converter is None
and _is_native(it)
):
if length > 1: # plain array register
cls = _ARRAY_REGISTER[reg.type](reg.address, length=length)
if reg.length is not None: # plain array register
cls = _ARRAY_REGISTER[reg.type](reg.address, length=reg.length)
cls.__name__ = cls.__qualname__ = class_name
return cls
return _new_class(class_name, (_SCALAR_REGISTER[reg.type],), {"address": reg.address})
Expand Down
11 changes: 6 additions & 5 deletions src/packages/harp-device/src/harp/device/schema/_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,9 +126,10 @@ class PayloadMember(BaseModel):
None,
description="The zero-based index at which encoding of this payload member starts.",
)
length: Optional[int] = Field(
None, description="The number of elements used to encode this payload member."
)
length: Annotated[
Optional[int],
Field(ge=1, description="The number of elements used to encode this payload member."),
] = None
description: Optional[str] = Field(
None, description="A summary description of this payload member."
)
Expand Down Expand Up @@ -164,8 +165,8 @@ class Register(BaseModel):
address: Annotated[int, Field(le=255, description="The unique 8-bit address of the register.")]
type: PayloadType = Field(..., description="The type of the register payload.")
length: Annotated[
Optional[int], Field(ge=1, default=1, description="The length of the register payload.")
]
Optional[int], Field(ge=1, description="The length of the register payload.")
] = None
access: Union[Access, List[Access]] = Field(
..., description="The expected use of the register."
)
Expand Down
14 changes: 7 additions & 7 deletions src/packages/harp-protocol/src/harp/protocol/_message.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
import struct
from typing import Any, ClassVar, Generic, Protocol, TypeVar, cast

import numpy as np
from typing_extensions import Sentinel

from ._builder import build_message_frame
Expand All @@ -18,6 +17,7 @@
_TIMESTAMPED_PAYLOAD_OFFSET,
)
from ._message_type import MessageType, _message_type_from_byte_safe
from ._payload import PayloadBase
from ._payload_type import PayloadType, decode_payload_type

P = TypeVar("P")
Expand All @@ -38,14 +38,14 @@ class PayloadDecoder(Protocol[_P_co]):
"""Reads a payload of type ``_P_co`` out of a message.

Structural rather than nominal, so a message never has to know about registers, and
anything declaring a payload type, a length and a ``parse`` satisfies it. Every
``RegisterBase`` does. ``length`` is the element count, or ``None`` for a single
value, and together with ``payload_type`` it fixes how many payload bytes the
decoder consumes.
anything declaring a payload type, a payload class and a ``parse`` satisfies it. Every
``RegisterBase`` does. ``payload_class.payload_dtype`` is what fixes how many payload
bytes the decoder consumes, and it is the same quantity ``parse`` reads the frame
with, so the two cannot disagree about the extent of a payload.
"""

payload_type: ClassVar["PayloadType"]
length: ClassVar[int | None]
payload_class: ClassVar[type[PayloadBase[Any]]]

@classmethod
def parse(cls, value: Any) -> _P_co: ...
Expand Down Expand Up @@ -187,7 +187,7 @@ def decode(self, decoder: type[PayloadDecoder[_P]]) -> "HarpMessage[_P]":
f"{decoder.__name__} declares {decoder.payload_type!r} but this "
f"message declares {self.payload_type!r}."
)
expected = (decoder.length or 1) * np.dtype(decoder.payload_type.value).itemsize
expected = decoder.payload_class.payload_dtype.itemsize
actual = len(self.payload_bytes)
if actual != expected:
raise HarpParseError(
Expand Down
54 changes: 31 additions & 23 deletions src/packages/harp-protocol/src/harp/protocol/_register.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from ._message import HarpMessage, HarpParseError
from ._message_type import MessageType, message_type_to_byte
from ._payload import (
AnonymousPayload,
Batch,
PayloadBase,
PayloadFloat,
Expand Down Expand Up @@ -148,14 +149,13 @@ class RegisterBase(ABC, Generic[U], metaclass=_RegisterBaseMeta):
or ``RegisterBase[ClockConfigurationFlags]`` for a whole-register
``BitMask`` or ``GroupMask``, even though each still has a ``payload_class``.

Subclasses must define ``address``, ``payload_type``, and
``payload_class`` as ``ClassVar``s.
Subclasses must define ``address``, ``payload_type``, and ``payload_class`` as
``ClassVar``s. The extent of a payload is always read from ``payload_class``.
"""

address: ClassVar[int]
payload_type: ClassVar[PayloadType]
payload_class: ClassVar[type[PayloadBase[Any]]]
length: ClassVar[int | None] = None

@classmethod
def parse(cls, value: HarpMessage | bytes | bytearray | memoryview) -> U:
Expand Down Expand Up @@ -435,32 +435,40 @@ class RegisterFloat(RegisterBase[np.float32], metaclass=_ScalarRegisterMeta):


class _ArrayRegisterMeta(_RegisterBaseMeta):
"""A base metaclass for array registers. Calling with address and length creates a concrete subclass: ``RegisterU16Array(0x28, length=3)``."""
"""A declared ``length`` sizes the payload, and calling a register base with an address
and a length creates a one-off subclass: ``RegisterU16Array(0x28, length=3)``.

def __call__(cls: "type[_AR]", address: int, *, length: int) -> "type[_AR]": # type: ignore[override, misc]
_require_no_address(cls)
base_payload = cls.payload_class # type: ignore[attr-defined]
# Anonymous payloads carry a plain (non-structured) dtype. The array
# variant uses a sub-dtype (inner_dtype, (length,)) so a single buffer
# element decodes directly to an ndarray of shape (length,).
inner = base_payload.payload_dtype
sub_dtype = np.dtype((inner, (length,)))
concrete_payload = type(
``length`` is declared here rather than on ``RegisterBase``, so only an array register
carries one. It is the element count, and nothing reads it to size a payload.
"""

length: int
payload_class: type[AnonymousPayload[Any]]

def __init__(
cls, name: str, bases: tuple[type, ...], namespace: dict[str, Any], **kwargs: Any
) -> None:
super().__init__(name, bases, namespace, **kwargs)
# The namespace holds this class body only, not inherited values, so a plain
# subclass reads None and keeps the payload already sized by its base.
length = namespace.get("length")
if length is None:
return
base_payload = cls.payload_class
if base_payload.payload_dtype.subdtype is not None:
raise TypeError(f"{name} redeclares a length already applied by its base class.")
# A sub-array dtype, so reading one buffer element gives an ndarray of that shape.
cls.payload_class = type(
f"{base_payload.__name__}_{length}",
(base_payload,),
{"payload_dtype": sub_dtype},
{"payload_dtype": np.dtype((base_payload.payload_dtype, (length,)))},
)

def __call__(cls: "type[_AR]", address: int, *, length: int) -> "type[_AR]": # type: ignore[override, misc]
_require_no_address(cls)
return cast(
"type[_AR]",
type(
f"{cls.__name__}_{address:#04x}",
(cls,),
{
"address": address,
"length": length,
"payload_class": concrete_payload,
},
),
type(f"{cls.__name__}_{address:#04x}", (cls,), {"address": address, "length": length}),
)


Expand Down
18 changes: 18 additions & 0 deletions tests/device/test_device.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,3 +176,21 @@ def test_error_reply_returned_when_not_raising():
reply = device.read(core.WhoAmI)
assert reply.has_error
assert int(reply.payload) == 7


def test_read_multi_element_register_returns_payload():
# read decodes the reply through the register, so a payload of several elements has to
# survive that step.
transport = _ScriptedTransport()
transport.on_write = lambda _: (
core.DeviceName.format("Behavior", message_type=MessageType.Read),
)
with _ShortTimeoutDevice(transport) as device:
assert device.read(core.DeviceName).payload == "Behavior"


def test_write_multi_element_register_returns_payload():
transport = _ScriptedTransport()
transport.on_write = lambda data: (data,) # a device echoing the write
with _ShortTimeoutDevice(transport) as device:
assert device.write(core.DeviceName, "Behavior").payload == "Behavior"
56 changes: 56 additions & 0 deletions tests/device/test_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,62 @@ def test_enum_names_match_generator_for_every_enum(device_registers):
}


# ---------------------------------------------------------------------------
# The declared length decides whether a register is an array
# ---------------------------------------------------------------------------

_LENGTH_YML = """
registers:
Absent: {address: 32, type: U16, access: Read}
One: {address: 33, type: U16, length: 1, access: Read}
Three: {address: 34, type: U16, length: 3, access: Read}
"""


def test_absent_length_emits_scalar_register():
# A register with no declared length holds a single value, and carries no length at all.
reg = create_registers(_LENGTH_YML)["Absent"]
assert not hasattr(reg, "length")
assert reg.payload_class.payload_dtype.shape == ()


@pytest.mark.parametrize("name, count", [("One", 1), ("Three", 3)])
def test_declared_length_emits_array_register(name, count):
# Any declared length means an array, 1 included. Reading a declared 1 as a single value
# gives the same type as declaring nothing, and the generator would then disagree.
reg = create_registers(_LENGTH_YML)[name]
assert reg.length == count
assert reg.payload_class.payload_dtype.shape == (count,)
values = np.arange(count, dtype=np.uint16)
np.testing.assert_array_equal(reg.parse(HarpMessage.parse(reg.format(values))), values)


def test_multi_element_struct_register_decodes_through_message():
# A struct register declares no length, so decode has to size its payload from the
# payload class.
reg = create_registers(
"registers:\n"
" Settings:\n"
" address: 32\n"
" type: U16\n"
" length: 3\n"
" access: Write\n"
" payloadSpec:\n"
" Gain: {offset: 0}\n"
" Offset: {offset: 1}\n"
" Threshold: {offset: 2}\n"
)["Settings"]
payload = reg.payload_class(gain=1, offset=2, threshold=3)
decoded = HarpMessage.parse(reg.format(payload)).decode(reg)
assert (int(decoded.payload.gain), int(decoded.payload.threshold)) == (1, 3)


def test_core_string_register_decodes_through_message():
# DeviceName spans 25 U8 elements, so the same defect reached every Harp device.
frame = core.DeviceName.format("Behavior")
assert HarpMessage.parse(frame).decode(core.DeviceName).payload == "Behavior"


# ---------------------------------------------------------------------------
# Core registers: the emitter and the generated package, from the same core.yml
# ---------------------------------------------------------------------------
Expand Down
Loading