diff --git a/s7commplus/server.py b/s7commplus/server.py index bd572982..68bb8ac0 100644 --- a/s7commplus/server.py +++ b/s7commplus/server.py @@ -604,23 +604,29 @@ def _process_request( request_data = payload[14:] if function_code == FunctionCode.INIT_SSL: - return self._handle_init_ssl(seq_num) + response = self._handle_init_ssl(seq_num) elif function_code == FunctionCode.CREATE_OBJECT: - return self._handle_create_object(seq_num, request_data) + response = self._handle_create_object(seq_num, request_data) elif function_code == FunctionCode.DELETE_OBJECT: - return self._handle_delete_object(seq_num, req_session_id) + response = self._handle_delete_object(seq_num, req_session_id) elif function_code == FunctionCode.EXPLORE: - return self._handle_explore(seq_num, req_session_id, request_data) + response = self._handle_explore(seq_num, req_session_id, request_data) elif function_code == FunctionCode.GET_MULTI_VARIABLES: - return self._handle_get_multi_variables(seq_num, req_session_id, request_data) + response = self._handle_get_multi_variables(seq_num, req_session_id, request_data) elif function_code == FunctionCode.SET_MULTI_VARIABLES: - return self._handle_set_multi_variables(seq_num, req_session_id, request_data) + response = self._handle_set_multi_variables(seq_num, req_session_id, request_data) elif function_code == FunctionCode.GET_VAR_SUBSTREAMED: - return self._handle_get_var_substreamed(seq_num, req_session_id, request_data) + response = self._handle_get_var_substreamed(seq_num, req_session_id, request_data) elif function_code == FunctionCode.SET_VAR_SUBSTREAMED: - return self._handle_set_var_substreamed(seq_num, req_session_id, request_data) + response = self._handle_set_var_substreamed(seq_num, req_session_id, request_data) else: - return self._build_error_response(seq_num, req_session_id, function_code) + response = self._build_error_response(seq_num, req_session_id, function_code) + + if self._protocol_version >= ProtocolVersion.V2 and req_session_id != 0: + integrity_id = integrity_id_read if function_code in READ_FUNCTION_CODES else integrity_id_write + response += encode_uint32_vlq(integrity_id) + + return response def _build_response_header(self, function_code: int, seq_num: int) -> bytes: """Build a 10-byte S7CommPlus data-response header. @@ -628,8 +634,8 @@ def _build_response_header(self, function_code: int, seq_num: int) -> bytes: Unlike requests (which carry a 4-byte SessionId, giving a 14-byte header), real S7-1500 *responses* omit the SessionId field, so the data header is 10 bytes: opcode + reserved + function + reserved + - seqnr + transport. For V2+, the IntegrityId travels at the *end* of - the payload (appended by the individual handlers), not in the header. + seqnr + transport. For V2+, _process_request appends the IntegrityId + at the *end* of the payload, not in the header. Args: function_code: Response function code @@ -895,8 +901,10 @@ def _handle_get_multi_variables(self, seq_num: int, session_id: int, request_dat # Terminate error list response += encode_uint32_vlq(0) - # IntegrityId - response += encode_uint32_vlq(0) + # V1 responses retain the legacy zero IntegrityId field. V2+ responses + # receive the current per-client counter in _process_request(). + if self._protocol_version < ProtocolVersion.V2: + response += encode_uint32_vlq(0) return bytes(response) @@ -918,7 +926,8 @@ def _handle_set_multi_variables(self, seq_num: int, session_id: int, request_dat logger.debug("SetMultiVariables: accepting session setup write") response += encode_uint64_vlq(0) # ReturnValue: success response += encode_uint32_vlq(0) # Empty error list - response += encode_uint32_vlq(0) # IntegrityId + if self._protocol_version < ProtocolVersion.V2: + response += encode_uint32_vlq(0) # Legacy V1 IntegrityId return bytes(response) # Parse request payload for DB writes @@ -944,8 +953,8 @@ def _handle_set_multi_variables(self, seq_num: int, session_id: int, request_dat # Terminate error list response += encode_uint32_vlq(0) - # IntegrityId - response += encode_uint32_vlq(0) + if self._protocol_version < ProtocolVersion.V2: + response += encode_uint32_vlq(0) # Legacy V1 IntegrityId return bytes(response) diff --git a/tests/test_s7_tls.py b/tests/test_s7_tls.py index e0e500ec..1c95ed85 100644 --- a/tests/test_s7_tls.py +++ b/tests/test_s7_tls.py @@ -213,6 +213,29 @@ async def test_integrity_id_tracking_enabled(self, tls_server: tuple[S7CommPlusS finally: await client.disconnect() + @pytest.mark.asyncio + async def test_sequential_reads_and_writes_keep_integrity_counters_in_sync( + self, tls_server: tuple[S7CommPlusServer, str, str] + ) -> None: + """Multiple V2 operations should succeed as both counters advance.""" + _, cert_path, _ = tls_server + + client = S7CommPlusAsyncClient() + await client.connect("127.0.0.1", port=TEST_PORT_V2_TLS, use_tls=True, tls_ca=cert_path) + + try: + initial_read_id = client._integrity_id_read + initial_write_id = client._integrity_id_write + + for value in (b"first", b"second", b"third"): + await client.db_write(1, 0, value) + assert await client.db_read(1, 0, len(value)) == value + + assert client._integrity_id_read == initial_read_id + 3 + assert client._integrity_id_write == initial_write_id + 3 + finally: + await client.disconnect() + @pytest.mark.asyncio async def test_protocol_version_is_v2(self, tls_server: tuple[S7CommPlusServer, str, str]) -> None: """V2 server should report protocol version 2.""" diff --git a/tests/test_s7_v2.py b/tests/test_s7_v2.py index a2299fd1..8a3411a4 100644 --- a/tests/test_s7_v2.py +++ b/tests/test_s7_v2.py @@ -5,23 +5,21 @@ """ import hashlib +import struct import pytest -from s7commplus.protocol import ( - FunctionCode, - LegitimationId, - ProtocolVersion, - READ_FUNCTION_CODES, -) +from s7commplus.codec import encode_header +from s7commplus.connection import S7CommPlusConnection from s7commplus.legitimation import ( LegitimationState, + _build_legitimation_payload, build_legacy_response, derive_legitimation_key, - _build_legitimation_payload, ) -from s7commplus.vlq import encode_uint32_vlq, decode_uint32_vlq -from s7commplus.connection import S7CommPlusConnection +from s7commplus.protocol import READ_FUNCTION_CODES, FunctionCode, LegitimationId, Opcode, ProtocolVersion +from s7commplus.server import S7CommPlusServer +from s7commplus.vlq import decode_uint32_vlq, encode_uint32_vlq class TestReadFunctionCodes: @@ -192,6 +190,55 @@ def test_protocol_version_default(self) -> None: assert conn.protocol_version == 0 +class TestServerResponseIntegrityId: + """Test V2 response IntegrityId selection and encoding.""" + + @staticmethod + def _request(function_code: int) -> bytes: + request = struct.pack( + ">BHHHHIB", + Opcode.REQUEST, + 0, + function_code, + 0, + 1, + 0x12345678, + 0x34, + ) + return encode_header(ProtocolVersion.V2, len(request)) + request + + @pytest.mark.parametrize( + ("function_code", "expected_integrity_id"), + [ + (FunctionCode.GET_MULTI_VARIABLES, 128), + (FunctionCode.EXPLORE, 128), + (FunctionCode.GET_VAR_SUBSTREAMED, 128), + (FunctionCode.SET_MULTI_VARIABLES, 16384), + (FunctionCode.SET_VAR_SUBSTREAMED, 16384), + (FunctionCode.DELETE_OBJECT, 16384), + ], + ) + def test_v2_response_appends_function_counter(self, function_code: int, expected_integrity_id: int) -> None: + server = S7CommPlusServer(protocol_version=ProtocolVersion.V2) + request = self._request(function_code) + + initial_response = server._process_request(request, 0x12345678) + advanced_response = server._process_request(request, 0x12345678, integrity_id_read=128, integrity_id_write=16384) + + assert initial_response is not None + assert advanced_response is not None + assert advanced_response == initial_response[:-1] + encode_uint32_vlq(expected_integrity_id) + + def test_v1_response_keeps_legacy_integrity_field(self) -> None: + server = S7CommPlusServer(protocol_version=ProtocolVersion.V1) + request = self._request(FunctionCode.GET_MULTI_VARIABLES) + + initial_response = server._process_request(request, 0x12345678) + advanced_response = server._process_request(request, 0x12345678, integrity_id_read=128) + + assert advanced_response == initial_response + + class TestIntegrityIdVlqEncoding: """Test VLQ encoding used for IntegrityId values."""