diff --git a/src/packages/harp-device/src/harp/device/schema/_emit.py b/src/packages/harp-device/src/harp/device/schema/_emit.py index 9411de9..ec525af 100644 --- a/src/packages/harp-device/src/harp/device/schema/_emit.py +++ b/src/packages/harp-device/src/harp/device/schema/_emit.py @@ -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 # register elements this value spans, at least one + length: Optional[int] # declared element count; None when the schema declares none 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, its element count times the element size.""" - return self.length * self.element_size + return (self.length or 1) * self.element_size @property def member_dtype(self) -> np.dtype: @@ -124,8 +124,8 @@ def member_dtype(self) -> np.dtype: @property def raw_dtype(self) -> np.dtype: - """Native passthrough dtype, a sub-array when the value spans more than one element.""" - if self.length > 1: + """Native passthrough dtype, a sub-array when the schema declares an element count.""" + if self.length is not None: return np.dtype((self.element.type, (self.length,))) return self.element @@ -319,7 +319,7 @@ def _resolve_converter(self, ctx: ConverterContext) -> Converter[Any]: if ctx.mask is not None: return IdentityConverter(ctx.member_dtype) # bit-field: native slice of the element if it is None: - if ctx.length > 1: + if ctx.length is not None: return ArrayConverter(ctx.element, ctx.length) # raw passthrough, sub-array return IdentityConverter(ctx.element) # raw passthrough, single element # A known primitive that didn't fit is re-interpreted per field (``{Name}Converter``); @@ -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 ctx.length > 1: + if _default_value is None or ctx.length is not None: return _NO_DEFAULT value = float(_default_value.root) group_mask = self._find_mask(type_name) @@ -376,7 +376,7 @@ def _build_field(self, key: str, member: PayloadMember, reg: Register) -> Any: name=key, interface_type=it, mask=member.mask, - length=member.length or 1, + length=member.length, element=elem, element_size=elem_size, ) @@ -451,7 +451,7 @@ def _new_payload(self, class_name: str, owner: str, reg: Register) -> type: name="__value__", interface_type=it, mask=None, - length=reg.length or 1, + length=reg.length, element=elem, element_size=elem_size, ) diff --git a/tests/assets/device.yml b/tests/assets/device.yml index 38b87cc..d882b3c 100644 --- a/tests/assets/device.yml +++ b/tests/assets/device.yml @@ -127,6 +127,30 @@ registers: PulseDO0: <<: *pulseDO address: 43 + MultiElementPayload: + address: 44 + type: U16 + length: 3 + access: Write + SingleElementPayload: + address: 45 + type: U16 + length: 1 + access: Write + MixedMemberLength: + address: 46 + type: U8 + length: 4 + access: Write + payloadSpec: + Absent: + offset: 0 + Single: + offset: 1 + length: 1 + Multiple: + offset: 2 + length: 2 StartPulse: address: 100 type: U16 diff --git a/tests/data/test_dataset.py b/tests/data/test_dataset.py index 24ad3cc..d0d8148 100644 --- a/tests/data/test_dataset.py +++ b/tests/data/test_dataset.py @@ -21,7 +21,7 @@ def _records(cls, n, seed): dtype = cls.payload_class.payload_dtype rng = np.random.default_rng(seed) raw = rng.integers(0, 128, size=n * dtype.itemsize, dtype=np.uint8) - return raw.view(dtype).copy() + return np.frombuffer(raw, dtype=dtype, count=n).copy() @pytest.fixture diff --git a/tests/device/converters.py b/tests/device/converters.py index 1957988..25b86dc 100644 --- a/tests/device/converters.py +++ b/tests/device/converters.py @@ -6,8 +6,8 @@ from harp.protocol import Converter -class DataConverter(Converter[int]): - """Maps two raw little-endian signed bytes to and from a Python int. +class DataConverter(Converter[np.int32]): + """Maps two raw little-endian signed bytes to and from a numpy int32. Models interfaceType: int over a two-byte sub-region of the CustomMemberConverter payload. """ @@ -18,8 +18,8 @@ def __init__(self) -> None: self._length = 2 self.dtype = np.dtype((np.uint8, (self._length,))) - def decode_scalar(self, view: np.generic) -> int: - return int.from_bytes(bytes(np.asarray(view).tolist()), "little", signed=True) + def decode_scalar(self, view: np.generic) -> np.int32: + return np.int32(int.from_bytes(bytes(np.asarray(view).tolist()), "little", signed=True)) def decode_batch(self, view: NDArray[np.generic]) -> Any: return np.array( @@ -30,7 +30,7 @@ def decode_batch(self, view: NDArray[np.generic]) -> Any: dtype=object, ) - def encode_into(self, view: NDArray[np.generic], value: int) -> None: + def encode_into(self, view: NDArray[np.generic], value: np.int32) -> None: view[...] = np.frombuffer( int(value).to_bytes(self._length, "little", signed=True), dtype=np.uint8 ) diff --git a/tests/device/expected_device.py b/tests/device/expected_device.py index 0f027f3..e67090b 100644 --- a/tests/device/expected_device.py +++ b/tests/device/expected_device.py @@ -8,6 +8,7 @@ from numpy.typing import NDArray from harp.protocol import ( AnonymousPayload, + ArrayConverter, BitMask, BoolConverter, Field, @@ -19,6 +20,7 @@ RegisterBase, RegisterS32, RegisterU16, + RegisterU16Array, RegisterU8, StringConverter, StructPayload, @@ -44,6 +46,7 @@ "CustomMemberConverterPayload", "BitmaskSplitterPayload", "PortDIOSetPayload", + "MixedMemberLengthPayload", "StartPulsePayload", "StartPulseTrainPayload", "EncoderModePayload", @@ -59,6 +62,9 @@ "PortDIOSet", "PulseDOPort0", "PulseDO0", + "MultiElementPayload", + "SingleElementPayload", + "MixedMemberLength", "StartPulse", "StartPulseTrain", "EncoderMode", @@ -100,9 +106,7 @@ class AnalogDataPayload(StructPayload[np.float32], length=6): analog0: np.float32 = Field(IdentityConverter(np.float32)) analog1: np.float32 = Field(IdentityConverter(np.float32), offset=1) analog2: np.float32 = Field(IdentityConverter(np.float32), offset=2) - accelerometer: NDArray[np.float32] = Field( - IdentityConverter(np.dtype((np.float32, (3,)))), offset=3 - ) + accelerometer: NDArray[np.float32] = Field(ArrayConverter(np.float32, 3), offset=3) class ComplexConfigurationPayload(StructPayload[np.uint8], length=17): @@ -122,9 +126,7 @@ class VersionPayload(StructPayload[np.uint8], length=32): firmware_version: HarpVersion = Field(HarpVersionConverter(np.uint8), offset=3) hardware_version: HarpVersion = Field(HarpVersionConverter(np.uint8), offset=6) core_id: str = Field(StringConverter(3), offset=9) - interface_hash: NDArray[np.uint8] = Field( - IdentityConverter(np.dtype((np.uint8, (20,)))), offset=12 - ) + interface_hash: NDArray[np.uint8] = Field(ArrayConverter(np.uint8, 20), offset=12) class CustomPayloadPayload(AnonymousPayload[np.uint32]): @@ -159,6 +161,14 @@ class PortDIOSetPayload(AnonymousPayload[np.uint8]): __value__: PortDigitalIOS = BitMask(enum=PortDigitalIOS) +class MixedMemberLengthPayload(StructPayload[np.uint8], length=4): + """Represents the payload of the MixedMemberLength register.""" + + absent: np.uint8 = Field(IdentityConverter(np.uint8)) + single: NDArray[np.uint8] = Field(ArrayConverter(np.uint8, 1), offset=1) + multiple: NDArray[np.uint8] = Field(ArrayConverter(np.uint8, 2), offset=2) + + class StartPulsePayload(StructPayload[np.uint16]): """Represents the payload of the StartPulse register.""" @@ -247,6 +257,22 @@ class PulseDO0(RegisterU16): address: ClassVar[int] = 43 +class MultiElementPayload(RegisterU16Array): + address: ClassVar[int] = 44 + length: int = 3 + + +class SingleElementPayload(RegisterU16Array): + address: ClassVar[int] = 45 + length: int = 1 + + +class MixedMemberLength(RegisterBase[MixedMemberLengthPayload]): + address: ClassVar[int] = 46 + payload_type: ClassVar[PayloadType] = PayloadType.U8 + payload_class = MixedMemberLengthPayload + + class StartPulse(RegisterBase[StartPulsePayload]): """Starts a PWM pulse.""" @@ -285,6 +311,9 @@ class EncoderMode(RegisterBase[EncoderModeMask]): 41: PortDIOSet, 42: PulseDOPort0, 43: PulseDO0, + 44: MultiElementPayload, + 45: SingleElementPayload, + 46: MixedMemberLength, 100: StartPulse, 101: StartPulseTrain, 103: EncoderMode, diff --git a/tests/device/test_emit.py b/tests/device/test_emit.py index d8ddd16..37a805e 100644 --- a/tests/device/test_emit.py +++ b/tests/device/test_emit.py @@ -574,7 +574,7 @@ def _random_records(dtype, n, seed): """ rng = np.random.default_rng(seed) raw = rng.integers(0, 128, size=n * dtype.itemsize, dtype=np.uint8) - return raw.view(dtype).copy() + return np.frombuffer(raw, dtype=dtype, count=n).copy() @pytest.mark.parametrize("name", sorted(_device_registers()))