mirror of
https://github.com/outbackdingo/patroni.git
synced 2026-08-25 14:53:37 +00:00
Explicitly shut down SSL connection before socket (#2425)
This is handled by calling the unwrap() method on SSLSocket. In addition to that simplify code that handles deferred handshakes. Close https://github.com/zalando/patroni/issues/2424
This commit is contained in:
+7
-19
@@ -825,32 +825,20 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread):
|
||||
else:
|
||||
logger.error('Bad value in the "restapi.verify_client": %s', verify_client)
|
||||
self.__ssl_serial_number = self.get_certificate_serial_number()
|
||||
self.socket = ctx.wrap_socket(self.socket, server_side=True)
|
||||
self.socket = ctx.wrap_socket(self.socket, server_side=True, do_handshake_on_connect=False)
|
||||
if reloading_config:
|
||||
self.start()
|
||||
|
||||
def process_request_thread(self, request, client_address):
|
||||
if isinstance(request, tuple):
|
||||
sock, newsock = request
|
||||
try:
|
||||
request = sock.context.wrap_socket(newsock, do_handshake_on_connect=sock.do_handshake_on_connect,
|
||||
suppress_ragged_eofs=sock.suppress_ragged_eofs, server_side=True)
|
||||
except socket.error:
|
||||
return
|
||||
enable_keepalive(request, 10, 3)
|
||||
if hasattr(request, 'context'): # SSLSocket
|
||||
request.do_handshake()
|
||||
super(RestApiServer, self).process_request_thread(request, client_address)
|
||||
|
||||
def get_request(self):
|
||||
sock = self.socket
|
||||
newsock, addr = socket.socket.accept(sock)
|
||||
enable_keepalive(newsock, 10, 3)
|
||||
if hasattr(sock, 'context'): # SSLSocket, we want to do the deferred handshake from a thread
|
||||
newsock = (sock, newsock)
|
||||
return newsock, addr
|
||||
|
||||
def shutdown_request(self, request):
|
||||
if isinstance(request, tuple):
|
||||
_, request = request # SSLSocket
|
||||
return super(RestApiServer, self).shutdown_request(request)
|
||||
if hasattr(request, 'context'): # SSLSocket
|
||||
request.unwrap()
|
||||
super(RestApiServer, self).shutdown_request(request)
|
||||
|
||||
def get_certificate_serial_number(self):
|
||||
if self.__ssl_options.get('certfile'):
|
||||
|
||||
+5
-22
@@ -12,6 +12,7 @@ from patroni.ha import _MemberStatus
|
||||
from patroni.utils import tzutc
|
||||
from six import BytesIO as IO
|
||||
from six.moves import BaseHTTPServer
|
||||
from six.moves.socketserver import ThreadingMixIn
|
||||
from . import psycopg_connect, MockCursor
|
||||
from .test_ha import get_cluster_initialized_without_leader
|
||||
|
||||
@@ -586,32 +587,14 @@ class TestRestApiServer(unittest.TestCase):
|
||||
def test_socket_error(self):
|
||||
self.assertRaises(socket.error, MockRestApiServer, Mock(), '', {'listen': '*:8008'})
|
||||
|
||||
@patch.object(MockRestApiServer, 'finish_request', Mock())
|
||||
@patch.object(ThreadingMixIn, 'process_request_thread', Mock())
|
||||
def test_process_request_thread(self):
|
||||
mock_socket = Mock()
|
||||
self.srv.process_request_thread((mock_socket, 1), '2')
|
||||
mock_socket.context.wrap_socket.side_effect = socket.error
|
||||
self.srv.process_request_thread((mock_socket, 1), '2')
|
||||
|
||||
@patch.object(socket.socket, 'accept')
|
||||
def test_get_request(self, mock_accept):
|
||||
newsock = Mock()
|
||||
mock_accept.return_value = (newsock, '2')
|
||||
self.srv.socket = Mock()
|
||||
self.assertEqual(self.srv.get_request(), ((self.srv.socket, newsock), '2'))
|
||||
self.srv.process_request_thread(Mock(), '2')
|
||||
|
||||
@patch.object(MockRestApiServer, 'process_request', Mock(side_effect=RuntimeError))
|
||||
@patch.object(MockRestApiServer, 'get_request', Mock(return_value=(Mock(), ('127.0.0.1', 55555))))
|
||||
def test_process_request_error(self):
|
||||
mock_address = ('127.0.0.1', 55555)
|
||||
mock_socket = Mock()
|
||||
mock_ssl_socket = (Mock(), Mock())
|
||||
for mock_request in (mock_socket, mock_ssl_socket):
|
||||
with patch.object(
|
||||
MockRestApiServer,
|
||||
'get_request',
|
||||
Mock(return_value=(mock_request, mock_address))
|
||||
):
|
||||
self.srv._handle_request_noblock()
|
||||
self.srv._handle_request_noblock()
|
||||
|
||||
@patch('ssl._ssl._test_decode_cert', Mock())
|
||||
def test_reload_local_certificate(self):
|
||||
|
||||
Reference in New Issue
Block a user