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
29 changes: 27 additions & 2 deletions cassandra/cluster.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,8 @@
BatchMessage, RESULT_KIND_PREPARED,
RESULT_KIND_SET_KEYSPACE, RESULT_KIND_ROWS,
RESULT_KIND_SCHEMA_CHANGE, ProtocolHandler,
RESULT_KIND_VOID, ProtocolException)
RESULT_KIND_VOID, ProtocolException,
StartupMessage)
from cassandra.metadata import Metadata, protect_name, murmur3, _NodeInfo
from cassandra.policies import (TokenAwarePolicy, DCAwareRoundRobinPolicy, SimpleConvictionPolicy,
ExponentialReconnectionPolicy, HostDistance,
Expand Down Expand Up @@ -972,6 +973,21 @@ def default_retry_policy(self, policy):
used for columns in this cluster.
"""

extra_startup_options: dict[str, str] | None = None
"""
A dict of extra options sent in the STARTUP message when a connection is established.

This is useful for sending custom startup options that are not supported by the driver.
For example, per `CASSANDRA-16378 <https://issues.apache.org/jira/browse/CASSANDRA-16378>`_,
custom application level options are exposed in client metrics::

Cluster(extra_startup_options={'APPLICATION_NAME': 'my-app'})

Options managed by the driver, such as ``CQL_VERSION``, ``COMPRESSION``, ``NO_COMPACT``,
``DRIVER_NAME`` and ``DRIVER_VERSION``, always take precedence. Extra options that collide
with them are ignored.
"""

@property
def schema_metadata_enabled(self):
"""
Expand Down Expand Up @@ -1069,7 +1085,8 @@ def __init__(self,
ssl_context=None,
endpoint_factory=None,
cloud=None,
column_encryption_policy=None):
column_encryption_policy=None,
extra_startup_options=None):
"""
``executor_threads`` defines the number of threads in a pool for handling asynchronous tasks such as
establishing connection pools or refreshing metadata.
Expand Down Expand Up @@ -1266,6 +1283,13 @@ def __init__(self,
self.connect_timeout = connect_timeout
self.prepare_on_all_hosts = prepare_on_all_hosts
self.reprepare_on_up = reprepare_on_up
self.extra_startup_options = dict(extra_startup_options) if extra_startup_options else {}
driver_managed_options = self.extra_startup_options.keys() & StartupMessage.KNOWN_OPTION_KEYS
if driver_managed_options:
log.warning("Ignoring extra startup option(s) %s: they are managed by the driver "
"and cannot be overridden", ', '.join(sorted(driver_managed_options)))
for key in driver_managed_options:
del self.extra_startup_options[key]

self._listeners = set()
self._listener_lock = Lock()
Expand Down Expand Up @@ -1556,6 +1580,7 @@ def _make_connection_kwargs(self, endpoint, kwargs_dict):
kwargs_dict.setdefault('user_type_map', self._user_types)
kwargs_dict.setdefault('allow_beta_protocol_version', self.allow_beta_protocol_version)
kwargs_dict.setdefault('no_compact', self.no_compact)
kwargs_dict.setdefault('extra_startup_options', self.extra_startup_options)

return kwargs_dict

Expand Down
6 changes: 4 additions & 2 deletions cassandra/connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -675,6 +675,7 @@ class Connection(object):

cql_version = None
no_compact = False
extra_startup_options = None
protocol_version = ProtocolVersion.MAX_SUPPORTED

keyspace = None
Expand Down Expand Up @@ -758,7 +759,7 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None,
ssl_options=None, sockopts=None, compression=True,
cql_version=None, protocol_version=ProtocolVersion.MAX_SUPPORTED, is_control_connection=False,
user_type_map=None, connect_timeout=None, allow_beta_protocol_version=False, no_compact=False,
ssl_context=None, on_orphaned_stream_released=None):
ssl_context=None, on_orphaned_stream_released=None, extra_startup_options=None):

# TODO next major rename host to endpoint and remove port kwarg.
self.endpoint = host if isinstance(host, EndPoint) else DefaultEndPoint(host, port)
Expand All @@ -782,6 +783,7 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None,
self._socket_writable = True
self.orphaned_request_ids = set()
self._on_orphaned_stream_released = on_orphaned_stream_released
self.extra_startup_options = extra_startup_options or {}

if ssl_options:
self.ssl_options.update(self.endpoint.ssl_options or {})
Expand Down Expand Up @@ -1403,7 +1405,7 @@ def _send_startup_message(self, compression=None, no_compact=False):
opts['COMPRESSION'] = compression
if no_compact:
opts['NO_COMPACT'] = 'true'
sm = StartupMessage(cqlversion=self.cql_version, options=opts)
sm = StartupMessage(cqlversion=self.cql_version, options=opts, extra_options=self.extra_startup_options)
self.send_msg(sm, self.get_request_id(), cb=self._handle_startup_response)
log.debug("Sent StartupMessage on %s", self)

Expand Down
10 changes: 7 additions & 3 deletions cassandra/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -400,15 +400,19 @@ class StartupMessage(_MessageType):
KNOWN_OPTION_KEYS = set((
'CQL_VERSION',
'COMPRESSION',
'NO_COMPACT'
'NO_COMPACT',
'DRIVER_NAME',
'DRIVER_VERSION'
))

def __init__(self, cqlversion, options):
def __init__(self, cqlversion, options, extra_options=None):

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nitpick: not sure if there is much use for distinguishing between options and extra_options in the StartupMessage class. Maybe they can be merged prior.

self.cqlversion = cqlversion
self.options = options
self.extra_options = extra_options

def send_body(self, f, protocol_version):
optmap = self.options.copy()
optmap: dict[str, str] = self.options.copy()
optmap.update(self.extra_options or {})
optmap['CQL_VERSION'] = self.cqlversion
write_stringmap(f, optmap)

Expand Down
2 changes: 2 additions & 0 deletions docs/api/cassandra/cluster.rst
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,8 @@

.. autoattribute:: cloud

.. autoattribute:: extra_startup_options

.. automethod:: connect

.. automethod:: shutdown
Expand Down
59 changes: 58 additions & 1 deletion tests/integration/standard/test_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,16 +23,19 @@
import threading
from threading import Thread, Event
import time
import uuid
from unittest import SkipTest

from cassandra import ConsistencyLevel, OperationTimedOut, DependencyException
from cassandra.cluster import NoHostAvailable, ConnectionShutdown, ExecutionProfile, EXEC_PROFILE_DEFAULT
from cassandra.connection import DRIVER_NAME, DRIVER_VERSION
from cassandra.protocol import QueryMessage
from cassandra.policies import HostFilterPolicy, RoundRobinPolicy, HostStateListener
from cassandra.pool import HostConnectionPool

from tests.integration import use_singledc, get_node, CASSANDRA_IP, local, \
requiresmallclockgranularity, greaterthancass20, TestCluster
requiresmallclockgranularity, greaterthancass20, greaterthanorequalcass41, \
requirecassandra, TestCluster

try:
import cassandra.io.asyncorereactor
Expand Down Expand Up @@ -465,3 +468,57 @@ def setUp(self):
def clean_global_loop(self):
cassandra.io.libevreactor._global_loop._cleanup()
cassandra.io.libevreactor._global_loop = None


@requirecassandra
@greaterthanorequalcass41
class ExtraStartupOptionsTest(unittest.TestCase):
"""
Ensures the extra startup options configured on a Cluster reach the server.

