Skip to content
Merged
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
3 changes: 3 additions & 0 deletions kafka/net/backend/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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: ...
Expand Down
17 changes: 7 additions & 10 deletions kafka/net/backend/selector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
212 changes: 179 additions & 33 deletions kafka/net/backend/transport.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from collections import deque
import copy
import enum
import logging
import selectors
import socket
Expand Down Expand Up @@ -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,
Expand All @@ -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):
Expand Down Expand Up @@ -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()}]>"
61 changes: 23 additions & 38 deletions test/net/backend/test_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down