diff --git a/patroni/api.py b/patroni/api.py index d718a6e9..5acd1fa1 100644 --- a/patroni/api.py +++ b/patroni/api.py @@ -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'): diff --git a/tests/test_api.py b/tests/test_api.py index 73f81412..641871af 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -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):