Skip to content
Open
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
41 changes: 25 additions & 16 deletions s7commplus/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -604,32 +604,38 @@ 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.

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
Expand Down Expand Up @@ -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)

Expand All @@ -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
Expand All @@ -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)

Expand Down
23 changes: 23 additions & 0 deletions tests/test_s7_tls.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
65 changes: 56 additions & 9 deletions tests/test_s7_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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."""

Expand Down
Loading