The options a client sent in its STARTUP message are exposed by the
``system_views.clients`` virtual table since Cassandra 4.1.
"""

def connect_and_get_client_options(self, **cluster_kwargs):
"""
Connects a cluster tagged with a unique APPLICATION_NAME and returns the
options the server recorded for it.
"""
application_name = f'app-{uuid.uuid4()}'
extra_startup_options = dict(cluster_kwargs.pop('extra_startup_options', {}),
APPLICATION_NAME=application_name)

cluster = TestCluster(extra_startup_options=extra_startup_options, **cluster_kwargs)
session = cluster.connect(wait_for_all_pools=True)
self.addCleanup(cluster.shutdown)

rows = session.execute("SELECT client_options FROM system_views.clients")
options = [row.client_options for row in rows
if row.client_options
and row.client_options.get('APPLICATION_NAME') == application_name]

self.assertGreater(len(options, 0))
return application_name, options

def test_extra_startup_options_are_sent_to_server(self):
application_name, options = self.connect_and_get_client_options(
extra_startup_options={'APPLICATION_VERSION': '1.2.3'})

for client_options in options:
self.assertEqual(client_options['APPLICATION_NAME'], application_name)
self.assertEqual(client_options['APPLICATION_VERSION'], '1.2.3')
self.assertEqual(client_options['DRIVER_NAME'], DRIVER_NAME)
self.assertEqual(client_options['DRIVER_VERSION'], DRIVER_VERSION)
self.assertIn('CQL_VERSION', client_options)

def test_extra_startup_options_do_not_override_driver_options(self):
_, options = self.connect_and_get_client_options(
extra_startup_options={'DRIVER_NAME': 'not the driver',
'DRIVER_VERSION': '0.0.0',
'CQL_VERSION': '2.0.0'})

for client_options in options:
self.assertEqual(client_options['DRIVER_NAME'], DRIVER_NAME)
self.assertEqual(client_options['DRIVER_VERSION'], DRIVER_VERSION)
self.assertNotEqual(client_options['CQL_VERSION'], '2.0.0')
43 changes: 42 additions & 1 deletion tests/unit/test_cluster.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,8 @@
InvalidRequest, Unauthorized, AuthenticationFailed, OperationTimedOut, UnsupportedOperation, RequestValidationException, ConfigurationException, ProtocolVersion
from cassandra.cluster import _Scheduler, Session, Cluster, default_lbp_factory, \
ExecutionProfile, _ConfigMode, EXEC_PROFILE_DEFAULT
from cassandra.connection import SniEndPoint, SniEndPointFactory
from cassandra.connection import SniEndPoint, SniEndPointFactory, DefaultEndPoint
from cassandra.protocol import StartupMessage
from cassandra.pool import Host
from cassandra.policies import HostDistance, RetryPolicy, RoundRobinPolicy, DowngradingConsistencyRetryPolicy, SimpleConvictionPolicy
from cassandra.query import SimpleStatement, named_tuple_factory, tuple_factory
Expand Down Expand Up @@ -150,6 +151,46 @@ def _mocked_proxy_dns_resolution(self):
# single SNI endpoint should be resolved to multiple unique IP addresses
self.assertEqual(len(addrs), len(set(addrs)))

def test_extra_startup_options_passed_to_connections(self):
"""
Ensures the extra startup options configured on a cluster reach the
connections it creates.
"""
extra_startup_options = {'EXTRA_OPTION': 'option-1'}
cluster = Cluster(extra_startup_options=extra_startup_options)

kwargs = cluster._make_connection_kwargs(DefaultEndPoint('127.0.0.1'), {})
self.assertEqual(kwargs['extra_startup_options'], extra_startup_options)

def test_extra_startup_options_default_to_empty(self):
kwargs = Cluster()._make_connection_kwargs(DefaultEndPoint('127.0.0.1'), {})
self.assertEqual(kwargs['extra_startup_options'], {})

def test_driver_managed_extra_startup_options_are_removed(self):
"""
Ensures the options managed by the driver are dropped from the extra
startup options, so they cannot be overridden.
"""
extra_startup_options = dict.fromkeys(StartupMessage.KNOWN_OPTION_KEYS, 'overridden')
extra_startup_options['EXTRA_OPTION'] = 'option-1'

with self.assertLogs('cassandra.cluster', level='WARNING') as logs:
cluster = Cluster(extra_startup_options=extra_startup_options)

self.assertEqual(cluster.extra_startup_options, {'EXTRA_OPTION': 'option-1'})

warnings = [line for line in logs.output if 'Ignoring extra startup option' in line]
self.assertEqual(len(warnings), 1)
for key in StartupMessage.KNOWN_OPTION_KEYS:
self.assertIn(key, warnings[0])
self.assertNotIn('EXTRA_OPTION', warnings[0])

def test_extra_startup_options_are_not_warned_about(self):
with self.assertNoLogs('cassandra.cluster', level='WARNING'):
cluster = Cluster(extra_startup_options={'EXTRA_OPTION': 'option-1'})

self.assertEqual(cluster.extra_startup_options, {'EXTRA_OPTION': 'option-1'})


class SchedulerTest(unittest.TestCase):
# TODO: this suite could be expanded; for now just adding a test covering a ticket
Expand Down
55 changes: 51 additions & 4 deletions tests/unit/test_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,18 +23,19 @@
from cassandra.cluster import Cluster
from cassandra.connection import (Connection, HEADER_DIRECTION_TO_CLIENT, ProtocolError,
locally_supported_compressions, ConnectionHeartbeat, _Frame, Timer, TimerManager,
ConnectionException, DefaultEndPoint)
ConnectionException, DefaultEndPoint, DRIVER_NAME, DRIVER_VERSION)
from cassandra.marshal import uint8_pack, uint32_pack, int32_pack
from cassandra.protocol import (write_stringmultimap, write_int, write_string,
SupportedMessage, ProtocolHandler)
SupportedMessage, ProtocolHandler, StartupMessage,
read_stringmap)

from tests.util import wait_until


class ConnectionTest(unittest.TestCase):

def make_connection(self):
c = Connection(DefaultEndPoint('1.2.3.4'))
def make_connection(self, **kwargs):
c = Connection(DefaultEndPoint('1.2.3.4'), **kwargs)
c._socket = Mock()
c._socket.send.side_effect = lambda x: len(x)
return c
Expand Down Expand Up @@ -245,6 +246,52 @@ def test_disable_compression(self, *args):

self.assertEqual(c.decompressor, None)

def send_startup_and_get_options(self, extra_startup_options, **startup_kwargs):
"""
Sends a StartupMessage on a connection configured with the given extra
options, and reads back the option map it would put on the wire.
"""
c = self.make_connection(cql_version='3.4.5', extra_startup_options=extra_startup_options)
c.send_msg = Mock()

c._send_startup_message(**startup_kwargs)

self.assertEqual(c.send_msg.call_count, 1)
message = c.send_msg.call_args[0][0]
self.assertIsInstance(message, StartupMessage)

buf = BytesIO()
message.send_body(buf, c.protocol_version)
buf.seek(0)
return read_stringmap(buf)

def test_extra_startup_options_are_sent(self):
"""
Ensures the extra options configured on a connection are sent in the
STARTUP message, alongside the options managed by the driver.
"""
options = self.send_startup_and_get_options({'EXTRA_OPTION_1': 'option-1', 'EXTRA_OPTION_2': 'option-2'},
compression='lz4', no_compact=True)

self.assertEqual(options, {
'CQL_VERSION': '3.4.5',
'DRIVER_NAME': DRIVER_NAME,
'DRIVER_VERSION': DRIVER_VERSION,
'COMPRESSION': 'lz4',
'NO_COMPACT': 'true',
'EXTRA_OPTION_1': 'option-1',
'EXTRA_OPTION_2': 'option-2'
})

def test_no_extra_startup_options(self):
"""
Ensures a connection without extra options still sends a valid STARTUP message.
"""
expected = {'CQL_VERSION': '3.4.5', 'DRIVER_NAME': DRIVER_NAME, 'DRIVER_VERSION': DRIVER_VERSION}
for extra_startup_options in (None, {}):
with self.subTest(extra_startup_options=extra_startup_options):
self.assertEqual(self.send_startup_and_get_options(extra_startup_options), expected)

def test_not_implemented(self):
"""
Ensure the following methods throw NIE's. If not, come back and test them.
Expand Down