Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
37 changes: 29 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,12 @@ 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
payload = data[9 : self.write_byte_count + 9]

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why make an extra split ? Line 154 could just use data directly.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done. The registers are now unpacked directly from data. I applied the same simplification to both 0x17 and 0x10.

self.write_registers = [
struct.unpack(">H", payload[i : i + 2])[0]
for i in range(0, len(payload) - 1, 2)
]

async def datastore_update(
self, context: ModbusServerContext, device_id: int
Expand All @@ -160,6 +163,14 @@ 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 +254,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 +266,25 @@ 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
payload = data[5 : self.byte_count + 5]
self.registers = [
struct.unpack(">H", payload[idx : idx + 2])[0]
for idx in range(0, len(payload) - 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