diff --git a/patroni/api.py b/patroni/api.py index 2c10ac22..3cd005d1 100644 --- a/patroni/api.py +++ b/patroni/api.py @@ -590,6 +590,23 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread): 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 + super(RestApiServer, self).process_request_thread(request, client_address) + + def get_request(self): + sock = self.socket + newsock, addr = socket.socket.accept(sock) + if hasattr(sock, 'context'): # SSLSocket, we want to do the deferred handshake from a thread + newsock = (sock, newsock) + return newsock, addr + def reload_config(self, config): if 'listen' not in config: # changing config in runtime raise ValueError('Can not find "restapi.listen" config') diff --git a/tests/test_api.py b/tests/test_api.py index 0819c6c9..928def0c 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -415,25 +415,27 @@ class TestRestApiHandler(unittest.TestCase): MockRestApiServer(RestApiHandler, post + '37\n\n{"candidate":"2","scheduled_at": "1"}') -@patch('ssl.SSLContext.load_cert_chain', Mock()) -@patch('ssl.SSLContext.wrap_socket', Mock(return_value=0)) -@patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock()) class TestRestApiServer(unittest.TestCase): + @patch('ssl.SSLContext.load_cert_chain', Mock()) + @patch('ssl.SSLContext.wrap_socket', Mock(return_value=0)) + @patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock()) + def setUp(self): + self.srv = MockRestApiServer(Mock(), '', {'listen': '*:8008', 'certfile': 'a', 'verify_client': 'required'}) + + @patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock()) def test_reload_config(self): bad_config = {'listen': 'foo'} self.assertRaises(ValueError, MockRestApiServer, None, '', bad_config) - srv = MockRestApiServer(Mock(), '', {'listen': '*:8008', 'certfile': 'a', 'verify_client': 'required'}) - self.assertRaises(ValueError, srv.reload_config, bad_config) - self.assertRaises(ValueError, srv.reload_config, {}) + self.assertRaises(ValueError, self.srv.reload_config, bad_config) + self.assertRaises(ValueError, self.srv.reload_config, {}) with patch.object(socket.socket, 'setsockopt', Mock(side_effect=socket.error)): - srv.reload_config({'listen': ':8008'}) + self.srv.reload_config({'listen': ':8008'}) def test_check_auth(self): - srv = MockRestApiServer(Mock(), '', {'listen': '*:8008', 'certfile': 'a', 'verify_client': 'required'}) mock_rh = Mock() mock_rh.request.getpeercert.return_value = None - self.assertIsNot(srv.check_auth(mock_rh), True) + self.assertIsNot(self.srv.check_auth(mock_rh), True) def test_handle_error(self): try: @@ -441,6 +443,18 @@ class TestRestApiServer(unittest.TestCase): except Exception: self.assertIsNone(MockRestApiServer.handle_error(None, ('127.0.0.1', 55555))) + @patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock(side_effect=socket.error)) def test_socket_error(self): - with patch.object(BaseHTTPServer.HTTPServer, '__init__', Mock(side_effect=socket.error)): - self.assertRaises(socket.error, MockRestApiServer, Mock(), '', {'listen': '*:8008'}) + self.assertRaises(socket.error, MockRestApiServer, Mock(), '', {'listen': '*:8008'}) + + @patch.object(MockRestApiServer, 'finish_request', 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', Mock(return_value=(1, '2'))) + def test_get_request(self): + self.srv.socket = Mock() + self.assertEqual(self.srv.get_request(), ((self.srv.socket, 1), '2'))