diff --git a/benchmarks/micro/bench_bind_no_encryption.py b/benchmarks/micro/bench_bind_no_encryption.py new file mode 100644 index 0000000000..6e3e23a677 --- /dev/null +++ b/benchmarks/micro/bench_bind_no_encryption.py @@ -0,0 +1,91 @@ +# 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: BoundStatement.bind() fast path without column encryption. + +Measures the improvement from skipping ColDesc namedtuple creation and +ce_policy checks when column_encryption_policy is None (the common case). + +Run: + python benchmarks/micro/bench_bind_no_encryption.py [iterations] +""" + +import datetime +import sys +import timeit + +from cassandra.protocol import ColumnMetadata +from cassandra.query import BoundStatement, PreparedStatement +from cassandra.cqltypes import ( + DateType, Int32Type, DoubleType, FloatType, UTF8Type, + BooleanType, LongType, +) + + +def make_prepared_statement(col_names, col_types): + """Build a real PreparedStatement (so _serializers is set up as in production).""" + col_meta = [ColumnMetadata('ks', 'metrics', name, ctype) + for name, ctype in zip(col_names, col_types)] + return PreparedStatement( + column_metadata=col_meta, query_id=None, routing_key_indexes=[], + query=None, keyspace='ks', protocol_version=4, + result_metadata=None, result_metadata_id=None) + + +def bench(): + schemas = [ + ( + "3-col (int, double, text)", + ['id', 'value', 'tag'], + [Int32Type, DoubleType, UTF8Type], + [42, 3.14159, 'sensor-001'], + ), + ( + "5-col time-series", + ['ts', 'sensor_id', 'value', 'quality', 'tag'], + [DateType, Int32Type, DoubleType, FloatType, UTF8Type], + [datetime.datetime(2025, 4, 5, 12, 0, 0, 123456), 42, 3.14, 0.95, 'alpha'], + ), + ( + "8-col wide row", + ['ts', 'id', 'v1', 'v2', 'v3', 'v4', 'flag', 'name'], + [DateType, LongType, DoubleType, DoubleType, FloatType, FloatType, BooleanType, UTF8Type], + [datetime.datetime(2025, 1, 1), 12345678, 1.1, 2.2, 3.3, 4.4, True, 'test-row'], + ), + ] + + n = int(sys.argv[1]) if len(sys.argv) > 1 else 200_000 + print(f"=== BoundStatement.bind() no-encryption fast path ({n:,} iters) ===\n") + + for label, col_names, col_types, row in schemas: + ps = make_prepared_statement(col_names, col_types) + + def do_bind(): + bs = BoundStatement(ps) + bs.bind(row) + + # Warmup + for _ in range(min(1000, n)): + do_bind() + + t = timeit.timeit(do_bind, number=n) + ns_per = t / n * 1e9 + print(f" {label}:") + print(f" {ns_per:.1f} ns/call ({n:,} iters)") + + +if __name__ == "__main__": + print(f"Python {sys.version}\n") + bench() diff --git a/benchmarks/micro/bench_timeseries.py b/benchmarks/micro/bench_timeseries.py new file mode 100644 index 0000000000..db5603e69f --- /dev/null +++ b/benchmarks/micro/bench_timeseries.py @@ -0,0 +1,184 @@ +#!/usr/bin/env python3 +""" +Microbenchmarks for time-series write and read hot paths. + +Covers: + - DateType.serialize / deserialize + - varint_pack / varint_unpack + - BoundStatement.bind() for a typical time-series schema + +All results in nanoseconds per call. Run with: + python benchmarks/micro/bench_timeseries.py +""" + +import datetime +import struct +import sys +import timeit +import uuid + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +WARMUP = 50_000 +ITERATIONS = 500_000 + + +def bench(label, stmt, setup="pass", number=ITERATIONS, warmup=WARMUP): + """Run *stmt* under *setup*, return ns/call and print a line.""" + globs = {} + exec(setup, globs) + # warmup + t_code = compile(stmt, "", "exec") + for _ in range(warmup): + exec(t_code, globs) + # measure + timer = timeit.Timer(stmt, setup, globals=globs) + raw = timer.timeit(number=number) + ns = raw / number * 1e9 + print(f" {label:.<60s} {ns:>9.1f} ns/call") + return ns + + +# --------------------------------------------------------------------------- +# DateType.serialize / deserialize +# --------------------------------------------------------------------------- + + +def bench_datetype(): + print("\n=== DateType.serialize ===") + setup = """\ +from cassandra.cqltypes import DateType +import datetime +dt_now = datetime.datetime(2025, 4, 5, 12, 0, 0, 123456) +dt_epoch = datetime.datetime(1970, 1, 1, 0, 0, 1, 0) +dt_far = datetime.datetime(2300, 1, 1, 0, 0, 0, 1000) +d_only = datetime.date(2025, 4, 5) +ts_int = 1712318400000 +""" + bench("serialize datetime (2025)", "DateType.serialize(dt_now, 4)", setup) + bench("serialize datetime (epoch)", "DateType.serialize(dt_epoch, 4)", setup) + bench("serialize datetime (2300)", "DateType.serialize(dt_far, 4)", setup) + bench("serialize date object", "DateType.serialize(d_only, 4)", setup) + bench("serialize raw int timestamp", "DateType.serialize(ts_int, 4)", setup) + + print("\n=== DateType.deserialize ===") + setup_deser = ( + setup + + """\ +packed_now = DateType.serialize(dt_now, 4) +packed_far = DateType.serialize(dt_far, 4) +""" + ) + bench("deserialize (2025)", "DateType.deserialize(packed_now, 4)", setup_deser) + bench("deserialize (2300)", "DateType.deserialize(packed_far, 4)", setup_deser) + + +# --------------------------------------------------------------------------- +# varint_pack / varint_unpack +# --------------------------------------------------------------------------- + + +def bench_varint(): + print("\n=== varint_pack ===") + setup = """\ +from cassandra.marshal import varint_pack, varint_unpack +small = 42 +medium = 2**62 +large = 2**127 +negative = -(2**62) +zero = 0 +""" + bench("varint_pack zero", "varint_pack(zero)", setup) + bench("varint_pack small", "varint_pack(small)", setup) + bench("varint_pack medium", "varint_pack(medium)", setup) + bench("varint_pack large", "varint_pack(large)", setup) + bench("varint_pack negative", "varint_pack(negative)", setup) + + print("\n=== varint_unpack ===") + setup_u = ( + setup + + """\ +packed_small = varint_pack(small) +packed_medium = varint_pack(medium) +packed_large = varint_pack(large) +packed_negative = varint_pack(negative) +packed_zero = varint_pack(zero) +""" + ) + bench("varint_unpack zero", "varint_unpack(packed_zero)", setup_u) + bench("varint_unpack small", "varint_unpack(packed_small)", setup_u) + bench("varint_unpack medium", "varint_unpack(packed_medium)", setup_u) + bench("varint_unpack large", "varint_unpack(packed_large)", setup_u) + bench("varint_unpack negative", "varint_unpack(packed_negative)", setup_u) + + +# --------------------------------------------------------------------------- +# BoundStatement.bind() — typical time-series schema +# --------------------------------------------------------------------------- + + +def bench_bind(): + print("\n=== BoundStatement.bind (time-series schema) ===") + setup = """\ +import datetime +from cassandra.query import BoundStatement, PreparedStatement +from cassandra.cqltypes import ( + DateType, Int32Type, DoubleType, FloatType, UTF8Type, +) +from cassandra.protocol import ProtocolVersion +from unittest.mock import MagicMock + +# Build a mock PreparedStatement with 5 columns: +# (ts timestamp, sensor_id int, value double, quality float, tag text) +col_types = [DateType, Int32Type, DoubleType, FloatType, UTF8Type] +col_names = ['ts', 'sensor_id', 'value', 'quality', 'tag'] + +col_meta = [] +for name, ctype in zip(col_names, col_types): + cm = MagicMock() + cm.name = name + cm.keyspace_name = 'ks' + cm.table_name = 'metrics' + cm.type = ctype + col_meta.append(cm) + +ps = MagicMock(spec=PreparedStatement) +ps.column_metadata = col_meta +ps.routing_key_indexes = None +ps.protocol_version = 4 +ps.column_encryption_policy = None +ps.serial_consistency_level = None +ps.retry_policy = None +ps.consistency_level = None +ps.fetch_size = None +ps.custom_payload = None +ps.is_idempotent = False + +dt = datetime.datetime(2025, 4, 5, 12, 0, 0, 123456) +row = [dt, 42, 3.14159, 0.95, 'sensor-alpha-001'] +""" + bench( + "bind 5-col time-series row", + """\ +bs = BoundStatement(ps) +bs.bind(row) +""", + setup, + ) + + +# --------------------------------------------------------------------------- +# Main +# --------------------------------------------------------------------------- + +if __name__ == "__main__": + print(f"Python {sys.version}") + print(f"Iterations per benchmark: {ITERATIONS:,}") + + bench_datetype() + bench_varint() + bench_bind() + + print("\nDone.") diff --git a/cassandra/cqltypes.py b/cassandra/cqltypes.py index 4d63ae5195..2fdac708df 100644 --- a/cassandra/cqltypes.py +++ b/cassandra/cqltypes.py @@ -31,6 +31,7 @@ from binascii import unhexlify import calendar from collections import namedtuple +import datetime as _datetime_mod from decimal import Decimal import io from itertools import chain @@ -62,6 +63,9 @@ _number_types = frozenset((int, float)) +_EPOCH_NAIVE = _datetime_mod.datetime(1970, 1, 1) +_EPOCH_DATE = _datetime_mod.date(1970, 1, 1) + def _name_from_hex_string(encoded_name): bin_str = unhexlify(encoded_name) @@ -649,13 +653,20 @@ def deserialize(byts, protocol_version): @staticmethod def serialize(v, protocol_version): try: - # v is datetime - timestamp_seconds = calendar.timegm(v.utctimetuple()) - timestamp = timestamp_seconds * 1000 + getattr(v, 'microsecond', 0) // 1000 + # v is a datetime; use integer arithmetic instead of + # calendar.timegm(v.utctimetuple()) to avoid allocating + # an intermediate struct_time object on every call. + utcoffset = v.utcoffset() + if utcoffset is not None: + v = v - utcoffset + v = v.replace(tzinfo=None) + td = v - _EPOCH_NAIVE + timestamp = (td.days * 86400 + td.seconds) * 1000 + td.microseconds // 1000 except AttributeError: try: - timestamp = calendar.timegm(v.timetuple()) * 1000 - except AttributeError: + td = v - _EPOCH_DATE + timestamp = td.days * 86400000 + except (AttributeError, TypeError): # Ints and floats are valid timestamps too if type(v) not in _number_types: raise TypeError('DateType arguments must be a datetime, date, or timestamp') diff --git a/cassandra/cython_marshal.pyx b/cassandra/cython_marshal.pyx index 0a926b6eef..e0992d6841 100644 --- a/cassandra/cython_marshal.pyx +++ b/cassandra/cython_marshal.pyx @@ -55,16 +55,7 @@ cdef varint_unpack(Buffer *term): """Unpack a variable-sized integer""" return varint_unpack_py3(to_bytes(term)) -# TODO: Optimize these two functions cdef varint_unpack_py3(bytes term): - val = int(''.join(["%02x" % i for i in term]), 16) - if (term[0] & 128) != 0: - shift = len(term) * 8 # * Note below - val -= 1 << shift - return val - -# * Note * -# '1 << (len(term) * 8)' Cython tries to do native -# integer shifts, which overflows. We need this to -# emulate Python shifting, which will expand the long -# to accommodate + if not term: + raise ValueError("Cannot unpack an empty varint") + return int.from_bytes(term, byteorder='big', signed=True) diff --git a/cassandra/encoder.py b/cassandra/encoder.py index b33be935df..d1422ca9cb 100644 --- a/cassandra/encoder.py +++ b/cassandra/encoder.py @@ -22,7 +22,6 @@ from binascii import hexlify from decimal import Decimal -import calendar import datetime import math import types @@ -32,6 +31,8 @@ from cassandra.util import (OrderedDict, OrderedMap, OrderedMapSerializedKey, sortedset, Time, Date, Point, LineString, Polygon) +_EPOCH_NAIVE = datetime.datetime(1970, 1, 1) + def cql_quote(term): if isinstance(term, str): @@ -140,8 +141,12 @@ def cql_encode_datetime(self, val): Converts a :class:`datetime.datetime` object to a (string) integer timestamp with millisecond precision. """ - timestamp = calendar.timegm(val.utctimetuple()) - return str(timestamp * 1000 + getattr(val, 'microsecond', 0) // 1000) + utcoffset = val.utcoffset() + if utcoffset is not None: + val = val - utcoffset + val = val.replace(tzinfo=None) + td = val - _EPOCH_NAIVE + return str((td.days * 86400 + td.seconds) * 1000 + td.microseconds // 1000) def cql_encode_date(self, val): """ diff --git a/cassandra/marshal.py b/cassandra/marshal.py index 413e1831d4..ccf84395d5 100644 --- a/cassandra/marshal.py +++ b/cassandra/marshal.py @@ -40,11 +40,9 @@ def _make_packer(format_string): def varint_unpack(term): - val = int(''.join("%02x" % i for i in term), 16) - if (term[0] & 128) != 0: - len_term = len(term) # pulling this out of the expression to avoid overflow in cython optimized code - val -= 1 << (len_term * 8) - return val + if not term: + raise ValueError('Cannot unpack an empty varint') + return int.from_bytes(term, byteorder='big', signed=True) def bit_length(n): @@ -52,21 +50,13 @@ def bit_length(n): def varint_pack(big): - pos = True if big == 0: return b'\x00' if big < 0: - bytelength = bit_length(abs(big) - 1) // 8 + 1 - big = (1 << bytelength * 8) + big - pos = False - revbytes = bytearray() - while big > 0: - revbytes.append(big & 0xff) - big >>= 8 - if pos and revbytes[-1] & 0x80: - revbytes.append(0) - revbytes.reverse() - return bytes(revbytes) + byte_length = (-big - 1).bit_length() // 8 + 1 + else: + byte_length = (big.bit_length() + 8) // 8 + return big.to_bytes(byte_length, byteorder='big', signed=True) point_be = struct.Struct('>dd') diff --git a/cassandra/query.py b/cassandra/query.py index 39b9fdb0ad..ab794beb71 100644 --- a/cassandra/query.py +++ b/cassandra/query.py @@ -684,28 +684,47 @@ def bind(self, values): self.raw_values = values self.values = [] - for value, col_spec in zip(values, col_meta): - if value is None: - self.values.append(None) - elif value is UNSET_VALUE: - if proto_version >= 4: - self._append_unset_value() + if ce_policy: + for value, col_spec in zip(values, col_meta): + if value is None: + self.values.append(None) + elif value is UNSET_VALUE: + if proto_version >= 4: + self._append_unset_value() + else: + raise ValueError("Attempt to bind UNSET_VALUE while using unsuitable protocol version (%d < 4)" % proto_version) else: - raise ValueError("Attempt to bind UNSET_VALUE while using unsuitable protocol version (%d < 4)" % proto_version) - else: - try: - col_desc = ColDesc(col_spec.keyspace_name, col_spec.table_name, col_spec.name) - uses_ce = ce_policy and ce_policy.contains_column(col_desc) - col_type = ce_policy.column_type(col_desc) if uses_ce else col_spec.type - col_bytes = col_type.serialize(value, proto_version) - if uses_ce: - col_bytes = ce_policy.encrypt(col_desc, col_bytes) - self.values.append(col_bytes) - except (TypeError, struct.error) as exc: - actual_type = type(value) - message = ('Received an argument of invalid type for column "%s". ' - 'Expected: %s, Got: %s; (%s)' % (col_spec.name, col_spec.type, actual_type, exc)) - raise TypeError(message) + try: + col_desc = ColDesc(col_spec.keyspace_name, col_spec.table_name, col_spec.name) + uses_ce = ce_policy.contains_column(col_desc) + col_type = ce_policy.column_type(col_desc) if uses_ce else col_spec.type + col_bytes = col_type.serialize(value, proto_version) + if uses_ce: + col_bytes = ce_policy.encrypt(col_desc, col_bytes) + self.values.append(col_bytes) + except (TypeError, struct.error) as exc: + actual_type = type(value) + message = ('Received an argument of invalid type for column "%s". ' + 'Expected: %s, Got: %s; (%s)' % (col_spec.name, col_spec.type, actual_type, exc)) + raise TypeError(message) + else: + # no column encryption: skip the per-column ColDesc allocation + for value, col_spec in zip(values, col_meta): + if value is None: + self.values.append(None) + elif value is UNSET_VALUE: + if proto_version >= 4: + self._append_unset_value() + else: + raise ValueError("Attempt to bind UNSET_VALUE while using unsuitable protocol version (%d < 4)" % proto_version) + else: + try: + self.values.append(col_spec.type.serialize(value, proto_version)) + except (TypeError, struct.error) as exc: + actual_type = type(value) + message = ('Received an argument of invalid type for column "%s". ' + 'Expected: %s, Got: %s; (%s)' % (col_spec.name, col_spec.type, actual_type, exc)) + raise TypeError(message) if proto_version >= 4: diff = col_meta_len - len(self.values) diff --git a/tests/unit/cython/test_types.py b/tests/unit/cython/test_types.py index 996be266c0..dc5b5ab38e 100644 --- a/tests/unit/cython/test_types.py +++ b/tests/unit/cython/test_types.py @@ -27,3 +27,7 @@ def test_datetype(self): @cythontest def test_date_side_by_side(self): types_testhelper.test_date_side_by_side() + + @cythontest + def test_decimal_empty_varint_raises(self): + types_testhelper.test_decimal_empty_varint_raises() diff --git a/tests/unit/cython/types_testhelper.pyx b/tests/unit/cython/types_testhelper.pyx index 81f9dca114..a7a6fe914a 100644 --- a/tests/unit/cython/types_testhelper.pyx +++ b/tests/unit/cython/types_testhelper.pyx @@ -106,3 +106,17 @@ def test_date_side_by_side(): for ms in range(1000): verify_time(x - ms) x //= 2 + + +def test_decimal_empty_varint_raises(): + # scale-only payload has an empty unscaled varint: must raise, not decode to 0 + from cassandra.cqltypes import DecimalType + cdef Deserializer des = find_deserializer(DecimalType) + cdef BytesIOReader reader = BytesIOReader(b'\x00\x00\x00\x04\x00\x00\x00\x05') + cdef Buffer buf + get_buf(reader, &buf) + try: + from_binary(des, &buf, 0) + except ValueError: + return + raise AssertionError("expected ValueError") diff --git a/tests/unit/test_marshalling.py b/tests/unit/test_marshalling.py index 02ca901abc..9485c14039 100644 --- a/tests/unit/test_marshalling.py +++ b/tests/unit/test_marshalling.py @@ -22,6 +22,7 @@ from uuid import UUID from cassandra.cqltypes import lookup_casstype, DecimalType, UTF8Type, DateType +from cassandra.marshal import varint_unpack from cassandra.util import OrderedMapSerializedKey, sortedset, Time, Date marshalled_value_pairs = ( @@ -134,3 +135,11 @@ def test_decimal(self): for n in converted_types: expected = Decimal(n) assert DecimalType.from_binary(DecimalType.to_binary(n, proto_ver), proto_ver) == expected + + def test_varint_unpack_empty_raises(self): + # empty varint is malformed wire data; must fail loudly, not decode to 0 + self.assertRaises(ValueError, varint_unpack, b'') + + def test_decimal_truncated_payload_raises(self): + # 4-byte payload (scale only, no unscaled varint) is a corrupt decimal + self.assertRaises(ValueError, DecimalType.deserialize, b'\x00\x00\x00\x05', 3)