From 2bd524f6c632c0bd4b8635eb272bcf45e6a06b6d Mon Sep 17 00:00:00 2001 From: Dana Powers Date: Sun, 19 Jul 2026 10:02:08 -0700 Subject: [PATCH] Move KafkaSSLTransport -> kafka.net.ssl --- kafka/net/backend/selector.py | 3 +- kafka/net/backend/transport.py | 249 ---------------------------- kafka/net/manager.py | 2 +- kafka/net/ssl.py | 256 +++++++++++++++++++++++++++++ test/net/backend/test_transport.py | 69 +------- test/net/test_ssl.py | 74 +++++++++ 6 files changed, 335 insertions(+), 318 deletions(-) create mode 100644 kafka/net/ssl.py create mode 100644 test/net/test_ssl.py diff --git a/kafka/net/backend/selector.py b/kafka/net/backend/selector.py index e35256003..1e23f3d04 100644 --- a/kafka/net/backend/selector.py +++ b/kafka/net/backend/selector.py @@ -12,7 +12,8 @@ import kafka.errors as Errors from kafka.future import Future from kafka.net.backend.inet import create_connection as _inet_create_connection -from kafka.net.backend.transport import KafkaSSLTransport, KafkaTCPTransport +from kafka.net.backend.transport import KafkaTCPTransport +from kafka.net.ssl import KafkaSSLTransport from kafka.version import __version__ diff --git a/kafka/net/backend/transport.py b/kafka/net/backend/transport.py index 72f0709fa..b876ca8a5 100644 --- a/kafka/net/backend/transport.py +++ b/kafka/net/backend/transport.py @@ -1,10 +1,8 @@ from collections import deque -import copy import enum import logging import selectors import socket -import ssl import time import kafka.errors as Errors @@ -253,250 +251,3 @@ def host_port(self): def __str__(self): state = ' (closed)' if self._closed else '' return f"<{self.__class__.__name__} [{self.host_port()}]{state}>" - - -class ConnectionState(enum.Enum): - HANDSHAKE = 'handshake' - CONNECTED = 'connected' - CLOSED = 'closed' - - -class KafkaSSLTransport: - DEFAULT_CONFIG = { - 'ssl_context': None, - 'ssl_check_hostname': True, - 'ssl_cafile': None, - 'ssl_certfile': None, - 'ssl_keyfile': None, - 'ssl_password': None, - 'ssl_crlfile': None, - } - def __init__(self, net, ssl_context, host=None): - self._net = net - self._state = None - self._connect_future = self._net.create_future() - self._ssl_context = ssl_context - self.host = host - server_hostname = host.rstrip('.') if host is not None else None - self._incoming = ssl.MemoryBIO() - self._outgoing = ssl.MemoryBIO() - self._ssl_object = self._ssl_context.wrap_bio( - self._incoming, self._outgoing, - server_hostname=server_hostname) - self._write_buffer = deque() # list of bytes that are pending ssl.send() - # recvs from transport, writes to protocol - self._transport = None - self._protocol = None - self._write = False - - @classmethod - def build_ssl_context(cls, configs): - config = copy.copy(cls.DEFAULT_CONFIG) - for key in config: - if key in configs: - config[key] = configs[key] - - if config['ssl_context'] is not None: - return config['ssl_context'] - ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ctx.minimum_version = ssl.TLSVersion.TLSv1_2 - ctx.check_hostname = config['ssl_check_hostname'] - if config['ssl_cafile']: - ctx.load_verify_locations(config['ssl_cafile']) - else: - ctx.load_default_certs() - if config['ssl_certfile']: - ctx.load_cert_chain( - certfile=config['ssl_certfile'], - keyfile=config['ssl_keyfile'], - password=config['ssl_password'], - ) - if config['ssl_crlfile']: - ctx.load_verify_locations(crl=config['ssl_crlfile']) - ctx.verify_flags |= ssl.VERIFY_CRL_CHECK_LEAF - return ctx - - def close(self, err=None): - self._state = ConnectionState.CLOSED - self._write = False - if self._protocol: - protocol, self._protocol = self._protocol, None - protocol.connection_lost(err) - if self._transport: - transport, self._transport = self._transport, None - if err: - transport.abort(err) - else: - transport.close() - if not self._connect_future.is_done: - self._connect_future.failure(err or Errors.Cancelled()) - - def abort(self, error): - self.close(error) - - def data_received(self, data): - # from underlying transport (tcp or proxy) - if self._state not in (ConnectionState.HANDSHAKE, ConnectionState.CONNECTED): - log.warning('%s: ignoring data_received %d bytes because not connected', self, len(data)) - return - log.debug('%s: data_received %d bytes', self, len(data)) - self._incoming.write(data) - if self._state == ConnectionState.HANDSHAKE: - self._do_handshake() - return - self._do_recv() - - def _do_recv(self): - data, err = self._ssl_recv() - if err: - self.close(err) - else: - self._process_outgoing() - self._protocol.data_received(data) - self._do_send() - - def write(self, data): - # from outer protocol (connection) - if self._state not in (ConnectionState.HANDSHAKE, ConnectionState.CONNECTED): - log.warning('%s: ignoring write %d bytes because not connected', self, len(data)) - return - log.debug('%s: write %d bytes', self, len(data)) - self._write_buffer.append(data) - if self._state == ConnectionState.HANDSHAKE: - self._do_handshake() - return - self._do_send() - - def _do_send(self): - nbytes, err = self._ssl_send() - if err: - self.close(err) - else: - self._process_outgoing() - - def _process_outgoing(self): - if not self._write: - return - data = self._outgoing.read() - if len(data): - self._transport.write(data) - - def set_protocol(self, protocol): - """Set a new protocol.""" - self._protocol = protocol - log.debug('%s: Set protocol %s', self, protocol) - - def get_protocol(self): - """Return the current protocol.""" - return self._protocol - - def _ssl_recv(self): - recvd = [] - err = None - while True: - try: - data = self._ssl_object.read(4096) - if not data: - log.error('%s: socket disconnected', self) - err = Errors.KafkaConnectionError('socket disconnected') - break - else: - recvd.append(data) - - except (ssl.SSLWantReadError, ssl.SSLWantWriteError): - break - except BaseException as e: - log.exception('%s: Error receiving ssl data' - ' closing transport', self) - err = Errors.KafkaConnectionError(e) - break - - recvd_data = b''.join(recvd) - return recvd_data, err - - def _ssl_send(self): - total_bytes = 0 - if self._state == ConnectionState.CLOSED: - return total_bytes, Errors.KafkaConnectionError('Connection closed') - while self._write_buffer: - next_chunk = self._write_buffer.popleft() - # Wrap in memoryview so partial-send slicing is O(1) instead of - # copying the unsent tail on every BlockingIOError / short write. - if not isinstance(next_chunk, memoryview): - next_chunk = memoryview(next_chunk) - while next_chunk: - try: - sent_bytes = self._ssl_object.write(next_chunk) - total_bytes += sent_bytes - next_chunk = next_chunk[sent_bytes:] - except (ssl.SSLWantReadError, ssl.SSLWantWriteError): - self._write_buffer.appendleft(next_chunk) - self._process_outgoing() - return total_bytes, None - except BaseException as e: - log.exception("%s: Error sending request data: %s", self, e) - return total_bytes, Errors.KafkaConnectionError(e) - return total_bytes, None - - def _do_handshake(self): - log.debug('%s: _do_handshake', self) - try: - self._ssl_object.do_handshake() - except (ssl.SSLWantReadError, ssl.SSLWantWriteError) as e: - log.debug('%s: %s', self, e) - self._process_outgoing() - pass - except BaseException as exc: - log.error("%s: Error during TLS Handshake: %s", self, exc) - self.close(exc) - else: - log.info('%s: connected', self) - self._state = ConnectionState.CONNECTED - self._connect_future.success(True) - self._do_send() - return - - async def handshake(self): - self._do_handshake() - await self._connect_future - - @property - def last_activity(self): - return self._transport.last_activity - - def is_closing(self): - return self._state is ConnectionState.CLOSED - - def pause_reading(self): - return self._transport.pause_reading() - - def resume_reading(self): - return self._transport.resume_reading() - - def pause_writing(self): - self._write = False - - def resume_writing(self): - self._write = True - self._process_outgoing() - - def host_port(self): - if self._transport: - return self._transport.host_port() - - def connection_made(self, transport): - self._transport = transport - self._transport.set_protocol(self) - self._state = ConnectionState.HANDSHAKE - self._transport.resume_reading() - self.resume_writing() - - def connection_lost(self, exc): - self.abort(exc) - - def get_peer(self): - if self._transport: - return self._transport.get_peer() - - def __str__(self): - return f"<{self.__class__.__name__} [{self.host_port()}]>" diff --git a/kafka/net/manager.py b/kafka/net/manager.py index 64eec5a02..8acc76e4d 100644 --- a/kafka/net/manager.py +++ b/kafka/net/manager.py @@ -10,7 +10,7 @@ from kafka.net.backend import resolve_backend from kafka.cluster import ClusterMetadata import kafka.errors as Errors -from kafka.net.backend.transport import KafkaSSLTransport +from kafka.net.ssl import KafkaSSLTransport from kafka.net.wakeup_notifier import WakeupNotifier from kafka.protocol.broker_version_data import BrokerVersionData from kafka.version import __version__ diff --git a/kafka/net/ssl.py b/kafka/net/ssl.py new file mode 100644 index 000000000..9749866b2 --- /dev/null +++ b/kafka/net/ssl.py @@ -0,0 +1,256 @@ +from collections import deque +import copy +import enum +import logging +import ssl + +import kafka.errors as Errors + +log = logging.getLogger(__name__) + + +class ConnectionState(enum.Enum): + HANDSHAKE = 'handshake' + CONNECTED = 'connected' + CLOSED = 'closed' + + +class KafkaSSLTransport: + DEFAULT_CONFIG = { + 'ssl_context': None, + 'ssl_check_hostname': True, + 'ssl_cafile': None, + 'ssl_certfile': None, + 'ssl_keyfile': None, + 'ssl_password': None, + 'ssl_crlfile': None, + } + def __init__(self, net, ssl_context, host=None): + self._net = net + self._state = None + self._connect_future = self._net.create_future() + self._ssl_context = ssl_context + self.host = host + server_hostname = host.rstrip('.') if host is not None else None + self._incoming = ssl.MemoryBIO() + self._outgoing = ssl.MemoryBIO() + self._ssl_object = self._ssl_context.wrap_bio( + self._incoming, self._outgoing, + server_hostname=server_hostname) + self._write_buffer = deque() # list of bytes that are pending ssl.send() + # recvs from transport, writes to protocol + self._transport = None + self._protocol = None + self._write = False + + @classmethod + def build_ssl_context(cls, configs): + config = copy.copy(cls.DEFAULT_CONFIG) + for key in config: + if key in configs: + config[key] = configs[key] + + if config['ssl_context'] is not None: + return config['ssl_context'] + ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + ctx.minimum_version = ssl.TLSVersion.TLSv1_2 + ctx.check_hostname = config['ssl_check_hostname'] + if config['ssl_cafile']: + ctx.load_verify_locations(config['ssl_cafile']) + else: + ctx.load_default_certs() + if config['ssl_certfile']: + ctx.load_cert_chain( + certfile=config['ssl_certfile'], + keyfile=config['ssl_keyfile'], + password=config['ssl_password'], + ) + if config['ssl_crlfile']: + ctx.load_verify_locations(crl=config['ssl_crlfile']) + ctx.verify_flags |= ssl.VERIFY_CRL_CHECK_LEAF + return ctx + + def close(self, err=None): + self._state = ConnectionState.CLOSED + self._write = False + if self._protocol: + protocol, self._protocol = self._protocol, None + protocol.connection_lost(err) + if self._transport: + transport, self._transport = self._transport, None + if err: + transport.abort(err) + else: + transport.close() + if not self._connect_future.is_done: + self._connect_future.failure(err or Errors.Cancelled()) + + def abort(self, error): + self.close(error) + + def data_received(self, data): + # from underlying transport (tcp or proxy) + if self._state not in (ConnectionState.HANDSHAKE, ConnectionState.CONNECTED): + log.warning('%s: ignoring data_received %d bytes because not connected', self, len(data)) + return + log.debug('%s: data_received %d bytes', self, len(data)) + self._incoming.write(data) + if self._state == ConnectionState.HANDSHAKE: + self._do_handshake() + return + self._do_recv() + + def _do_recv(self): + data, err = self._ssl_recv() + if err: + self.close(err) + else: + self._process_outgoing() + self._protocol.data_received(data) + self._do_send() + + def write(self, data): + # from outer protocol (connection) + if self._state not in (ConnectionState.HANDSHAKE, ConnectionState.CONNECTED): + log.warning('%s: ignoring write %d bytes because not connected', self, len(data)) + return + log.debug('%s: write %d bytes', self, len(data)) + self._write_buffer.append(data) + if self._state == ConnectionState.HANDSHAKE: + self._do_handshake() + return + self._do_send() + + def _do_send(self): + nbytes, err = self._ssl_send() + if err: + self.close(err) + else: + self._process_outgoing() + + def _process_outgoing(self): + if not self._write: + return + data = self._outgoing.read() + if len(data): + self._transport.write(data) + + def set_protocol(self, protocol): + """Set a new protocol.""" + self._protocol = protocol + log.debug('%s: Set protocol %s', self, protocol) + + def get_protocol(self): + """Return the current protocol.""" + return self._protocol + + def _ssl_recv(self): + recvd = [] + err = None + while True: + try: + data = self._ssl_object.read(4096) + if not data: + log.error('%s: socket disconnected', self) + err = Errors.KafkaConnectionError('socket disconnected') + break + else: + recvd.append(data) + + except (ssl.SSLWantReadError, ssl.SSLWantWriteError): + break + except BaseException as e: + log.exception('%s: Error receiving ssl data' + ' closing transport', self) + err = Errors.KafkaConnectionError(e) + break + + recvd_data = b''.join(recvd) + return recvd_data, err + + def _ssl_send(self): + total_bytes = 0 + if self._state == ConnectionState.CLOSED: + return total_bytes, Errors.KafkaConnectionError('Connection closed') + while self._write_buffer: + next_chunk = self._write_buffer.popleft() + # Wrap in memoryview so partial-send slicing is O(1) instead of + # copying the unsent tail on every BlockingIOError / short write. + if not isinstance(next_chunk, memoryview): + next_chunk = memoryview(next_chunk) + while next_chunk: + try: + sent_bytes = self._ssl_object.write(next_chunk) + total_bytes += sent_bytes + next_chunk = next_chunk[sent_bytes:] + except (ssl.SSLWantReadError, ssl.SSLWantWriteError): + self._write_buffer.appendleft(next_chunk) + self._process_outgoing() + return total_bytes, None + except BaseException as e: + log.exception("%s: Error sending request data: %s", self, e) + return total_bytes, Errors.KafkaConnectionError(e) + return total_bytes, None + + def _do_handshake(self): + log.debug('%s: _do_handshake', self) + try: + self._ssl_object.do_handshake() + except (ssl.SSLWantReadError, ssl.SSLWantWriteError) as e: + log.debug('%s: %s', self, e) + self._process_outgoing() + pass + except BaseException as exc: + log.error("%s: Error during TLS Handshake: %s", self, exc) + self.close(exc) + else: + log.info('%s: connected', self) + self._state = ConnectionState.CONNECTED + self._connect_future.success(True) + self._do_send() + return + + async def handshake(self): + self._do_handshake() + await self._connect_future + + @property + def last_activity(self): + return self._transport.last_activity + + def is_closing(self): + return self._state is ConnectionState.CLOSED + + def pause_reading(self): + return self._transport.pause_reading() + + def resume_reading(self): + return self._transport.resume_reading() + + def pause_writing(self): + self._write = False + + def resume_writing(self): + self._write = True + self._process_outgoing() + + def host_port(self): + if self._transport: + return self._transport.host_port() + + def connection_made(self, transport): + self._transport = transport + self._transport.set_protocol(self) + self._state = ConnectionState.HANDSHAKE + self._transport.resume_reading() + self.resume_writing() + + def connection_lost(self, exc): + self.abort(exc) + + def get_peer(self): + if self._transport: + return self._transport.get_peer() + + def __str__(self): + return f"<{self.__class__.__name__} [{self.host_port()}]>" diff --git a/test/net/backend/test_transport.py b/test/net/backend/test_transport.py index 77d83ad7a..26ebc36fc 100644 --- a/test/net/backend/test_transport.py +++ b/test/net/backend/test_transport.py @@ -9,7 +9,8 @@ from kafka.future import Future from kafka.net.backend import NetTransport, NetProtocol from kafka.net.backend.selector import NetworkSelector, TaskState -from kafka.net.backend.transport import KafkaSSLTransport, KafkaTCPTransport +from kafka.net.backend.transport import KafkaTCPTransport +from kafka.net.ssl import KafkaSSLTransport @pytest.fixture @@ -301,48 +302,6 @@ def test_str_closed(self, net): s = str(t) assert 'closed' in s - -class TestKafkaSSLTransport: - """Regression tests for https://github.com/dpkp/kafka-python/issues/3113: - TLS SNI (server_hostname) must be sent regardless of ssl_check_hostname so - that SNI-routed clusters (nginx/Istio/Strimzi ingress) remain reachable when - hostname verification is disabled. - """ - - def test_sni_sent_when_check_hostname_true(self, net): - ctx = MagicMock() - ctx.check_hostname = True - KafkaSSLTransport(net, ctx, host='broker.example.com') - _, kwargs = ctx.wrap_bio.call_args - assert kwargs['server_hostname'] == 'broker.example.com' - - def test_sni_sent_when_check_hostname_false(self, net): - # The bug: SNI used to be suppressed when verification was disabled. - ctx = MagicMock() - ctx.check_hostname = False - KafkaSSLTransport(net, ctx, host='broker.example.com') - _, kwargs = ctx.wrap_bio.call_args - assert kwargs['server_hostname'] == 'broker.example.com' - - def test_sni_strips_trailing_dot(self, net): - # A trailing dot is a valid FQDN but illegal in the SNI extension. - ctx = MagicMock() - ctx.check_hostname = False - KafkaSSLTransport(net, ctx, host='broker.example.com.') - _, kwargs = ctx.wrap_bio.call_args - assert kwargs['server_hostname'] == 'broker.example.com' - - def test_sni_none_when_host_missing(self, net): - ctx = MagicMock() - KafkaSSLTransport(net, ctx, host=None) - _, kwargs = ctx.wrap_bio.call_args - assert kwargs['server_hostname'] is None - - def test_provided_ssl_context_is_used(self, net): - ctx = MagicMock() - t = KafkaSSLTransport(net, ctx, host='broker.example.com') - assert t._ssl_context is ctx - def test_ssl_wrapper_transport(self, net, socketpair): rsock, wsock = socketpair t = KafkaTCPTransport(net, wsock) @@ -352,30 +311,6 @@ def test_ssl_wrapper_transport(self, net, socketpair): assert isinstance(ssl_wrapper, NetProtocol) -class TestBuildSSLContext: - def test_returns_provided_context(self): - ctx = MagicMock() - config = dict(ssl_context=ctx) - assert KafkaSSLTransport.build_ssl_context(config) is ctx - - def test_check_hostname_propagates_to_context(self): - config = dict(ssl_check_hostname=False) - ctx = KafkaSSLTransport.build_ssl_context(config) - assert ctx.check_hostname is False - - def test_check_hostname_true_requires_verification(self): - config = dict(ssl_check_hostname=True) - ctx = KafkaSSLTransport.build_ssl_context(config) - assert ctx.check_hostname is True - - def test_empty_config_builds_default_context(self): - # Missing keys fall back to DEFAULT_CONFIG: a real TLS-client context - # with hostname checking on. This is the default-build path. - ctx = KafkaSSLTransport.build_ssl_context({}) - assert isinstance(ctx, ssl.SSLContext) - assert ctx.check_hostname is True - - class TestTransportWaiterCleanup: """Regression: a locally-initiated close()/abort() must reclaim the socket read/write coroutine tasks parked in the event loop. diff --git a/test/net/test_ssl.py b/test/net/test_ssl.py new file mode 100644 index 000000000..88462ad40 --- /dev/null +++ b/test/net/test_ssl.py @@ -0,0 +1,74 @@ +import ssl + +from unittest.mock import MagicMock + +from kafka.net.backend import NetTransport, NetProtocol +from kafka.net.backend.transport import KafkaTCPTransport + +from kafka.net.ssl import KafkaSSLTransport + + +class TestKafkaSSLTransport: + """Regression tests for https://github.com/dpkp/kafka-python/issues/3113: + TLS SNI (server_hostname) must be sent regardless of ssl_check_hostname so + that SNI-routed clusters (nginx/Istio/Strimzi ingress) remain reachable when + hostname verification is disabled. + """ + + def test_sni_sent_when_check_hostname_true(self, net): + ctx = MagicMock() + ctx.check_hostname = True + KafkaSSLTransport(net, ctx, host='broker.example.com') + _, kwargs = ctx.wrap_bio.call_args + assert kwargs['server_hostname'] == 'broker.example.com' + + def test_sni_sent_when_check_hostname_false(self, net): + # The bug: SNI used to be suppressed when verification was disabled. + ctx = MagicMock() + ctx.check_hostname = False + KafkaSSLTransport(net, ctx, host='broker.example.com') + _, kwargs = ctx.wrap_bio.call_args + assert kwargs['server_hostname'] == 'broker.example.com' + + def test_sni_strips_trailing_dot(self, net): + # A trailing dot is a valid FQDN but illegal in the SNI extension. + ctx = MagicMock() + ctx.check_hostname = False + KafkaSSLTransport(net, ctx, host='broker.example.com.') + _, kwargs = ctx.wrap_bio.call_args + assert kwargs['server_hostname'] == 'broker.example.com' + + def test_sni_none_when_host_missing(self, net): + ctx = MagicMock() + KafkaSSLTransport(net, ctx, host=None) + _, kwargs = ctx.wrap_bio.call_args + assert kwargs['server_hostname'] is None + + def test_provided_ssl_context_is_used(self, net): + ctx = MagicMock() + t = KafkaSSLTransport(net, ctx, host='broker.example.com') + assert t._ssl_context is ctx + + +class TestBuildSSLContext: + def test_returns_provided_context(self): + ctx = MagicMock() + config = dict(ssl_context=ctx) + assert KafkaSSLTransport.build_ssl_context(config) is ctx + + def test_check_hostname_propagates_to_context(self): + config = dict(ssl_check_hostname=False) + ctx = KafkaSSLTransport.build_ssl_context(config) + assert ctx.check_hostname is False + + def test_check_hostname_true_requires_verification(self): + config = dict(ssl_check_hostname=True) + ctx = KafkaSSLTransport.build_ssl_context(config) + assert ctx.check_hostname is True + + def test_empty_config_builds_default_context(self): + # Missing keys fall back to DEFAULT_CONFIG: a real TLS-client context + # with hostname checking on. This is the default-build path. + ctx = KafkaSSLTransport.build_ssl_context({}) + assert isinstance(ctx, ssl.SSLContext) + assert ctx.check_hostname is True