From 3dbe6a542a0e4d7e928e1b2452b7a973061b6ef1 Mon Sep 17 00:00:00 2001 From: Michael Todorovic Date: Mon, 29 Mar 2021 08:07:48 +0200 Subject: [PATCH] fix: reload api if certificate changed on disk (#1887) This PR fixes #1886. We get the certificate serial number on server startup and store it in `api.__ssl_serial_number` On reload, we get again the serial number from disk and compare it to the one stored in `api.__ssl_serial_number`: if different, then the api will be reloaded (even if the config file didn't change) --- patroni/__init__.py | 1 + patroni/api.py | 23 ++++++++++++++++++++++- 2 files changed, 23 insertions(+), 1 deletion(-) diff --git a/patroni/__init__.py b/patroni/__init__.py index 4d3463a5..dc5823fc 100644 --- a/patroni/__init__.py +++ b/patroni/__init__.py @@ -74,6 +74,7 @@ class Patroni(AbstractPatroniDaemon): if local: self.tags = self.get_tags() self.request.reload_config(self.config) + if local or self.api.reload_local_certificate(): self.api.reload_config(self.config['restapi']) self.watchdog.reload_config(self.config) self.postgresql.reload_config(self.config['postgresql'], sighup) diff --git a/patroni/api.py b/patroni/api.py index 7bc9fc63..4fbc19e5 100644 --- a/patroni/api.py +++ b/patroni/api.py @@ -616,6 +616,8 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread): self.http_extra_headers = {} self.reload_config(config) self.daemon = True + self.__ssl_serial_number = None + self._received_new_cert = False def query(self, sql, *params): cursor = None @@ -701,6 +703,7 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread): self.__listen = listen self.__ssl_options = ssl_options + self._received_new_cert = False # reset to False after reload_config() self.__httpserver_init(host, port) Thread.__init__(self, target=self.serve_forever) @@ -723,6 +726,7 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread): ctx.verify_mode = modes[verify_client] 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) if reloading_config: self.start() @@ -750,6 +754,23 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread): _, request = request # SSLSocket return super(RestApiServer, self).shutdown_request(request) + def get_certificate_serial_number(self): + if self.__ssl_options.get('certfile'): + import ssl + try: + crt = ssl._ssl._test_decode_cert(self.__ssl_options['certfile']) + return crt.get('serialNumber') + except ssl.SSLError as e: + logger.error('Failed to get serial number from certificate %s: %r', self.__ssl_options['certfile'], e) + + def reload_local_certificate(self): + if self.__protocol == 'https': + on_disk_cert_serial_number = self.get_certificate_serial_number() + if on_disk_cert_serial_number != self.__ssl_serial_number: + self._received_new_cert = True + self.__ssl_serial_number = on_disk_cert_serial_number + return True + def reload_config(self, config): if 'listen' not in config: # changing config in runtime raise ValueError('Can not find "restapi.listen" config') @@ -763,7 +784,7 @@ class RestApiServer(ThreadingMixIn, HTTPServer, Thread): if isinstance(config.get('verify_client'), six.string_types): ssl_options['verify_client'] = config['verify_client'].lower() - if self.__listen != config['listen'] or self.__ssl_options != ssl_options: + if self.__listen != config['listen'] or self.__ssl_options != ssl_options or self._received_new_cert: self.__initialize(config['listen'], ssl_options) self.__auth_key = base64.b64encode(config['auth'].encode('utf-8')) if 'auth' in config else None