Implement simple asynchronos dns-resolve cache (#360)

This commit is contained in:
Alexander Kukushkin
2016-12-07 13:16:26 +01:00
committed by GitHub
parent b38d98a6a3
commit ec78777778
2 changed files with 114 additions and 31 deletions
+94 -24
View File
@@ -14,8 +14,10 @@ from patroni.exceptions import DCSError
from patroni.utils import Retry, RetryFailedError
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
from six.moves.urllib_parse import urlparse, urlunparse
from threading import Thread
logger = logging.getLogger(__name__)
@@ -24,9 +26,64 @@ class EtcdError(DCSError):
pass
class DnsCachingResolver(Thread):
def __init__(self, cache_time=600.0, cache_fail_time=30.0):
super(DnsCachingResolver, self).__init__()
self._cache = {}
self._cache_time = cache_time
self._cache_fail_time = cache_fail_time
self._resolve_queue = Queue()
self.daemon = True
self.start()
def run(self):
while True:
hostname, attempt = self._resolve_queue.get()
ips = self._do_resolve(hostname)
if ips:
self._cache[hostname] = (time.time(), ips)
else:
if attempt < 10:
self.resolve_async(hostname, attempt + 1)
time.sleep(1)
def resolve(self, hostname):
current_time = time.time()
cached_time, ips = self._cache.get(hostname, (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
def resolve_async(self, hostname, attempt=0):
self._resolve_queue.put((hostname, attempt))
@staticmethod
def _do_resolve(hostname):
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)
except socket.gaierror:
logger.warning('failed to resolve host %s', hostname)
return []
class Client(etcd.Client):
def __init__(self, config):
def __init__(self, config, cache_ttl=300):
self._dns_resolver = DnsCachingResolver()
self._base_uri_unresolved = None
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',
'cert', 'ca_cert') if config.get(p)}
super(Client, self).__init__(read_timeout=config['retry_timeout'], **args)
@@ -51,6 +108,9 @@ class Client(etcd.Client):
kwargs['timeout'] = self.read_timeout/float(kwargs['retries'] + 1)/etcd_nodes
return kwargs
def set_machines_cache_ttl(self, cache_ttl):
self._machines_cache_ttl = cache_ttl
@property
def machines(self):
"""Original `machines` method(property) of `etcd.Client` class raise exception
@@ -71,20 +131,24 @@ class Client(etcd.Client):
machines = [n.strip() for n in self._handle_server_response(response).data.decode('utf-8').split(',')]
logger.debug("Retrieved list of machines: %s", machines)
random.shuffle(machines)
for url in machines:
r = urlparse(url)
self._dns_resolver.resolve_async(r.hostname)
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)
if self._machines_cache:
self._base_uri = self._machines_cache.pop(0)
try:
self._base_uri = self._next_server()
logger.info("Retrying on %s", self._base_uri)
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 []
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 []
def set_read_timeout(self, timeout):
self._read_timeout = timeout
@@ -105,6 +169,16 @@ 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 /')
@@ -120,7 +194,7 @@ class Client(etcd.Client):
raise etcd.EtcdException('HTTP method {0} not supported'.format(method))
# Update machines_cache if previous attempt of update has failed
if self._update_machines_cache:
if self._update_machines_cache or time.time() - self._machines_cache_updated > self._machines_cache_ttl:
self._load_machines_cache()
kwargs.update(self._build_request_parameters())
@@ -139,8 +213,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 in self._machines_cache:
self._machines_cache.remove(self._base_uri)
if self._base_uri_unresolved in self._machines_cache:
self._machines_cache.remove(self._base_uri_unresolved)
except etcd.EtcdConnectionFailed:
self._update_machines_cache = True
if not response:
@@ -184,13 +258,7 @@ 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 = []
try:
for r in set(socket.getaddrinfo(host, port, socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP)):
ret.append('{0}://{1}:{2}'.format(self.protocol, *r[4]))
except socket.error:
logger.exception('Can not resolve %s', host)
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)]
def _load_machines_cache(self):
@@ -220,13 +288,14 @@ class Client(etcd.Client):
raise etcd.EtcdException
# After filling up initial list of machines_cache we should ask etcd-cluster about actual list
self._base_uri = self._machines_cache.pop(0)
self._base_uri = self._next_server()
self._machines_cache = self.machines
if self._base_uri in self._machines_cache:
self._machines_cache.remove(self._base_uri)
if self._base_uri_unresolved in self._machines_cache:
self._machines_cache.remove(self._base_uri_unresolved)
self._update_machines_cache = False
self._machines_cache_updated = time.time()
def catch_etcd_errors(func):
@@ -295,6 +364,7 @@ class Etcd(AbstractDCS):
ttl = int(ttl)
self.__do_not_watch = self._ttl != ttl
self._ttl = ttl
self._client.set_machines_cache_ttl(ttl*10)
def set_retry_timeout(self, retry_timeout):
self._retry.deadline = retry_timeout
@@ -427,7 +497,7 @@ class Etcd(AbstractDCS):
self._client.http.clear()
return False
except etcd.EtcdException:
logging.exception('watch')
logger.exception('watch')
timeout = end_time - time.time()
+20 -7
View File
@@ -6,7 +6,7 @@ import unittest
from dns.exception import DNSException
from mock import Mock, patch
from patroni.dcs.etcd import AbstractDCS, Client, Cluster, Etcd, EtcdError
from patroni.dcs.etcd import AbstractDCS, Client, Cluster, Etcd, EtcdError, DnsCachingResolver
from patroni.exceptions import DCSError
from urllib3.exceptions import ReadTimeoutError
@@ -128,29 +128,40 @@ def dns_query(name, _):
def socket_getaddrinfo(*args):
if args[0] == 'ok':
return [(2, 1, 6, '', ('127.0.0.1', 2379)), (2, 1, 6, '', ('127.0.0.1', 2379))]
raise socket.error
if args[0] in ('ok', 'localhost', '127.0.0.1'):
return [(2, 1, 6, '', ('127.0.0.1', 0)), (10, 1, 6, '', ('::1', 0))]
raise socket.gaierror
def http_request(method, url, **kwargs):
if url == 'http://localhost:2379/timeout':
if url in ('http://127.0.0.1:2379/timeout', 'http://[::1]:2379/timeout'):
raise ReadTimeoutError(None, None, None)
if url == 'http://localhost:2379/v2/machines':
if url in ('http://127.0.0.1:2379/v2/machines', 'http://[::1]:2379/v2/machines'):
ret = MockResponse()
ret.content = 'http://localhost:2379,http://localhost:4001'
return ret
if url == 'http://localhost:2379/':
if url in ('http://127.0.0.1:2379/', 'http://[::1]:2379/'):
return MockResponse()
raise socket.error
class TestDnsCachingResolver(unittest.TestCase):
@patch('time.sleep', Mock(side_effect=SleepException))
@patch('socket.gethostbyname_ex', Mock(side_effect=socket.gaierror))
def test_run(self):
r = DnsCachingResolver()
self.assertIsNone(r.resolve_async(''))
r.join()
@patch('dns.resolver.query', dns_query)
@patch('socket.getaddrinfo', socket_getaddrinfo)
@patch('requests.get', requests_get)
class TestClient(unittest.TestCase):
@patch('dns.resolver.query', dns_query)
@patch('socket.getaddrinfo', socket_getaddrinfo)
@patch('requests.get', requests_get)
def setUp(self):
with patch.object(Client, 'machines') as mock_machines:
@@ -209,11 +220,13 @@ class TestClient(unittest.TestCase):
@patch('requests.get', requests_get)
@patch('socket.getaddrinfo', socket_getaddrinfo)
@patch.object(etcd.Client, 'write', etcd_write)
@patch.object(etcd.Client, 'read', etcd_read)
@patch.object(etcd.Client, 'delete', Mock(side_effect=etcd.EtcdException))
class TestEtcd(unittest.TestCase):
@patch('socket.getaddrinfo', socket_getaddrinfo)
def setUp(self):
with patch.object(Client, 'machines') as mock_machines:
mock_machines.__get__ = Mock(return_value=['http://localhost:2379', 'http://localhost:4001'])