codekingpro/portable-devtools
114k
1import asyncio
2import contextlib
3import datetime
4import functools
5import socket
6import traceback
7import typing
8import unittest
9
10from tornado.concurrent import Future
11from tornado import gen
12from tornado.httpclient import HTTPError, HTTPRequest
13from tornado.locks import Event
14from tornado.log import gen_log, app_log
15from tornado.netutil import Resolver
16from tornado.simple_httpclient import SimpleAsyncHTTPClient
17from tornado.template import DictLoader
18from tornado.test.util import abstract_base_test, ignore_deprecation
19from tornado.testing import AsyncHTTPTestCase, gen_test, bind_unused_port, ExpectLog
20from tornado.web import Application, RequestHandler
21
22try:
23 import tornado.websocket # noqa: F401
24 from tornado.util import _websocket_mask_python
25except ImportError:
26 # The unittest module presents misleading errors on ImportError
27 # (it acts as if websocket_test could not be found, hiding the underlying
28 # error). If we get an ImportError here (which could happen due to
29 # TORNADO_EXTENSION=1), print some extra information before failing.
30 traceback.print_exc()
31 raise
32
33from tornado.websocket import (
34 WebSocketHandler,
35 websocket_connect,
36 WebSocketError,
37 WebSocketClosedError,
38)
39
40try:
41 from tornado import speedups
42except ImportError:
43 speedups = None # type: ignore
44
45
46class TestWebSocketHandler(WebSocketHandler):
47 """Base class for testing handlers that exposes the on_close event.
48
49 This allows for tests to see the close code and reason on the
50 server side.
51
52 """
53
54 def initialize(self, close_future=None, compression_options=None):
55 self.close_future = close_future
56 self.compression_options = compression_options
57
58 def get_compression_options(self):
59 return self.compression_options
60
61 def on_close(self):
62 if self.close_future is not None:
63 self.close_future.set_result((self.close_code, self.close_reason))
64
65
66class EchoHandler(TestWebSocketHandler):
67 @gen.coroutine
68 def on_message(self, message):
69 try:
70 yield self.write_message(message, isinstance(message, bytes))
71 except asyncio.CancelledError:
72 pass
73 except WebSocketClosedError:
74 pass
75
76
77class ErrorInOnMessageHandler(TestWebSocketHandler):
78 def on_message(self, message):
79 1 / 0
80
81
82class HeaderHandler(TestWebSocketHandler):
83 def open(self):
84 methods_to_test = [
85 functools.partial(self.write, "This should not work"),
86 functools.partial(self.redirect, "http://localhost/elsewhere"),
87 functools.partial(self.set_header, "X-Test", ""),
88 functools.partial(self.set_cookie, "Chocolate", "Chip"),
89 functools.partial(self.set_status, 503),
90 self.flush,
91 self.finish,
92 ]
93 for method in methods_to_test:
94 try:
95 # In a websocket context, many RequestHandler methods
96 # raise RuntimeErrors.
97 method() # type: ignore
98 raise Exception("did not get expected exception")
99 except RuntimeError:
100 pass
101 self.write_message(self.request.headers.get("X-Test", ""))
102
103
104class HeaderEchoHandler(TestWebSocketHandler):
105 def set_default_headers(self):
106 self.set_header("X-Extra-Response-Header", "Extra-Response-Value")
107
108 def prepare(self):
109 for k, v in self.request.headers.get_all():
110 if k.lower().startswith("x-test"):
111 self.set_header(k, v)
112
113
114class NonWebSocketHandler(RequestHandler):
115 def get(self):
116 self.write("ok")
117
118
119class RedirectHandler(RequestHandler):
120 def get(self):
121 self.redirect("/echo")
122
123
124class CloseReasonHandler(TestWebSocketHandler):
125 def open(self):
126 self.on_close_called = False
127 self.close(1001, "goodbye")
128
129
130class AsyncPrepareHandler(TestWebSocketHandler):
131 @gen.coroutine
132 def prepare(self):
133 yield gen.moment
134
135 def on_message(self, message):
136 self.write_message(message)
137
138
139class PathArgsHandler(TestWebSocketHandler):
140 def open(self, arg):
141 self.write_message(arg)
142
143
144class CoroutineOnMessageHandler(TestWebSocketHandler):
145 def initialize(self, **kwargs):
146 super().initialize(**kwargs)
147 self.sleeping = 0
148
149 @gen.coroutine
150 def on_message(self, message):
151 if self.sleeping > 0:
152 self.write_message("another coroutine is already sleeping")
153 self.sleeping += 1
154 yield gen.sleep(0.01)
155 self.sleeping -= 1
156 self.write_message(message)
157
158
159class RenderMessageHandler(TestWebSocketHandler):
160 def on_message(self, message):
161 self.write_message(self.render_string("message.html", message=message))
162
163
164class SubprotocolHandler(TestWebSocketHandler):
165 def initialize(self, **kwargs):
166 super().initialize(**kwargs)
167 self.select_subprotocol_called = False
168
169 def select_subprotocol(self, subprotocols):
170 if self.select_subprotocol_called:
171 raise Exception("select_subprotocol called twice")
172 self.select_subprotocol_called = True
173 if "goodproto" in subprotocols:
174 return "goodproto"
175 return None
176
177 def open(self):
178 if not self.select_subprotocol_called:
179 raise Exception("select_subprotocol not called")
180 self.write_message("subprotocol=%s" % self.selected_subprotocol)
181
182
183class OpenCoroutineHandler(TestWebSocketHandler):
184 def initialize(self, test, **kwargs):
185 super().initialize(**kwargs)
186 self.test = test
187 self.open_finished = False
188
189 @gen.coroutine
190 def open(self):
191 yield self.test.message_sent.wait()
192 yield gen.sleep(0.010)
193 self.open_finished = True
194
195 def on_message(self, message):
196 if not self.open_finished:
197 raise Exception("on_message called before open finished")
198 self.write_message("ok")
199
200
201class ErrorInOpenHandler(TestWebSocketHandler):
202 def open(self):
203 raise Exception("boom")
204
205
206class ErrorInAsyncOpenHandler(TestWebSocketHandler):
207 async def open(self):
208 await asyncio.sleep(0)
209 raise Exception("boom")
210
211
212class NoDelayHandler(TestWebSocketHandler):
213 def open(self):
214 self.set_nodelay(True)
215 self.write_message("hello")
216
217
218class WebSocketBaseTestCase(AsyncHTTPTestCase):
219 def setUp(self):
220 super().setUp()
221 self.conns_to_close = []
222
223 def tearDown(self):
224 for conn in self.conns_to_close:
225 conn.close()
226 super().tearDown()
227
228 @gen.coroutine
229 def ws_connect(self, path, **kwargs):
230 ws = yield websocket_connect(
231 "ws://127.0.0.1:%d%s" % (self.get_http_port(), path), **kwargs
232 )
233 self.conns_to_close.append(ws)
234 raise gen.Return(ws)
235
236
237class WebSocketTest(WebSocketBaseTestCase):
238 def get_app(self):
239 self.close_future = Future() # type: Future[None]
240 return Application(
241 [
242 ("/echo", EchoHandler, dict(close_future=self.close_future)),
243 ("/non_ws", NonWebSocketHandler),
244 ("/redirect", RedirectHandler),
245 ("/header", HeaderHandler, dict(close_future=self.close_future)),
246 (
247 "/header_echo",
248 HeaderEchoHandler,
249 dict(close_future=self.close_future),
250 ),
251 (
252 "/close_reason",
253 CloseReasonHandler,
254 dict(close_future=self.close_future),
255 ),
256 (
257 "/error_in_on_message",
258 ErrorInOnMessageHandler,
259 dict(close_future=self.close_future),
260 ),
261 (
262 "/async_prepare",
263 AsyncPrepareHandler,
264 dict(close_future=self.close_future),
265 ),
266 (
267 "/path_args/(.*)",
268 PathArgsHandler,
269 dict(close_future=self.close_future),
270 ),
271 (
272 "/coroutine",
273 CoroutineOnMessageHandler,
274 dict(close_future=self.close_future),
275 ),
276 ("/render", RenderMessageHandler, dict(close_future=self.close_future)),
277 (
278 "/subprotocol",
279 SubprotocolHandler,
280 dict(close_future=self.close_future),
281 ),
282 (
283 "/open_coroutine",
284 OpenCoroutineHandler,
285 dict(close_future=self.close_future, test=self),
286 ),
287 ("/error_in_open", ErrorInOpenHandler),
288 ("/error_in_async_open", ErrorInAsyncOpenHandler),
289 ("/nodelay", NoDelayHandler),
290 ],
291 template_loader=DictLoader({"message.html": "<b>{{ message }}</b>"}),
292 )
293
294 def get_http_client(self):
295 # These tests require HTTP/1; force the use of SimpleAsyncHTTPClient.
296 return SimpleAsyncHTTPClient()
297
298 def tearDown(self):
299 super().tearDown()
300 RequestHandler._template_loaders.clear()
301
302 def test_http_request(self):
303 # WS server, HTTP client.
304 response = self.fetch("/echo")
305 self.assertEqual(response.code, 400)
306
307 def test_missing_websocket_key(self):
308 response = self.fetch(
309 "/echo",
310 headers={
311 "Connection": "Upgrade",
312 "Upgrade": "WebSocket",
313 "Sec-WebSocket-Version": "13",
314 },
315 )
316 self.assertEqual(response.code, 400)
317
318 def test_bad_websocket_version(self):
319 response = self.fetch(
320 "/echo",
321 headers={
322 "Connection": "Upgrade",
323 "Upgrade": "WebSocket",
324 "Sec-WebSocket-Version": "12",
325 },
326 )
327 self.assertEqual(response.code, 426)
328
329 @gen_test
330 def test_websocket_gen(self):
331 ws = yield self.ws_connect("/echo")
332 yield ws.write_message("hello")
333 response = yield ws.read_message()
334 self.assertEqual(response, "hello")
335
336 def test_websocket_callbacks(self):
337 with ignore_deprecation():
338 websocket_connect(
339 "ws://127.0.0.1:%d/echo" % self.get_http_port(), callback=self.stop
340 )
341 ws = self.wait().result()
342 ws.write_message("hello")
343 ws.read_message(self.stop)
344 response = self.wait().result()
345 self.assertEqual(response, "hello")
346 self.close_future.add_done_callback(lambda f: self.stop())
347 ws.close()
348 self.wait()
349
350 @gen_test
351 def test_binary_message(self):
352 ws = yield self.ws_connect("/echo")
353 ws.write_message(b"hello \xe9", binary=True)
354 response = yield ws.read_message()
355 self.assertEqual(response, b"hello \xe9")
356
357 @gen_test
358 def test_unicode_message(self):
359 ws = yield self.ws_connect("/echo")
360 ws.write_message("hello \u00e9")
361 response = yield ws.read_message()
362 self.assertEqual(response, "hello \u00e9")
363
364 @gen_test
365 def test_error_in_closed_client_write_message(self):
366 ws = yield self.ws_connect("/echo")
367 ws.close()
368 with self.assertRaises(WebSocketClosedError):
369 ws.write_message("hello \u00e9")
370
371 @gen_test
372 def test_render_message(self):
373 ws = yield self.ws_connect("/render")
374 ws.write_message("hello")
375 response = yield ws.read_message()
376 self.assertEqual(response, "<b>hello</b>")
377
378 @gen_test
379 def test_error_in_on_message(self):
380 ws = yield self.ws_connect("/error_in_on_message")
381 ws.write_message("hello")
382 with ExpectLog(app_log, "Uncaught exception"):
383 response = yield ws.read_message()
384 self.assertIsNone(response)
385
386 @gen_test
387 def test_websocket_http_fail(self):
388 with self.assertRaises(HTTPError) as cm:
389 yield self.ws_connect("/notfound")
390 self.assertEqual(cm.exception.code, 404)
391
392 @gen_test
393 def test_websocket_http_success(self):
394 with self.assertRaises(WebSocketError):
395 yield self.ws_connect("/non_ws")
396
397 @gen_test
398 def test_websocket_http_redirect(self):
399 with self.assertRaises(HTTPError):
400 yield self.ws_connect("/redirect")
401
402 @gen_test
403 def test_websocket_network_fail(self):
404 sock, port = bind_unused_port()
405 sock.close()
406 with self.assertRaises(IOError):
407 with ExpectLog(gen_log, ".*", required=False):
408 yield websocket_connect(
409 "ws://127.0.0.1:%d/" % port, connect_timeout=3600
410 )
411
412 @gen_test
413 def test_websocket_close_buffered_data(self):
414 with contextlib.closing(
415 (yield websocket_connect("ws://127.0.0.1:%d/echo" % self.get_http_port()))
416 ) as ws:
417 ws.write_message("hello")
418 ws.write_message("world")
419 # Close the underlying stream.
420 ws.stream.close()
421
422 @gen_test
423 def test_websocket_headers(self):
424 # Ensure that arbitrary headers can be passed through websocket_connect.
425 with contextlib.closing(
426 (
427 yield websocket_connect(
428 HTTPRequest(
429 "ws://127.0.0.1:%d/header" % self.get_http_port(),
430 headers={"X-Test": "hello"},
431 )
432 )
433 )
434 ) as ws:
435 response = yield ws.read_message()
436 self.assertEqual(response, "hello")
437
438 @gen_test
439 def test_websocket_header_echo(self):
440 # Ensure that headers can be returned in the response.
441 # Specifically, that arbitrary headers passed through websocket_connect
442 # can be returned.
443 with contextlib.closing(
444 (
445 yield websocket_connect(
446 HTTPRequest(
447 "ws://127.0.0.1:%d/header_echo" % self.get_http_port(),
448 headers={"X-Test-Hello": "hello"},
449 )
450 )
451 )
452 ) as ws:
453 self.assertEqual(ws.headers.get("X-Test-Hello"), "hello")
454 self.assertEqual(
455 ws.headers.get("X-Extra-Response-Header"), "Extra-Response-Value"
456 )
457
458 @gen_test
459 def test_server_close_reason(self):
460 ws = yield self.ws_connect("/close_reason")
461 msg = yield ws.read_message()
462 # A message of None means the other side closed the connection.
463 self.assertIs(msg, None)
464 self.assertEqual(ws.close_code, 1001)
465 self.assertEqual(ws.close_reason, "goodbye")
466 # The on_close callback is called no matter which side closed.
467 code, reason = yield self.close_future
468 # The client echoed the close code it received to the server,
469 # so the server's close code (returned via close_future) is
470 # the same.
471 self.assertEqual(code, 1001)
472
473 @gen_test
474 def test_client_close_reason(self):
475 ws = yield self.ws_connect("/echo")
476 ws.close(1001, "goodbye")
477 code, reason = yield self.close_future
478 self.assertEqual(code, 1001)
479 self.assertEqual(reason, "goodbye")
480
481 @gen_test
482 def test_write_after_close(self):
483 ws = yield self.ws_connect("/close_reason")
484 msg = yield ws.read_message()
485 self.assertIs(msg, None)
486 with self.assertRaises(WebSocketClosedError):
487 ws.write_message("hello")
488
489 @gen_test
490 def test_async_prepare(self):
491 # Previously, an async prepare method triggered a bug that would
492 # result in a timeout on test shutdown (and a memory leak).
493 ws = yield self.ws_connect("/async_prepare")
494 ws.write_message("hello")
495 res = yield ws.read_message()
496 self.assertEqual(res, "hello")
497
498 @gen_test
499 def test_path_args(self):
500 ws = yield self.ws_connect("/path_args/hello")
501 res = yield ws.read_message()
502 self.assertEqual(res, "hello")
503
504 @gen_test
505 def test_coroutine(self):
506 ws = yield self.ws_connect("/coroutine")
507 # Send both messages immediately, coroutine must process one at a time.
508 yield ws.write_message("hello1")
509 yield ws.write_message("hello2")
510 res = yield ws.read_message()
511 self.assertEqual(res, "hello1")
512 res = yield ws.read_message()
513 self.assertEqual(res, "hello2")
514
515 @gen_test
516 def test_check_origin_valid_no_path(self):
517 port = self.get_http_port()
518
519 url = "ws://127.0.0.1:%d/echo" % port
520 headers = {"Origin": "http://127.0.0.1:%d" % port}
521
522 with contextlib.closing(
523 (yield websocket_connect(HTTPRequest(url, headers=headers)))
524 ) as ws:
525 ws.write_message("hello")
526 response = yield ws.read_message()
527 self.assertEqual(response, "hello")
528
529 @gen_test
530 def test_check_origin_valid_with_path(self):
531 port = self.get_http_port()
532
533 url = "ws://127.0.0.1:%d/echo" % port
534 headers = {"Origin": "http://127.0.0.1:%d/something" % port}
535
536 with contextlib.closing(
537 (yield websocket_connect(HTTPRequest(url, headers=headers)))
538 ) as ws:
539 ws.write_message("hello")
540 response = yield ws.read_message()
541 self.assertEqual(response, "hello")
542
543 @gen_test
544 def test_check_origin_invalid_partial_url(self):
545 port = self.get_http_port()
546
547 url = "ws://127.0.0.1:%d/echo" % port
548 headers = {"Origin": "127.0.0.1:%d" % port}
549
550 with self.assertRaises(HTTPError) as cm:
551 yield websocket_connect(HTTPRequest(url, headers=headers))
552 self.assertEqual(cm.exception.code, 403)
553
554 @gen_test
555 def test_check_origin_invalid(self):
556 port = self.get_http_port()
557
558 url = "ws://127.0.0.1:%d/echo" % port
559 # Host is 127.0.0.1, which should not be accessible from some other
560 # domain
561 headers = {"Origin": "http://somewhereelse.com"}
562
563 with self.assertRaises(HTTPError) as cm:
564 yield websocket_connect(HTTPRequest(url, headers=headers))
565
566 self.assertEqual(cm.exception.code, 403)
567
568 @gen_test
569 def test_check_origin_invalid_subdomains(self):
570 port = self.get_http_port()
571
572 # CaresResolver may return ipv6-only results for localhost, but our
573 # server is only running on ipv4. Test for this edge case and skip
574 # the test if it happens.
575 addrinfo = yield Resolver().resolve("localhost", port)
576 families = {addr[0] for addr in addrinfo}
577 if socket.AF_INET not in families:
578 self.skipTest("localhost does not resolve to ipv4")
579 return
580
581 url = "ws://localhost:%d/echo" % port
582 # Subdomains should be disallowed by default. If we could pass a
583 # resolver to websocket_connect we could test sibling domains as well.
584 headers = {"Origin": "http://subtenant.localhost"}
585
586 with self.assertRaises(HTTPError) as cm:
587 yield websocket_connect(HTTPRequest(url, headers=headers))
588
589 self.assertEqual(cm.exception.code, 403)
590
591 @gen_test
592 def test_subprotocols(self):
593 ws = yield self.ws_connect(
594 "/subprotocol", subprotocols=["badproto", "goodproto"]
595 )
596 self.assertEqual(ws.selected_subprotocol, "goodproto")
597 res = yield ws.read_message()
598 self.assertEqual(res, "subprotocol=goodproto")
599
600 @gen_test
601 def test_subprotocols_not_offered(self):
602 ws = yield self.ws_connect("/subprotocol")
603 self.assertIs(ws.selected_subprotocol, None)
604 res = yield ws.read_message()
605 self.assertEqual(res, "subprotocol=None")
606
607 @gen_test
608 def test_open_coroutine(self):
609 self.message_sent = Event()
610 ws = yield self.ws_connect("/open_coroutine")
611 yield ws.write_message("hello")
612 self.message_sent.set()
613 res = yield ws.read_message()
614 self.assertEqual(res, "ok")
615
616 @gen_test
617 def test_error_in_open(self):
618 with ExpectLog(app_log, "Uncaught exception"):
619 ws = yield self.ws_connect("/error_in_open")
620 res = yield ws.read_message()
621 self.assertIsNone(res)
622
623 @gen_test
624 def test_error_in_async_open(self):
625 with ExpectLog(app_log, "Uncaught exception"):
626 ws = yield self.ws_connect("/error_in_async_open")
627 res = yield ws.read_message()
628 self.assertIsNone(res)
629
630 @gen_test
631 def test_nodelay(self):
632 ws = yield self.ws_connect("/nodelay")
633 res = yield ws.read_message()
634 self.assertEqual(res, "hello")
635
636
637class NativeCoroutineOnMessageHandler(TestWebSocketHandler):
638 def initialize(self, **kwargs):
639 super().initialize(**kwargs)
640 self.sleeping = 0
641
642 async def on_message(self, message):
643 if self.sleeping > 0:
644 self.write_message("another coroutine is already sleeping")
645 self.sleeping += 1
646 await gen.sleep(0.01)
647 self.sleeping -= 1
648 self.write_message(message)
649
650
651class WebSocketNativeCoroutineTest(WebSocketBaseTestCase):
652 def get_app(self):
653 return Application([("/native", NativeCoroutineOnMessageHandler)])
654
655 @gen_test
656 def test_native_coroutine(self):
657 ws = yield self.ws_connect("/native")
658 # Send both messages immediately, coroutine must process one at a time.
659 yield ws.write_message("hello1")
660 yield ws.write_message("hello2")
661 res = yield ws.read_message()
662 self.assertEqual(res, "hello1")
663 res = yield ws.read_message()
664 self.assertEqual(res, "hello2")
665
666
667@abstract_base_test
668class CompressionTestMixin(WebSocketBaseTestCase):
669 MESSAGE = "Hello world. Testing 123 123"
670
671 def get_app(self):
672 class LimitedHandler(TestWebSocketHandler):
673 @property
674 def max_message_size(self):
675 return 1024
676
677 def on_message(self, message):
678 self.write_message(str(len(message)))
679
680 return Application(
681 [
682 (
683 "/echo",
684 EchoHandler,
685 dict(compression_options=self.get_server_compression_options()),
686 ),
687 (
688 "/limited",
689 LimitedHandler,
690 dict(compression_options=self.get_server_compression_options()),
691 ),
692 ]
693 )
694
695 def get_server_compression_options(self):
696 return None
697
698 def get_client_compression_options(self):
699 return None
700
701 def verify_wire_bytes(self, bytes_in: int, bytes_out: int) -> None:
702 raise NotImplementedError()
703
704 @gen_test
705 def test_message_sizes(self):
706 ws = yield self.ws_connect(
707 "/echo", compression_options=self.get_client_compression_options()
708 )
709 # Send the same message three times so we can measure the
710 # effect of the context_takeover options.
711 for i in range(3):
712 ws.write_message(self.MESSAGE)
713 response = yield ws.read_message()
714 self.assertEqual(response, self.MESSAGE)
715 self.assertEqual(ws.protocol._message_bytes_out, len(self.MESSAGE) * 3)
716 self.assertEqual(ws.protocol._message_bytes_in, len(self.MESSAGE) * 3)
717 self.verify_wire_bytes(ws.protocol._wire_bytes_in, ws.protocol._wire_bytes_out)
718
719 @gen_test
720 def test_size_limit(self):
721 ws = yield self.ws_connect(
722 "/limited", compression_options=self.get_client_compression_options()
723 )
724 # Small messages pass through.
725 ws.write_message("a" * 128)
726 response = yield ws.read_message()
727 self.assertEqual(response, "128")
728 # This message is too big after decompression, but it compresses
729 # down to a size that will pass the initial checks.
730 ws.write_message("a" * 2048)
731 response = yield ws.read_message()
732 self.assertIsNone(response)
733
734
735@abstract_base_test
736class UncompressedTestMixin(CompressionTestMixin):
737 """Specialization of CompressionTestMixin when we expect no compression."""
738
739 def verify_wire_bytes(self, bytes_in, bytes_out):
740 # Bytes out includes the 4-byte mask key per message.
741 self.assertEqual(bytes_out, 3 * (len(self.MESSAGE) + 6))
742 self.assertEqual(bytes_in, 3 * (len(self.MESSAGE) + 2))
743
744
745class NoCompressionTest(UncompressedTestMixin):
746 pass
747
748
749# If only one side tries to compress, the extension is not negotiated.
750class ServerOnlyCompressionTest(UncompressedTestMixin):
751 def get_server_compression_options(self):
752 return {}
753
754
755class ClientOnlyCompressionTest(UncompressedTestMixin):
756 def get_client_compression_options(self):
757 return {}
758
759
760class DefaultCompressionTest(CompressionTestMixin):
761 def get_server_compression_options(self):
762 return {}
763
764 def get_client_compression_options(self):
765 return {}
766
767 def verify_wire_bytes(self, bytes_in, bytes_out):
768 self.assertLess(bytes_out, 3 * (len(self.MESSAGE) + 6))
769 self.assertLess(bytes_in, 3 * (len(self.MESSAGE) + 2))
770 # Bytes out includes the 4 bytes mask key per message.
771 self.assertEqual(bytes_out, bytes_in + 12)
772
773
774@abstract_base_test
775class MaskFunctionMixin(unittest.TestCase):
776 # Subclasses should define self.mask(mask, data)
777 def mask(self, mask: bytes, data: bytes) -> bytes:
778 raise NotImplementedError()
779
780 def test_mask(self: typing.Any):
781 self.assertEqual(self.mask(b"abcd", b""), b"")
782 self.assertEqual(self.mask(b"abcd", b"b"), b"\x03")
783 self.assertEqual(self.mask(b"abcd", b"54321"), b"TVPVP")
784 self.assertEqual(self.mask(b"ZXCV", b"98765432"), b"c`t`olpd")
785 # Include test cases with \x00 bytes (to ensure that the C
786 # extension isn't depending on null-terminated strings) and
787 # bytes with the high bit set (to smoke out signedness issues).
788 self.assertEqual(
789 self.mask(b"\x00\x01\x02\x03", b"\xff\xfb\xfd\xfc\xfe\xfa"),
790 b"\xff\xfa\xff\xff\xfe\xfb",
791 )
792 self.assertEqual(
793 self.mask(b"\xff\xfb\xfd\xfc", b"\x00\x01\x02\x03\x04\x05"),
794 b"\xff\xfa\xff\xff\xfb\xfe",
795 )
796
797
798class PythonMaskFunctionTest(MaskFunctionMixin):
799 def mask(self, mask, data):
800 return _websocket_mask_python(mask, data)
801
802
803@unittest.skipIf(speedups is None, "tornado.speedups module not present")
804class CythonMaskFunctionTest(MaskFunctionMixin):
805 def mask(self, mask, data):
806 return speedups.websocket_mask(mask, data)
807
808
809class ServerPeriodicPingTest(WebSocketBaseTestCase):
810 def get_app(self):
811 class PingHandler(TestWebSocketHandler):
812 def on_pong(self, data):
813 self.write_message("got pong")
814
815 return Application(
816 [("/", PingHandler)],
817 websocket_ping_interval=0.01,
818 websocket_ping_timeout=0,
819 )
820
821 @gen_test
822 def test_server_ping(self):
823 ws = yield self.ws_connect("/")
824 for i in range(3):
825 response = yield ws.read_message()
826 self.assertEqual(response, "got pong")
827 # TODO: test that the connection gets closed if ping responses stop.
828
829
830class ClientPeriodicPingTest(WebSocketBaseTestCase):
831 def get_app(self):
832 class PingHandler(TestWebSocketHandler):
833 def on_ping(self, data):
834 self.write_message("got ping")
835
836 return Application([("/", PingHandler)])
837
838 @gen_test
839 def test_client_ping(self):
840 ws = yield self.ws_connect("/", ping_interval=0.01, ping_timeout=0)
841 for i in range(3):
842 response = yield ws.read_message()
843 self.assertEqual(response, "got ping")
844 ws.close()
845
846
847class ServerPingTimeoutTest(WebSocketBaseTestCase):
848 def get_app(self):
849 self.handlers: list[WebSocketHandler] = []
850 test = self
851
852 class PingHandler(TestWebSocketHandler):
853 def initialize(self, close_future=None, compression_options=None):
854 self.handlers = test.handlers
855 # capture the handler instance so we can interrogate it later
856 self.handlers.append(self)
857 return super().initialize(
858 close_future=close_future, compression_options=compression_options
859 )
860
861 app = Application([("/", PingHandler)])
862 return app
863
864 @staticmethod
865 def install_hook(ws):
866 """Optionally suppress the client's "pong" response."""
867
868 ws.drop_pongs = False
869 ws.pongs_received = 0
870
871 def wrapper(fcn):
872 def _inner(opcode: int, data: bytes):
873 if opcode == 0xA: # NOTE: 0x9=ping, 0xA=pong
874 ws.pongs_received += 1
875 if ws.drop_pongs:
876 # prevent pong responses
877 return
878 # leave all other responses unchanged
879 return fcn(opcode, data)
880
881 return _inner
882
883 ws.protocol._handle_message = wrapper(ws.protocol._handle_message)
884
885 @gen_test
886 def test_client_ping_timeout(self):
887 # websocket client
888 interval = 0.2
889 ws = yield self.ws_connect(
890 "/", ping_interval=interval, ping_timeout=interval / 4
891 )
892 self.install_hook(ws)
893
894 # websocket handler (server side)
895 handler = self.handlers[0]
896
897 for _ in range(5):
898 # wait for the ping period
899 yield gen.sleep(interval)
900
901 # connection should still be open from the server end
902 self.assertIsNone(handler.close_code)
903 self.assertIsNone(handler.close_reason)
904
905 # connection should still be open from the client end
906 assert ws.protocol.close_code is None
907
908 # Check that our hook is intercepting messages; allow for
909 # some variance in timing (due to e.g. cpu load)
910 self.assertGreaterEqual(ws.pongs_received, 4)
911
912 # suppress the pong response message
913 ws.drop_pongs = True
914
915 # give the server time to register this
916 yield gen.sleep(interval * 1.5)
917
918 # connection should be closed from the server side
919 self.assertEqual(handler.close_code, 1000)
920 self.assertEqual(handler.close_reason, "ping timed out")
921
922 # client should have received a close operation
923 self.assertEqual(ws.protocol.close_code, 1000)
924
925
926class PingCalculationTest(unittest.TestCase):
927 def test_ping_sleep_time(self):
928 from tornado.websocket import WebSocketProtocol13
929
930 now = datetime.datetime(2025, 1, 1, 12, 0, 0, tzinfo=datetime.timezone.utc)
931 interval = 10 # seconds
932 last_ping_time = datetime.datetime(
933 2025, 1, 1, 11, 59, 54, tzinfo=datetime.timezone.utc
934 )
935 sleep_time = WebSocketProtocol13.ping_sleep_time(
936 last_ping_time=last_ping_time.timestamp(),
937 interval=interval,
938 now=now.timestamp(),
939 )
940 self.assertEqual(sleep_time, 4)
941
942
943class ManualPingTest(WebSocketBaseTestCase):
944 def get_app(self):
945 class PingHandler(TestWebSocketHandler):
946 def on_ping(self, data):
947 self.write_message(data, binary=isinstance(data, bytes))
948
949 return Application([("/", PingHandler)])
950
951 @gen_test
952 def test_manual_ping(self):
953 ws = yield self.ws_connect("/")
954
955 self.assertRaises(ValueError, ws.ping, "a" * 126)
956
957 ws.ping("hello")
958 resp = yield ws.read_message()
959 # on_ping always sees bytes.
960 self.assertEqual(resp, b"hello")
961
962 ws.ping(b"binary hello")
963 resp = yield ws.read_message()
964 self.assertEqual(resp, b"binary hello")
965
966
967class MaxMessageSizeTest(WebSocketBaseTestCase):
968 def get_app(self):
969 return Application([("/", EchoHandler)], websocket_max_message_size=1024)
970
971 @gen_test
972 def test_large_message(self):
973 ws = yield self.ws_connect("/")
974
975 # Write a message that is allowed.
976 msg = "a" * 1024
977 ws.write_message(msg)
978 resp = yield ws.read_message()
979 self.assertEqual(resp, msg)
980
981 # Write a message that is too large.
982 ws.write_message(msg + "b")
983 resp = yield ws.read_message()
984 # A message of None means the other side closed the connection.
985 self.assertIs(resp, None)
986 self.assertEqual(ws.close_code, 1009)
987 self.assertEqual(ws.close_reason, "message too big")
988 # TODO: Needs tests of messages split over multiple
989 # continuation frames.
990 