[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:
tsm-odoo
2022-08-23 17:55:09 +02:00
parent cbb5805479
commit e06bb9a42d
10 changed files with 936 additions and 10 deletions
+1
View File
@@ -1,3 +1,4 @@
# -*- coding: utf-8 -*-
from . import models
from . import controllers
from . import websocket
+1
View File
@@ -1,2 +1,3 @@
# -*- coding: utf-8 -*-
from . import main
from . import websocket
+14
View File
@@ -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)
+2
View File
@@ -1,2 +1,4 @@
from . import common
from . import test_assetsbundle
from . import test_health
from . import test_websocket_caryall
+94
View File
@@ -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)
+679
View File
@@ -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
View File
@@ -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
View File
@@ -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
+1
View File
@@ -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.