Skip to content
Closed
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
1 change: 1 addition & 0 deletions src/packages/harp-device/src/harp/device/schema/_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
},
)

Expand Down
47 changes: 46 additions & 1 deletion src/packages/harp-protocol/src/harp/protocol/_register.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,14 +149,59 @@ 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]
payload_type: ClassVar[PayloadType]
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.
Expand Down
64 changes: 64 additions & 0 deletions tests/device/test_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
101 changes: 101 additions & 0 deletions tests/protocol/test_register.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
RegisterS32,
RegisterS64,
RegisterU8,
RegisterU8Array,
RegisterU16,
RegisterU16Array,
RegisterU32,
Expand Down Expand Up @@ -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("<u4"))
Expand Down