diff --git a/s7commplus/async_client.py b/s7commplus/async_client.py index 93f4b4aa..7ff53c1a 100644 --- a/s7commplus/async_client.py +++ b/s7commplus/async_client.py @@ -436,10 +436,18 @@ async def db_read(self, db_number: int, start: int, size: int) -> bytes: async def db_write(self, db_number: int, start: int, data: bytes) -> None: """Write raw bytes to a data block.""" - payload = _build_write_payload([(db_number, start, data)], self._protocol_version) + await self.db_write_multi([(db_number, start, data)]) + + async def db_write_multi(self, items: list[tuple[int, int, bytes]]) -> None: + """Write multiple data block regions in a single request.""" + payload = _build_write_payload(items, self._protocol_version) response = await self._send_request(FunctionCode.SET_MULTI_VARIABLES, payload) _parse_write_response(response) + async def write_multi(self, items: list[tuple[int, int, bytes]]) -> None: + """Alias for :meth:`db_write_multi`.""" + await self.db_write_multi(items) + async def db_read_multi(self, items: list[tuple[int, int, int]]) -> list[bytes]: """Read multiple data block regions in a single request.""" payload = _build_read_payload(items, self._protocol_version) diff --git a/s7commplus/client.py b/s7commplus/client.py index 7a47c819..6bc9050e 100644 --- a/s7commplus/client.py +++ b/s7commplus/client.py @@ -207,17 +207,30 @@ def db_write(self, db_number: int, start: int, data: bytes) -> None: start: Start byte offset data: Bytes to write """ + self.db_write_multi([(db_number, start, data)]) + + def db_write_multi(self, items: list[tuple[int, int, bytes]]) -> None: + """Write multiple data block regions in a single request. + + Args: + items: List of ``(db_number, start_offset, data)`` tuples. + """ if self._connection is None: raise RuntimeError("Not connected") if self._connection.requires_substreamed: - self._db_write_substreamed(db_number, start, data) + for db_number, start, data in items: + self._db_write_substreamed(db_number, start, data) return - payload = _build_write_payload([(db_number, start, data)], self._connection.protocol_version) + payload = _build_write_payload(items, self._connection.protocol_version) response = self._connection.send_request(FunctionCode.SET_MULTI_VARIABLES, payload) _parse_write_response(response) + def write_multi(self, items: list[tuple[int, int, bytes]]) -> None: + """Alias for :meth:`db_write_multi`.""" + self.db_write_multi(items) + def _db_write_substreamed(self, db_number: int, start: int, data: bytes) -> None: assert self._connection is not None access_area = Ids.DB_ACCESS_AREA_BASE + (db_number & 0xFFFF) diff --git a/tests/test_s7_server.py b/tests/test_s7_server.py index 6b114562..d498dd97 100644 --- a/tests/test_s7_server.py +++ b/tests/test_s7_server.py @@ -208,6 +208,24 @@ def test_multi_read(self, server: S7CommPlusServer) -> None: finally: client.disconnect() + def test_multi_write(self, server: S7CommPlusServer) -> None: + client = S7CommPlusClient() + client.connect("127.0.0.1", port=TEST_PORT) + try: + client.db_write_multi( + [ + (1, 0, b"first"), + (1, 10, b"second"), + (2, 20, b"third"), + ] + ) + + assert client.db_read(1, 0, 5) == b"first" + assert client.db_read(1, 10, 6) == b"second" + assert client.db_read(2, 20, 5) == b"third" + finally: + client.disconnect() + def test_explore(self, server: S7CommPlusServer) -> None: client = S7CommPlusClient() client.connect("127.0.0.1", port=TEST_PORT) @@ -305,6 +323,21 @@ async def test_multi_read(self, server: S7CommPlusServer) -> None: temp = struct.unpack(">f", results[0])[0] assert abs(temp - 23.5) < 0.1 # May be modified by earlier test + async def test_multi_write(self, server: S7CommPlusServer) -> None: + async with S7CommPlusAsyncClient() as client: + await client.connect("127.0.0.1", port=TEST_PORT) + await client.write_multi( + [ + (1, 0, b"alpha"), + (1, 10, b"beta"), + (2, 20, b"gamma"), + ] + ) + + assert await client.db_read(1, 0, 5) == b"alpha" + assert await client.db_read(1, 10, 4) == b"beta" + assert await client.db_read(2, 20, 5) == b"gamma" + async def test_explore(self, server: S7CommPlusServer) -> None: async with S7CommPlusAsyncClient() as client: await client.connect("127.0.0.1", port=TEST_PORT) diff --git a/tests/test_s7_unit.py b/tests/test_s7_unit.py index 7466b967..9acb272e 100644 --- a/tests/test_s7_unit.py +++ b/tests/test_s7_unit.py @@ -1,6 +1,8 @@ """Unit tests for S7CommPlus client payload builders, connection parsing, and error paths.""" import struct +from unittest.mock import MagicMock, call + import pytest from s7commplus.client import ( @@ -15,12 +17,13 @@ _build_area_write_payload, _build_symbolic_read_payload, _build_symbolic_write_payload, + _build_substreamed_write_payload, ) from s7commplus.connection import S7CommPlusConnection, _strip_paom_string_in_session_version from s7commplus.codec import encode_object_qualifier, encode_pvalue_blob from s7commplus.codec import _pvalue_element_size as _element_size from s7commplus.codec import skip_typed_value, parse_server_session_version -from s7commplus.protocol import DataType, ElementID, ObjectId +from s7commplus.protocol import DataType, ElementID, FunctionCode, Ids, ObjectId from s7commplus.vlq import ( encode_uint32_vlq, encode_uint64_vlq, @@ -553,6 +556,42 @@ def test_db_read_multi_not_connected(self) -> None: with pytest.raises(RuntimeError, match="Not connected"): client.db_read_multi([(1, 0, 4)]) + def test_db_write_multi_not_connected(self) -> None: + client = S7CommPlusClient() + with pytest.raises(RuntimeError, match="Not connected"): + client.db_write_multi([(1, 0, b"data")]) + + def test_write_multi_not_connected(self) -> None: + client = S7CommPlusClient() + with pytest.raises(RuntimeError, match="Not connected"): + client.write_multi([(1, 0, b"data")]) + + def test_db_write_multi_uses_one_substreamed_request_per_item(self) -> None: + client = S7CommPlusClient() + connection = MagicMock() + connection.requires_substreamed = True + connection.session_id = 0x70000001 + client._connection = connection + items = [(1, 0, b"first"), (2, 10, b"second")] + + client.db_write_multi(items) + + connection.send_request.assert_has_calls( + [ + call( + FunctionCode.SET_VAR_SUBSTREAMED, + _build_substreamed_write_payload( + connection.session_id, + Ids.DB_ACCESS_AREA_BASE + db_number, + Ids.DB_VALUE_ACTUAL, + [start + 1, len(data)], + data, + ), + ) + for db_number, start, data in items + ] + ) + def test_explore_not_connected(self) -> None: client = S7CommPlusClient() with pytest.raises(RuntimeError, match="Not connected"):