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)