diff --git a/addons/bus/__init__.py b/addons/bus/__init__.py index 3b38916015c..9ab5f1778cd 100644 --- a/addons/bus/__init__.py +++ b/addons/bus/__init__.py @@ -1,3 +1,4 @@ # -*- coding: utf-8 -*- from . import models from . import controllers +from . import websocket diff --git a/addons/bus/controllers/__init__.py b/addons/bus/controllers/__init__.py index 757b12a1f17..7c4fc4891f6 100644 --- a/addons/bus/controllers/__init__.py +++ b/addons/bus/controllers/__init__.py @@ -1,2 +1,3 @@ # -*- coding: utf-8 -*- from . import main +from . import websocket diff --git a/addons/bus/controllers/websocket.py b/addons/bus/controllers/websocket.py new file mode 100644 index 00000000000..49403c3537c --- /dev/null +++ b/addons/bus/controllers/websocket.py @@ -0,0 +1,14 @@ +# Part of Odoo. See LICENSE file for full copyright and licensing details. + +from odoo.http import Controller, request, route +from ..websocket import WebsocketConnectionHandler + + +class WebsocketController(Controller): + @route('/websocket', type="http", auth="public", cors='*', websocket=True) + def websocket(self): + """ + Handle the websocket handshake, upgrade the connection if + successfull. + """ + return WebsocketConnectionHandler.open_connection(request) diff --git a/addons/bus/tests/__init__.py b/addons/bus/tests/__init__.py index 7cd3a9d428b..c8f25a068a1 100644 --- a/addons/bus/tests/__init__.py +++ b/addons/bus/tests/__init__.py @@ -1,2 +1,4 @@ +from . import common from . import test_assetsbundle from . import test_health +from . import test_websocket_caryall diff --git a/addons/bus/tests/common.py b/addons/bus/tests/common.py new file mode 100644 index 00000000000..9fdbafbbf67 --- /dev/null +++ b/addons/bus/tests/common.py @@ -0,0 +1,94 @@ +# Part of Odoo. See LICENSE file for full copyright and licensing details. + +import struct +from threading import Event +import unittest +from unittest.mock import patch + +try: + import websocket +except ImportError: + websocket = None + +import odoo.tools +from odoo.tests import HOST, common +from odoo.addons.bus.websocket import CloseCode, WebsocketConnectionHandler + + +class WebsocketCase(common.HttpCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + if websocket is None: + cls._logger.warning("websocket-client module is not installed") + raise unittest.SkipTest("websocket-client module is not installed") + cls._WEBSOCKET_URL = f"ws://{HOST}:{odoo.tools.config['http_port']}/websocket" + + def setUp(self): + super().setUp() + self._websockets = set() + # Used to ensure websocket connections have been closed + # properly. + self._websocket_events = set() + original_serve_forever = WebsocketConnectionHandler._serve_forever + + def _mocked_serve_forever(*args): + websocket_closed_event = Event() + self._websocket_events.add(websocket_closed_event) + original_serve_forever(*args) + websocket_closed_event.set() + + self._serve_forever_patch = patch.object( + WebsocketConnectionHandler, + '_serve_forever', + wraps=_mocked_serve_forever + ) + self._serve_forever_patch.start() + self.addCleanup(self._serve_forever_patch.stop) + + def tearDown(self): + self._close_websockets() + super().tearDown() + + def _close_websockets(self): + """ + Close all the connected websockets and wait for the connection + to terminate. + """ + for ws in self._websockets: + if ws.connected: + ws.close(CloseCode.CLEAN) + self.wait_remaining_websocket_connections() + + def websocket_connect(self, *args, **kwargs): + """ + Connect a websocket. If no cookie is given, the connection is + opened with a default session. The created websocket is closed + at the end of the test. + """ + if 'cookie' not in kwargs: + self.session = self.authenticate(None, None) + kwargs['cookie'] = f'session_id={self.session.sid}' + if 'timeout' not in kwargs: + kwargs['timeout'] = 5 + ws = websocket.create_connection( + type(self)._WEBSOCKET_URL, *args, **kwargs + ) + self._websockets.add(ws) + return ws + + def wait_remaining_websocket_connections(self): + """ Wait for the websocket connections to terminate. """ + for event in self._websocket_events: + event.wait(5) + + def assert_close_with_code(self, websocket, expected_code): + """ + Assert that the websocket is closed with the expected_code. + """ + opcode, payload = websocket.recv_data() + # ensure it's a close frame + self.assertEqual(opcode, 8) + code = struct.unpack('!H', payload[:2])[0] + # ensure the close code is the one we expected + self.assertEqual(code, expected_code) diff --git a/addons/bus/tests/test_websocket_caryall.py b/addons/bus/tests/test_websocket_caryall.py new file mode 100644 index 00000000000..16d50346143 --- /dev/null +++ b/addons/bus/tests/test_websocket_caryall.py @@ -0,0 +1,80 @@ +# Part of Odoo. See LICENSE file for full copyright and licensing details. + +import gc +from datetime import timedelta +from freezegun import freeze_time + +from odoo.tests import common +from .common import WebsocketCase +from ..websocket import ( + CloseCode, + Frame, + Opcode, + TimeoutManager, + TimeoutReason, + Websocket +) + + +@common.tagged('post_install', '-at_install') +class TestWebsocketCaryall(WebsocketCase): + def test_instances_weak_set(self): + gc.collect() + first_ws = self.websocket_connect() + second_ws = self.websocket_connect() + self.assertEqual(len(Websocket._instances), 2) + first_ws.close(CloseCode.CLEAN) + second_ws.close(CloseCode.CLEAN) + self.wait_remaining_websocket_connections() + # serve_forever_patch prevent websocket instances from being + # collected. Stop it now. + self._serve_forever_patch.stop() + gc.collect() + self.assertEqual(len(Websocket._instances), 0) + + def test_timeout_manager_no_response_timeout(self): + with freeze_time('2022-08-19') as frozen_time: + timeout_manager = TimeoutManager() + # A PING frame was just sent, if no pong has been received + # within TIMEOUT seconds, the connection should have timed out. + timeout_manager.acknowledge_frame_sent(Frame(Opcode.PING)) + self.assertEqual(timeout_manager._awaited_opcode, Opcode.PONG) + frozen_time.tick(delta=timedelta(seconds=TimeoutManager.TIMEOUT / 2)) + self.assertFalse(timeout_manager.has_timed_out()) + frozen_time.tick(delta=timedelta(seconds=TimeoutManager.TIMEOUT / 2)) + self.assertTrue(timeout_manager.has_timed_out()) + self.assertEqual(timeout_manager.timeout_reason, TimeoutReason.NO_RESPONSE) + + timeout_manager = TimeoutManager() + # A CLOSE frame was just sent, if no close has been received + # within TIMEOUT seconds, the connection should have timed out. + timeout_manager.acknowledge_frame_sent(Frame(Opcode.CLOSE)) + self.assertEqual(timeout_manager._awaited_opcode, Opcode.CLOSE) + frozen_time.tick(delta=timedelta(seconds=TimeoutManager.TIMEOUT / 2)) + self.assertFalse(timeout_manager.has_timed_out()) + frozen_time.tick(delta=timedelta(seconds=TimeoutManager.TIMEOUT / 2)) + self.assertTrue(timeout_manager.has_timed_out()) + self.assertEqual(timeout_manager.timeout_reason, TimeoutReason.NO_RESPONSE) + + def test_timeout_manager_keep_alive_timeout(self): + with freeze_time('2022-08-19') as frozen_time: + timeout_manager = TimeoutManager() + frozen_time.tick(delta=timedelta(seconds=TimeoutManager.KEEP_ALIVE_TIMEOUT / 2)) + self.assertFalse(timeout_manager.has_timed_out()) + frozen_time.tick(delta=timedelta(seconds=TimeoutManager.KEEP_ALIVE_TIMEOUT / 2)) + self.assertTrue(timeout_manager.has_timed_out()) + self.assertEqual(timeout_manager.timeout_reason, TimeoutReason.KEEP_ALIVE) + + def test_timeout_manager_reset_wait_for(self): + timeout_manager = TimeoutManager() + # PING frame + timeout_manager.acknowledge_frame_sent(Frame(Opcode.PING)) + self.assertEqual(timeout_manager._awaited_opcode, Opcode.PONG) + timeout_manager.acknowledge_frame_receipt(Frame(Opcode.PONG)) + self.assertIsNone(timeout_manager._awaited_opcode) + + # CLOSE frame + timeout_manager.acknowledge_frame_sent(Frame(Opcode.CLOSE)) + self.assertEqual(timeout_manager._awaited_opcode, Opcode.CLOSE) + timeout_manager.acknowledge_frame_receipt(Frame(Opcode.CLOSE)) + self.assertIsNone(timeout_manager._awaited_opcode) diff --git a/addons/bus/websocket.py b/addons/bus/websocket.py new file mode 100644 index 00000000000..4fe39af53fd --- /dev/null +++ b/addons/bus/websocket.py @@ -0,0 +1,679 @@ +import base64 +import functools +import hashlib +import json +import logging +import queue +import socket +import struct +import selectors +import threading +import time +from contextlib import suppress +from enum import IntEnum +from itertools import count +from weakref import WeakSet + +from werkzeug.exceptions import BadRequest, HTTPException + +from odoo.http import Response +from odoo.service.server import CommonServer +from odoo.tools import config + +_logger = logging.getLogger(__name__) + + +# Idea taken from the python cookbook: +# https://github.com/dabeaz/python-cookbook/blob/6e46b78e5644b3e5bf7426d900e2203b7cc630da/src/12/polling_multiple_thread_queues/pqueue.py +class PollablePriorityQueue(queue.PriorityQueue): + """ A custom PriorityQueue than can be polled """ + # This class allow this queue to be used with select. + def __init__(self): + super().__init__() + self._putsocket, self._getsocket = socket.socketpair() + + def fileno(self): + return self._getsocket.fileno() + + def put(self, item, **kwargs): + super().put(item, **kwargs) + self._putsocket.send(b'x') + + def get(self, **kwargs): + self._getsocket.recv(1) + return super().get(**kwargs) + + +# ------------------------------------------------------ +# EXCEPTIONS +# ------------------------------------------------------ + +class UpgradeRequired(HTTPException): + code = 426 + description = "Wrong websocket version was given during the handshake" + + def get_headers(self, environ=None): + headers = super().get_headers(environ) + headers.append(( + 'Sec-WebSocket-Version', + '; '.join(WebsocketConnectionHandler.SUPPORTED_VERSIONS) + )) + return headers + + +class WebsocketException(Exception): + """ Base class for all websockets exceptions """ + + +class ConnectionClosed(WebsocketException): + """ + Raised when the other end closes the socket without performing + the closing handshake. + """ + + +class InvalidCloseCodeException(WebsocketException): + def __init__(self, code): + super().__init__(f"Invalid close code: {code}") + + +class InvalidStateException(WebsocketException): + """ + Raised when an operation is forbidden in the current state. + """ + + +class PayloadTooLargeException(WebsocketException): + """ + Raised when a websocket message is too large. + """ + + +class ProtocolError(WebsocketException): + """ + Raised when a frame format doesn't match expectations. + """ + + +# ------------------------------------------------------ +# WEBSOCKET +# ------------------------------------------------------ + + +class Opcode(IntEnum): + CONTINUE = 0x00 + TEXT = 0x01 + BINARY = 0x02 + CLOSE = 0x08 + PING = 0x09 + PONG = 0x0A + + +class CloseCode(IntEnum): + CLEAN = 1000 + GOING_AWAY = 1001 + PROTOCOL_ERROR = 1002 + INCORRECT_DATA = 1003 + ABNORMAL_CLOSURE = 1006 + INCONSISTENT_DATA = 1007 + MESSAGE_VIOLATING_POLICY = 1008 + MESSAGE_TOO_BIG = 1009 + EXTENSION_NEGOTIATION_FAILED = 1010 + SERVER_ERROR = 1011 + RESTART = 1012 + TRY_LATER = 1013 + BAD_GATEWAY = 1014 + KEEP_ALIVE_TIMEOUT = 4002 + + +class ConnectionState(IntEnum): + OPEN = 0 + CLOSING = 1 + CLOSED = 2 + + +DATA_OP = {Opcode.TEXT, Opcode.BINARY} +CTRL_OP = {Opcode.CLOSE, Opcode.PING, Opcode.PONG} +HEARTBEAT_OP = {Opcode.PING, Opcode.PONG} + +VALID_CLOSE_CODES = { + code for code in CloseCode if code is not CloseCode.ABNORMAL_CLOSURE +} +CLEAN_CLOSE_CODES = {CloseCode.CLEAN, CloseCode.GOING_AWAY, CloseCode.RESTART} +RESERVED_CLOSE_CODES = range(3000, 5000) + +_XOR_TABLE = [bytes(a ^ b for a in range(256)) for b in range(256)] + + +class Frame: + # This class implements the `__lt__` method in order for frames to + # be stored in a `PriorityQueue`: ping/pong frames are prioritary. + _frames_sent = count(0) + + def __init__( + self, + opcode, + payload=b'', + fin=True, + rsv1=False, + rsv2=False, + rsv3=False + ): + self._send_order = next(self._frames_sent) + self.opcode = opcode + self.payload = payload + self.fin = fin + self.rsv1 = rsv1 + self.rsv2 = rsv2 + self.rsv3 = rsv3 + + def __lt__(self, other): + if not isinstance(other, Frame): + return NotImplemented + if ( + self.opcode in HEARTBEAT_OP and + other.opcode in HEARTBEAT_OP or + self.opcode in DATA_OP and other.opcode in DATA_OP + ): + return self._send_order < other._send_order + return self.opcode in HEARTBEAT_OP + + +class CloseFrame(Frame): + def __init__(self, code, reason): + if code not in VALID_CLOSE_CODES and code not in RESERVED_CLOSE_CODES: + raise InvalidCloseCodeException(code) + payload = struct.pack('!H', code) + if reason: + payload += reason.encode('utf-8') + self.code = code + self.reason = reason + super().__init__(Opcode.CLOSE, payload) + + +class Websocket: + _instances = WeakSet() + # Maximum size for a message in bytes, whether it is sent as one + # frame or many fragmented ones. + MESSAGE_MAX_SIZE = 2 ** 20 + # Proxies usually close a connection after 1 minute of inactivity. + # Therefore, a PING frame have to be sent if no frame is either sent + # or received within CONNECTION_TIMEOUT - 15 seconds. + CONNECTION_TIMEOUT = 60 + INACTIVITY_TIMEOUT = CONNECTION_TIMEOUT - 15 + + def __init__(self, socket): + self._socket = socket + self._close_sent = False + self._close_received = False + self._outgoing_frame_queue = PollablePriorityQueue() + self._selector = selectors.DefaultSelector() + self._selector.register(self._socket, selectors.EVENT_READ) + self._selector.register(self._outgoing_frame_queue, selectors.EVENT_READ) + self._timeout_manager = TimeoutManager() + self.state = ConnectionState.OPEN + type(self)._instances.add(self) + + # ------------------------------------------------------ + # PUBLIC METHODS + # ------------------------------------------------------ + + def get_messages(self): + while self.state is not ConnectionState.CLOSED: + try: + readables = { + selector_key[0].fileobj for selector_key in + self._selector.select(type(self).INACTIVITY_TIMEOUT) + } + if self._timeout_manager.has_timed_out(): + self.disconnect( + CloseCode.ABNORMAL_CLOSURE + if self._timeout_manager.timeout_reason is TimeoutReason.NO_RESPONSE + else CloseCode.KEEP_ALIVE_TIMEOUT + ) + break + if not readables: + self._enqueue_ping_frame() + continue + if self._outgoing_frame_queue in readables: + self._send_next_frame() + if self.state is ConnectionState.CLOSED: + break + if self._socket in readables: + message = self._process_next_message() + if message is not None: + yield message + except Exception as exc: + self._handle_transport_error(exc) + + def send(self, message): + if self.state is not ConnectionState.OPEN: + raise InvalidStateException( + "Trying to send a frame on a closed socket" + ) + opcode = Opcode.BINARY + if not isinstance(message, (bytes, bytearray)): + opcode = Opcode.TEXT + self._outgoing_frame_queue.put(Frame(opcode, message)) + + def disconnect(self, code, reason=None): + """ + Initiate the closing handshake that is, send a close frame + to the other end which will then send us back an + acknowledgment. Upon the reception of this acknowledgment, + the `_terminate` method will be called to perform an + orderly shutdown. Note that we don't need to wait for the + acknowledgment if the connection was failed beforewards. + """ + if code is not CloseCode.ABNORMAL_CLOSURE: + self._enqueue_close_frame(code, reason) + else: + self._terminate() + + # ------------------------------------------------------ + # PRIVATE METHODS + # ------------------------------------------------------ + + def _get_next_frame(self): + # 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 + # +-+-+-+-+-------+-+-------------+-------------------------------+ + # |F|R|R|R| opcode|M| Payload len | Extended payload length | + # |I|S|S|S| (4) |A| (7) | (16/64) | + # |N|V|V|V| |S| | (if payload len==126/127) | + # | |1|2|3| |K| | | + # +-+-+-+-+-------+-+-------------+ - - - - - - - - - - - - - - - + + # | Extended payload length continued, if payload len == 127 | + # + - - - - - - - - - - - - - - - +-------------------------------+ + # | |Masking-key, if MASK set to 1 | + # +-------------------------------+-------------------------------+ + # | Masking-key (continued) | Payload Data | + # +-------------------------------- - - - - - - - - - - - - - - - + + # : Payload Data continued ... : + # + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + # | Payload Data continued ... | + # +---------------------------------------------------------------+ + def recv_bytes(n): + """ Pull n bytes from the socket """ + data = bytearray() + while len(data) < n: + received_data = self._socket.recv(n - len(data)) + if not received_data: + raise ConnectionClosed() + data.extend(received_data) + return data + + def is_bit_set(byte, n): + """ + Check whether nth bit of byte is set or not (from left + to right). + """ + return byte & (1 << (7 - n)) + + def apply_mask(payload, mask): + # see: https://www.willmcgugan.com/blog/tech/post/speeding-up-websockets-60x/ + a, b, c, d = (_XOR_TABLE[n] for n in mask) + payload[::4] = payload[::4].translate(a) + payload[1::4] = payload[1::4].translate(b) + payload[2::4] = payload[2::4].translate(c) + payload[3::4] = payload[3::4].translate(d) + return payload + + first_byte, second_byte = recv_bytes(2) + fin, rsv1, rsv2, rsv3 = (is_bit_set(first_byte, n) for n in range(4)) + try: + opcode = Opcode(first_byte & 0b00001111) + except ValueError as exc: + raise ProtocolError(exc) + payload_length = second_byte & 0b01111111 + + if rsv1 or rsv2 or rsv3: + raise ProtocolError("Reserved bits must be unset") + if not is_bit_set(second_byte, 0): + raise ProtocolError("Frame must be masked") + if opcode in CTRL_OP: + if not fin: + raise ProtocolError("Control frames cannot be fragmented") + if payload_length > 125: + raise ProtocolError( + "Control frames payload must be smaller than 126" + ) + if payload_length == 126: + payload_length = struct.unpack('!H', recv_bytes(2))[0] + elif payload_length == 127: + payload_length = struct.unpack('!Q', recv_bytes(8))[0] + if payload_length > type(self).MESSAGE_MAX_SIZE: + raise PayloadTooLargeException() + + mask = recv_bytes(4) + payload = apply_mask(recv_bytes(payload_length), mask) + frame = Frame(opcode, bytes(payload), fin, rsv1, rsv2, rsv3) + self._timeout_manager.acknowledge_frame_receipt(frame) + return frame + + def _send_next_frame(self): + """ Send the next frame available in the outgoing frame queue. """ + frame = self._outgoing_frame_queue.get_nowait() + self._send_frame(frame) + if not isinstance(frame, CloseFrame): + return + self.state = ConnectionState.CLOSING + self._close_sent = True + if frame.code not in CLEAN_CLOSE_CODES or self._close_received: + return self._terminate() + # After sending a control frame indicating the connection + # should be closed, a peer does not send any further data. + self._selector.unregister(self._outgoing_frame_queue) + + def _process_next_message(self): + """ + Process the next message coming throught the socket. If a + data message can be extracted, return its decoded payload. + As per the RFC, only control frames will be processed once + the connection reaches the closing state. + """ + frame = self._get_next_frame() + if frame.opcode in CTRL_OP: + return self._handle_control_frame(frame) + if self.state is not ConnectionState.OPEN: + # After receiving a control frame indicating the connection + # should be closed, a peer discards any further data + # received. + return + if frame.opcode is Opcode.CONTINUE: + raise ProtocolError("Unexpected continuation frame") + message = frame.payload + if not frame.fin: + message = self._recover_fragmented_message(frame) + return ( + message.decode('utf-8') + if message is not None and frame.opcode is Opcode.TEXT else message + ) + + def _recover_fragmented_message(self, initial_frame): + message_fragments = bytearray(initial_frame.payload) + while True: + frame = self._get_next_frame() + if frame.opcode in CTRL_OP: + # Control frames can be received in the middle of a + # fragmented message, process them as soon as possible. + self._handle_control_frame(frame) + if self.state is not ConnectionState.OPEN: + return + continue + if frame.opcode is not Opcode.CONTINUE: + raise ProtocolError("A continuation frame was expected") + message_fragments.extend(frame.payload) + if len(message_fragments) > type(self).MESSAGE_MAX_SIZE: + raise PayloadTooLargeException() + if frame.fin: + return bytes(message_fragments) + + def _send_frame(self, frame): + if frame.opcode in CTRL_OP and len(frame.payload) > 125: + raise ProtocolError( + "Control frames should have a payload length smaller than 126" + ) + if isinstance(frame.payload, str): + frame.payload = frame.payload.encode('utf-8') + elif not isinstance(frame.payload, (bytes, bytearray)): + frame.payload = json.dumps(frame.payload).encode('utf-8') + + output = bytearray() + first_byte = ( + (0b10000000 if frame.fin else 0) + | (0b01000000 if frame.rsv1 else 0) + | (0b00100000 if frame.rsv2 else 0) + | (0b00010000 if frame.rsv3 else 0) + | frame.opcode + ) + payload_length = len(frame.payload) + if payload_length < 126: + output.extend( + struct.pack('!BB', first_byte, payload_length) + ) + elif payload_length < 65536: + output.extend( + struct.pack('!BBH', first_byte, 126, payload_length) + ) + else: + output.extend( + struct.pack('!BBQ', first_byte, 127, payload_length) + ) + output.extend(frame.payload) + self._socket.sendall(output) + self._timeout_manager.acknowledge_frame_sent(frame) + + def _enqueue_close_frame(self, code, reason=None): + """ Put a close frame in the outgoing frame queue. """ + self._outgoing_frame_queue.put(CloseFrame(code, reason)) + + def _enqueue_ping_frame(self): + """ Put a ping frame in the outgoing frame queue. """ + self._outgoing_frame_queue.put(Frame(Opcode.PING)) + + def _enqueue_pong_frame(self, payload): + """ Put a pong frame in the outgoing frame queue. """ + self._outgoing_frame_queue.put(Frame(Opcode.PONG, payload)) + + def _terminate(self): + """ Close the underlying TCP socket. """ + with suppress(OSError, TimeoutError): + self._socket.shutdown(socket.SHUT_WR) + # Call recv until obtaining a return value of 0 indicating + # the other end has performed an orderly shutdown. A timeout + # is set to ensure the connection will be closed even if + # the other end does not close the socket properly. + self._socket.settimeout(1) + while self._socket.recv(4096): + pass + self._selector.unregister(self._socket) + self._selector.close() + self._socket.close() + self.state = ConnectionState.CLOSED + + def _handle_control_frame(self, frame): + if frame.opcode is Opcode.PING: + self._enqueue_pong_frame(frame.payload) + elif frame.opcode is Opcode.CLOSE: + self.state = ConnectionState.CLOSING + self._close_received = True + code, reason = CloseCode.CLEAN, None + if len(frame.payload) >= 2: + code = struct.unpack('!H', frame.payload[:2])[0] + reason = frame.payload[2:].decode('utf-8') + elif frame.payload: + raise ProtocolError("Malformed closing frame") + if not self._close_sent: + self._enqueue_close_frame(code, reason) + else: + self._terminate() + + def _handle_transport_error(self, exc): + """ + Find out which close code should be sent according to given + exception and call `self.disconnect` in order to close the + connection cleanly. + """ + code, reason = CloseCode.SERVER_ERROR, str(exc) + if isinstance(exc, (ConnectionClosed, OSError)): + code = CloseCode.ABNORMAL_CLOSURE + elif isinstance(exc, (ProtocolError, InvalidCloseCodeException)): + code = CloseCode.PROTOCOL_ERROR + elif isinstance(exc, UnicodeDecodeError): + code = CloseCode.INCONSISTENT_DATA + elif isinstance(exc, PayloadTooLargeException): + code = CloseCode.MESSAGE_TOO_BIG + if code is CloseCode.SERVER_ERROR: + reason = None + _logger.error(exc, exc_info=True) + self.disconnect(code, reason) + + @classmethod + def _kick_all(cls): + """ Disconnect all the websocket instances. """ + for websocket in cls._instances: + if websocket.state is ConnectionState.OPEN: + websocket.disconnect(CloseCode.GOING_AWAY) + + +class TimeoutReason(IntEnum): + KEEP_ALIVE = 0 + NO_RESPONSE = 1 + + +class TimeoutManager: + """ + This class handles the Websocket timeouts. If no response to a + PING/CLOSE frame is received after `TIMEOUT` seconds or if the + connection is opened for more than `KEEP_ALIVE_TIMEOUT` seconds, the + connection is considered to have timed out. To determine if the + connection has timed out, use the `has_timed_out` method. + """ + TIMEOUT = 15 + # Timeout specifying how many seconds the connection should be kept + # alive. + KEEP_ALIVE_TIMEOUT = int(config['websocket_keep_alive_timeout']) + + def __init__(self): + super().__init__() + self._awaited_opcode = None + # Time in which the connection was opened. + self._opened_at = time.time() + self.timeout_reason = None + # Start time recorded when we started awaiting an answer to a + # PING/CLOSE frame. + self._waiting_start_time = None + + def acknowledge_frame_receipt(self, frame): + if self._awaited_opcode is frame.opcode: + self._awaited_opcode = None + self._waiting_start_time = None + + def acknowledge_frame_sent(self, frame): + """ + Acknowledge a frame was sent. If this frame is a PING/CLOSE + frame, start waiting for an answer. + """ + if self.has_timed_out(): + return + if frame.opcode is Opcode.PING: + self._awaited_opcode = Opcode.PONG + elif frame.opcode is Opcode.CLOSE: + self._awaited_opcode = Opcode.CLOSE + if self._awaited_opcode is not None: + self._waiting_start_time = time.time() + + def has_timed_out(self): + """ + Determine whether the connection has timed out or not. The + connection times out when the answer to a CLOSE/PING frame + is not received within `TIMEOUT` seconds or if the connection + is opened for more than `KEEP_ALIVE_TIMEOUT` seconds. + """ + now = time.time() + if now - self._opened_at >= type(self).KEEP_ALIVE_TIMEOUT: + self.timeout_reason = TimeoutReason.KEEP_ALIVE + return True + if self._awaited_opcode and now - self._waiting_start_time >= type(self).TIMEOUT: + self.timeout_reason = TimeoutReason.NO_RESPONSE + return True + return False + + +# ------------------------------------------------------ +# WEBSOCKET SERVING +# ------------------------------------------------------ + + +class WebsocketConnectionHandler: + SUPPORTED_VERSIONS = {'13'} + # Given by the RFC in order to generate Sec-WebSocket-Accept from + # Sec-WebSocket-Key value. + _HANDSHAKE_GUID = '258EAFA5-E914-47DA-95CA-C5AB0DC85B11' + _REQUIRED_HANDSHAKE_HEADERS = { + 'connection', 'host', 'sec-websocket-key', + 'sec-websocket-version', 'upgrade', + } + + @classmethod + def open_connection(cls, request): + """ + Open a websocket connection if the handshake is successfull. + :return: Response indicating the server performed a connection + upgrade. + :raise: UpgradeRequired if there is no intersection between the + versions the client supports and those we support. + :raise: BadRequest if the handshake data is incorrect. + """ + response = cls._get_handshake_response(request.httprequest.headers) + response.call_on_close(functools.partial( + cls._serve_forever, + Websocket(request.httprequest.environ['socket']), + )) + return response + + @classmethod + def _get_handshake_response(cls, headers): + """ + :return: Response indicating the server performed a connection + upgrade. + :raise: BadRequest + :raise: UpgradeRequired + """ + cls._assert_handshake_validity(headers) + # sha-1 is used as it is required by + # https://datatracker.ietf.org/doc/html/rfc6455#page-7 + accept_header = hashlib.sha1( + (headers['sec-websocket-key'] + cls._HANDSHAKE_GUID).encode()).digest() + accept_header = base64.b64encode(accept_header) + return Response(status=101, headers={ + 'Upgrade': 'websocket', + 'Connection': 'Upgrade', + 'Sec-WebSocket-Accept': accept_header, + }) + + @classmethod + def _assert_handshake_validity(cls, headers): + """ + :raise: UpgradeRequired if there is no intersection between + the version the client supports and those we support. + :raise: BadRequest in case of invalid handshake. + """ + missing_or_empty_headers = { + header for header in cls._REQUIRED_HANDSHAKE_HEADERS + if header not in headers + } + if missing_or_empty_headers: + raise BadRequest( + f"""Empty or missing header(s): {', '.join(missing_or_empty_headers)}""" + ) + + if headers['upgrade'].lower() != 'websocket': + raise BadRequest('Invalid upgrade header') + if 'upgrade' not in headers['connection'].lower(): + raise BadRequest('Invalid connection header') + if headers['sec-websocket-version'] not in cls.SUPPORTED_VERSIONS: + raise UpgradeRequired() + + key = headers['sec-websocket-key'] + try: + decoded_key = base64.b64decode(key, validate=True) + except ValueError: + raise BadRequest("Sec-WebSocket-Key should be b64 encoded") + if len(decoded_key) != 16: + raise BadRequest( + "Sec-WebSocket-Key should be of length 16 once decoded" + ) + + @classmethod + def _serve_forever(cls, websocket): + """ + Process incoming messages and dispatch them to the application. + """ + current_thread = threading.current_thread() + current_thread.type = 'websocket' + for message in websocket.get_messages(): + pass + + +CommonServer.on_stop(Websocket._kick_all) diff --git a/odoo/http.py b/odoo/http.py index 17f58cb7d14..b367a76e9ba 100644 --- a/odoo/http.py +++ b/odoo/http.py @@ -164,8 +164,8 @@ from .exceptions import UserError, AccessError, AccessDenied from .modules.module import get_manifest from .modules.registry import Registry from .service import security, model as service_model -from .tools import (config, consteq, date_utils, file_path, profiler, - resolve_attr, submap, unique, ustr,) +from .tools import (config, consteq, date_utils, file_path, parse_version, + profiler, submap, unique, ustr,) from .tools.geoipresolver import GeoIPResolver from .tools.func import filter_kwargs, lazy_property from .tools.mimetypes import guess_mimetype @@ -257,6 +257,14 @@ ROUTING_KEYS = { 'alias', 'host', 'methods', } +if parse_version(werkzeug.__version__) >= parse_version('2.0.2'): + # Werkzeug 2.0.2 adds the websocket option. If a websocket request + # (ws/wss) is trying to access an HTTP route, a WebsocketMismatch + # exception is raised. On the other hand, Werkzeug 0.16 does not + # support the websocket routing key. In order to bypass this issue, + # let's add the websocket key only when appropriate. + ROUTING_KEYS.add('websocket') + # The duration of a user session before it is considered expired, # three months. SESSION_LIFETIME = 60 * 60 * 24 * 90 diff --git a/odoo/service/server.py b/odoo/service/server.py index 8991bf33500..83643705791 100644 --- a/odoo/service/server.py +++ b/odoo/service/server.py @@ -126,6 +126,21 @@ class RequestHandler(werkzeug.serving.WSGIRequestHandler): me = threading.current_thread() me.name = 'odoo.service.http.request.%s' % (me.ident,) + def send_response(self, code, message=None): + # Since the upgrade header is introduced in version 1.1, Firefox + # won't accept a websocket connection if the version is set to + # 1.0. + if self.environ.get('REQUEST_URI') == '/websocket': + self.protocol_version = "HTTP/1.1" + return super().send_response(code, message=message) + + def make_environ(self): + environ = super().make_environ() + # Add the TCP socket to environ in order for the websocket + # connections to use it. + environ['socket'] = self.connection + return environ + class ThreadedWSGIServerReloadable(LoggingBaseWSGIServerMixIn, werkzeug.serving.ThreadedWSGIServer): """ werkzeug Threaded WSGI Server patched to allow reusing a listen socket @@ -304,9 +319,10 @@ class FSWatcherInotify(FSWatcherBase): #---------------------------------------------------------- class CommonServer(object): + _on_stop_funcs = [] + def __init__(self, app): self.app = app - self._on_stop_funcs = [] # config self.interface = config['http_interface'] or '0.0.0.0' self.port = config['http_port'] @@ -334,12 +350,13 @@ class CommonServer(object): raise sock.close() - def on_stop(self, func): + @classmethod + def on_stop(cls, func): """ Register a cleanup function to be executed when the server stops """ - self._on_stop_funcs.append(func) + cls._on_stop_funcs.append(func) def stop(self): - for func in self._on_stop_funcs: + for func in type(self)._on_stop_funcs: try: _logger.debug("on_close call %s", func) func() @@ -388,9 +405,10 @@ class ThreadedServer(CommonServer): self.limits_reached_threads.add(threading.current_thread()) for thread in threading.enumerate(): - if not thread.daemon or getattr(thread, 'type', None) == 'cron': + thread_type = getattr(thread, 'type', None) + if not thread.daemon and thread_type != 'websocket' or thread_type == 'cron': # We apply the limits on cron threads and HTTP requests, - # longpolling requests excluded. + # websocket requests excluded. if getattr(thread, 'start_time', None): thread_execution_time = time.time() - thread.start_time thread_limit_time_real = config['limit_time_real'] @@ -614,11 +632,11 @@ class GeventServer(CommonServer): def process_limits(self): restart = False if self.ppid != os.getppid(): - _logger.warning("LongPolling Parent changed: %s", self.pid) + _logger.warning("Gevent Parent changed: %s", self.pid) restart = True memory = memory_info(psutil.Process(self.pid)) if config['limit_memory_soft'] and memory > config['limit_memory_soft']: - _logger.warning('LongPolling virtual memory limit reached: %s', memory) + _logger.warning('Gevent virtual memory limit reached: %s', memory) restart = True if restart: # suicide !! @@ -645,6 +663,13 @@ class GeventServer(CommonServer): Derived from werzeug.serving.WSGIRequestHandler.log / werzeug.serving.WSGIRequestHandler.address_string """ + def _connection_upgrade_requested(self): + if self.headers.get('Connection', '').lower() == 'upgrade': + return True + if self.headers.get('Upgrade', '').lower() == 'websocket': + return True + return False + def format_request(self): old_address = self.client_address if getattr(self, 'environ', None): @@ -657,6 +682,27 @@ class GeventServer(CommonServer): finally: self.client_address = old_address + def finalize_headers(self): + # We need to make gevent.pywsgi stop dealing with chunks when the connection + # Is being upgraded. see https://github.com/gevent/gevent/issues/1712 + super().finalize_headers() + if self.code == 101: + # Switching Protocols. Disable chunked writes. + self.response_use_chunked = False + + def get_environ(self): + # Add the TCP socket to environ in order for the websocket + # connections to use it. + environ = super().get_environ() + environ['socket'] = self.socket + # Disable support for HTTP chunking on reads which cause + # an issue when the connection is being upgraded, see + # https://github.com/gevent/gevent/issues/1712 + if self._connection_upgrade_requested(): + environ['wsgi.input'] = self.rfile + environ['wsgi.input_terminated'] = False + return environ + set_limit_memory_hard() if os.name == 'posix': # Set process memory limit as an extra safeguard diff --git a/odoo/tools/config.py b/odoo/tools/config.py index 02d988a06d5..10da83acc3b 100644 --- a/odoo/tools/config.py +++ b/odoo/tools/config.py @@ -77,6 +77,7 @@ class configmanager(object): 'publisher_warranty_url': 'http://services.openerp.com/publisher-warranty/', 'reportgz': False, 'root_path': None, + 'websocket_keep_alive_timeout': 600, } # Not exposed in the configuration file.