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
+85 -57
View File
@@ -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)
+1
View File
@@ -1,3 +1,4 @@
urllib3>=1.9
boto
psycopg2>=2.6.1
PyYAML
+15 -5
View File
@@ -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)