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
28 changes: 4 additions & 24 deletions kafka/net/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
import inspect
import random
import socket
import ssl
import time

from .inet import create_connection
Expand Down Expand Up @@ -205,35 +204,16 @@ def close_idle_connections(self):
def ssl_enabled(self):
return self.config['security_protocol'] in ('SSL', 'SASL_SSL')

def _build_ssl_context(self):
if self.config['ssl_context'] is not None:
return self.config['ssl_context']
ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
ctx.minimum_version = ssl.TLSVersion.TLSv1_2
ctx.check_hostname = self.config['ssl_check_hostname']
if self.config['ssl_cafile']:
ctx.load_verify_locations(self.config['ssl_cafile'])
else:
ctx.load_default_certs()
if self.config['ssl_certfile']:
ctx.load_cert_chain(
certfile=self.config['ssl_certfile'],
keyfile=self.config['ssl_keyfile'],
password=self.config['ssl_password'],
)
if self.config['ssl_crlfile']:
ctx.load_verify_locations(crl=self.config['ssl_crlfile'])
ctx.verify_flags |= ssl.VERIFY_CRL_CHECK_LEAF
return ctx

async def _build_transport(self, node, timeout_at=None):
sock = await create_connection(self._net, node.host, node.port,
self.config['socket_options'],
proxy_url=self.config['proxy_url'],
timeout_at=timeout_at)
if self.ssl_enabled:
transport = KafkaSSLTransport(self._net, sock, self._build_ssl_context(),
host=node.host, ssl_check_hostname=self.config['ssl_check_hostname'])
ssl_configs = {key: value
for key, value in self.config.items()
if key.startswith('ssl_')}
transport = KafkaSSLTransport(self._net, sock, host=node.host, **ssl_configs)
else:
transport = KafkaTCPTransport(self._net, sock, host=node.host)

Expand Down
47 changes: 42 additions & 5 deletions kafka/net/transport.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from collections import deque
import copy
import logging
import selectors
import socket
Expand Down Expand Up @@ -371,13 +372,49 @@ def __str__(self):


class KafkaSSLTransport(KafkaTCPTransport):
def __init__(self, net, sock, ssl_context, host=None, ssl_check_hostname=False):
self._ssl_context = ssl_context
server_hostname = host if ssl_check_hostname else None
sock = ssl_context.wrap_socket(
sock, server_hostname=server_hostname, do_handshake_on_connect=False)
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, sock, host=None, **configs):
self.ssl_config = copy.copy(self.DEFAULT_CONFIG)
for key in self.ssl_config:
if key in configs:
self.ssl_config[key] = configs[key]
self._ssl_context = self._build_ssl_context(self.ssl_config)
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)

@staticmethod
def _build_ssl_context(config):
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

async def handshake(self):
while True:
try:
Expand Down
89 changes: 88 additions & 1 deletion test/net/test_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import kafka.errors as Errors
from kafka.future import Future
from kafka.net.selector import NetworkSelector, TaskState
from kafka.net.transport import KafkaTCPTransport
from kafka.net.transport import KafkaSSLTransport, KafkaTCPTransport


@pytest.fixture
Expand Down Expand Up @@ -367,6 +367,93 @@ def test_str_closed(self, net):
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 _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()
KafkaSSLTransport(net, sock, host='broker.example.com',
ssl_context=ctx, ssl_check_hostname=True)
_, kwargs = ctx.wrap_socket.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()
KafkaSSLTransport(net, sock, host='broker.example.com',
ssl_context=ctx, ssl_check_hostname=False)
_, kwargs = ctx.wrap_socket.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()
KafkaSSLTransport(net, sock, host='broker.example.com.',
ssl_context=ctx, ssl_check_hostname=False)
_, kwargs = ctx.wrap_socket.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, host=None, ssl_context=ctx)
_, kwargs = ctx.wrap_socket.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, host='broker.example.com', ssl_context=ctx)
_, 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, host='broker.example.com',
ssl_context=ctx)
assert t._ssl_context is ctx
assert t._sock is wrapped

def test_config_defaults_and_overrides(self, net):
sock, ctx, _ = self._make_ssl_sock()
t = KafkaSSLTransport(net, sock, host='h', ssl_context=ctx,
ssl_check_hostname=False)
# Explicitly-passed ssl_* keys land in ssl_config...
assert t.ssl_config['ssl_check_hostname'] is False
assert t.ssl_config['ssl_context'] is ctx
# ...unspecified keys keep their defaults.
assert t.ssl_config['ssl_cafile'] is None
assert t.ssl_config['ssl_check_hostname'] is not None


class TestBuildSSLContext:
def test_returns_provided_context(self):
ctx = MagicMock()
config = dict(KafkaSSLTransport.DEFAULT_CONFIG, ssl_context=ctx)
assert KafkaSSLTransport._build_ssl_context(config) is ctx

def test_check_hostname_propagates_to_context(self):
config = dict(KafkaSSLTransport.DEFAULT_CONFIG, 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(KafkaSSLTransport.DEFAULT_CONFIG, ssl_check_hostname=True)
ctx = KafkaSSLTransport._build_ssl_context(config)
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.
Expand Down