[ADD] bus: add websocket implementation
This commit is the first commit of the websocket integration in Odoo.
It focuses on the implementation of the websocket protocol as per RFC6455.
The implementation is tested thanks to the autobahn test suite.
A config parameter is available to customize the websocket connection:
- websocket_keep_alive_timeout (default 600): Integer specifying how
many seconds a websocket connection should be kept alive
Part-of: odoo/odoo#75510
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from . import models
|
||||
from . import controllers
|
||||
from . import websocket
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from . import main
|
||||
from . import websocket
|
||||
|
||||
@@ -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)
|
||||
@@ -1,2 +1,4 @@
|
||||
from . import common
|
||||
from . import test_assetsbundle
|
||||
from . import test_health
|
||||
from . import test_websocket_caryall
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
+10
-2
@@ -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
|
||||
|
||||
+54
-8
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user