diff --git a/README.md b/README.md index c477b05..afaaa0f 100644 --- a/README.md +++ b/README.md @@ -166,6 +166,27 @@ For all entity definitions: - E.g., `sum_scale: [1, 10000]` for two registers starting at `address: 5` uses r1=5, r2=6, calculating r1 * 1 + r2 * 10000. - `shift_bits`: Bit shift right (integer). - `bits`: Bit mask length (integer). + - On a **writable** entity (`control: number`, `control: select` or `control: switch`), + `bits` and `shift_bits` make the entity address a *bit field*: the write becomes a + read-modify-write, so the other bits of the register keep their values. The read and the + write are issued under the client lock, so a poll cannot interleave between them. + - This lets several independent controls share one register. E.g. two switches in register 0, + one on bit 1 and one on bit 2, where toggling either leaves the other untouched: + ```yaml + heating: + address: 0 + bits: 1 + shift_bits: 1 + control: switch + hot_water: + address: 0 + bits: 1 + shift_bits: 2 + control: switch + ``` + - `signed` and `sum_scale` are rejected on a *writable* bit field — neither has a meaningful + inverse when merging a value back into part of a register. They remain valid on read-only + entities. - `multiplier`: Scaling factor (float). - `offset`: Adds an offset (float). - **Display**: diff --git a/custom_components/modbus_local_gateway/conversion.py b/custom_components/modbus_local_gateway/conversion.py index 346efc0..4d43b03 100644 --- a/custom_components/modbus_local_gateway/conversion.py +++ b/custom_components/modbus_local_gateway/conversion.py @@ -200,27 +200,99 @@ def _apply_conversion_operations( return num - def _convert_from_decimal( - self, num: float, desc: ModbusEntityDescription - ) -> list[int]: - """Convert from a decimal to registers""" + def _descale(self, num: float, desc: ModbusEntityDescription) -> int: + """Reverse multiplier/offset and round to the raw register value. + + The inverse of `_apply_conversion_operations`. + """ if desc.conv_offset: num -= desc.conv_offset if desc.conv_multiplier is not None: num = num / desc.conv_multiplier - if desc.conv_bits: - raise NotSupportedError("Setting of bit fields is not supported") - if desc.conv_shift_bits: - raise NotSupportedError("Setting of bit fields is not supported") + return int(round(num)) + + def field_geometry(self, desc: ModbusEntityDescription) -> tuple[int, int]: + """Return (shift, mask) for the bit field described by `desc`. + + The geometry comes from `desc.conv_shift_bits` and `desc.conv_bits` - + the `shift_bits` and `bits` keys of a device YAML. With no `conv_bits` + the field runs from the shift to the top of the register span, which is + how the read path already treats it. + """ + span: int = 16 * (desc.register_count or 1) + shift: int = desc.conv_shift_bits or 0 + width: int = desc.conv_bits if desc.conv_bits is not None else span - shift + return shift, (1 << width) - 1 + + def _convert_from_decimal( + self, num: float, desc: ModbusEntityDescription + ) -> list[int]: + """Convert from a decimal to registers""" + if desc.conv_bits or desc.conv_shift_bits: + raise NotSupportedError( + "Bit fields must be written with merge_into_registers" + ) if desc.conv_sum_scale: raise NotSupportedError("Setting of scaled sums is not supported") registers: list[int] = self.client.convert_to_registers( - int(round(num)), + self._descale(num, desc), data_type=self._get_number_data_type(desc), ) return registers + def merge_into_registers( + self, + desc: ModbusEntityDescription, + value: float, + current_registers: list[int], + ) -> list[int]: + """Merge a bit-field value into the register(s) currently on the device. + + Writing a field declared with `bits` / `shift_bits` means reading the + whole register, replacing that field and writing it back. FC 0x16 (Mask + Write Register) would do this on the device, but it is optional and + cannot span a multi-register field. + + The caller must hold the client lock across the read and the write. + """ + if desc.conv_sum_scale: + raise NotSupportedError("Setting of scaled sums is not supported") + + data_type = self._get_number_data_type(desc) + # Same _swap_registers call is used here and on the return: it is its own + # inverse, so it un-swaps into native order here and re-swaps back into + # device order below. + raw = self.client.convert_from_registers( + self._swap_registers(current_registers, desc), data_type=data_type + ) + if not isinstance(raw, int): + raise InvalidDataTypeError( + f"Invalid data type for bit field merge: {type(raw).__name__}" + ) + + shift, mask = self.field_geometry(desc) + field: int = self._descale(value, desc) + if not 0 <= field <= mask: + raise ValueError( + f"Value {value} maps to {field}, which does not fit the " + f"{mask.bit_length()}-bit field at shift {shift} of {desc.key}" + ) + + merged: int = (raw & ~(mask << shift)) | (field << shift) + _LOGGER.debug( + "Merging %s into %s: 0x%04X -> 0x%04X (mask 0x%X << %d)", + field, + desc.key, + raw, + merged, + mask, + shift, + ) + return self._swap_registers( + self.client.convert_to_registers(merged, data_type=data_type), desc + ) + def convert_from_response( self, desc: ModbusEntityDescription, response: ModbusPDU ) -> str | float | int | bool | None: diff --git a/custom_components/modbus_local_gateway/entity_management/base.py b/custom_components/modbus_local_gateway/entity_management/base.py index 6c4b7bb..c09ebe4 100644 --- a/custom_components/modbus_local_gateway/entity_management/base.py +++ b/custom_components/modbus_local_gateway/entity_management/base.py @@ -101,6 +101,80 @@ def validate(self) -> bool: return False if not self._validate_scan_interval(): return False + if not self._validate_bitfield(): + return False + return True + + def _validate_bitfield(self) -> bool: + """Check constraints for writable bit fields. + + The merge assumes an unsigned value: a signed field has no well-defined + representation once masked into part of a register. + + The geometry must also fit the register span, and a coil is already a + single bit so the options mean nothing there. + + Limited to the number, switch and select controls - applying it to + sensors would stop already-working entities from being created. + """ + if self.conv_bits is None and self.conv_shift_bits is None: + return True + if self.control_type not in ( + ControlType.NUMBER, + ControlType.SWITCH, + ControlType.SELECT, + ): + return True + + if self.is_signed: + _LOGGER.warning( + "Unable to create entity for %s: %s cannot be combined with " + "%s or %s on a writable entity", + self.key, + IS_SIGNED, + CONV_BITS, + CONV_SHIFT_BITS, + ) + return False + if self.conv_sum_scale: + _LOGGER.warning( + "Unable to create entity for %s: %s cannot be combined with " + "%s or %s on a writable entity", + self.key, + CONV_SUM_SCALE, + CONV_BITS, + CONV_SHIFT_BITS, + ) + return False + return self._validate_bitfield_geometry() + + def _validate_bitfield_geometry(self) -> bool: + """The field must be a real run of bits inside the registers it names.""" + if self.data_type == ModbusDataType.COIL: + _LOGGER.warning( + "Unable to create entity for %s: %s and %s have no meaning on a " + "coil, which is already a single bit", + self.key, + CONV_BITS, + CONV_SHIFT_BITS, + ) + return False + + span: int = 16 * (self.register_count or 1) + shift: int = self.conv_shift_bits or 0 + width: int = self.conv_bits if self.conv_bits is not None else span - shift + if shift < 0 or width <= 0 or shift + width > span: + _LOGGER.warning( + "Unable to create entity for %s: %s %s / %s %s does not fit the " + "%s bits it addresses", + self.key, + CONV_SHIFT_BITS, + shift, + CONV_BITS, + width, + span, + ) + return False return True def _validate_scan_interval(self) -> bool: diff --git a/custom_components/modbus_local_gateway/entity_management/modbus_device_info.py b/custom_components/modbus_local_gateway/entity_management/modbus_device_info.py index d043aec..aba4938 100644 --- a/custom_components/modbus_local_gateway/entity_management/modbus_device_info.py +++ b/custom_components/modbus_local_gateway/entity_management/modbus_device_info.py @@ -345,10 +345,6 @@ def _handle_switch_description( ) -> None | type[ModbusSwitchEntityDescription]: """Handle switch description specific logic""" switch_data = _data.get("switch", {}) - if _data.get(CONV_BITS) or _data.get(CONV_SHIFT_BITS): - _LOGGER.warning("bits / shift bits cannot be set for Switches") - return None - if not isinstance(switch_data, dict): _LOGGER.warning( "Switch configuration for %s should be a dictionary", entity diff --git a/custom_components/modbus_local_gateway/tcp_client.py b/custom_components/modbus_local_gateway/tcp_client.py index 83542db..93c311b 100644 --- a/custom_components/modbus_local_gateway/tcp_client.py +++ b/custom_components/modbus_local_gateway/tcp_client.py @@ -213,6 +213,32 @@ async def _write_registers_individually( return _LOGGER.debug("All individual writes successful using fallback") + async def _read_current_registers(self, entity: ModbusContext) -> list[int]: + """Read the register(s) backing a bit field, for a read-modify-write. + + Must be called with the client lock held, so the read and the write it + feeds cannot be interleaved with a poll. + + The span is read in one transaction rather than in `max_register_read` + chunks: a field split across two reads could tear if the device changed + in between. A failed read raises, abandoning the write - merging onto a + guess would clear the field's neighbours. + """ + response: ModbusPDU | None = await self.read_data( + func=self.read_holding_registers, + address=entity.desc.register_address, + count=entity.desc.register_count, + device_id=entity.device_id, + max_read_size=entity.desc.register_count, + ) + if response is None or response.isError(): + raise ModbusException( + "Unable to read current value of " + f"{entity.desc.key} at {entity.desc.register_address} - " + "aborting bit field write" + ) + return response.registers + async def write_data(self, entity: ModbusContext, value: Any) -> ModbusPDU | None: """Writes data to Holding Registers or Coils""" pdu: ModbusPDU | None = None @@ -236,9 +262,18 @@ async def write_data(self, entity: ModbusContext, value: Any) -> ModbusPDU | Non ) if entity.desc.data_type == ModbusDataType.HOLDING_REGISTER: - registers = Conversion(type(self)).convert_to_registers( - entity.desc, value - ) + conversion = Conversion(type(self)) + if entity.desc.conv_bits or entity.desc.conv_shift_bits: + # No dependable device-side bit write (FC 0x16 is optional): + # read the register, replace this field, write it back. Still + # inside `self.lock`, so nothing lands in between. + registers = conversion.merge_into_registers( + entity.desc, + value, + await self._read_current_registers(entity), + ) + else: + registers = conversion.convert_to_registers(entity.desc, value) _LOGGER.debug( "Raw value after conversion to registers: %s (type: %s)", registers, diff --git a/tests/test_client.py b/tests/test_client.py index 7ae49fa..83ddf96 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -934,3 +934,131 @@ async def test_write_data_invalid_coil_value_type() -> None: pytest.raises(TypeError, match="Value for COIL must be boolean, got int"), ): await client.write_data(entity, value=123) + + +def _bitfield_entity() -> ModbusContext: + """A switch on bit 4 of a holding register.""" + return ModbusContext( + device_id=1, + desc=ModbusSwitchEntityDescription( + key="bitfield", + register_address=1, + register_count=1, + control_type="switch", + data_type=ModbusDataType.HOLDING_REGISTER, + conv_bits=1, + conv_shift_bits=4, + on=1, + off=0, + ), + ) + + +@pytest.mark.asyncio +async def test_write_data_bitfield_read_modify_write() -> None: + """Writing a bit field reads the register and merges into it.""" + client = AsyncModbusTcpClientGateway(host="localhost") + client.connect = AsyncMock() + client.write_register = AsyncMock(return_value=ModbusPDU()) + client.write_register.return_value.isError = lambda: False + + with ( + patch.object( + AsyncModbusTcpClientGateway, "connected", PropertyMock(return_value=True) + ), + patch.object( + AsyncModbusTcpClientGateway, + "read_data", + AsyncMock( + return_value=ReadHoldingRegistersResponse(registers=[0b0000_0011]) + ), + ), + ): + await client.write_data(_bitfield_entity(), value=1) + + # bit 4 set, the two bits already on are untouched + client.write_register.assert_called_once_with( + address=1, + value=0b0001_0011, + device_id=1, + ) + + +@pytest.mark.asyncio +async def test_write_data_bitfield_read_failure_aborts_write() -> None: + """A failed read must abort - merging onto a guess would clear the field's + neighbours, which is worse than not writing at all.""" + client = AsyncModbusTcpClientGateway(host="localhost") + client.connect = AsyncMock() + client.write_register = AsyncMock(return_value=ModbusPDU()) + + with ( + patch.object( + AsyncModbusTcpClientGateway, "connected", PropertyMock(return_value=True) + ), + patch.object( + AsyncModbusTcpClientGateway, "read_data", AsyncMock(return_value=None) + ), + pytest.raises(ModbusException, match="aborting bit field write"), + ): + await client.write_data(_bitfield_entity(), value=1) + + client.write_register.assert_not_called() + + +@pytest.mark.asyncio +async def test_write_data_bitfield_error_response_aborts_write() -> None: + """An error PDU from the read is a failed read, not a value of zero.""" + client = AsyncModbusTcpClientGateway(host="localhost") + client.connect = AsyncMock() + client.write_register = AsyncMock(return_value=ModbusPDU()) + + error_response = ReadHoldingRegistersResponse(registers=[0]) + error_response.isError = lambda: True # type: ignore[method-assign] + + with ( + patch.object( + AsyncModbusTcpClientGateway, "connected", PropertyMock(return_value=True) + ), + patch.object( + AsyncModbusTcpClientGateway, + "read_data", + AsyncMock(return_value=error_response), + ), + pytest.raises(ModbusException, match="aborting bit field write"), + ): + await client.write_data(_bitfield_entity(), value=1) + + client.write_register.assert_not_called() + + +@pytest.mark.asyncio +async def test_write_data_non_bitfield_does_not_read_first() -> None: + """Plain registers keep the single-transaction write they always had.""" + client = AsyncModbusTcpClientGateway(host="localhost") + client.connect = AsyncMock() + client.write_register = AsyncMock(return_value=ModbusPDU()) + client.write_register.return_value.isError = lambda: False + + entity = ModbusContext( + device_id=1, + desc=ModbusEntityDescription( + key="plain", + register_address=1, + register_count=1, + data_type=ModbusDataType.HOLDING_REGISTER, + ), + ) + + with ( + patch.object( + AsyncModbusTcpClientGateway, "connected", PropertyMock(return_value=True) + ), + patch.object( + AsyncModbusTcpClientGateway, "read_data", AsyncMock() + ) as read_data, + ): + await client.write_data(entity, value=123) + + read_data.assert_not_called() + client.write_register.assert_called_once() diff --git a/tests/test_conversion.py b/tests/test_conversion.py index 6d33b8a..a532629 100644 --- a/tests/test_conversion.py +++ b/tests/test_conversion.py @@ -4,15 +4,21 @@ import pytest from pymodbus.client.mixin import ModbusClientMixin from pymodbus.pdu.bit_message import ReadCoilsResponse, ReadDiscreteInputsResponse -from pymodbus.pdu.register_message import ReadInputRegistersResponse +from pymodbus.pdu.register_message import ( + ReadHoldingRegistersResponse, + ReadInputRegistersResponse, +) from custom_components.modbus_local_gateway.conversion import ( Conversion, InvalidDataTypeError, + NotSupportedError, ) from custom_components.modbus_local_gateway.entity_management.base import ( ModbusDataType, + ModbusNumberEntityDescription, ModbusSensorEntityDescription, + ModbusSwitchEntityDescription, ) from custom_components.modbus_local_gateway.entity_management.const import SwapType from custom_components.modbus_local_gateway.tcp_client import AsyncModbusTcpClient @@ -667,3 +673,159 @@ def test_get_float_type(size: int, expected: ModbusClientMixin.DATATYPE) -> None else: result: ModbusClientMixin.DATATYPE = conversion._get_float_data_type(desc) assert result == expected + + +def _switch_desc(**kwargs) -> ModbusSwitchEntityDescription: + """A writable bit-field switch description.""" + return ModbusSwitchEntityDescription( + register_address=1, + key="test", + control_type="switch", + data_type=ModbusDataType.HOLDING_REGISTER, + **kwargs, + ) + + +def _number_desc(**kwargs) -> ModbusNumberEntityDescription: + """A writable bit-field number description.""" + return ModbusNumberEntityDescription( + register_address=1, + key="test", + control_type="number", + data_type=ModbusDataType.HOLDING_REGISTER, + min=0, + max=255, + **kwargs, + ) + + +@pytest.mark.parametrize( + ("current", "value", "expected"), + [ + (0b0000_0000_0001_0011, 1, 0b0000_0000_0001_0011), # already set, no change + (0b0000_0000_0000_0011, 1, 0b0000_0000_0001_0011), # set, neighbours kept + (0b1111_1111_1111_1111, 0, 0b1111_1111_1110_1111), # clear, neighbours kept + (0b0000_0000_0001_0000, 0, 0b0000_0000_0000_0000), # clear the only bit + ], +) +def test_merge_single_bit(current: int, value: int, expected: int) -> None: + """Setting or clearing one bit must leave every other bit untouched.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _switch_desc(conv_bits=1, conv_shift_bits=4, on=1, off=0) + + assert conversion.merge_into_registers(desc, value, [current]) == [expected] + + +def test_merge_low_byte_preserves_high_byte() -> None: + """Two 8-bit fields in one register: writing the low byte must not zero the high byte. + + A register holding 30 in the high byte and 6 in the low byte; writing 45 + to the low byte must leave the high byte at 30. + """ + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_bits=8, conv_shift_bits=0) + + result = conversion.merge_into_registers(desc, 45, [(30 << 8) | 6]) + + assert result == [(30 << 8) | 45] + assert result[0] >> 8 == 30 + + +def test_merge_high_byte_preserves_low_byte() -> None: + """The mirror case: writing the high byte must not disturb the low byte.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_bits=8, conv_shift_bits=8) + + result = conversion.merge_into_registers(desc, 30, [(12 << 8) | 45]) + + assert result == [(30 << 8) | 45] + assert result[0] & 0xFF == 45 + + +def test_merge_across_two_registers() -> None: + """A field straddling the boundary between the two registers of a 32-bit entity.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(register_count=2, conv_bits=8, conv_shift_bits=12) + + current = AsyncModbusTcpClient.convert_to_registers( + 0xABCD_1234, data_type=AsyncModbusTcpClient.DATATYPE.UINT32 + ) + result = conversion.merge_into_registers(desc, 0xFF, current) + + merged = AsyncModbusTcpClient.convert_from_registers( + result, data_type=AsyncModbusTcpClient.DATATYPE.UINT32 + ) + assert merged == 0xABCF_F234 + + +@pytest.mark.parametrize( + "swap", [None, SwapType.BYTE, SwapType.WORD, SwapType.WORD_BYTE] +) +def test_merge_round_trips_through_swap(swap) -> None: + """Merging then reading back must return the value that was written. + + `_swap_registers` is its own inverse, which is what lets the merge use the + same call to un-swap and re-swap. + """ + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc( + register_count=2, conv_bits=8, conv_shift_bits=8, conv_swap=swap + ) + + current = AsyncModbusTcpClient.convert_to_registers( + 0x0000_0000, data_type=AsyncModbusTcpClient.DATATYPE.UINT32 + ) + merged = conversion.merge_into_registers(desc, 0x5A, current) + + read_back = conversion.convert_from_response( + desc=desc, + response=ReadHoldingRegistersResponse(registers=merged), + ) + assert read_back == 0x5A + + +def test_merge_applies_multiplier_and_offset() -> None: + """A scaled bit field descales before it is packed.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_bits=8, conv_shift_bits=0, conv_multiplier=0.5) + + # 21.5 degrees / 0.5 == 43 raw, merged into the low byte + assert conversion.merge_into_registers(desc, 21.5, [0xFF00]) == [0xFF00 | 43] + + +@pytest.mark.parametrize("value", [256, -1]) +def test_merge_rejects_value_that_does_not_fit(value: int) -> None: + """Overflowing the field would corrupt the neighbouring controls.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_bits=8, conv_shift_bits=0) + + with pytest.raises(ValueError, match="does not fit"): + conversion.merge_into_registers(desc, value, [0x1234]) + + +def test_merge_width_defaults_to_top_of_register() -> None: + """With `shift_bits` but no `bits`, the field runs to the top of the span.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_shift_bits=12) + + assert conversion.field_geometry(desc) == (12, 0xF) + assert conversion.merge_into_registers(desc, 0xA, [0x5678]) == [0xA678] + + +def test_convert_to_registers_still_refuses_bit_fields() -> None: + """The plain (non read-modify-write) path must not silently zero a field.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_bits=8, conv_shift_bits=0) + + with pytest.raises(NotSupportedError, match="merge_into_registers"): + conversion.convert_to_registers(desc, 5) + + +def test_merge_refuses_sum_scale() -> None: + """`sum_scale` has no inverse. Validation rejects it at load, but + merge_into_registers is public, so it guards too.""" + conversion = Conversion(client=AsyncModbusTcpClient) + desc = _number_desc(conv_bits=8, conv_sum_scale=[1.0, 0.1]) + + with pytest.raises(NotSupportedError, match="scaled sums"): + conversion.merge_into_registers(desc, 5, [0x1234]) diff --git a/tests/test_entity_management_base.py b/tests/test_entity_management_base.py index c3c2922..b1acf91 100644 --- a/tests/test_entity_management_base.py +++ b/tests/test_entity_management_base.py @@ -4,9 +4,15 @@ # pylint: disable=unexpected-keyword-arg, protected-access from unittest.mock import patch +import pytest + from custom_components.modbus_local_gateway.entity_management.base import ( ModbusEntityDescription, ) +from custom_components.modbus_local_gateway.entity_management.const import ( + ControlType, + ModbusDataType, +) def test_validate_both_float_and_string( @@ -106,3 +112,137 @@ def test_validate_valid_entity( ) as mock_warning: assert entity.validate() mock_warning.assert_not_called() + + +@pytest.mark.parametrize( + "control_type", [ControlType.NUMBER, ControlType.SWITCH, ControlType.SELECT] +) +def test_validate_signed_bitfield_rejected_when_writable( + valid_entity_description: ModbusEntityDescription, control_type: str +) -> None: + """A writable bit field cannot be signed - the merge assumes unsigned.""" + entity: ModbusEntityDescription = valid_entity_description + entity = entity.__class__( + **{ + **entity.__dict__, + "conv_bits": 8, + "is_signed": True, + "control_type": control_type, + } + ) + with patch( + "custom_components.modbus_local_gateway.entity_management.base._LOGGER.warning" + ): + assert not entity.validate() + + +def test_validate_signed_bitfield_allowed_on_sensor( + valid_entity_description: ModbusEntityDescription, +) -> None: + """Read-only entities keep working: tightening them would delete entities + from configs that are valid today.""" + entity: ModbusEntityDescription = valid_entity_description + entity = entity.__class__( + **{ + **entity.__dict__, + "conv_bits": 8, + "is_signed": True, + "control_type": ControlType.SENSOR, + } + ) + assert entity.validate() + + +def test_validate_sum_scale_bitfield_rejected_when_writable( + valid_entity_description: ModbusEntityDescription, +) -> None: + """`sum_scale` has no meaningful inverse, so it cannot be written.""" + entity: ModbusEntityDescription = valid_entity_description + entity = entity.__class__( + **{ + **entity.__dict__, + "conv_shift_bits": 4, + "conv_sum_scale": [1.0, 0.1], + "control_type": ControlType.NUMBER, + } + ) + with patch( + "custom_components.modbus_local_gateway.entity_management.base._LOGGER.warning" + ): + assert not entity.validate() + + +def test_validate_unsigned_bitfield_accepted_when_writable( + valid_entity_description: ModbusEntityDescription, +) -> None: + """The ordinary case still validates.""" + entity: ModbusEntityDescription = valid_entity_description + entity = entity.__class__( + **{ + **entity.__dict__, + "conv_bits": 1, + "conv_shift_bits": 4, + "control_type": ControlType.SWITCH, + } + ) + assert entity.validate() + + +def _writable_bitfield( + entity: ModbusEntityDescription, **overrides +) -> ModbusEntityDescription: + """Build a writable bit-field description from the valid fixture.""" + return entity.__class__( + **{ + **entity.__dict__, + "control_type": ControlType.SWITCH, + **overrides, + } + ) + + +@pytest.mark.parametrize( + ("overrides", "reason"), + [ + ( + {"register_count": 1, "conv_shift_bits": 17}, + "shift past the end of the span gives a negative width, and the " + "first write would raise from 1 << width", + ), + ( + {"register_count": 1, "conv_bits": 8, "conv_shift_bits": 12}, + "field starts inside the span but runs off the end", + ), + ( + {"conv_bits": 0, "conv_shift_bits": 0}, + "an explicit bits: 0 is a mistake, not a request for the whole register", + ), + ( + {"data_type": ModbusDataType.COIL, "conv_bits": 1, "conv_shift_bits": 4}, + "a coil is already one bit and the write path ignores both options", + ), + ], + ids=["shift_past_span", "width_overflows_span", "explicit_zero_width", "coil"], +) +def test_validate_bitfield_geometry_rejected( + valid_entity_description: ModbusEntityDescription, + overrides: dict, + reason: str, +) -> None: + """Bad geometry is refused at load, rather than failing at the first write.""" + entity = _writable_bitfield(valid_entity_description, **overrides) + with patch( + "custom_components.modbus_local_gateway.entity_management.base._LOGGER.warning" + ) as mock_warning: + assert not entity.validate(), reason + mock_warning.assert_called_once() + + +def test_validate_bitfield_spanning_two_registers_accepted( + valid_entity_description: ModbusEntityDescription, +) -> None: + """The geometry check must not reject a field that legitimately spans registers.""" + entity = _writable_bitfield( + valid_entity_description, register_count=2, conv_bits=8, conv_shift_bits=12 + ) + assert entity.validate() diff --git a/tests/test_modbus_device_info.py b/tests/test_modbus_device_info.py index 3b4a0ea..9c4c99b 100644 --- a/tests/test_modbus_device_info.py +++ b/tests/test_modbus_device_info.py @@ -32,14 +32,15 @@ def test_entity_load() -> None: name: Read-Write Entity address: 1 - entity_invalid_bits: - name: Entity Invalid Bits + entity_bitfield_switch: + name: Entity Bitfield Switch address: 2 control: switch bits: 1 + shift_bits: 4 - entity_invalid_shift: - name: Entity Invalid Shift + entity_shifted_switch: + name: Entity Shifted Switch address: 3 control: switch shift_bits: 1 @@ -78,9 +79,17 @@ def test_entity_load() -> None: device = ModbusDeviceInfo("test.yaml") entities = device.entity_descriptions - assert len(entities) == 5 + assert len(entities) == 7 assert device.model == "Model" assert device.manufacturer == "Manufacturer" + # A switch on a bit field is valid: writing it is a read-modify-write. + assert any( + e.key == "entity_bitfield_switch" + and e.conv_bits == 1 + and e.conv_shift_bits == 4 + for e in entities + ) + assert any(e.key == "entity_shifted_switch" for e in entities) assert any( e.key == "entity_rw" and e.data_type == ModbusDataType.HOLDING_REGISTER for e in entities