Skip to content
Merged
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
32 changes: 24 additions & 8 deletions pymodbus/pdu/register_message.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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
Expand All @@ -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
)
Expand Down Expand Up @@ -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."""
Expand All @@ -253,17 +262,24 @@ 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
) -> ModbusPDU:
"""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
)
Expand Down
40 changes: 40 additions & 0 deletions test/pdu/test_register_read_messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
40 changes: 40 additions & 0 deletions test/pdu/test_register_write_messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()):
Expand Down
Loading