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:
Alexander Kukushkin
2017-01-18 13:46:02 +01:00
committed by GitHub
parent 2b5c08d17d
commit c6252bc004
3 changed files with 101 additions and 62 deletions
+80 -52
View File
@@ -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
View File
@@ -1,3 +1,4 @@
urllib3>=1.9
boto boto
psycopg2>=2.6.1 psycopg2>=2.6.1
PyYAML PyYAML
+15 -5
View File
@@ -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)