5555
5656log = logging .getLogger (__name__ )
5757
58+
5859segment_codec_no_compression = SegmentCodec ()
5960segment_codec_lz4 = None
6061
@@ -130,6 +131,81 @@ def decompress(byts):
130131frame_header_v3 = struct .Struct ('>BhBi' )
131132
132133
134+ def _default_pyopenssl_ssl_method (ssl_module ):
135+ for method_name in ('TLS_CLIENT_METHOD' , 'TLS_METHOD' , 'TLSv1_2_METHOD' ):
136+ method = getattr (ssl_module , method_name , None )
137+ if method is not None :
138+ return method
139+ raise ImportError ('pyOpenSSL does not expose a secure TLS client method' )
140+
141+
142+ def _pyopenssl_ssl_method_from_stdlib (ssl_module , ssl_version ):
143+ if ssl_version is None :
144+ return _default_pyopenssl_ssl_method (ssl_module )
145+
146+ protocol_method_names = (
147+ ('PROTOCOL_TLS_CLIENT' , ('TLS_CLIENT_METHOD' , 'TLS_METHOD' , 'TLSv1_2_METHOD' )),
148+ ('PROTOCOL_TLS' , ('TLS_METHOD' , 'TLS_CLIENT_METHOD' , 'TLSv1_2_METHOD' )),
149+ ('PROTOCOL_SSLv23' , ('TLS_METHOD' , 'TLS_CLIENT_METHOD' , 'TLSv1_2_METHOD' )),
150+ ('PROTOCOL_TLSv1_2' , ('TLSv1_2_METHOD' ,)),
151+ ('PROTOCOL_TLSv1_1' , ('TLSv1_1_METHOD' ,)),
152+ ('PROTOCOL_TLSv1' , ('TLSv1_METHOD' ,)),
153+ )
154+ for protocol_name , method_names in protocol_method_names :
155+ protocol = getattr (ssl , protocol_name , None )
156+ if (protocol is not None and
157+ ssl_version .__class__ is protocol .__class__ and
158+ ssl_version == protocol ):
159+ for method_name in method_names :
160+ method = getattr (ssl_module , method_name , None )
161+ if method is not None :
162+ return method
163+ raise ImportError ('pyOpenSSL does not expose a method for %s' % (protocol_name ,))
164+
165+ return ssl_version
166+
167+
168+ def _pyopenssl_verify_mode_from_cert_reqs (ssl_module , cert_reqs ):
169+ if cert_reqs is None :
170+ return None
171+ if cert_reqs .__class__ is not type (ssl .CERT_REQUIRED ):
172+ return cert_reqs
173+ if cert_reqs == ssl .CERT_NONE :
174+ return ssl_module .VERIFY_NONE
175+ if cert_reqs in (ssl .CERT_OPTIONAL , ssl .CERT_REQUIRED ):
176+ return ssl_module .VERIFY_PEER
177+ return cert_reqs
178+
179+
180+ def _build_pyopenssl_context_from_options (ssl_module , ssl_options ):
181+ ssl_options = ssl_options or {}
182+ context = ssl_module .Context (
183+ _pyopenssl_ssl_method_from_stdlib (ssl_module , ssl_options .get ('ssl_version' , None ))
184+ )
185+ if 'certfile' in ssl_options :
186+ context .use_certificate_file (ssl_options ['certfile' ])
187+ if 'keyfile' in ssl_options :
188+ context .use_privatekey_file (ssl_options ['keyfile' ])
189+ if 'ca_certs' in ssl_options :
190+ context .load_verify_locations (ssl_options ['ca_certs' ])
191+ cert_reqs = _pyopenssl_verify_mode_from_cert_reqs (
192+ ssl_module , ssl_options .get ('cert_reqs' , None ))
193+ if cert_reqs is None :
194+ cert_reqs = (ssl_module .VERIFY_PEER
195+ if (ssl_options .get ('ca_certs' , None ) or ssl_options .get ('check_hostname' , False ))
196+ else ssl_module .VERIFY_NONE )
197+ context .set_verify (
198+ cert_reqs ,
199+ callback = lambda _connection , _x509 , _errnum , _errdepth , ok : ok
200+ )
201+ ciphers = ssl_options .get ('ciphers' , None )
202+ if ciphers :
203+ if isinstance (ciphers , str ):
204+ ciphers = ciphers .encode ('ascii' )
205+ context .set_cipher_list (ciphers )
206+ return context
207+
208+
133209class EndPoint (object ):
134210 """
135211 Represents the information to connect to a cassandra node.
@@ -803,6 +879,7 @@ class Connection(object):
803879 endpoint = None
804880 ssl_options = None
805881 ssl_context = None
882+ _ssl_options_explicit = False
806883 last_error = None
807884
808885 # The current number of operations that are in flight. More precisely,
@@ -885,7 +962,11 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None,
885962 self .endpoint = host if isinstance (host , EndPoint ) else DefaultEndPoint (host , port )
886963
887964 self .authenticator = authenticator
888- self .ssl_options = ssl_options .copy () if ssl_options else {}
965+ endpoint_ssl_options = self .endpoint .ssl_options
966+ # Explicit ssl_options={} enables SSL with default options; omitted
967+ # ssl_options=None leaves SSL disabled unless an endpoint supplies options.
968+ self ._ssl_options_explicit = ssl_options is not None
969+ self .ssl_options = ssl_options .copy () if ssl_options is not None else {}
889970 self .ssl_context = ssl_context
890971 self .sockopts = sockopts
891972 self .compression = compression
@@ -905,10 +986,13 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None,
905986 self ._on_orphaned_stream_released = on_orphaned_stream_released
906987 self ._application_info = application_info
907988
908- if ssl_options :
909- self .ssl_options .update (self .endpoint .ssl_options or {})
910- elif self .endpoint .ssl_options :
911- self .ssl_options = self .endpoint .ssl_options
989+ if ssl_options is not None :
990+ self .ssl_options .update (endpoint_ssl_options or {})
991+ elif endpoint_ssl_options is not None :
992+ self ._ssl_options_explicit = True
993+ self .ssl_options = endpoint_ssl_options
994+ self ._check_hostname = bool (self .ssl_options .get ('check_hostname' , False ) or
995+ getattr (self .ssl_context , 'check_hostname' , False ))
912996
913997 # PYTHON-1331
914998 #
@@ -918,7 +1002,7 @@ def __init__(self, host='127.0.0.1', port=9042, authenticator=None,
9181002 #
9191003 # Note the use of pop() here; we are very deliberately removing these params from ssl_options if they're present. After this
9201004 # operation ssl_options should contain only args needed for the ssl_context.wrap_socket() call.
921- if not self .ssl_context and self .ssl_options :
1005+ if self .ssl_context is None and self ._ssl_options_explicit :
9221006 self .ssl_context = self ._build_ssl_context_from_options ()
9231007
9241008 self .max_request_id = min (self .max_in_flight - 1 , (2 ** 15 ) - 1 )
@@ -942,6 +1026,10 @@ def host(self):
9421026 def port (self ):
9431027 return self .endpoint .port
9441028
1029+ @property
1030+ def _ssl_enabled (self ):
1031+ return self .ssl_context is not None or self ._ssl_options_explicit
1032+
9451033 @classmethod
9461034 def initialize_reactor (cls ):
9471035 """
@@ -1000,10 +1088,15 @@ def _build_ssl_context_from_options(self):
10001088 # Python >= 3.10 requires either PROTOCOL_TLS_CLIENT or PROTOCOL_TLS_SERVER so we'll get ahead of things by always
10011089 # being explicit
10021090 ssl_version = opts .get ('ssl_version' , None ) or ssl .PROTOCOL_TLS_CLIENT
1003- cert_reqs = opts .get ('cert_reqs' , None ) or ssl .CERT_REQUIRED
1091+ cert_reqs = opts .get ('cert_reqs' , None )
1092+ if cert_reqs is None :
1093+ cert_reqs = (ssl .CERT_REQUIRED
1094+ if (opts .get ('ca_certs' , None ) or opts .get ('check_hostname' , False ))
1095+ else ssl .CERT_NONE )
10041096 rv = ssl .SSLContext (protocol = int (ssl_version ))
1097+ rv .check_hostname = False
1098+ rv .verify_mode = cert_reqs
10051099 rv .check_hostname = bool (opts .get ('check_hostname' , False ))
1006- rv .options = int (cert_reqs )
10071100
10081101 certfile = opts .get ('certfile' , None )
10091102 keyfile = opts .get ('keyfile' , None )
0 commit comments