Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
websocket_test.py990 linesDownload Raw Back to test
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 
codekingpro/portable-devtools · Team Ai