diff --git a/pymodbus/pdu/register_message.py b/pymodbus/pdu/register_message.py index 00eb331ee..de18f05ac 100644 --- a/pymodbus/pdu/register_message.py +++ b/pymodbus/pdu/register_message.py @@ -119,6 +119,7 @@ def __init__( self.write_registers = write_registers self.write_count = len(self.write_registers) self.write_byte_count = self.write_count * 2 + self._payload_byte_count: int | None = None def encode(self) -> bytes: """Encode the request packet.""" @@ -147,10 +148,11 @@ def decode(self, data: bytes) -> None: self.write_count, self.write_byte_count, ) = struct.unpack(">HHHHB", data[:9]) - self.write_registers = [] - for i in range(9, self.write_byte_count + 9, 2): - register = struct.unpack(">H", data[i : i + 2])[0] - self.write_registers.append(register) + self._payload_byte_count = len(data) - 9 + self.write_registers = [ + struct.unpack(">H", data[i : i + 2])[0] + for i in range(9, min(len(data), self.write_byte_count + 9) - 1, 2) + ] async def datastore_update( self, context: ModbusServerContext, device_id: int @@ -160,6 +162,11 @@ async def datastore_update( return ExceptionResponse(self.function_code, ExcCodes.ILLEGAL_VALUE) if not 1 <= self.write_count <= 0x079: return ExceptionResponse(self.function_code, ExcCodes.ILLEGAL_VALUE) + if self.write_byte_count != self.write_count * 2 or ( + self._payload_byte_count is not None + and self._payload_byte_count != self.write_byte_count + ): + return ExceptionResponse(self.function_code, ExcCodes.ILLEGAL_VALUE) rc = await context.async_setValues( device_id, self.function_code, self.write_address, self.write_registers ) @@ -243,6 +250,8 @@ class WriteMultipleRegistersRequest(ModbusPDU): function_code = 16 rtu_byte_count_pos = 6 _pdu_length = 5 # func + adress1 + adress2 + outputQuant1 + outputQuant2 + byte_count: int | None = None + _payload_byte_count: int | None = None def encode(self) -> bytes: """Encode a write single register packet packet request.""" @@ -253,10 +262,12 @@ def encode(self) -> bytes: def decode(self, data: bytes) -> None: """Decode a write single register packet packet request.""" - self.address, self.count, _byte_count = struct.unpack(">HHB", data[:5]) - self.registers = [] - for idx in range(5, (self.count * 2) + 5, 2): - self.registers.append(struct.unpack(">H", data[idx : idx + 2])[0]) + self.address, self.count, self.byte_count = struct.unpack(">HHB", data[:5]) + self._payload_byte_count = len(data) - 5 + self.registers = [ + struct.unpack(">H", data[idx : idx + 2])[0] + for idx in range(5, min(len(data), self.byte_count + 5) - 1, 2) + ] async def datastore_update( self, context: ModbusServerContext, device_id: int @@ -264,6 +275,11 @@ async def datastore_update( """Update diagnostic request on the given device.""" if not 1 <= self.count <= 0x07B: return ExceptionResponse(self.function_code, ExcCodes.ILLEGAL_VALUE) + if self.byte_count is not None and ( + self.byte_count != self.count * 2 + or self._payload_byte_count != self.byte_count + ): + return ExceptionResponse(self.function_code, ExcCodes.ILLEGAL_VALUE) rc = await context.async_setValues( device_id, self.function_code, self.address, self.registers ) diff --git a/test/pdu/test_register_read_messages.py b/test/pdu/test_register_read_messages.py index 6d606efe8..fe4ffb154 100644 --- a/test/pdu/test_register_read_messages.py +++ b/test/pdu/test_register_read_messages.py @@ -141,6 +141,46 @@ async def test_read_write_multiple_registers_request(self, mock_server_context): response = await request.datastore_update(context, 0) assert request.function_code == response.function_code + @pytest.mark.parametrize( + "frame", + [ + b"\x00\x01\x00\x01\x00\x02\x00\x02\x02\x00\x0a", + b"\x00\x01\x00\x01\x00\x02\x00\x01\x04\x00\x0a\x00\x0b", + b"\x00\x01\x00\x01\x00\x02\x00\x01\x02\x00", + b"\x00\x01\x00\x01\x00\x02\x00\x01\x02\x00\x0a\x00", + b"\x00\x01\x00\x01\x00\x02\x00\x01\x01\x00", + ], + ) + async def test_read_write_multiple_registers_rejects_invalid_byte_count( + self, frame, mock_server_context + ): + """Test inconsistent write byte counts are rejected before writing.""" + request = ReadWriteMultipleRegistersRequest() + request.decode(frame) + context = mock_server_context() + context.async_setValues = mock.AsyncMock() + + result = await request.datastore_update(context, 1) + + assert result.exception_code == ExcCodes.ILLEGAL_VALUE + context.async_setValues.assert_not_awaited() + + async def test_read_write_multiple_registers_accepts_valid_byte_count( + self, mock_server_context + ): + """Test a consistent write byte count reaches the datastore.""" + request = ReadWriteMultipleRegistersRequest() + request.decode(b"\x00\x01\x00\x01\x00\x02\x00\x02\x04\x00\x0a\x00\x0b") + context = mock_server_context() + context.async_setValues = mock.AsyncMock(return_value=0) + + result = await request.datastore_update(context, 1) + + assert result.function_code == request.function_code + context.async_setValues.assert_awaited_once_with( + 1, request.function_code, 2, [0x0A, 0x0B] + ) + async def test_read_write_multiple_registers_verify(self, mock_server_context): """Test read/write multiple registers.""" context = mock_server_context() diff --git a/test/pdu/test_register_write_messages.py b/test/pdu/test_register_write_messages.py index 0be86721b..61134205a 100644 --- a/test/pdu/test_register_write_messages.py +++ b/test/pdu/test_register_write_messages.py @@ -73,6 +73,46 @@ def test_invalid_write_multiple_registers_request(self): request = WriteMultipleRegistersRequest(address=0, registers=None) assert not request.registers + @pytest.mark.parametrize( + "frame", + [ + b"\x00\x01\x00\x02\x02\x00\x0a", + b"\x00\x01\x00\x01\x04\x00\x0a\x00\x0b", + b"\x00\x01\x00\x01\x02\x00", + b"\x00\x01\x00\x01\x02\x00\x0a\x00", + b"\x00\x01\x00\x01\x01\x00", + ], + ) + async def test_write_multiple_registers_rejects_invalid_byte_count( + self, frame, mock_server_context + ): + """Test inconsistent byte counts are rejected before writing.""" + request = WriteMultipleRegistersRequest() + request.decode(frame) + context = mock_server_context() + context.async_setValues = mock.AsyncMock() + + result = await request.datastore_update(context, 1) + + assert result.exception_code == ExcCodes.ILLEGAL_VALUE + context.async_setValues.assert_not_awaited() + + async def test_write_multiple_registers_accepts_valid_byte_count( + self, mock_server_context + ): + """Test a consistent byte count reaches the datastore.""" + request = WriteMultipleRegistersRequest() + request.decode(b"\x00\x01\x00\x02\x04\x00\x0a\x00\x0b") + context = mock_server_context() + context.async_setValues = mock.AsyncMock(return_value=0) + + result = await request.datastore_update(context, 1) + + assert result.count == 2 + context.async_setValues.assert_awaited_once_with( + 1, request.function_code, 1, [0x0A, 0x0B] + ) + def test_serializing_to_string(self): """Test serializing to string.""" for request in iter(self.write.keys()):