mirror of
https://github.com/outbackdingo/patroni.git
synced 2026-08-25 14:53:37 +00:00
Implement simple asynchronos dns-resolve cache (#360)
This commit is contained in:
committed by
GitHub
parent
b38d98a6a3
commit
ec78777778
+94
-24
@@ -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
@@ -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'])
|
||||
|
||||
Reference in New Issue
Block a user