diff --git a/kafka/net/backend/abstract.py b/kafka/net/backend/abstract.py index 04521bb94..183de2632 100644 --- a/kafka/net/backend/abstract.py +++ b/kafka/net/backend/abstract.py @@ -115,6 +115,9 @@ class NetTransport(Protocol): """ # Monotonic timestamp of the last read/write; used for idle-sweeping last_activity: float + # User-supplied hostname (before DNS resolution); read by _sasl_authenticate + # so mechanisms like GSSAPI build principals against the configured name. + host: Optional[str] def write(self, data: bytes) -> None: ... def close(self) -> None: ... diff --git a/kafka/net/backend/selector.py b/kafka/net/backend/selector.py index ecb7b3f58..e35256003 100644 --- a/kafka/net/backend/selector.py +++ b/kafka/net/backend/selector.py @@ -575,19 +575,16 @@ async def create_connection(self, protocol, host, port, *, ssl=None, """ sock = await _inet_create_connection(self, host, port, socket_options, proxy_url=proxy_url, timeout_at=timeout_at) + transport = KafkaTCPTransport(self, sock, host=host) if ssl is not None: - transport = KafkaSSLTransport(self, sock, ssl, host=host) - else: - transport = KafkaTCPTransport(self, sock, host=host) - try: - await transport.handshake() - except Exception as e: - transport.close() - raise Errors.KafkaConnectionError('Handshake failed: %s' % e) + ssl_wrapper = KafkaSSLTransport(self, ssl, host=host) + ssl_wrapper.connection_made(transport) + await ssl_wrapper.handshake() + transport = ssl_wrapper try: protocol.connection_made(transport) - except Exception: - transport.close() + except Exception as e: + transport.abort(e) raise def sleep(self, delay): diff --git a/kafka/net/backend/transport.py b/kafka/net/backend/transport.py index 8846b848d..72f0709fa 100644 --- a/kafka/net/backend/transport.py +++ b/kafka/net/backend/transport.py @@ -1,5 +1,6 @@ from collections import deque import copy +import enum import logging import selectors import socket @@ -254,7 +255,13 @@ def __str__(self): return f"<{self.__class__.__name__} [{self.host_port()}]{state}>" -class KafkaSSLTransport(KafkaTCPTransport): +class ConnectionState(enum.Enum): + HANDSHAKE = 'handshake' + CONNECTED = 'connected' + CLOSED = 'closed' + + +class KafkaSSLTransport: DEFAULT_CONFIG = { 'ssl_context': None, 'ssl_check_hostname': True, @@ -264,13 +271,23 @@ class KafkaSSLTransport(KafkaTCPTransport): 'ssl_password': None, 'ssl_crlfile': None, } - def __init__(self, net, sock, ssl_context, host=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 - sock = self._ssl_context.wrap_socket( - sock, server_hostname=server_hostname, - do_handshake_on_connect=False) - super().__init__(net, sock, host=host) + 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): @@ -299,58 +316,187 @@ def build_ssl_context(cls, configs): ctx.verify_flags |= ssl.VERIFY_CRL_CHECK_LEAF return ctx - async def handshake(self): - while True: - try: - self._sock.do_handshake() - log.info('%s: connected to %s', self, self._sock) - return - except ssl.SSLWantReadError: - await self._net.wait_read(self._sock) - except ssl.SSLWantWriteError: - await self._net.wait_write(self._sock) + 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 _sock_recv(self): + 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._sock.recv(4096) + 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 (BlockingIOError, InterruptedError, - ssl.SSLWantReadError, ssl.SSLWantWriteError): + + except (ssl.SSLWantReadError, ssl.SSLWantWriteError): break except BaseException as e: - log.exception('%s: Error receiving network data' - ' closing socket', self) + 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 _sock_send(self): + def _ssl_send(self): total_bytes = 0 - err = None - if self._sock is None: - return total_bytes, Errors.KafkaConnectionError('Connection closed during send') + 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._sock.send(next_chunk) + sent_bytes = self._ssl_object.write(next_chunk) total_bytes += sent_bytes next_chunk = next_chunk[sent_bytes:] - except (BlockingIOError, InterruptedError, - ssl.SSLWantReadError, ssl.SSLWantWriteError): + except (ssl.SSLWantReadError, ssl.SSLWantWriteError): self._write_buffer.appendleft(next_chunk) - return total_bytes, err + self._process_outgoing() + return total_bytes, None except BaseException as e: log.exception("%s: Error sending request data: %s", self, e) - err = Errors.KafkaConnectionError(e) - return total_bytes, err - return total_bytes, err + 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 b746a6f8e..77d83ad7a 100644 --- a/test/net/backend/test_transport.py +++ b/test/net/backend/test_transport.py @@ -7,6 +7,7 @@ import kafka.errors as Errors 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 @@ -308,63 +309,47 @@ class TestKafkaSSLTransport: hostname verification is disabled. """ - def _make_ssl_sock(self): - # wrap_socket returns a wrapped socket; give it the peer/name accessors - # that KafkaTCPTransport.__init__ pokes at via str()/repr helpers. - wrapped = _make_mock_sock() - sock = _make_mock_sock() - ctx = MagicMock() - ctx.wrap_socket.return_value = wrapped - return sock, ctx, wrapped - def test_sni_sent_when_check_hostname_true(self, net): - sock, ctx, _ = self._make_ssl_sock() + ctx = MagicMock() ctx.check_hostname = True - KafkaSSLTransport(net, sock, ctx, host='broker.example.com') - _, kwargs = ctx.wrap_socket.call_args + 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. - sock, ctx, _ = self._make_ssl_sock() + ctx = MagicMock() ctx.check_hostname = False - KafkaSSLTransport(net, sock, ctx, host='broker.example.com') - _, kwargs = ctx.wrap_socket.call_args + 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. - sock, ctx, _ = self._make_ssl_sock() + ctx = MagicMock() ctx.check_hostname = False - KafkaSSLTransport(net, sock, ctx, host='broker.example.com.') - _, kwargs = ctx.wrap_socket.call_args + 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): - sock, ctx, _ = self._make_ssl_sock() - KafkaSSLTransport(net, sock, ctx, host=None) - _, kwargs = ctx.wrap_socket.call_args + ctx = MagicMock() + KafkaSSLTransport(net, ctx, host=None) + _, kwargs = ctx.wrap_bio.call_args assert kwargs['server_hostname'] is None - def test_handshake_not_done_on_connect(self, net): - sock, ctx, _ = self._make_ssl_sock() - KafkaSSLTransport(net, sock, ctx, host='broker.example.com') - _, kwargs = ctx.wrap_socket.call_args - assert kwargs['do_handshake_on_connect'] is False - def test_provided_ssl_context_is_used(self, net): - sock, ctx, wrapped = self._make_ssl_sock() - t = KafkaSSLTransport(net, sock, ctx, host='broker.example.com') + ctx = MagicMock() + t = KafkaSSLTransport(net, ctx, host='broker.example.com') assert t._ssl_context is ctx - assert t._sock is wrapped - - def test_ssl_context_is_required(self, net): - # The transport no longer builds a context itself; callers must pass - # a pre-built one (via build_ssl_context). Omitting it is a TypeError, - # not a silently-default context. - sock, _, _ = self._make_ssl_sock() - with pytest.raises(TypeError): - KafkaSSLTransport(net, sock, host='broker.example.com') # pylint: disable=E1120 + + def test_ssl_wrapper_transport(self, net, socketpair): + rsock, wsock = socketpair + t = KafkaTCPTransport(net, wsock) + ssl_wrapper = KafkaSSLTransport(net, KafkaSSLTransport.build_ssl_context({})) + ssl_wrapper.connection_made(t) + assert isinstance(ssl_wrapper, NetTransport) + assert isinstance(ssl_wrapper, NetProtocol) class TestBuildSSLContext: