mirror of
https://github.com/outbackdingo/patroni.git
synced 2026-08-25 14:53:37 +00:00
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.
This commit is contained in:
committed by
GitHub
parent
2b5c08d17d
commit
c6252bc004
+80
-52
@@ -2,6 +2,7 @@ from __future__ import absolute_import
|
|||||||
import etcd
|
import etcd
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import urllib3.util.connection
|
||||||
import random
|
import random
|
||||||
import requests
|
import requests
|
||||||
import socket
|
import socket
|
||||||
@@ -16,7 +17,7 @@ from urllib3.exceptions import HTTPError, ReadTimeoutError
|
|||||||
from requests.exceptions import RequestException
|
from requests.exceptions import RequestException
|
||||||
from six.moves.queue import Queue
|
from six.moves.queue import Queue
|
||||||
from six.moves.http_client import HTTPException
|
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
|
from threading import Thread
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -39,49 +40,42 @@ class DnsCachingResolver(Thread):
|
|||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
while True:
|
while True:
|
||||||
hostname, attempt = self._resolve_queue.get()
|
(host, port), attempt = self._resolve_queue.get()
|
||||||
ips = self._do_resolve(hostname)
|
response = self._do_resolve(host, port)
|
||||||
if ips:
|
if response:
|
||||||
self._cache[hostname] = (time.time(), ips)
|
self._cache[(host, port)] = (time.time(), response)
|
||||||
else:
|
else:
|
||||||
if attempt < 10:
|
if attempt < 10:
|
||||||
self.resolve_async(hostname, attempt + 1)
|
self.resolve_async(host, port, attempt + 1)
|
||||||
time.sleep(1)
|
time.sleep(1)
|
||||||
|
|
||||||
def resolve(self, hostname):
|
def resolve(self, host, port):
|
||||||
current_time = time.time()
|
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
|
time_passed = current_time - cached_time
|
||||||
if time_passed > self._cache_time or (not ips and time_passed > self._cache_fail_time):
|
if time_passed > self._cache_time or (not response and time_passed > self._cache_fail_time):
|
||||||
new_ips = self._do_resolve(hostname)
|
new_response = self._do_resolve(host, port)
|
||||||
if new_ips:
|
if new_response:
|
||||||
self._cache[hostname] = (current_time, new_ips)
|
self._cache[(host, port)] = (current_time, new_response)
|
||||||
ips = new_ips
|
response = new_response
|
||||||
return ips
|
return response
|
||||||
|
|
||||||
def resolve_async(self, hostname, attempt=0):
|
def resolve_async(self, host, port, attempt=0):
|
||||||
self._resolve_queue.put((hostname, attempt))
|
self._resolve_queue.put(((host, port), attempt))
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _do_resolve(hostname):
|
def _do_resolve(host, port):
|
||||||
try:
|
try:
|
||||||
ret = set()
|
return socket.getaddrinfo(host, port, 0, socket.SOCK_STREAM, socket.IPPROTO_TCP)
|
||||||
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:
|
except socket.gaierror:
|
||||||
logger.warning('failed to resolve host %s', hostname)
|
logger.warning('failed to resolve host %s', host)
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
class Client(etcd.Client):
|
class Client(etcd.Client):
|
||||||
|
|
||||||
def __init__(self, config, cache_ttl=300):
|
def __init__(self, config, dns_resolver, cache_ttl=300):
|
||||||
self._dns_resolver = DnsCachingResolver()
|
self._dns_resolver = dns_resolver
|
||||||
self._base_uri_unresolved = None
|
|
||||||
self.set_machines_cache_ttl(cache_ttl)
|
self.set_machines_cache_ttl(cache_ttl)
|
||||||
self._machines_cache_updated = 0
|
self._machines_cache_updated = 0
|
||||||
args = {p: config.get(p) for p in ('host', 'port', 'protocol', 'use_proxies', 'username', 'password',
|
args = {p: config.get(p) for p in ('host', 'port', 'protocol', 'use_proxies', 'username', 'password',
|
||||||
@@ -133,17 +127,17 @@ class Client(etcd.Client):
|
|||||||
random.shuffle(machines)
|
random.shuffle(machines)
|
||||||
for url in machines:
|
for url in machines:
|
||||||
r = urlparse(url)
|
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
|
return machines
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# We can't get the list of machines, if one server is in the
|
# We can't get the list of machines, if one server is in the
|
||||||
# machines cache, try on it
|
# machines cache, try on it
|
||||||
logger.error("Failed to get list of machines from %s%s: %r", self._base_uri, self.version_prefix, e)
|
logger.error("Failed to get list of machines from %s%s: %r", self._base_uri, self.version_prefix, e)
|
||||||
try:
|
if self._machines_cache:
|
||||||
self._base_uri = self._next_server()
|
self._base_uri = self._machines_cache.pop(0)
|
||||||
logger.info("Retrying on %s", self._base_uri)
|
logger.info("Retrying on %s", self._base_uri)
|
||||||
except etcd.EtcdConnectionFailed:
|
elif self._update_machines_cache:
|
||||||
if self._update_machines_cache:
|
|
||||||
raise etcd.EtcdException("Could not get the list of servers, "
|
raise etcd.EtcdException("Could not get the list of servers, "
|
||||||
"maybe you provided the wrong "
|
"maybe you provided the wrong "
|
||||||
"host(s) to connect to?")
|
"host(s) to connect to?")
|
||||||
@@ -169,16 +163,6 @@ class Client(etcd.Client):
|
|||||||
response = False
|
response = False
|
||||||
return response
|
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):
|
def api_execute(self, path, method, params=None, timeout=None):
|
||||||
if not path.startswith('/'):
|
if not path.startswith('/'):
|
||||||
raise ValueError('Path does not start with /')
|
raise ValueError('Path does not start with /')
|
||||||
@@ -198,8 +182,8 @@ class Client(etcd.Client):
|
|||||||
self._load_machines_cache()
|
self._load_machines_cache()
|
||||||
elif time.time() - self._machines_cache_updated > self._machines_cache_ttl:
|
elif time.time() - self._machines_cache_updated > self._machines_cache_ttl:
|
||||||
self._machines_cache = self.machines
|
self._machines_cache = self.machines
|
||||||
if self._base_uri_unresolved in self._machines_cache:
|
if self._base_uri in self._machines_cache:
|
||||||
self._machines_cache.remove(self._base_uri_unresolved)
|
self._machines_cache.remove(self._base_uri)
|
||||||
self._machines_cache_updated = time.time()
|
self._machines_cache_updated = time.time()
|
||||||
|
|
||||||
kwargs.update(self._build_request_parameters())
|
kwargs.update(self._build_request_parameters())
|
||||||
@@ -218,8 +202,8 @@ class Client(etcd.Client):
|
|||||||
some_request_failed = True
|
some_request_failed = True
|
||||||
if some_request_failed and not self._use_proxies:
|
if some_request_failed and not self._use_proxies:
|
||||||
self._machines_cache = self.machines
|
self._machines_cache = self.machines
|
||||||
if self._base_uri_unresolved in self._machines_cache:
|
if self._base_uri in self._machines_cache:
|
||||||
self._machines_cache.remove(self._base_uri_unresolved)
|
self._machines_cache.remove(self._base_uri)
|
||||||
except etcd.EtcdConnectionFailed:
|
except etcd.EtcdConnectionFailed:
|
||||||
self._update_machines_cache = True
|
self._update_machines_cache = True
|
||||||
if not response:
|
if not response:
|
||||||
@@ -264,8 +248,16 @@ class Client(etcd.Client):
|
|||||||
|
|
||||||
def _get_machines_cache_from_dns(self, host, port):
|
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"""
|
"""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))]
|
if self.protocol == 'http':
|
||||||
return list(set(ret)) if ret else ['{0}://{1}:{2}'.format(self.protocol, host, port)]
|
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):
|
def _load_machines_cache(self):
|
||||||
"""This method should fill up `_machines_cache` from scratch.
|
"""This method should fill up `_machines_cache` from scratch.
|
||||||
@@ -297,8 +289,8 @@ class Client(etcd.Client):
|
|||||||
self._base_uri = self._next_server()
|
self._base_uri = self._next_server()
|
||||||
self._machines_cache = self.machines
|
self._machines_cache = self.machines
|
||||||
|
|
||||||
if self._base_uri_unresolved in self._machines_cache:
|
if self._base_uri in self._machines_cache:
|
||||||
self._machines_cache.remove(self._base_uri_unresolved)
|
self._machines_cache.remove(self._base_uri)
|
||||||
|
|
||||||
self._update_machines_cache = False
|
self._update_machines_cache = False
|
||||||
self._machines_cache_updated = time.time()
|
self._machines_cache_updated = time.time()
|
||||||
@@ -357,10 +349,46 @@ class Etcd(AbstractDCS):
|
|||||||
for p in ('discovery_srv', 'srv_domain'):
|
for p in ('discovery_srv', 'srv_domain'):
|
||||||
if p in config:
|
if p in config:
|
||||||
config['srv'] = config.pop(p)
|
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
|
client = None
|
||||||
while not client:
|
while not client:
|
||||||
try:
|
try:
|
||||||
client = Client(config)
|
client = Client(config, dns_resolver)
|
||||||
except etcd.EtcdException:
|
except etcd.EtcdException:
|
||||||
logger.info('waiting on etcd')
|
logger.info('waiting on etcd')
|
||||||
time.sleep(5)
|
time.sleep(5)
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
urllib3>=1.9
|
||||||
boto
|
boto
|
||||||
psycopg2>=2.6.1
|
psycopg2>=2.6.1
|
||||||
PyYAML
|
PyYAML
|
||||||
|
|||||||
+15
-5
@@ -1,5 +1,6 @@
|
|||||||
import etcd
|
import etcd
|
||||||
import json
|
import json
|
||||||
|
import urllib3.util.connection
|
||||||
import requests
|
import requests
|
||||||
import socket
|
import socket
|
||||||
import unittest
|
import unittest
|
||||||
@@ -134,13 +135,13 @@ def socket_getaddrinfo(*args):
|
|||||||
|
|
||||||
|
|
||||||
def http_request(method, url, **kwargs):
|
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)
|
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 = MockResponse()
|
||||||
ret.content = 'http://localhost:2379,http://localhost:4001'
|
ret.content = 'http://localhost:2379,http://localhost:4001'
|
||||||
return ret
|
return ret
|
||||||
if url in ('http://127.0.0.1:2379/', 'http://[::1]:2379/'):
|
if url == 'http://localhost:2379/':
|
||||||
return MockResponse()
|
return MockResponse()
|
||||||
raise socket.error
|
raise socket.error
|
||||||
|
|
||||||
@@ -151,7 +152,7 @@ class TestDnsCachingResolver(unittest.TestCase):
|
|||||||
@patch('socket.getaddrinfo', Mock(side_effect=socket.gaierror))
|
@patch('socket.getaddrinfo', Mock(side_effect=socket.gaierror))
|
||||||
def test_run(self):
|
def test_run(self):
|
||||||
r = DnsCachingResolver()
|
r = DnsCachingResolver()
|
||||||
self.assertIsNone(r.resolve_async(''))
|
self.assertIsNone(r.resolve_async('', 0))
|
||||||
r.join()
|
r.join()
|
||||||
|
|
||||||
|
|
||||||
@@ -166,7 +167,7 @@ class TestClient(unittest.TestCase):
|
|||||||
def setUp(self):
|
def setUp(self):
|
||||||
with patch.object(Client, 'machines') as mock_machines:
|
with patch.object(Client, 'machines') as mock_machines:
|
||||||
mock_machines.__get__ = Mock(return_value=['http://localhost:2379', 'http://localhost:4001'])
|
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 = http_request
|
||||||
self.client.http.request_encode_body = http_request
|
self.client.http.request_encode_body = http_request
|
||||||
|
|
||||||
@@ -221,6 +222,15 @@ class TestClient(unittest.TestCase):
|
|||||||
self.client._config = {'srv': 'blabla'}
|
self.client._config = {'srv': 'blabla'}
|
||||||
self.assertRaises(etcd.EtcdException, self.client._load_machines_cache)
|
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('requests.get', requests_get)
|
||||||
@patch('socket.getaddrinfo', socket_getaddrinfo)
|
@patch('socket.getaddrinfo', socket_getaddrinfo)
|
||||||
|
|||||||
Reference in New Issue
Block a user