codekingpro/portable-devtools
115k
1from tornado import gen, netutil
2from tornado.escape import (
3 json_decode,
4 json_encode,
5 utf8,
6 _unicode,
7 recursive_unicode,
8 native_str,
9)
10from tornado.http1connection import HTTP1Connection
11from tornado.httpclient import HTTPError
12from tornado.httpserver import HTTPServer
13from tornado.httputil import (
14 HTTPHeaders,
15 HTTPMessageDelegate,
16 HTTPServerConnectionDelegate,
17 ResponseStartLine,
18)
19from tornado.iostream import IOStream
20from tornado.locks import Event
21from tornado.log import gen_log, app_log
22from tornado.simple_httpclient import SimpleAsyncHTTPClient
23from tornado.testing import (
24 AsyncHTTPTestCase,
25 AsyncHTTPSTestCase,
26 AsyncTestCase,
27 ExpectLog,
28 gen_test,
29)
30from tornado.test.util import abstract_base_test
31from tornado.web import Application, RequestHandler, stream_request_body
32
33from contextlib import closing, contextmanager
34import datetime
35import gzip
36import logging
37import os
38import shutil
39import socket
40import ssl
41import sys
42import tempfile
43import textwrap
44import unittest
45import urllib.parse
46import uuid
47from io import BytesIO
48
49import typing
50
51if typing.TYPE_CHECKING:
52 from typing import Dict, List # noqa: F401
53
54
55async def read_stream_body(stream):
56 """Reads an HTTP response from `stream` and returns a tuple of its
57 start_line, headers and body."""
58 chunks = []
59
60 class Delegate(HTTPMessageDelegate):
61 def headers_received(self, start_line, headers):
62 self.headers = headers
63 self.start_line = start_line
64
65 def data_received(self, chunk):
66 chunks.append(chunk)
67
68 def finish(self):
69 conn.detach() # type: ignore
70
71 conn = HTTP1Connection(stream, True)
72 delegate = Delegate()
73 await conn.read_response(delegate)
74 return delegate.start_line, delegate.headers, b"".join(chunks)
75
76
77class HandlerBaseTestCase(AsyncHTTPTestCase):
78 Handler = None
79
80 def get_app(self):
81 return Application([("/", self.__class__.Handler)])
82
83 def fetch_json(self, *args, **kwargs):
84 response = self.fetch(*args, **kwargs)
85 response.rethrow()
86 return json_decode(response.body)
87
88
89class HelloWorldRequestHandler(RequestHandler):
90 def initialize(self, protocol="http"):
91 self.expected_protocol = protocol
92
93 def get(self):
94 if self.request.protocol != self.expected_protocol:
95 raise Exception("unexpected protocol")
96 self.finish("Hello world")
97
98 def post(self):
99 self.finish("Got %d bytes in POST" % len(self.request.body))
100
101
102class SSLTest(AsyncHTTPSTestCase):
103 def get_app(self):
104 return Application([("/", HelloWorldRequestHandler, dict(protocol="https"))])
105
106 def get_ssl_options(self):
107 return dict(
108 ssl_version=ssl.PROTOCOL_TLS_SERVER,
109 **AsyncHTTPSTestCase.default_ssl_options(),
110 )
111
112 def test_ssl(self):
113 response = self.fetch("/")
114 self.assertEqual(response.body, b"Hello world")
115
116 def test_large_post(self):
117 response = self.fetch("/", method="POST", body="A" * 5000)
118 self.assertEqual(response.body, b"Got 5000 bytes in POST")
119
120 def test_non_ssl_request(self):
121 # Make sure the server closes the connection when it gets a non-ssl
122 # connection, rather than waiting for a timeout or otherwise
123 # misbehaving.
124 with ExpectLog(gen_log, "(SSL Error|uncaught exception)"):
125 with ExpectLog(gen_log, "Uncaught exception", required=False):
126 with self.assertRaises((IOError, HTTPError)): # type: ignore
127 self.fetch(
128 self.get_url("/").replace("https:", "http:"),
129 request_timeout=3600,
130 connect_timeout=3600,
131 raise_error=True,
132 )
133
134 def test_error_logging(self):
135 # No stack traces are logged for SSL errors.
136 with ExpectLog(gen_log, "SSL Error") as expect_log:
137 with self.assertRaises((IOError, HTTPError)): # type: ignore
138 self.fetch(
139 self.get_url("/").replace("https:", "http:"), raise_error=True
140 )
141 self.assertFalse(expect_log.logged_stack)
142
143
144class BadSSLOptionsTest(unittest.TestCase):
145 def test_missing_arguments(self):
146 application = Application()
147 self.assertRaises(
148 KeyError,
149 HTTPServer,
150 application,
151 ssl_options={"keyfile": "/__missing__.crt"},
152 )
153
154 def test_missing_key(self):
155 """A missing SSL key should cause an immediate exception."""
156
157 application = Application()
158 module_dir = os.path.dirname(__file__)
159 existing_certificate = os.path.join(module_dir, "test.crt")
160 existing_key = os.path.join(module_dir, "test.key")
161
162 self.assertRaises(
163 (ValueError, IOError),
164 HTTPServer,
165 application,
166 ssl_options={"certfile": "/__mising__.crt"},
167 )
168 self.assertRaises(
169 (ValueError, IOError),
170 HTTPServer,
171 application,
172 ssl_options={
173 "certfile": existing_certificate,
174 "keyfile": "/__missing__.key",
175 },
176 )
177
178 # This actually works because both files exist
179 HTTPServer(
180 application,
181 ssl_options={"certfile": existing_certificate, "keyfile": existing_key},
182 )
183
184
185class MultipartTestHandler(RequestHandler):
186 def post(self):
187 self.finish(
188 {
189 "header": self.request.headers["X-Header-Encoding-Test"],
190 "argument": self.get_argument("argument"),
191 "filename": self.request.files["files"][0].filename,
192 "filebody": _unicode(self.request.files["files"][0]["body"]),
193 }
194 )
195
196
197# This test is also called from wsgi_test
198class HTTPConnectionTest(AsyncHTTPTestCase):
199 def get_handlers(self):
200 return [
201 ("/multipart", MultipartTestHandler),
202 ("/hello", HelloWorldRequestHandler),
203 ]
204
205 def get_app(self):
206 return Application(self.get_handlers())
207
208 def raw_fetch(self, headers, body, newline=b"\r\n"):
209 with closing(IOStream(socket.socket())) as stream:
210 self.io_loop.run_sync(
211 lambda: stream.connect(("127.0.0.1", self.get_http_port()))
212 )
213 stream.write(
214 newline.join(headers + [utf8("Content-Length: %d" % len(body))])
215 + newline
216 + newline
217 + body
218 )
219 start_line, headers, body = self.io_loop.run_sync(
220 lambda: read_stream_body(stream)
221 )
222 return body
223
224 def test_multipart_form(self):
225 # Encodings here are tricky: Headers are latin1, bodies can be
226 # anything (we use utf8 by default).
227 response = self.raw_fetch(
228 [
229 b"POST /multipart HTTP/1.0",
230 b"Content-Type: multipart/form-data; boundary=1234567890",
231 b"X-Header-encoding-test: \xe9",
232 ],
233 b"\r\n".join(
234 [
235 b"Content-Disposition: form-data; name=argument",
236 b"",
237 "\u00e1".encode(),
238 b"--1234567890",
239 'Content-Disposition: form-data; name="files"; filename="\u00f3"'.encode(),
240 b"",
241 "\u00fa".encode(),
242 b"--1234567890--",
243 b"",
244 ]
245 ),
246 )
247 data = json_decode(response)
248 self.assertEqual("\u00e9", data["header"])
249 self.assertEqual("\u00e1", data["argument"])
250 self.assertEqual("\u00f3", data["filename"])
251 self.assertEqual("\u00fa", data["filebody"])
252
253 def test_newlines(self):
254 # We support both CRLF and bare LF as line separators.
255 for newline in (b"\r\n", b"\n"):
256 response = self.raw_fetch([b"GET /hello HTTP/1.0"], b"", newline=newline)
257 self.assertEqual(response, b"Hello world")
258
259 @gen_test
260 def test_100_continue(self):
261 # Run through a 100-continue interaction by hand:
262 # When given Expect: 100-continue, we get a 100 response after the
263 # headers, and then the real response after the body.
264 stream = IOStream(socket.socket())
265 yield stream.connect(("127.0.0.1", self.get_http_port()))
266 yield stream.write(
267 b"\r\n".join(
268 [
269 b"POST /hello HTTP/1.1",
270 b"Host: 127.0.0.1",
271 b"Content-Length: 1024",
272 b"Expect: 100-continue",
273 b"Connection: close",
274 b"\r\n",
275 ]
276 )
277 )
278 data = yield stream.read_until(b"\r\n\r\n")
279 self.assertTrue(data.startswith(b"HTTP/1.1 100 "), data)
280 stream.write(b"a" * 1024)
281 first_line = yield stream.read_until(b"\r\n")
282 self.assertTrue(first_line.startswith(b"HTTP/1.1 200"), first_line)
283 header_data = yield stream.read_until(b"\r\n\r\n")
284 headers = HTTPHeaders.parse(native_str(header_data.decode("latin1")))
285 body = yield stream.read_bytes(int(headers["Content-Length"]))
286 self.assertEqual(body, b"Got 1024 bytes in POST")
287 stream.close()
288
289
290class EchoHandler(RequestHandler):
291 def get(self):
292 self.write(recursive_unicode(self.request.arguments))
293
294 def post(self):
295 self.write(recursive_unicode(self.request.arguments))
296
297
298class TypeCheckHandler(RequestHandler):
299 def prepare(self):
300 self.errors = {} # type: Dict[str, str]
301 fields = [
302 ("method", str),
303 ("uri", str),
304 ("version", str),
305 ("remote_ip", str),
306 ("protocol", str),
307 ("host", str),
308 ("path", str),
309 ("query", str),
310 ]
311 for field, expected_type in fields:
312 self.check_type(field, getattr(self.request, field), expected_type)
313
314 self.check_type("header_key", list(self.request.headers.keys())[0], str)
315 self.check_type("header_value", list(self.request.headers.values())[0], str)
316
317 self.check_type("cookie_key", list(self.request.cookies.keys())[0], str)
318 self.check_type(
319 "cookie_value", list(self.request.cookies.values())[0].value, str
320 )
321 # secure cookies
322
323 self.check_type("arg_key", list(self.request.arguments.keys())[0], str)
324 self.check_type("arg_value", list(self.request.arguments.values())[0][0], bytes)
325
326 def post(self):
327 self.check_type("body", self.request.body, bytes)
328 self.write(self.errors)
329
330 def get(self):
331 self.write(self.errors)
332
333 def check_type(self, name, obj, expected_type):
334 actual_type = type(obj)
335 if expected_type != actual_type:
336 self.errors[name] = f"expected {expected_type}, got {actual_type}"
337
338
339class PostEchoHandler(RequestHandler):
340 def post(self, *path_args):
341 self.write(dict(echo=self.get_argument("data")))
342
343
344class PostEchoGBKHandler(PostEchoHandler):
345 def decode_argument(self, value, name=None):
346 try:
347 return value.decode("gbk")
348 except Exception:
349 raise HTTPError(400, "invalid gbk bytes: %r" % value)
350
351
352class HTTPServerTest(AsyncHTTPTestCase):
353 def get_app(self):
354 return Application(
355 [
356 ("/echo", EchoHandler),
357 ("/typecheck", TypeCheckHandler),
358 ("//doubleslash", EchoHandler),
359 ("/post_utf8", PostEchoHandler),
360 ("/post_gbk", PostEchoGBKHandler),
361 ]
362 )
363
364 def test_query_string_encoding(self):
365 response = self.fetch("/echo?foo=%C3%A9")
366 data = json_decode(response.body)
367 self.assertEqual(data, {"foo": ["\u00e9"]})
368
369 def test_empty_query_string(self):
370 response = self.fetch("/echo?foo=&foo=")
371 data = json_decode(response.body)
372 self.assertEqual(data, {"foo": ["", ""]})
373
374 def test_empty_post_parameters(self):
375 response = self.fetch("/echo", method="POST", body="foo=&bar=")
376 data = json_decode(response.body)
377 self.assertEqual(data, {"foo": [""], "bar": [""]})
378
379 def test_types(self):
380 headers = {"Cookie": "foo=bar"}
381 response = self.fetch("/typecheck?foo=bar", headers=headers)
382 data = json_decode(response.body)
383 self.assertEqual(data, {})
384
385 response = self.fetch(
386 "/typecheck", method="POST", body="foo=bar", headers=headers
387 )
388 data = json_decode(response.body)
389 self.assertEqual(data, {})
390
391 def test_double_slash(self):
392 # urlparse.urlsplit (which tornado.httpserver used to use
393 # incorrectly) would parse paths beginning with "//" as
394 # protocol-relative urls.
395 response = self.fetch("//doubleslash")
396 self.assertEqual(200, response.code)
397 self.assertEqual(json_decode(response.body), {})
398
399 def test_post_encodings(self):
400 headers = {"Content-Type": "application/x-www-form-urlencoded"}
401 uni_text = "chinese: \u5f20\u4e09"
402 for enc in ("utf8", "gbk"):
403 for quote in (True, False):
404 with self.subTest(enc=enc, quote=quote):
405 bin_text = uni_text.encode(enc)
406 if quote:
407 bin_text = urllib.parse.quote(bin_text).encode("ascii")
408 response = self.fetch(
409 "/post_" + enc,
410 method="POST",
411 headers=headers,
412 body=(b"data=" + bin_text),
413 )
414 self.assertEqual(json_decode(response.body), {"echo": uni_text})
415
416
417class HTTPServerRawTest(AsyncHTTPTestCase):
418 def get_app(self):
419 return Application([("/echo", EchoHandler)])
420
421 def setUp(self):
422 super().setUp()
423 self.stream = IOStream(socket.socket())
424 self.io_loop.run_sync(
425 lambda: self.stream.connect(("127.0.0.1", self.get_http_port()))
426 )
427
428 def tearDown(self):
429 self.stream.close()
430 super().tearDown()
431
432 def test_empty_request(self):
433 self.stream.close()
434 self.io_loop.add_timeout(datetime.timedelta(seconds=0.001), self.stop)
435 self.wait()
436
437 def test_malformed_first_line_response(self):
438 with ExpectLog(gen_log, ".*Malformed HTTP request line", level=logging.INFO):
439 self.stream.write(b"asdf\r\n\r\n")
440 start_line, headers, response = self.io_loop.run_sync(
441 lambda: read_stream_body(self.stream)
442 )
443 self.assertEqual("HTTP/1.1", start_line.version)
444 self.assertEqual(400, start_line.code)
445 self.assertEqual("Bad Request", start_line.reason)
446
447 def test_malformed_first_line_log(self):
448 with ExpectLog(gen_log, ".*Malformed HTTP request line", level=logging.INFO):
449 self.stream.write(b"asdf\r\n\r\n")
450 # TODO: need an async version of ExpectLog so we don't need
451 # hard-coded timeouts here.
452 self.io_loop.add_timeout(datetime.timedelta(seconds=0.05), self.stop)
453 self.wait()
454
455 def test_malformed_headers(self):
456 with ExpectLog(
457 gen_log,
458 ".*Malformed HTTP message.*no colon in header line",
459 level=logging.INFO,
460 ):
461 self.stream.write(b"GET / HTTP/1.0\r\nasdf\r\n\r\n")
462 self.io_loop.add_timeout(datetime.timedelta(seconds=0.05), self.stop)
463 self.wait()
464
465 def test_invalid_host_header_with_whitespace(self):
466 with ExpectLog(
467 gen_log, ".*Malformed HTTP message.*Invalid Host header", level=logging.INFO
468 ):
469 self.stream.write(b"GET / HTTP/1.0\r\nHost: foo bar\r\n\r\n")
470 start_line, headers, response = self.io_loop.run_sync(
471 lambda: read_stream_body(self.stream)
472 )
473 self.assertEqual("HTTP/1.1", start_line.version)
474 self.assertEqual(400, start_line.code)
475 self.assertEqual("Bad Request", start_line.reason)
476
477 def test_chunked_request_body(self):
478 # Chunked requests are not widely supported and we don't have a way
479 # to generate them in AsyncHTTPClient, but HTTPServer will read them.
480 self.stream.write(
481 b"""\
482POST /echo HTTP/1.1
483Host: 127.0.0.1
484Transfer-Encoding: chunked
485Content-Type: application/x-www-form-urlencoded
486
4874
488foo=
4893
490bar
4910
492
493""".replace(
494 b"\n", b"\r\n"
495 )
496 )
497 start_line, headers, response = self.io_loop.run_sync(
498 lambda: read_stream_body(self.stream)
499 )
500 self.assertEqual(json_decode(response), {"foo": ["bar"]})
501
502 def test_chunked_request_uppercase(self):
503 # As per RFC 2616 section 3.6, "Transfer-Encoding" header's value is
504 # case-insensitive.
505 self.stream.write(
506 b"""\
507POST /echo HTTP/1.1
508Host: 127.0.0.1
509Transfer-Encoding: Chunked
510Content-Type: application/x-www-form-urlencoded
511
5124
513foo=
5143
515bar
5160
517
518""".replace(
519 b"\n", b"\r\n"
520 )
521 )
522 start_line, headers, response = self.io_loop.run_sync(
523 lambda: read_stream_body(self.stream)
524 )
525 self.assertEqual(json_decode(response), {"foo": ["bar"]})
526
527 def test_chunked_request_body_invalid_size(self):
528 # Only hex digits are allowed in chunk sizes. Python's int() function
529 # also accepts underscores, so make sure we reject them here.
530 self.stream.write(
531 b"""\
532POST /echo HTTP/1.1
533Host: 127.0.0.1
534Transfer-Encoding: chunked
535
5361_a
5371234567890abcdef1234567890
5380
539
540""".replace(
541 b"\n", b"\r\n"
542 )
543 )
544 with ExpectLog(gen_log, ".*invalid chunk size", level=logging.INFO):
545 start_line, headers, response = self.io_loop.run_sync(
546 lambda: read_stream_body(self.stream)
547 )
548 self.assertEqual(400, start_line.code)
549
550 def test_chunked_request_body_duplicate_header(self):
551 # Repeated Transfer-Encoding headers should be an error (and not confuse
552 # the chunked-encoding detection to mess up framing).
553 self.stream.write(
554 b"""\
555POST /echo HTTP/1.1
556Host: 127.0.0.1
557Transfer-Encoding: chunked
558Transfer-encoding: chunked
559
5602
561ok
5620
563
564"""
565 )
566 with ExpectLog(
567 gen_log,
568 ".*Unsupported Transfer-Encoding chunked,chunked",
569 level=logging.INFO,
570 ):
571 start_line, headers, response = self.io_loop.run_sync(
572 lambda: read_stream_body(self.stream)
573 )
574 self.assertEqual(400, start_line.code)
575
576 def test_chunked_request_body_unsupported_transfer_encoding(self):
577 # We don't support transfer-encodings other than chunked.
578 self.stream.write(
579 b"""\
580POST /echo HTTP/1.1
581Host: 127.0.0.1
582Transfer-Encoding: gzip, chunked
583
5842
585ok
5860
587
588"""
589 )
590 with ExpectLog(
591 gen_log, ".*Unsupported Transfer-Encoding gzip, chunked", level=logging.INFO
592 ):
593 start_line, headers, response = self.io_loop.run_sync(
594 lambda: read_stream_body(self.stream)
595 )
596 self.assertEqual(400, start_line.code)
597
598 def test_chunked_request_body_transfer_encoding_and_content_length(self):
599 # Transfer-encoding and content-length are mutually exclusive
600 self.stream.write(
601 b"""\
602POST /echo HTTP/1.1
603Host: 127.0.0.1
604Transfer-Encoding: chunked
605Content-Length: 2
606
6072
608ok
6090
610
611"""
612 )
613 with ExpectLog(
614 gen_log,
615 ".*Message with both Transfer-Encoding and Content-Length",
616 level=logging.INFO,
617 ):
618 start_line, headers, response = self.io_loop.run_sync(
619 lambda: read_stream_body(self.stream)
620 )
621 self.assertEqual(400, start_line.code)
622
623 @gen_test
624 def test_invalid_content_length(self):
625 # HTTP only allows decimal digits in content-length. Make sure we don't
626 # accept anything else, with special attention to things accepted by the
627 # python int() function (leading plus signs and internal underscores).
628 test_cases = [
629 ("alphabetic", "foo"),
630 ("leading plus", "+10"),
631 ("internal underscore", "1_0"),
632 ]
633 for name, value in test_cases:
634 with self.subTest(name=name), closing(IOStream(socket.socket())) as stream:
635 with ExpectLog(
636 gen_log,
637 ".*Only integer Content-Length is allowed",
638 level=logging.INFO,
639 ):
640 yield stream.connect(("127.0.0.1", self.get_http_port()))
641 stream.write(
642 utf8(
643 textwrap.dedent(
644 f"""\
645 POST /echo HTTP/1.1
646 Host: 127.0.0.1
647 Content-Length: {value}
648 Connection: close
649
650 1234567890
651 """
652 ).replace("\n", "\r\n")
653 )
654 )
655 yield stream.read_until_close()
656
657 @gen_test
658 def test_invalid_methods(self):
659 # RFC 9110 distinguishes between syntactically invalid methods and those that are
660 # valid but unknown. The former must give a 400 status code, while the latter should
661 # give a 405.
662 test_cases = [
663 ("FOO", 405, None),
664 ("FOO,BAR", 400, ".*Malformed HTTP request line"),
665 ]
666 for method, code, log_msg in test_cases:
667 if log_msg is not None:
668 expect_log = ExpectLog(gen_log, log_msg, level=logging.INFO)
669 else:
670
671 @contextmanager
672 def noop_context():
673 yield
674
675 expect_log = noop_context() # type: ignore
676 with (
677 self.subTest(method=method),
678 closing(IOStream(socket.socket())) as stream,
679 expect_log,
680 ):
681 yield stream.connect(("127.0.0.1", self.get_http_port()))
682 stream.write(utf8(f"{method} /echo HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n"))
683 resp = yield stream.read_until(b"\r\n\r\n")
684 self.assertTrue(
685 resp.startswith(b"HTTP/1.1 %d" % code),
686 f"expected status code {code} in {resp!r}",
687 )
688
689
690class XHeaderTest(HandlerBaseTestCase):
691 class Handler(RequestHandler):
692 def get(self):
693 self.set_header("request-version", self.request.version)
694 self.write(
695 dict(
696 remote_ip=self.request.remote_ip,
697 remote_protocol=self.request.protocol,
698 )
699 )
700
701 def get_httpserver_options(self):
702 return dict(xheaders=True, trusted_downstream=["5.5.5.5"])
703
704 def test_ip_headers(self):
705 self.assertEqual(self.fetch_json("/")["remote_ip"], "127.0.0.1")
706
707 valid_ipv4 = {"X-Real-IP": "4.4.4.4"}
708 self.assertEqual(
709 self.fetch_json("/", headers=valid_ipv4)["remote_ip"], "4.4.4.4"
710 )
711
712 valid_ipv4_list = {"X-Forwarded-For": "127.0.0.1, 4.4.4.4"}
713 self.assertEqual(
714 self.fetch_json("/", headers=valid_ipv4_list)["remote_ip"], "4.4.4.4"
715 )
716
717 valid_ipv6 = {"X-Real-IP": "2620:0:1cfe:face:b00c::3"}
718 self.assertEqual(
719 self.fetch_json("/", headers=valid_ipv6)["remote_ip"],
720 "2620:0:1cfe:face:b00c::3",
721 )
722
723 valid_ipv6_list = {"X-Forwarded-For": "::1, 2620:0:1cfe:face:b00c::3"}
724 self.assertEqual(
725 self.fetch_json("/", headers=valid_ipv6_list)["remote_ip"],
726 "2620:0:1cfe:face:b00c::3",
727 )
728
729 invalid_chars = {"X-Real-IP": "4.4.4.4<script>"}
730 self.assertEqual(
731 self.fetch_json("/", headers=invalid_chars)["remote_ip"], "127.0.0.1"
732 )
733
734 invalid_chars_list = {"X-Forwarded-For": "4.4.4.4, 5.5.5.5<script>"}
735 self.assertEqual(
736 self.fetch_json("/", headers=invalid_chars_list)["remote_ip"], "127.0.0.1"
737 )
738
739 invalid_host = {"X-Real-IP": "www.google.com"}
740 self.assertEqual(
741 self.fetch_json("/", headers=invalid_host)["remote_ip"], "127.0.0.1"
742 )
743
744 def test_trusted_downstream(self):
745 valid_ipv4_list = {"X-Forwarded-For": "127.0.0.1, 4.4.4.4, 5.5.5.5"}
746 resp = self.fetch("/", headers=valid_ipv4_list)
747 if resp.headers["request-version"].startswith("HTTP/2"):
748 # This is a hack - there's nothing that fundamentally requires http/1
749 # here but tornado_http2 doesn't support it yet.
750 self.skipTest("requires HTTP/1.x")
751 result = json_decode(resp.body)
752 self.assertEqual(result["remote_ip"], "4.4.4.4")
753
754 def test_scheme_headers(self):
755 self.assertEqual(self.fetch_json("/")["remote_protocol"], "http")
756
757 https_scheme = {"X-Scheme": "https"}
758 self.assertEqual(
759 self.fetch_json("/", headers=https_scheme)["remote_protocol"], "https"
760 )
761
762 https_forwarded = {"X-Forwarded-Proto": "https"}
763 self.assertEqual(
764 self.fetch_json("/", headers=https_forwarded)["remote_protocol"], "https"
765 )
766
767 https_multi_forwarded = {"X-Forwarded-Proto": "https , http"}
768 self.assertEqual(
769 self.fetch_json("/", headers=https_multi_forwarded)["remote_protocol"],
770 "http",
771 )
772
773 http_multi_forwarded = {"X-Forwarded-Proto": "http,https"}
774 self.assertEqual(
775 self.fetch_json("/", headers=http_multi_forwarded)["remote_protocol"],
776 "https",
777 )
778
779 bad_forwarded = {"X-Forwarded-Proto": "unknown"}
780 self.assertEqual(
781 self.fetch_json("/", headers=bad_forwarded)["remote_protocol"], "http"
782 )
783
784
785class SSLXHeaderTest(AsyncHTTPSTestCase, HandlerBaseTestCase):
786 def get_app(self):
787 return Application([("/", XHeaderTest.Handler)])
788
789 def get_httpserver_options(self):
790 output = super().get_httpserver_options()
791 output["xheaders"] = True
792 return output
793
794 def test_request_without_xprotocol(self):
795 self.assertEqual(self.fetch_json("/")["remote_protocol"], "https")
796
797 http_scheme = {"X-Scheme": "http"}
798 self.assertEqual(
799 self.fetch_json("/", headers=http_scheme)["remote_protocol"], "http"
800 )
801
802 bad_scheme = {"X-Scheme": "unknown"}
803 self.assertEqual(
804 self.fetch_json("/", headers=bad_scheme)["remote_protocol"], "https"
805 )
806
807
808class ManualProtocolTest(HandlerBaseTestCase):
809 class Handler(RequestHandler):
810 def get(self):
811 self.write(dict(protocol=self.request.protocol))
812
813 def get_httpserver_options(self):
814 return dict(protocol="https")
815
816 def test_manual_protocol(self):
817 self.assertEqual(self.fetch_json("/")["protocol"], "https")
818
819
820@abstract_base_test
821class UnixSocketTest(AsyncTestCase):
822 """HTTPServers can listen on Unix sockets too.
823
824 Why would you want to do this? Nginx can proxy to backends listening
825 on unix sockets, for one thing (and managing a namespace for unix
826 sockets can be easier than managing a bunch of TCP port numbers).
827
828 Unfortunately, there's no way to specify a unix socket in a url for
829 an HTTP client, so we have to test this by hand.
830 """
831
832 address = ""
833
834 def setUp(self):
835 super().setUp()
836 app = Application([("/hello", HelloWorldRequestHandler)])
837 self.server = HTTPServer(app)
838 self.server.add_socket(netutil.bind_unix_socket(self.address))
839
840 def tearDown(self):
841 self.io_loop.run_sync(self.server.close_all_connections)
842 self.server.stop()
843 super().tearDown()
844
845 @gen_test
846 def test_unix_socket(self):
847 with closing(IOStream(socket.socket(socket.AF_UNIX))) as stream:
848 stream.connect(self.address)
849 stream.write(b"GET /hello HTTP/1.0\r\n\r\n")
850 response = yield stream.read_until(b"\r\n")
851 self.assertEqual(response, b"HTTP/1.1 200 OK\r\n")
852 header_data = yield stream.read_until(b"\r\n\r\n")
853 headers = HTTPHeaders.parse(header_data.decode("latin1"))
854 body = yield stream.read_bytes(int(headers["Content-Length"]))
855 self.assertEqual(body, b"Hello world")
856
857 @gen_test
858 def test_unix_socket_bad_request(self):
859 # Unix sockets don't have remote addresses so they just return an
860 # empty string.
861 with ExpectLog(gen_log, "Malformed HTTP message from", level=logging.INFO):
862 with closing(IOStream(socket.socket(socket.AF_UNIX))) as stream:
863 stream.connect(self.address)
864 stream.write(b"garbage\r\n\r\n")
865 response = yield stream.read_until_close()
866 self.assertEqual(response, b"HTTP/1.1 400 Bad Request\r\n\r\n")
867
868
869@unittest.skipIf(
870 not hasattr(socket, "AF_UNIX") or sys.platform == "cygwin",
871 "unix sockets not supported on this platform",
872)
873class UnixSocketTestFile(UnixSocketTest):
874 def setUp(self):
875 self.tmpdir = tempfile.mkdtemp()
876 self.address = os.path.join(self.tmpdir, "test.sock")
877 super().setUp()
878
879 def tearDown(self):
880 super().tearDown()
881 shutil.rmtree(self.tmpdir)
882
883
884@unittest.skipIf(
885 not (hasattr(socket, "AF_UNIX") and sys.platform.startswith("linux")),
886 "abstract namespace unix sockets not supported on this platform",
887)
888class UnixSocketTestAbstract(UnixSocketTest):
889 def setUp(self):
890 self.address = "\0" + uuid.uuid4().hex
891 super().setUp()
892
893
894class KeepAliveTest(AsyncHTTPTestCase):
895 """Tests various scenarios for HTTP 1.1 keep-alive support.
896
897 These tests don't use AsyncHTTPClient because we want to control
898 connection reuse and closing.
899 """
900
901 def get_app(self):
902 class HelloHandler(RequestHandler):
903 def get(self):
904 self.finish("Hello world")
905
906 def post(self):
907 self.finish("Hello world")
908
909 class LargeHandler(RequestHandler):
910 def get(self):
911 # 512KB should be bigger than the socket buffers so it will
912 # be written out in chunks.
913 self.write("".join(chr(i % 256) * 1024 for i in range(512)))
914
915 class TransferEncodingChunkedHandler(RequestHandler):
916 @gen.coroutine
917 def head(self):
918 self.write("Hello world")
919 yield self.flush()
920
921 class FinishOnCloseHandler(RequestHandler):
922 def initialize(self, cleanup_event):
923 self.cleanup_event = cleanup_event
924
925 @gen.coroutine
926 def get(self):
927 self.flush()
928 yield self.cleanup_event.wait()
929
930 def on_connection_close(self):
931 # This is not very realistic, but finishing the request
932 # from the close callback has the right timing to mimic
933 # some errors seen in the wild.
934 self.finish("closed")
935
936 self.cleanup_event = Event()
937 return Application(
938 [
939 ("/", HelloHandler),
940 ("/large", LargeHandler),
941 ("/chunked", TransferEncodingChunkedHandler),
942 (
943 "/finish_on_close",
944 FinishOnCloseHandler,
945 dict(cleanup_event=self.cleanup_event),
946 ),
947 ]
948 )
949
950 def setUp(self):
951 super().setUp()
952 self.http_version = b"HTTP/1.1"
953
954 def tearDown(self):
955 # We just closed the client side of the socket; let the IOLoop run
956 # once to make sure the server side got the message.
957 self.io_loop.add_timeout(datetime.timedelta(seconds=0.001), self.stop)
958 self.wait()
959
960 if hasattr(self, "stream"):
961 self.stream.close()
962 super().tearDown()
963
964 # The next few methods are a crude manual http client
965 @gen.coroutine
966 def connect(self):
967 self.stream = IOStream(socket.socket())
968 yield self.stream.connect(("127.0.0.1", self.get_http_port()))
969
970 @gen.coroutine
971 def read_headers(self):
972 first_line = yield self.stream.read_until(b"\r\n")
973 self.assertTrue(first_line.startswith(b"HTTP/1.1 200"), first_line)
974 header_bytes = yield self.stream.read_until(b"\r\n\r\n")
975 headers = HTTPHeaders.parse(header_bytes.decode("latin1"))
976 raise gen.Return(headers)
977
978 @gen.coroutine
979 def read_response(self):
980 self.headers = yield self.read_headers()
981 body = yield self.stream.read_bytes(int(self.headers["Content-Length"]))
982 self.assertEqual(b"Hello world", body)
983
984 def close(self):
985 self.stream.close()
986 del self.stream
987
988 @gen_test
989 def test_two_requests(self):
990 yield self.connect()
991 self.stream.write(b"GET / HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n")
992 yield self.read_response()
993 self.stream.write(b"GET / HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n")
994 yield self.read_response()
995 self.close()
996
997 @gen_test
998 def test_request_close(self):
999 yield self.connect()
1000 self.stream.write(
1001 b"GET / HTTP/1.1\r\nHost:127.0.0.1\r\nConnection: close\r\n\r\n"
1002 )
1003 yield self.read_response()
1004 data = yield self.stream.read_until_close()
1005 self.assertTrue(not data)
1006 self.assertEqual(self.headers["Connection"], "close")
1007 self.close()
1008
1009 # keepalive is supported for http 1.0 too, but it's opt-in
1010 @gen_test
1011 def test_http10(self):
1012 self.http_version = b"HTTP/1.0"
1013 yield self.connect()
1014 self.stream.write(b"GET / HTTP/1.0\r\n\r\n")
1015 yield self.read_response()
1016 data = yield self.stream.read_until_close()
1017 self.assertFalse(data)
1018 self.assertNotIn("Connection", self.headers)
1019 self.close()
1020
1021 @gen_test
1022 def test_http10_keepalive(self):
1023 self.http_version = b"HTTP/1.0"
1024 yield self.connect()
1025 self.stream.write(b"GET / HTTP/1.0\r\nConnection: keep-alive\r\n\r\n")
1026 yield self.read_response()
1027 self.assertEqual(self.headers["Connection"], "Keep-Alive")
1028 self.stream.write(b"GET / HTTP/1.0\r\nConnection: keep-alive\r\n\r\n")
1029 yield self.read_response()
1030 self.assertEqual(self.headers["Connection"], "Keep-Alive")
1031 self.close()
1032
1033 @gen_test
1034 def test_http10_keepalive_extra_crlf(self):
1035 self.http_version = b"HTTP/1.0"
1036 yield self.connect()
1037 self.stream.write(b"GET / HTTP/1.0\r\nConnection: keep-alive\r\n\r\n\r\n")
1038 yield self.read_response()
1039 self.assertEqual(self.headers["Connection"], "Keep-Alive")
1040 self.stream.write(b"GET / HTTP/1.0\r\nConnection: keep-alive\r\n\r\n")
1041 yield self.read_response()
1042 self.assertEqual(self.headers["Connection"], "Keep-Alive")
1043 self.close()
1044
1045 @gen_test
1046 def test_pipelined_requests(self):
1047 yield self.connect()
1048 self.stream.write(
1049 b"GET / HTTP/1.1\r\nHost:127.0.0.1\r\n\r\nGET / HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n"
1050 )
1051 yield self.read_response()
1052 yield self.read_response()
1053 self.close()
1054
1055 @gen_test
1056 def test_pipelined_cancel(self):
1057 yield self.connect()
1058 self.stream.write(
1059 b"GET / HTTP/1.1\r\nHost:127.0.0.1\r\n\r\nGET / HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n"
1060 )
1061 # only read once
1062 yield self.read_response()
1063 self.close()
1064
1065 @gen_test
1066 def test_cancel_during_download(self):
1067 yield self.connect()
1068 self.stream.write(b"GET /large HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n")
1069 yield self.read_headers()
1070 yield self.stream.read_bytes(1024)
1071 self.close()
1072
1073 @gen_test
1074 def test_finish_while_closed(self):
1075 yield self.connect()
1076 self.stream.write(b"GET /finish_on_close HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n")
1077 yield self.read_headers()
1078 self.close()
1079 # Let the hanging coroutine clean up after itself
1080 self.cleanup_event.set()
1081
1082 @gen_test
1083 def test_keepalive_chunked(self):
1084 self.http_version = b"HTTP/1.0"
1085 yield self.connect()
1086 self.stream.write(
1087 b"POST / HTTP/1.0\r\n"
1088 b"Connection: keep-alive\r\n"
1089 b"Transfer-Encoding: chunked\r\n"
1090 b"\r\n"
1091 b"0\r\n"
1092 b"\r\n"
1093 )
1094 yield self.read_response()
1095 self.assertEqual(self.headers["Connection"], "Keep-Alive")
1096 self.stream.write(b"GET / HTTP/1.0\r\nConnection: keep-alive\r\n\r\n")
1097 yield self.read_response()
1098 self.assertEqual(self.headers["Connection"], "Keep-Alive")
1099 self.close()
1100
1101 @gen_test
1102 def test_keepalive_chunked_head_no_body(self):
1103 yield self.connect()
1104 self.stream.write(b"HEAD /chunked HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n")
1105 yield self.read_headers()
1106
1107 self.stream.write(b"HEAD /chunked HTTP/1.1\r\nHost:127.0.0.1\r\n\r\n")
1108 yield self.read_headers()
1109 self.close()
1110
1111
1112class GzipBaseTest(AsyncHTTPTestCase):
1113 def get_app(self):
1114 return Application([("/", EchoHandler)])
1115
1116 def post_gzip(self, body):
1117 bytesio = BytesIO()
1118 gzip_file = gzip.GzipFile(mode="w", fileobj=bytesio)
1119 gzip_file.write(utf8(body))
1120 gzip_file.close()
1121 compressed_body = bytesio.getvalue()
1122 return self.fetch(
1123 "/",
1124 method="POST",
1125 body=compressed_body,
1126 headers={"Content-Encoding": "gzip"},
1127 )
1128
1129 def test_uncompressed(self):
1130 response = self.fetch("/", method="POST", body="foo=bar")
1131 self.assertEqual(json_decode(response.body), {"foo": ["bar"]})
1132
1133
1134class GzipTest(GzipBaseTest, AsyncHTTPTestCase):
1135 def get_httpserver_options(self):
1136 return dict(decompress_request=True)
1137
1138 def test_gzip(self):
1139 response = self.post_gzip("foo=bar")
1140 self.assertEqual(json_decode(response.body), {"foo": ["bar"]})
1141
1142 def test_gzip_case_insensitive(self):
1143 # https://datatracker.ietf.org/doc/html/rfc7231#section-3.1.2.1
1144 bytesio = BytesIO()
1145 gzip_file = gzip.GzipFile(mode="w", fileobj=bytesio)
1146 gzip_file.write(utf8("foo=bar"))
1147 gzip_file.close()
1148 compressed_body = bytesio.getvalue()
1149 response = self.fetch(
1150 "/",
1151 method="POST",
1152 body=compressed_body,
1153 headers={"Content-Encoding": "GZIP"},
1154 )
1155 self.assertEqual(json_decode(response.body), {"foo": ["bar"]})
1156
1157
1158class GzipUnsupportedTest(GzipBaseTest, AsyncHTTPTestCase):
1159 def test_gzip_unsupported(self):
1160 # Gzip support is opt-in; without it the server fails to parse
1161 # the body (but parsing form bodies is currently just a log message,
1162 # not a fatal error).
1163 with ExpectLog(gen_log, ".*Unsupported Content-Encoding"):
1164 response = self.post_gzip("foo=bar")
1165 self.assertEqual(response.code, 400)
1166
1167
1168class StreamingChunkSizeTest(AsyncHTTPTestCase):
1169 # 50 characters long, and repetitive so it can be compressed.
1170 BODY = b"01234567890123456789012345678901234567890123456789"
1171 CHUNK_SIZE = 16
1172
1173 def get_http_client(self):
1174 # body_producer doesn't work on curl_httpclient, so override the
1175 # configured AsyncHTTPClient implementation.
1176 return SimpleAsyncHTTPClient()
1177
1178 def get_httpserver_options(self):
1179 return dict(chunk_size=self.CHUNK_SIZE, decompress_request=True)
1180
1181 class MessageDelegate(HTTPMessageDelegate):
1182 def __init__(self, connection):
1183 self.connection = connection
1184
1185 def headers_received(self, start_line, headers):
1186 self.chunk_lengths = [] # type: List[int]
1187
1188 def data_received(self, chunk):
1189 self.chunk_lengths.append(len(chunk))
1190
1191 def finish(self):
1192 response_body = utf8(json_encode(self.chunk_lengths))
1193 self.connection.write_headers(
1194 ResponseStartLine("HTTP/1.1", 200, "OK"),
1195 HTTPHeaders({"Content-Length": str(len(response_body))}),
1196 )
1197 self.connection.write(response_body)
1198 self.connection.finish()
1199
1200 def get_app(self):
