From c6252bc004ce56fd718df19e39aae8b938b1d8e2 Mon Sep 17 00:00:00 2001 From: Alexander Kukushkin Date: Wed, 18 Jan 2017 13:46:02 +0100 Subject: [PATCH] Don't resolve url hostnames manualy but mokey patch urllib3 (#385) Change hostnames by ip addresses was causing certificate verification to fail. Instead of doing it we will better monkey patch urllib3 functionality which does name resolution. It should work without problems even for https connection. --- patroni/dcs/etcd.py | 142 ++++++++++++++++++++++++++------------------ requirements.txt | 1 + tests/test_etcd.py | 20 +++++-- 3 files changed, 101 insertions(+), 62 deletions(-) diff --git a/patroni/dcs/etcd.py b/patroni/dcs/etcd.py index 91ad420c..4cb1f391 100644 --- a/patroni/dcs/etcd.py +++ b/patroni/dcs/etcd.py @@ -2,6 +2,7 @@ from __future__ import absolute_import import etcd import logging import os +import urllib3.util.connection import random import requests import socket @@ -16,7 +17,7 @@ from urllib3.exceptions import HTTPError, ReadTimeoutError from requests.exceptions import RequestException from six.moves.queue import Queue from six.moves.http_client import HTTPException -from six.moves.urllib_parse import urlparse, urlunparse +from six.moves.urllib_parse import urlparse from threading import Thread logger = logging.getLogger(__name__) @@ -39,49 +40,42 @@ class DnsCachingResolver(Thread): def run(self): while True: - hostname, attempt = self._resolve_queue.get() - ips = self._do_resolve(hostname) - if ips: - self._cache[hostname] = (time.time(), ips) + (host, port), attempt = self._resolve_queue.get() + response = self._do_resolve(host, port) + if response: + self._cache[(host, port)] = (time.time(), response) else: if attempt < 10: - self.resolve_async(hostname, attempt + 1) + self.resolve_async(host, port, attempt + 1) time.sleep(1) - def resolve(self, hostname): + def resolve(self, host, port): current_time = time.time() - cached_time, ips = self._cache.get(hostname, (0, [])) + cached_time, response = self._cache.get((host, port), (0, [])) time_passed = current_time - cached_time - if time_passed > self._cache_time or (not ips and time_passed > self._cache_fail_time): - new_ips = self._do_resolve(hostname) - if new_ips: - self._cache[hostname] = (current_time, new_ips) - ips = new_ips - return ips + if time_passed > self._cache_time or (not response and time_passed > self._cache_fail_time): + new_response = self._do_resolve(host, port) + if new_response: + self._cache[(host, port)] = (current_time, new_response) + response = new_response + return response - def resolve_async(self, hostname, attempt=0): - self._resolve_queue.put((hostname, attempt)) + def resolve_async(self, host, port, attempt=0): + self._resolve_queue.put(((host, port), attempt)) @staticmethod - def _do_resolve(hostname): + def _do_resolve(host, port): try: - ret = set() - for r in socket.getaddrinfo(hostname, 0, 0, 0, socket.IPPROTO_TCP): - if r[0] == socket.AF_INET6: - ret.add('[{0}]'.format(r[4][0])) - else: - ret.add(r[4][0]) - return list(ret) + return socket.getaddrinfo(host, port, 0, socket.SOCK_STREAM, socket.IPPROTO_TCP) except socket.gaierror: - logger.warning('failed to resolve host %s', hostname) + logger.warning('failed to resolve host %s', host) return [] class Client(etcd.Client): - def __init__(self, config, cache_ttl=300): - self._dns_resolver = DnsCachingResolver() - self._base_uri_unresolved = None + def __init__(self, config, dns_resolver, cache_ttl=300): + self._dns_resolver = dns_resolver self.set_machines_cache_ttl(cache_ttl) self._machines_cache_updated = 0 args = {p: config.get(p) for p in ('host', 'port', 'protocol', 'use_proxies', 'username', 'password', @@ -133,22 +127,22 @@ class Client(etcd.Client): random.shuffle(machines) for url in machines: r = urlparse(url) - self._dns_resolver.resolve_async(r.hostname) + port = r.port or (443 if r.scheme == 'https' else 80) + self._dns_resolver.resolve_async(r.hostname, port) return machines except Exception as e: # We can't get the list of machines, if one server is in the # machines cache, try on it logger.error("Failed to get list of machines from %s%s: %r", self._base_uri, self.version_prefix, e) - try: - self._base_uri = self._next_server() + if self._machines_cache: + self._base_uri = self._machines_cache.pop(0) logger.info("Retrying on %s", self._base_uri) - except etcd.EtcdConnectionFailed: - if self._update_machines_cache: - raise etcd.EtcdException("Could not get the list of servers, " - "maybe you provided the wrong " - "host(s) to connect to?") - else: - return [] + elif self._update_machines_cache: + raise etcd.EtcdException("Could not get the list of servers, " + "maybe you provided the wrong " + "host(s) to connect to?") + else: + return [] def set_read_timeout(self, timeout): self._read_timeout = timeout @@ -169,16 +163,6 @@ class Client(etcd.Client): response = False return response - def _next_server(self, cause=None): - while True: - url = super(Client, self)._next_server(cause) - r = urlparse(url) - ips = self._dns_resolver.resolve(r.hostname) - if ips: - netloc = '{0}:{1}'.format(random.choice(ips), r.port) - self._base_uri_unresolved = url - return urlunparse((r.scheme, netloc, r.path, r.params, r.query, r.fragment)) - def api_execute(self, path, method, params=None, timeout=None): if not path.startswith('/'): raise ValueError('Path does not start with /') @@ -198,8 +182,8 @@ class Client(etcd.Client): self._load_machines_cache() elif time.time() - self._machines_cache_updated > self._machines_cache_ttl: self._machines_cache = self.machines - if self._base_uri_unresolved in self._machines_cache: - self._machines_cache.remove(self._base_uri_unresolved) + if self._base_uri in self._machines_cache: + self._machines_cache.remove(self._base_uri) self._machines_cache_updated = time.time() kwargs.update(self._build_request_parameters()) @@ -218,8 +202,8 @@ class Client(etcd.Client): some_request_failed = True if some_request_failed and not self._use_proxies: self._machines_cache = self.machines - if self._base_uri_unresolved in self._machines_cache: - self._machines_cache.remove(self._base_uri_unresolved) + if self._base_uri in self._machines_cache: + self._machines_cache.remove(self._base_uri) except etcd.EtcdConnectionFailed: self._update_machines_cache = True if not response: @@ -264,8 +248,16 @@ class Client(etcd.Client): def _get_machines_cache_from_dns(self, host, port): """One host might be resolved into multiple ip addresses. We will make list out of it""" - ret = ['{0}://{1}:{2}'.format(self.protocol, ip, port) for ip in set(self._dns_resolver.resolve(host))] - return list(set(ret)) if ret else ['{0}://{1}:{2}'.format(self.protocol, host, port)] + if self.protocol == 'http': + ret = [] + for af, _, _, _, sa in self._dns_resolver.resolve(host, port): + host, port = sa[:2] + if af == socket.AF_INET6: + host = '[{0}]'.format(host) + ret.append('{0}://{1}:{2}'.format(self.protocol, host, port)) + if ret: + return list(set(ret)) + return ['{0}://{1}:{2}'.format(self.protocol, host, port)] def _load_machines_cache(self): """This method should fill up `_machines_cache` from scratch. @@ -297,8 +289,8 @@ class Client(etcd.Client): self._base_uri = self._next_server() self._machines_cache = self.machines - if self._base_uri_unresolved in self._machines_cache: - self._machines_cache.remove(self._base_uri_unresolved) + if self._base_uri in self._machines_cache: + self._machines_cache.remove(self._base_uri) self._update_machines_cache = False self._machines_cache_updated = time.time() @@ -357,10 +349,46 @@ class Etcd(AbstractDCS): for p in ('discovery_srv', 'srv_domain'): if p in config: config['srv'] = config.pop(p) + + dns_resolver = DnsCachingResolver() + + def create_connection_patched(address, timeout=socket._GLOBAL_DEFAULT_TIMEOUT, + source_address=None, socket_options=None): + host, port = address + if host.startswith('['): + host = host.strip('[]') + err = None + for af, socktype, proto, _, sa in dns_resolver.resolve(host, port): + sock = None + try: + sock = socket.socket(af, socktype, proto) + if socket_options: + for opt in socket_options: + sock.setsockopt(*opt) + if timeout is not socket._GLOBAL_DEFAULT_TIMEOUT: + sock.settimeout(timeout) + if source_address: + sock.bind(source_address) + sock.connect(sa) + return sock + + except socket.error as e: + err = e + if sock is not None: + sock.close() + sock = None + + if err is not None: + raise err + + raise socket.error("getaddrinfo returns an empty list") + + urllib3.util.connection.create_connection = create_connection_patched + client = None while not client: try: - client = Client(config) + client = Client(config, dns_resolver) except etcd.EtcdException: logger.info('waiting on etcd') time.sleep(5) diff --git a/requirements.txt b/requirements.txt index 7835ed13..7063351b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,4 @@ +urllib3>=1.9 boto psycopg2>=2.6.1 PyYAML diff --git a/tests/test_etcd.py b/tests/test_etcd.py index 825cab0d..3dd93da6 100644 --- a/tests/test_etcd.py +++ b/tests/test_etcd.py @@ -1,5 +1,6 @@ import etcd import json +import urllib3.util.connection import requests import socket import unittest @@ -134,13 +135,13 @@ def socket_getaddrinfo(*args): def http_request(method, url, **kwargs): - if url in ('http://127.0.0.1:2379/timeout', 'http://[::1]:2379/timeout'): + if url == 'http://localhost:2379/timeout': raise ReadTimeoutError(None, None, None) - if url in ('http://127.0.0.1:2379/v2/machines', 'http://[::1]:2379/v2/machines'): + if url == 'http://localhost:2379/v2/machines': ret = MockResponse() ret.content = 'http://localhost:2379,http://localhost:4001' return ret - if url in ('http://127.0.0.1:2379/', 'http://[::1]:2379/'): + if url == 'http://localhost:2379/': return MockResponse() raise socket.error @@ -151,7 +152,7 @@ class TestDnsCachingResolver(unittest.TestCase): @patch('socket.getaddrinfo', Mock(side_effect=socket.gaierror)) def test_run(self): r = DnsCachingResolver() - self.assertIsNone(r.resolve_async('')) + self.assertIsNone(r.resolve_async('', 0)) r.join() @@ -166,7 +167,7 @@ class TestClient(unittest.TestCase): def setUp(self): with patch.object(Client, 'machines') as mock_machines: mock_machines.__get__ = Mock(return_value=['http://localhost:2379', 'http://localhost:4001']) - self.client = Client({'srv': 'test', 'retry_timeout': 3}) + self.client = Client({'srv': 'test', 'retry_timeout': 3}, DnsCachingResolver()) self.client.http.request = http_request self.client.http.request_encode_body = http_request @@ -221,6 +222,15 @@ class TestClient(unittest.TestCase): self.client._config = {'srv': 'blabla'} self.assertRaises(etcd.EtcdException, self.client._load_machines_cache) + @patch.object(socket.socket, 'connect') + def test_create_connection_patched(self, mock_connect): + self.assertRaises(socket.error, urllib3.util.connection.create_connection, ('fail', 2379)) + urllib3.util.connection.create_connection(('[localhost]', 2379)) + mock_connect.side_effect = socket.error + self.assertRaises(socket.error, urllib3.util.connection.create_connection, ('[localhost]', 2379), + timeout=1, source_address=('localhost', 53333), + socket_options=[(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)]) + @patch('requests.get', requests_get) @patch('socket.getaddrinfo', socket_getaddrinfo)