Tornado update
This commit is contained in:
@@ -1,14 +1,14 @@
|
||||
#!/usr/bin/env python
|
||||
from __future__ import absolute_import, division, with_statement
|
||||
from __future__ import absolute_import, division, print_function, with_statement
|
||||
|
||||
from tornado.escape import utf8, _unicode, native_str
|
||||
from tornado.httpclient import HTTPRequest, HTTPResponse, HTTPError, AsyncHTTPClient, main, _RequestProxy
|
||||
from tornado.httputil import HTTPHeaders
|
||||
from tornado.iostream import IOStream, SSLIOStream
|
||||
from tornado.netutil import Resolver
|
||||
from tornado.netutil import Resolver, OverrideResolver
|
||||
from tornado.log import gen_log
|
||||
from tornado import stack_context
|
||||
from tornado.util import b, GzipDecompressor
|
||||
from tornado.util import GzipDecompressor
|
||||
|
||||
import base64
|
||||
import collections
|
||||
@@ -17,9 +17,8 @@ import functools
|
||||
import os.path
|
||||
import re
|
||||
import socket
|
||||
import ssl
|
||||
import sys
|
||||
import time
|
||||
import urlparse
|
||||
|
||||
try:
|
||||
from io import BytesIO # python 3
|
||||
@@ -27,9 +26,9 @@ except ImportError:
|
||||
from cStringIO import StringIO as BytesIO # python 2
|
||||
|
||||
try:
|
||||
import ssl # python 2.6+
|
||||
import urlparse # py2
|
||||
except ImportError:
|
||||
ssl = None
|
||||
import urllib.parse as urlparse # py3
|
||||
|
||||
_DEFAULT_CA_CERTS = os.path.dirname(__file__) + '/ca-certificates.crt'
|
||||
|
||||
@@ -45,12 +44,8 @@ class SimpleAsyncHTTPClient(AsyncHTTPClient):
|
||||
supported. In particular, proxies are not supported, connections
|
||||
are not reused, and callers cannot select the network interface to be
|
||||
used.
|
||||
|
||||
Python 2.6 or higher is required for HTTPS support. Users of Python 2.5
|
||||
should use the curl-based AsyncHTTPClient if HTTPS support is required.
|
||||
|
||||
"""
|
||||
def initialize(self, io_loop=None, max_clients=10,
|
||||
def initialize(self, io_loop, max_clients=10,
|
||||
hostname_mapping=None, max_buffer_size=104857600,
|
||||
resolver=None, defaults=None):
|
||||
"""Creates a AsyncHTTPClient.
|
||||
@@ -72,32 +67,24 @@ class SimpleAsyncHTTPClient(AsyncHTTPClient):
|
||||
max_buffer_size is the number of bytes that can be read by IOStream. It
|
||||
defaults to 100mb.
|
||||
"""
|
||||
self.io_loop = io_loop
|
||||
super(SimpleAsyncHTTPClient, self).initialize(io_loop,
|
||||
defaults=defaults)
|
||||
self.max_clients = max_clients
|
||||
self.queue = collections.deque()
|
||||
self.active = {}
|
||||
self.hostname_mapping = hostname_mapping
|
||||
self.max_buffer_size = max_buffer_size
|
||||
self.resolver = resolver or Resolver(io_loop=io_loop)
|
||||
self.defaults = dict(HTTPRequest._DEFAULTS)
|
||||
if defaults is not None:
|
||||
self.defaults.update(defaults)
|
||||
if hostname_mapping is not None:
|
||||
self.resolver = OverrideResolver(resolver=self.resolver,
|
||||
mapping=hostname_mapping)
|
||||
|
||||
def fetch(self, request, callback, **kwargs):
|
||||
if not isinstance(request, HTTPRequest):
|
||||
request = HTTPRequest(url=request, **kwargs)
|
||||
# We're going to modify this (to add Host, Accept-Encoding, etc),
|
||||
# so make sure we don't modify the caller's object. This is also
|
||||
# where normal dicts get converted to HTTPHeaders objects.
|
||||
request.headers = HTTPHeaders(request.headers)
|
||||
request = _RequestProxy(request, self.defaults)
|
||||
callback = stack_context.wrap(callback)
|
||||
def fetch_impl(self, request, callback):
|
||||
self.queue.append((request, callback))
|
||||
self._process_queue()
|
||||
if self.queue:
|
||||
gen_log.debug("max_clients limit reached, request queued. "
|
||||
"%d active, %d queued requests." % (
|
||||
len(self.active), len(self.queue)))
|
||||
len(self.active), len(self.queue)))
|
||||
|
||||
def _process_queue(self):
|
||||
with stack_context.NullContext():
|
||||
@@ -108,7 +95,7 @@ class SimpleAsyncHTTPClient(AsyncHTTPClient):
|
||||
_HTTPConnection(self.io_loop, self, request,
|
||||
functools.partial(self._release_fetch, key),
|
||||
callback,
|
||||
self.max_buffer_size)
|
||||
self.max_buffer_size, self.resolver)
|
||||
|
||||
def _release_fetch(self, key):
|
||||
del self.active[key]
|
||||
@@ -119,7 +106,7 @@ class _HTTPConnection(object):
|
||||
_SUPPORTED_METHODS = set(["GET", "HEAD", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"])
|
||||
|
||||
def __init__(self, io_loop, client, request, release_callback,
|
||||
final_callback, max_buffer_size):
|
||||
final_callback, max_buffer_size, resolver):
|
||||
self.start_time = io_loop.time()
|
||||
self.io_loop = io_loop
|
||||
self.client = client
|
||||
@@ -127,6 +114,7 @@ class _HTTPConnection(object):
|
||||
self.release_callback = release_callback
|
||||
self.final_callback = final_callback
|
||||
self.max_buffer_size = max_buffer_size
|
||||
self.resolver = resolver
|
||||
self.code = None
|
||||
self.headers = None
|
||||
self.chunks = None
|
||||
@@ -135,9 +123,6 @@ class _HTTPConnection(object):
|
||||
self._timeout = None
|
||||
with stack_context.ExceptionStackContext(self._handle_exception):
|
||||
self.parsed = urlparse.urlsplit(_unicode(self.request.url))
|
||||
if ssl is None and self.parsed.scheme == "https":
|
||||
raise ValueError("HTTPS requires either python2.6+ or "
|
||||
"curl_httpclient")
|
||||
if self.parsed.scheme not in ("http", "https"):
|
||||
raise ValueError("Unsupported url scheme: %s" %
|
||||
self.request.url)
|
||||
@@ -157,8 +142,6 @@ class _HTTPConnection(object):
|
||||
# raw ipv6 addresses in urls are enclosed in brackets
|
||||
host = host[1:-1]
|
||||
self.parsed_hostname = host # save final host for _on_connect
|
||||
if self.client.hostname_mapping is not None:
|
||||
host = self.client.hostname_mapping.get(host, host)
|
||||
|
||||
if request.allow_ipv6:
|
||||
af = socket.AF_UNSPEC
|
||||
@@ -167,7 +150,7 @@ class _HTTPConnection(object):
|
||||
# so restrict to ipv4 by default.
|
||||
af = socket.AF_INET
|
||||
|
||||
self.client.resolver.getaddrinfo(
|
||||
self.resolver.getaddrinfo(
|
||||
host, port, af, socket.SOCK_STREAM, 0, 0,
|
||||
callback=self._on_resolve)
|
||||
|
||||
@@ -220,31 +203,29 @@ class _HTTPConnection(object):
|
||||
self.start_time + timeout,
|
||||
stack_context.wrap(self._on_timeout))
|
||||
self.stream.set_close_callback(self._on_close)
|
||||
self.stream.connect(sockaddr, self._on_connect)
|
||||
# ipv6 addresses are broken (in self.parsed.hostname) until
|
||||
# 2.7, here is correctly parsed value calculated in __init__
|
||||
self.stream.connect(sockaddr, self._on_connect,
|
||||
server_hostname=self.parsed_hostname)
|
||||
|
||||
def _on_timeout(self):
|
||||
self._timeout = None
|
||||
if self.final_callback is not None:
|
||||
raise HTTPError(599, "Timeout")
|
||||
|
||||
def _on_connect(self):
|
||||
def _remove_timeout(self):
|
||||
if self._timeout is not None:
|
||||
self.io_loop.remove_timeout(self._timeout)
|
||||
self._timeout = None
|
||||
|
||||
def _on_connect(self):
|
||||
self._remove_timeout()
|
||||
if self.request.request_timeout:
|
||||
self._timeout = self.io_loop.add_timeout(
|
||||
self.start_time + self.request.request_timeout,
|
||||
stack_context.wrap(self._on_timeout))
|
||||
if (self.request.validate_cert and
|
||||
isinstance(self.stream, SSLIOStream)):
|
||||
match_hostname(self.stream.socket.getpeercert(),
|
||||
# ipv6 addresses are broken (in
|
||||
# self.parsed.hostname) until 2.7, here is
|
||||
# correctly parsed value calculated in
|
||||
# __init__
|
||||
self.parsed_hostname)
|
||||
if (self.request.method not in self._SUPPORTED_METHODS and
|
||||
not self.request.allow_nonstandard_methods):
|
||||
not self.request.allow_nonstandard_methods):
|
||||
raise KeyError("unknown method %s" % self.request.method)
|
||||
for key in ('network_interface',
|
||||
'proxy_host', 'proxy_port',
|
||||
@@ -265,8 +246,8 @@ class _HTTPConnection(object):
|
||||
username = self.request.auth_username
|
||||
password = self.request.auth_password or ''
|
||||
if username is not None:
|
||||
auth = utf8(username) + b(":") + utf8(password)
|
||||
self.request.headers["Authorization"] = (b("Basic ") +
|
||||
auth = utf8(username) + b":" + utf8(password)
|
||||
self.request.headers["Authorization"] = (b"Basic " +
|
||||
base64.b64encode(auth))
|
||||
if self.request.user_agent:
|
||||
self.request.headers["User-Agent"] = self.request.user_agent
|
||||
@@ -277,25 +258,25 @@ class _HTTPConnection(object):
|
||||
assert self.request.body is None
|
||||
if self.request.body is not None:
|
||||
self.request.headers["Content-Length"] = str(len(
|
||||
self.request.body))
|
||||
self.request.body))
|
||||
if (self.request.method == "POST" and
|
||||
"Content-Type" not in self.request.headers):
|
||||
"Content-Type" not in self.request.headers):
|
||||
self.request.headers["Content-Type"] = "application/x-www-form-urlencoded"
|
||||
if self.request.use_gzip:
|
||||
self.request.headers["Accept-Encoding"] = "gzip"
|
||||
req_path = ((self.parsed.path or '/') +
|
||||
(('?' + self.parsed.query) if self.parsed.query else ''))
|
||||
(('?' + self.parsed.query) if self.parsed.query else ''))
|
||||
request_lines = [utf8("%s %s HTTP/1.1" % (self.request.method,
|
||||
req_path))]
|
||||
for k, v in self.request.headers.get_all():
|
||||
line = utf8(k) + b(": ") + utf8(v)
|
||||
if b('\n') in line:
|
||||
line = utf8(k) + b": " + utf8(v)
|
||||
if b'\n' in line:
|
||||
raise ValueError('Newline in header: ' + repr(line))
|
||||
request_lines.append(line)
|
||||
self.stream.write(b("\r\n").join(request_lines) + b("\r\n\r\n"))
|
||||
self.stream.write(b"\r\n".join(request_lines) + b"\r\n\r\n")
|
||||
if self.request.body is not None:
|
||||
self.stream.write(self.request.body)
|
||||
self.stream.read_until_regex(b("\r?\n\r?\n"), self._on_headers)
|
||||
self.stream.read_until_regex(b"\r?\n\r?\n", self._on_headers)
|
||||
|
||||
def _release(self):
|
||||
if self.release_callback is not None:
|
||||
@@ -312,10 +293,11 @@ class _HTTPConnection(object):
|
||||
|
||||
def _handle_exception(self, typ, value, tb):
|
||||
if self.final_callback:
|
||||
self._remove_timeout()
|
||||
gen_log.warning("uncaught exception", exc_info=(typ, value, tb))
|
||||
self._run_callback(HTTPResponse(self.request, 599, error=value,
|
||||
request_time=self.io_loop.time() - self.start_time,
|
||||
))
|
||||
request_time=self.io_loop.time() - self.start_time,
|
||||
))
|
||||
|
||||
if hasattr(self, "stream"):
|
||||
self.stream.close()
|
||||
@@ -334,19 +316,22 @@ class _HTTPConnection(object):
|
||||
message = str(self.stream.error)
|
||||
raise HTTPError(599, message)
|
||||
|
||||
def _handle_1xx(self, code):
|
||||
self.stream.read_until_regex(b"\r?\n\r?\n", self._on_headers)
|
||||
|
||||
def _on_headers(self, data):
|
||||
data = native_str(data.decode("latin1"))
|
||||
first_line, _, header_data = data.partition("\n")
|
||||
match = re.match("HTTP/1.[01] ([0-9]+) ([^\r]*)", first_line)
|
||||
assert match
|
||||
code = int(match.group(1))
|
||||
self.headers = HTTPHeaders.parse(header_data)
|
||||
if 100 <= code < 200:
|
||||
self.stream.read_until_regex(b("\r?\n\r?\n"), self._on_headers)
|
||||
self._handle_1xx(code)
|
||||
return
|
||||
else:
|
||||
self.code = code
|
||||
self.reason = match.group(2)
|
||||
self.headers = HTTPHeaders.parse(header_data)
|
||||
|
||||
if "Content-Length" in self.headers:
|
||||
if "," in self.headers["Content-Length"]:
|
||||
@@ -372,38 +357,36 @@ class _HTTPConnection(object):
|
||||
if self.request.method == "HEAD" or self.code == 304:
|
||||
# HEAD requests and 304 responses never have content, even
|
||||
# though they may have content-length headers
|
||||
self._on_body(b(""))
|
||||
self._on_body(b"")
|
||||
return
|
||||
if 100 <= self.code < 200 or self.code == 204:
|
||||
# These response codes never have bodies
|
||||
# http://www.w3.org/Protocols/rfc2616/rfc2616-sec4.html#sec4.3
|
||||
if ("Transfer-Encoding" in self.headers or
|
||||
content_length not in (None, 0)):
|
||||
content_length not in (None, 0)):
|
||||
raise ValueError("Response with code %d should not have body" %
|
||||
self.code)
|
||||
self._on_body(b(""))
|
||||
self._on_body(b"")
|
||||
return
|
||||
|
||||
if (self.request.use_gzip and
|
||||
self.headers.get("Content-Encoding") == "gzip"):
|
||||
self.headers.get("Content-Encoding") == "gzip"):
|
||||
self._decompressor = GzipDecompressor()
|
||||
if self.headers.get("Transfer-Encoding") == "chunked":
|
||||
self.chunks = []
|
||||
self.stream.read_until(b("\r\n"), self._on_chunk_length)
|
||||
self.stream.read_until(b"\r\n", self._on_chunk_length)
|
||||
elif content_length is not None:
|
||||
self.stream.read_bytes(content_length, self._on_body)
|
||||
else:
|
||||
self.stream.read_until_close(self._on_body)
|
||||
|
||||
def _on_body(self, data):
|
||||
if self._timeout is not None:
|
||||
self.io_loop.remove_timeout(self._timeout)
|
||||
self._timeout = None
|
||||
self._remove_timeout()
|
||||
original_request = getattr(self.request, "original_request",
|
||||
self.request)
|
||||
if (self.request.follow_redirects and
|
||||
self.request.max_redirects > 0 and
|
||||
self.code in (301, 302, 303, 307)):
|
||||
self.code in (301, 302, 303, 307)):
|
||||
assert isinstance(self.request, _RequestProxy)
|
||||
new_request = copy.copy(self.request.request)
|
||||
new_request.url = urlparse.urljoin(self.request.url,
|
||||
@@ -472,13 +455,13 @@ class _HTTPConnection(object):
|
||||
# all the data has been decompressed, so we don't need to
|
||||
# decompress again in _on_body
|
||||
self._decompressor = None
|
||||
self._on_body(b('').join(self.chunks))
|
||||
self._on_body(b''.join(self.chunks))
|
||||
else:
|
||||
self.stream.read_bytes(length + 2, # chunk ends with \r\n
|
||||
self._on_chunk_data)
|
||||
self._on_chunk_data)
|
||||
|
||||
def _on_chunk_data(self, data):
|
||||
assert data[-2:] == b("\r\n")
|
||||
assert data[-2:] == b"\r\n"
|
||||
chunk = data[:-2]
|
||||
if self._decompressor:
|
||||
chunk = self._decompressor.decompress(chunk)
|
||||
@@ -486,69 +469,9 @@ class _HTTPConnection(object):
|
||||
self.request.streaming_callback(chunk)
|
||||
else:
|
||||
self.chunks.append(chunk)
|
||||
self.stream.read_until(b("\r\n"), self._on_chunk_length)
|
||||
self.stream.read_until(b"\r\n", self._on_chunk_length)
|
||||
|
||||
|
||||
# match_hostname was added to the standard library ssl module in python 3.2.
|
||||
# The following code was backported for older releases and copied from
|
||||
# https://bitbucket.org/brandon/backports.ssl_match_hostname
|
||||
class CertificateError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def _dnsname_to_pat(dn):
|
||||
pats = []
|
||||
for frag in dn.split(r'.'):
|
||||
if frag == '*':
|
||||
# When '*' is a fragment by itself, it matches a non-empty dotless
|
||||
# fragment.
|
||||
pats.append('[^.]+')
|
||||
else:
|
||||
# Otherwise, '*' matches any dotless fragment.
|
||||
frag = re.escape(frag)
|
||||
pats.append(frag.replace(r'\*', '[^.]*'))
|
||||
return re.compile(r'\A' + r'\.'.join(pats) + r'\Z', re.IGNORECASE)
|
||||
|
||||
|
||||
def match_hostname(cert, hostname):
|
||||
"""Verify that *cert* (in decoded format as returned by
|
||||
SSLSocket.getpeercert()) matches the *hostname*. RFC 2818 rules
|
||||
are mostly followed, but IP addresses are not accepted for *hostname*.
|
||||
|
||||
CertificateError is raised on failure. On success, the function
|
||||
returns nothing.
|
||||
"""
|
||||
if not cert:
|
||||
raise ValueError("empty or no certificate")
|
||||
dnsnames = []
|
||||
san = cert.get('subjectAltName', ())
|
||||
for key, value in san:
|
||||
if key == 'DNS':
|
||||
if _dnsname_to_pat(value).match(hostname):
|
||||
return
|
||||
dnsnames.append(value)
|
||||
if not san:
|
||||
# The subject is only checked when subjectAltName is empty
|
||||
for sub in cert.get('subject', ()):
|
||||
for key, value in sub:
|
||||
# XXX according to RFC 2818, the most specific Common Name
|
||||
# must be used.
|
||||
if key == 'commonName':
|
||||
if _dnsname_to_pat(value).match(hostname):
|
||||
return
|
||||
dnsnames.append(value)
|
||||
if len(dnsnames) > 1:
|
||||
raise CertificateError("hostname %r "
|
||||
"doesn't match either of %s"
|
||||
% (hostname, ', '.join(map(repr, dnsnames))))
|
||||
elif len(dnsnames) == 1:
|
||||
raise CertificateError("hostname %r "
|
||||
"doesn't match %r"
|
||||
% (hostname, dnsnames[0]))
|
||||
else:
|
||||
raise CertificateError("no appropriate commonName or "
|
||||
"subjectAltName fields were found")
|
||||
|
||||
if __name__ == "__main__":
|
||||
AsyncHTTPClient.configure(SimpleAsyncHTTPClient)
|
||||
main()
|
||||
|
||||
Reference in New Issue
Block a user