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 4923e5c..43dd452 100644 --- a/src/packages/harp-device/src/harp/device/schema/_emit.py +++ b/src/packages/harp-device/src/harp/device/schema/_emit.py @@ -489,6 +489,7 @@ def _build_register(self, name: str, class_name: str, reg: Register) -> type[Reg "address": reg.address, "payload_type": ProtoPayloadType[reg.type.name], "payload_class": payload_cls, + "length": reg.length, }, ) diff --git a/src/packages/harp-protocol/src/harp/protocol/_register.py b/src/packages/harp-protocol/src/harp/protocol/_register.py index 31109e1..a6c6882 100644 --- a/src/packages/harp-protocol/src/harp/protocol/_register.py +++ b/src/packages/harp-protocol/src/harp/protocol/_register.py @@ -149,7 +149,8 @@ class RegisterBase(ABC, Generic[U], metaclass=_RegisterBaseMeta): ``BitMask`` or ``GroupMask``, even though each still has a ``payload_class``. Subclasses must define ``address``, ``payload_type``, and - ``payload_class`` as ``ClassVar``s. + ``payload_class`` as ``ClassVar``s. ``length`` need not be declared: it is derived + from ``payload_class``, and a subclass that does declare it keeps what it declared. """ address: ClassVar[int] @@ -157,6 +158,50 @@ class RegisterBase(ABC, Generic[U], metaclass=_RegisterBaseMeta): payload_class: ClassVar[type[PayloadBase[Any]]] length: ClassVar[int | None] = None + def __init_subclass__(cls, **kwargs: Any) -> None: + """Set ``length`` from ``payload_class``, so the two can never disagree.""" + super().__init_subclass__(**kwargs) + has_payload_class = getattr(cls, "payload_class", None) is not None + has_payload_type = getattr(cls, "payload_type", None) is not None + if not (has_payload_class and has_payload_type): + # We may want to remove this in the future (or guard against it). + # It is here to allow the creation of intermediate abstract base + # classes that don't declare a `payload_class` or `payload_type`. + return + cls.length = cls._derive_length() + + @classmethod + def _derive_length(cls) -> int | None: + """How many ``payload_type`` elements ``payload_class`` spans, ``None`` for one. + + The element width comes from ``payload_type`` rather than the payload's own + ``_elem_dtype``, since that is what ``HarpMessage.decode`` measures the frame with + and the two need not agree. A register declaring its own ``length`` has it checked + against the payload but not otherwise trusted: the payload decides, and a + sub-array payload counts as an array however few elements it spans, so arrayness + comes from the layout rather than from anyone remembering to declare it. + """ + payload_class: type[PayloadBase[Any]] = cls.payload_class + payload_dtype: np.dtype = payload_class.payload_dtype + itemsize: int = payload_dtype.itemsize + element_size: int = np.dtype(cls.payload_type.value).itemsize + count, remainder = divmod(itemsize, element_size) + if remainder: + raise TypeError( + f"{cls.__name__}: payload {payload_class.__name__} spans {itemsize} bytes, " + f"which is not a whole number of {cls.payload_type!r} elements " + f"({element_size} byte(s))." + ) + declared: int | None = cls.__dict__.get("length") + if declared == 1: + declared = None + if declared is not None and declared != count: + raise TypeError( + f"{cls.__name__} declares length={declared!r} but payload " + f"{payload_class.__name__} spans {count} {cls.payload_type!r} element(s)." + ) + return count if count > 1 or payload_dtype.shape else None + @classmethod def parse(cls, value: HarpMessage | bytes | bytearray | memoryview) -> U: """Parse a single message into the user-facing payload value. diff --git a/tests/device/test_emit.py b/tests/device/test_emit.py index f51c4b4..47d3ac2 100644 --- a/tests/device/test_emit.py +++ b/tests/device/test_emit.py @@ -198,6 +198,70 @@ def test_struct_masked_members_roundtrip(device_registers): assert int(parsed.pulse_width) == 300 +def test_multi_element_register_decodes_through_a_message(): + # https://github.com/harp-tech/python/issues/36: an emitted register reported no + # length, so decode -- the Device.read/write path -- rejected its own frames. + regs = 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" + ) + reg = regs["Settings"] + assert reg.length == 3 + 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_register_spanning_elements_without_length_decodes_through_a_message(): + # Same failure without a declared length: the payloadSpec offsets set the width. + regs = create_registers( + "registers:\n" + " Combo:\n" + " address: 33\n" + " type: U16\n" + " access: Read\n" + " payloadSpec:\n" + " Alpha: {offset: 0}\n" + " Beta: {offset: 1}\n" + ) + reg = regs["Combo"] + assert reg.length == 2 + payload = reg.payload_class(alpha=7, beta=9) + decoded = HarpMessage.parse(reg.format(payload)).decode(reg) + assert (int(decoded.payload.alpha), int(decoded.payload.beta)) == (7, 9) + + +def test_declared_length_of_one_collapses_to_none(): + base = ( + "registers:\n" + " R:\n" + " address: 40\n" + " type: U8\n" + " access: Read\n" + ) # fmt: skip + omitted = create_registers(base)["R"] + declared = create_registers(base + " length: 1\n")["R"] + assert omitted.length is None + assert declared.length is None + assert declared.payload_class.payload_dtype == omitted.payload_class.payload_dtype + assert declared.payload_class.payload_dtype.shape == () # scalar, not a 1-elem array + + +def test_core_multi_element_register_decodes_through_a_message(): + # DeviceName carries 25 U8 elements, so this broke on every Harp device. + assert core.DeviceName.length == 25 + decoded = HarpMessage.parse(core.DeviceName.format("Behavior")).decode(core.DeviceName) + assert decoded.payload == "Behavior" + + def test_custom_converter_roundtrip(device_registers): reg = device_registers["CustomMemberConverter"] payload_cls = reg.payload_class diff --git a/tests/protocol/test_register.py b/tests/protocol/test_register.py index ccd4a88..db13dfc 100644 --- a/tests/protocol/test_register.py +++ b/tests/protocol/test_register.py @@ -31,6 +31,7 @@ RegisterS32, RegisterS64, RegisterU8, + RegisterU8Array, RegisterU16, RegisterU16Array, RegisterU32, @@ -284,6 +285,106 @@ def test_array_register_factory_different_lengths_independent(): assert r1.payload_class is not r2.payload_class +def test_multi_element_struct_register_derives_length(): + # https://github.com/harp-tech/python/issues/36: a register reporting no length for + # a multi-element payload parses fine, but HarpMessage.decode rejects its own frames. + assert AnalogData.length == 3 + sample = np.array([(100, 512, -200)], dtype=AnalogDataPayload.payload_dtype) + decoded = _parse_frame(AnalogData.format(sample)).decode(AnalogData) + assert int(decoded.payload.encoder) == 512 + + +def test_multi_byte_single_value_register_derives_length(): + from harp.protocol._payload import AnonymousPayload + from harp.protocol._payload_converters import StringConverter + + class PayloadDeviceName(AnonymousPayload[np.uint8]): + __value__: str = Field(StringConverter(25)) + + class DeviceName(RegisterBase[str]): + address: ClassVar[int] = 12 + payload_type: ClassVar[PayloadType] = PayloadType.U8 + payload_class = PayloadDeviceName + + assert DeviceName.length == 25 + decoded = _parse_frame(DeviceName.format("Behavior")).decode(DeviceName) + assert decoded.payload == "Behavior" + + +def test_scalar_register_derives_no_length(): + # A single element stays None, the "one value" contract decode() reads. + assert RegisterU16(0x20).length is None + assert DigitalOutputSet.length is None + + +def test_array_register_length_survives_derivation(): + # Array registers declare their own length; it agrees with the sub-array dtype. + reg = RegisterU32Array(0x08, length=3) + assert reg.length == 3 + assert reg.payload_class.payload_dtype.itemsize == 3 * 4 + + +def test_one_element_array_register_reports_a_length_from_its_dtype(): + # A one-element array and a scalar span the same byte, so the sub-array dtype is + # what tells them apart rather than the count. + reg = RegisterU8Array(0x20, length=1) + scalar = RegisterU8(0x20) + assert reg.payload_class.payload_dtype.itemsize == scalar.payload_class.payload_dtype.itemsize + assert reg.payload_class.payload_dtype.shape == (1,) + assert scalar.payload_class.payload_dtype.shape == () + assert reg.length == 1 + assert scalar.length is None + # decode() reads (length or 1), so either way the frame round-trips + frame = reg.format(np.array([7], dtype=np.uint8)) + np.testing.assert_array_equal(_parse_frame(frame).decode(reg).payload, [7]) + + +def test_subclassing_a_one_element_array_register_keeps_it_an_array(): + # The subclass declares nothing, so the dtype has to carry the arrayness. + reg = RegisterU8Array(0x20, length=1) + sub = type("SubArrayRegister", (reg,), {}) + assert sub.length == 1 + + +def test_declared_length_of_one_on_a_scalar_payload_normalizes_to_none(): + # A declared 1 says nothing a single-element payload does not; see the TODO in + # RegisterBase._derive_length. + class DeclaredOne(RegisterBase[np.uint8]): + address: ClassVar[int] = 35 + payload_type: ClassVar[PayloadType] = PayloadType.U8 + payload_class = PayloadU8 + length: ClassVar[int | None] = 1 + + assert DeclaredOne.length is None + + +def test_declared_length_must_be_a_positive_count(): + with pytest.raises(TypeError, match="declares length=0"): + + class ZeroLength(RegisterBase[np.uint8]): + address: ClassVar[int] = 35 + payload_type: ClassVar[PayloadType] = PayloadType.U8 + payload_class = PayloadU8 + length: ClassVar[int | None] = 0 + + +def test_derive_length_measures_in_payload_type_elements(): + assert AnalogDataPayload._elem_dtype == np.dtype(np.uint8) + assert AnalogDataPayload.payload_dtype.itemsize == 6 + assert AnalogData._derive_length() == 3 + + +def test_declared_length_contradicting_payload_is_rejected(): + # A wrong declaration fails at class creation, not at the first decode. + with pytest.raises(TypeError, match="declares length=2 but payload"): + + class Mismatched(RegisterBase[AnalogDataPayload]): + address: ClassVar[int] = 34 + payload_type: ClassVar[PayloadType] = PayloadType.S16 + payload_class = AnalogDataPayload + length: ClassVar[int | None] = 2 + + def test_array_register_format_write(): reg = RegisterU32Array(0x08, length=3) values = np.array([10, 20, 30], dtype=np.dtype("