diff --git a/benchmarks/micro/bench_checksumming_inline.py b/benchmarks/micro/bench_checksumming_inline.py new file mode 100644 index 0000000000..965b144197 --- /dev/null +++ b/benchmarks/micro/bench_checksumming_inline.py @@ -0,0 +1,57 @@ +# Copyright ScyllaDB, Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Micro-benchmark: inline checksumming check vs classmethod call. + +Measures the overhead of ProtocolVersion.has_checksumming_support() +classmethod call versus an inline integer comparison on the +encode/decode hot path. + +Run: + python benchmarks/micro/bench_checksumming_inline.py +""" + +import sys +import timeit + +from cassandra import ProtocolVersion +from cassandra.protocol import _CHECKSUMMING_MIN_VERSION, _CHECKSUMMING_MAX_VERSION_EXCLUSIVE + + +def bench(): + protocol_version = ProtocolVersion.V4 + + def via_classmethod(): + return ProtocolVersion.has_checksumming_support(protocol_version) + + def via_inline(): + return _CHECKSUMMING_MIN_VERSION <= protocol_version < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE + + n = 5_000_000 + t_classmethod = timeit.timeit(via_classmethod, number=n) + t_inline = timeit.timeit(via_inline, number=n) + + saving_ns = (t_classmethod - t_inline) / n * 1e9 + speedup = t_classmethod / t_inline if t_inline > 0 else float('inf') + + print(f"=== has_checksumming_support ({n:,} iters) ===") + print(f" classmethod call: {t_classmethod / n * 1e9:.1f} ns") + print(f" inline compare: {t_inline / n * 1e9:.1f} ns") + print(f" saving: {saving_ns:.1f} ns/call ({speedup:.1f}x)") + + +if __name__ == "__main__": + print(f"Python {sys.version}") + bench() diff --git a/cassandra/protocol.py b/cassandra/protocol.py index 9dfdbf3022..223b3a667a 100644 --- a/cassandra/protocol.py +++ b/cassandra/protocol.py @@ -72,6 +72,28 @@ class InternalError(Exception): _UNSET_VALUE = object() +# Inline constants mirroring ProtocolVersion.has_checksumming_support(), to +# avoid the classmethod call overhead (~94 ns per call) on the encode/decode +# hot path: +# +# ProtocolVersion.has_checksumming_support(v) == ( +# _CHECKSUMMING_MIN_VERSION <= v < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE) +# +# _CHECKSUMMING_MAX_VERSION_EXCLUSIVE is an *exclusive* upper bound -- DSE_V1 +# itself is deliberately NOT considered checksumming-capable. DSE_V1/DSE_V2 +# are private DSE protocol extensions that do not carry the checksumming +# feature introduced with the (Cassandra) V5 native protocol, even though +# their numeric values (0x41/0x42) are greater than V5/V6. This matches the +# other V5-only feature gates on ProtocolVersion (uses_prepare_flags, +# uses_prepared_metadata, uses_keyspace_flag), which all explicitly exclude +# DSE_V1 too. See ProtocolVersion.has_checksumming_support(), which remains +# the canonical definition (still used as-is in connection.py to decide +# whether to enable frame-level checksumming for a connection); if it ever +# changes, these constants must be updated to match -- test_protocol.py +# asserts the two stay in sync. +_CHECKSUMMING_MIN_VERSION = ProtocolVersion.V5 +_CHECKSUMMING_MAX_VERSION_EXCLUSIVE = ProtocolVersion.DSE_V1 + def register_class(cls): _message_types_by_opcode[cls.opcode] = cls @@ -1150,40 +1172,52 @@ def encode_message(cls, msg, stream_id, protocol_version, compressor, allow_beta flags |= USE_BETA_FLAG buff = io.BytesIO() - buff.seek(9) # With checksumming, the compression is done at the segment frame encoding - if (compressor and not ProtocolVersion.has_checksumming_support(protocol_version)): - body = io.BytesIO() + if (compressor and not (_CHECKSUMMING_MIN_VERSION <= protocol_version < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE)): if msg.custom_payload: - write_bytesmap(body, msg.custom_payload) - msg.send_body(body, protocol_version, protocol_features) - body = body.getvalue() + write_bytesmap(buff, msg.custom_payload) + msg.send_body(buff, protocol_version, protocol_features) + body = buff.getvalue() if len(body) > 0: body = compressor(body) flags |= COMPRESSED_FLAG - buff.write(body) length = len(body) + # Same header layout as the non-compression path below: both go + # through _pack_header so the two can never silently diverge. + return cls._pack_header(protocol_version, flags, stream_id, msg.opcode, length) + body else: + buff.seek(9) + if msg.custom_payload: write_bytesmap(buff, msg.custom_payload) msg.send_body(buff, protocol_version, protocol_features) length = buff.tell() - 9 - buff.seek(0) - cls._write_header(buff, protocol_version, flags, stream_id, msg.opcode, length) - return buff.getvalue() + buff.seek(0) + cls._write_header(buff, protocol_version, flags, stream_id, msg.opcode, length) + return buff.getvalue() @staticmethod - def _write_header(f, version, flags, stream_id, opcode, length): + def _pack_header(version, flags, stream_id, opcode, length): + """ + Pack a CQL protocol frame header into bytes. + + This is the single source of truth for the frame header layout; + ``_write_header`` and the compressed encode path both go through + this method so they cannot silently diverge. + """ + return v3_header_pack(version, flags, stream_id, opcode) + int32_pack(length) + + @classmethod + def _write_header(cls, f, version, flags, stream_id, opcode, length): """ Write a CQL protocol frame header. """ - f.write(v3_header_pack(version, flags, stream_id, opcode)) - write_int(f, length) + f.write(cls._pack_header(version, flags, stream_id, opcode, length)) @classmethod def decode_message(cls, protocol_version, protocol_features, user_type_map, stream_id, flags, opcode, body, @@ -1200,7 +1234,7 @@ def decode_message(cls, protocol_version, protocol_features, user_type_map, stre :param decompressor: optional decompression function to inflate the body :return: a message decoded from the body and frame attributes """ - if (not ProtocolVersion.has_checksumming_support(protocol_version) and + if (not (_CHECKSUMMING_MIN_VERSION <= protocol_version < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE) and flags & COMPRESSED_FLAG): if decompressor is None: raise RuntimeError("No de-compressor available for compressed frame!") diff --git a/tests/unit/test_protocol.py b/tests/unit/test_protocol.py index 75dc69bca5..2dbf34880d 100644 --- a/tests/unit/test_protocol.py +++ b/tests/unit/test_protocol.py @@ -17,14 +17,15 @@ import unittest from typing import ClassVar -from unittest.mock import Mock +from unittest.mock import Mock, patch from cassandra import ConsistencyLevel, ProtocolVersion, UnsupportedOperation from cassandra.protocol import ( PrepareMessage, QueryMessage, ExecuteMessage, BatchMessage, StartupMessage, OptionsMessage, RegisterMessage, - AuthResponseMessage, ProtocolHandler, _MessageType, - ResultMessage, RESULT_KIND_ROWS + AuthResponseMessage, ProtocolHandler, _ProtocolHandler, _MessageType, + ResultMessage, RESULT_KIND_ROWS, COMPRESSED_FLAG, ReadyMessage, + _CHECKSUMMING_MIN_VERSION, _CHECKSUMMING_MAX_VERSION_EXCLUSIVE, ) from cassandra.protocol_features import ProtocolFeatures from cassandra.query import BatchType @@ -555,3 +556,181 @@ def test_frames_without_features(self): def test_frames_with_default_features(self): self._assert_frames(ProtocolFeatures()) + + +class ChecksummingBoundaryTest(unittest.TestCase): + """ + _CHECKSUMMING_MIN_VERSION / _CHECKSUMMING_MAX_VERSION_EXCLUSIVE in + cassandra/protocol.py are hand-inlined for performance, duplicating + ProtocolVersion.has_checksumming_support()'s ``V5 <= v < DSE_V1`` range + check (see cassandra/__init__.py). DSE_V1/DSE_V2 are private DSE + protocol extensions that do not carry the V5 native-protocol + checksumming feature, even though their numeric values (0x41/0x42) are + greater than V5/V6 -- so the upper bound is deliberately *exclusive* of + DSE_V1 itself. This mirrors the other V5-only feature gates on + ProtocolVersion (uses_prepare_flags, uses_prepared_metadata, + uses_keyspace_flag), which also explicitly carve out DSE_V1, and it is + the same boundary connection.py already relies on (unchanged by this + module) to decide whether to enable frame-level checksumming at all. + + These tests guard against the two definitions drifting apart, and + directly exercise the encode_message/decode_message boundary at and + around DSE_V1 -- the exact case a naive ``< _CHECKSUMMING_MAX_VERSION`` + read could get backwards. + """ + + # A representative span of versions straddling both boundaries. + ALL_VERSIONS = ( + ProtocolVersion.V3, ProtocolVersion.V4, + ProtocolVersion.V5, ProtocolVersion.V6, + ProtocolVersion.DSE_V1, ProtocolVersion.DSE_V2, + ) + + def test_inline_constants_match_has_checksumming_support(self): + for version in self.ALL_VERSIONS: + inline_result = _CHECKSUMMING_MIN_VERSION <= version < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE + canonical_result = ProtocolVersion.has_checksumming_support(version) + assert inline_result == canonical_result, ( + "inline checksumming range check disagrees with " + "ProtocolVersion.has_checksumming_support for version %r" % (version,) + ) + + def test_dse_v1_itself_is_excluded(self): + # The crux of the boundary: DSE_V1 must NOT be treated as + # checksumming-capable, despite being the constant's value. + assert not (_CHECKSUMMING_MIN_VERSION <= ProtocolVersion.DSE_V1 < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE) + assert not ProtocolVersion.has_checksumming_support(ProtocolVersion.DSE_V1) + + def test_v5_and_v6_are_included(self): + for version in (ProtocolVersion.V5, ProtocolVersion.V6): + assert _CHECKSUMMING_MIN_VERSION <= version < _CHECKSUMMING_MAX_VERSION_EXCLUSIVE + assert ProtocolVersion.has_checksumming_support(version) + + def _encode_with_spy_compressor(self, protocol_version): + calls = [] + + def spy_compressor(body): + calls.append(body) + return body + + msg = RegisterMessage(["TOPOLOGY_CHANGE", "STATUS_CHANGE"]) + frame = ProtocolHandler.encode_message( + msg, stream_id=0, protocol_version=protocol_version, compressor=spy_compressor, + allow_beta_protocol_version=False, protocol_features=ProtocolFeatures()) + flags_byte = frame[1] + return calls, flags_byte + + def test_encode_message_compresses_at_dse_v1(self): + """ + DSE_V1 has no checksumming, so encode_message must fall back to + message-level compression: the compressor must run and + COMPRESSED_FLAG must be set. + """ + calls, flags_byte = self._encode_with_spy_compressor(ProtocolVersion.DSE_V1) + assert len(calls) == 1, "compressor should have been invoked for DSE_V1" + assert flags_byte & COMPRESSED_FLAG + + def test_encode_message_does_not_compress_at_v5(self): + """ + V5 has checksumming support, so encode_message must NOT apply + message-level compression (that happens at the segment level + instead): the compressor must not run and COMPRESSED_FLAG must not + be set. + """ + calls, flags_byte = self._encode_with_spy_compressor(ProtocolVersion.V5) + assert len(calls) == 0, "compressor should not have been invoked for V5" + assert not (flags_byte & COMPRESSED_FLAG) + + def test_encode_message_compresses_at_v4(self): + # Sanity baseline below the checksumming range. + calls, flags_byte = self._encode_with_spy_compressor(ProtocolVersion.V4) + assert len(calls) == 1 + assert flags_byte & COMPRESSED_FLAG + + def _decode_with_spy_decompressor(self, protocol_version): + calls = [] + + def spy_decompressor(body): + calls.append(body) + return body + + ProtocolHandler.decode_message( + protocol_version=protocol_version, protocol_features=ProtocolFeatures(), + user_type_map={}, stream_id=0, flags=COMPRESSED_FLAG, opcode=ReadyMessage.opcode, + body=b'ignored-by-readymessage', decompressor=spy_decompressor, result_metadata=None) + return calls + + def test_decode_message_decompresses_at_dse_v1(self): + calls = self._decode_with_spy_decompressor(ProtocolVersion.DSE_V1) + assert len(calls) == 1, "decompressor should have been invoked for DSE_V1" + + def test_decode_message_does_not_decompress_at_v5(self): + calls = self._decode_with_spy_decompressor(ProtocolVersion.V5) + assert len(calls) == 0, "decompressor should not have been invoked for V5" + + +class HeaderConsistencyTest(unittest.TestCase): + """ + encode_message has two branches: the compression branch (which avoids a + second BytesIO allocation by building the header directly as bytes) and + the non-compression branch (which writes into a BytesIO via + _write_header). Both must produce byte-identical headers for the same + (version, flags, stream_id, opcode, length) -- if the two ever diverged + (e.g. because _write_header grew version-specific behavior that the + compression branch's header-building didn't replicate), messages would + be silently corrupted for whichever protocol version triggered the + difference. + """ + + def test_both_branches_use_the_shared_header_packer(self): + """ + The compression branch must not rebuild the header inline; it must + go through the same _pack_header single source of truth that + _write_header uses. + """ + with patch.object(_ProtocolHandler, '_pack_header', + wraps=_ProtocolHandler._pack_header) as spy: + msg = RegisterMessage(["TOPOLOGY_CHANGE"]) + + # Compression branch (checksumming not active at V4). + ProtocolHandler.encode_message( + msg, stream_id=1, protocol_version=ProtocolVersion.V4, compressor=lambda b: b, + allow_beta_protocol_version=False, protocol_features=ProtocolFeatures()) + assert spy.call_count == 1, ( + "compression branch must build its header via _pack_header, " + "not by duplicating v3_header_pack(...) + int32_pack(...) inline" + ) + + # Non-compression branch, via _write_header. + ProtocolHandler.encode_message( + msg, stream_id=1, protocol_version=ProtocolVersion.V4, compressor=None, + allow_beta_protocol_version=False, protocol_features=ProtocolFeatures()) + assert spy.call_count == 2, "non-compression branch must also go through _pack_header" + + def test_compressed_and_uncompressed_headers_are_identical_modulo_compressed_flag(self): + """ + With an identity compressor (output same length as input), the + compression and non-compression branches encode the same message at + the same version with only the COMPRESSED_FLAG bit differing in the + header -- version, stream_id, opcode and length must match exactly. + """ + msg = RegisterMessage(["TOPOLOGY_CHANGE", "STATUS_CHANGE"]) + + compressed_frame = ProtocolHandler.encode_message( + msg, stream_id=42, protocol_version=ProtocolVersion.V4, compressor=lambda b: b, + allow_beta_protocol_version=False, protocol_features=ProtocolFeatures()) + plain_frame = ProtocolHandler.encode_message( + msg, stream_id=42, protocol_version=ProtocolVersion.V4, compressor=None, + allow_beta_protocol_version=False, protocol_features=ProtocolFeatures()) + + # 9-byte v3 frame header: version, flags, stream_id (2 bytes), opcode, length (4 bytes) + compressed_header = bytearray(compressed_frame[:9]) + plain_header = bytearray(plain_frame[:9]) + + assert compressed_header[1] & COMPRESSED_FLAG + assert not (plain_header[1] & COMPRESSED_FLAG) + + # Mask out the COMPRESSED_FLAG bit and compare the rest byte-for-byte. + compressed_header[1] &= ~COMPRESSED_FLAG + plain_header[1] &= ~COMPRESSED_FLAG + assert bytes(compressed_header) == bytes(plain_header)