diff --git a/addons/bus/models/bus.py b/addons/bus/models/bus.py index a0c2cd549cc..fdc62b4dca0 100644 --- a/addons/bus/models/bus.py +++ b/addons/bus/models/bus.py @@ -6,15 +6,11 @@ import random import selectors import threading import time -from contextlib import suppress import odoo from odoo import api, fields, models -from odoo.http import root -from odoo.service.security import check_session from odoo.tools.misc import DEFAULT_SERVER_DATETIME_FORMAT from odoo.tools import date_utils -from ..websocket import CloseCode, InvalidStateException, Websocket _logger = logging.getLogger(__name__) @@ -118,33 +114,8 @@ class BusSubscription: class ImDispatch(threading.Thread): def __init__(self): super().__init__(daemon=True, name=f'{__name__}.Bus') - self._ws_to_subscription = {} self._channels_to_ws = {} - def _dispatch_notifications(self, websocket): - """ - Dispatch notifications available for the given websocket. If the - session is expired, close the connection with the `SESSION_EXPIRED` - close code. - """ - subscription = self._ws_to_subscription.get(websocket) - if not subscription: - return - session = root.session_store.get(websocket._session.sid) - if not session: - return websocket.disconnect(CloseCode.SESSION_EXPIRED) - with odoo.registry(session.db).cursor() as cr: - env = api.Environment(cr, session.uid, session.context) - if session.uid is not None and not check_session(session, env): - return websocket.disconnect(CloseCode.SESSION_EXPIRED) - notifications = env['bus.bus']._poll( - subscription.channels, subscription.last_notification_id) - if not notifications: - return - with suppress(InvalidStateException): - subscription.last_notification_id = notifications[-1]['id'] - websocket.send(notifications) - def subscribe(self, channels, last, db, websocket): """ Subcribe to bus notifications. Every notification related to the @@ -152,23 +123,16 @@ class ImDispatch(threading.Thread): is already present, overwrite it. """ channels = {hashable(channel_with_db(db, c)) for c in channels} - subscription = self._ws_to_subscription.get(websocket) - if subscription: - outdated_channels = subscription.channels - channels - self._clear_outdated_channels(websocket, outdated_channels) for channel in channels: self._channels_to_ws.setdefault(channel, set()).add(websocket) - self._ws_to_subscription[websocket] = BusSubscription(channels, last) + outdated_channels = websocket._channels - channels + self._clear_outdated_channels(websocket, outdated_channels) + websocket.subscribe(channels, last) if not self.is_alive(): self.start() - # Dispatch past notifications if there are any. - self._dispatch_notifications(websocket) def unsubscribe(self, websocket): - websocket_subscription = self._ws_to_subscription.pop(websocket, None) - if not websocket_subscription: - return - self._clear_outdated_channels(websocket, websocket_subscription.channels) + self._clear_outdated_channels(websocket, websocket._channels) def _clear_outdated_channels(self, websocket, outdated_channels): """ Remove channels from channel to websocket map. """ @@ -187,20 +151,18 @@ class ImDispatch(threading.Thread): conn = cr._cnx sel.register(conn, selectors.EVENT_READ) while True: - sel.select(TIMEOUT) - conn.poll() - channels = [] - while conn.notifies: - channels.extend(json.loads(conn.notifies.pop().payload)) - # relay notifications to websockets that have - # subscribed to the corresponding channels. - websockets = set() - for channel in channels: - websockets.update( - self._channels_to_ws.get(hashable(channel), []) - ) - for websocket in websockets: - self._dispatch_notifications(websocket) + if sel.select(TIMEOUT): + conn.poll() + channels = [] + while conn.notifies: + channels.extend(json.loads(conn.notifies.pop().payload)) + # relay notifications to websockets that have + # subscribed to the corresponding channels. + websockets = set() + for channel in channels: + websockets.update(self._channels_to_ws.get(hashable(channel), [])) + for websocket in websockets: + websocket.trigger_notification_dispatching() def run(self): while True: @@ -214,8 +176,3 @@ dispatch = None if not odoo.multi_process or odoo.evented: # We only use the event dispatcher in threaded and gevent mode dispatch = ImDispatch() - - -@Websocket.onclose -def _unsubscribe(env, websocket): - dispatch.unsubscribe(websocket) diff --git a/addons/bus/static/src/services/bus_service.js b/addons/bus/static/src/services/bus_service.js index 33e25320de2..4bc35e8ac38 100644 --- a/addons/bus/static/src/services/bus_service.js +++ b/addons/bus/static/src/services/bus_service.js @@ -6,7 +6,11 @@ import { browser } from "@web/core/browser/browser"; import { registry } from '@web/core/registry'; const { EventBus } = owl; -const NO_POPUP_CLOSE_CODES = [WEBSOCKET_CLOSE_CODES.SESSION_EXPIRED, WEBSOCKET_CLOSE_CODES.KEEP_ALIVE_TIMEOUT]; +const NO_POPUP_CLOSE_CODES = [ + WEBSOCKET_CLOSE_CODES.SESSION_EXPIRED, + WEBSOCKET_CLOSE_CODES.KEEP_ALIVE_TIMEOUT, + WEBSOCKET_CLOSE_CODES.TRY_LATER, +]; /** * Communicate with a SharedWorker in order to provide a single websocket diff --git a/addons/bus/static/src/workers/websocket_worker.js b/addons/bus/static/src/workers/websocket_worker.js index a2fc412c0ad..f7a6e99ba8e 100644 --- a/addons/bus/static/src/workers/websocket_worker.js +++ b/addons/bus/static/src/workers/websocket_worker.js @@ -43,7 +43,7 @@ export class WebsocketWorker { constructor(websocketURL) { this.websocketURL = websocketURL; this.channelsByClient = new Map(); - this.connectRetryDelay = 0; + this.connectRetryDelay = 1000; this.connectTimeout = null; this.isReconnecting = false; this.lastChannelSubscription = null; @@ -216,6 +216,10 @@ export class WebsocketWorker { // WebSocket was not closed cleanly, let's try to reconnect. this.broadcast('reconnecting', { closeCode: code }); this.isReconnecting = true; + if (code === WEBSOCKET_CLOSE_CODES.KEEP_ALIVE_TIMEOUT) { + // Don't wait to reconnect on keep alive timeout. + this.connectRetryDelay = 0; + } this._onWebsocketError(); } diff --git a/addons/bus/tests/test_websocket_caryall.py b/addons/bus/tests/test_websocket_caryall.py index d3aaa2b883d..018bf0cd928 100644 --- a/addons/bus/tests/test_websocket_caryall.py +++ b/addons/bus/tests/test_websocket_caryall.py @@ -127,16 +127,19 @@ class TestWebsocketCaryall(WebsocketCase): def test_user_logout_outgoing_message(self): subscribe_done_event = Event() - original_subscribe = dispatch.subscribe + original_subscribe = Websocket.subscribe + odoo_ws = None - def patched_subscribe(*args): - original_subscribe(*args) + def patched_subscribe(self, *args): + nonlocal odoo_ws + odoo_ws = self + original_subscribe(self, *args) subscribe_done_event.set() new_test_user(self.env, login='test_user', password='Password!1') user_session = self.authenticate('test_user', 'Password!1') websocket = self.websocket_connect(cookie=f'session_id={user_session.sid};') - with patch.object(dispatch, 'subscribe', patched_subscribe): + with patch.object(Websocket, 'subscribe', patched_subscribe): websocket.send(json.dumps({ 'event_name': 'subscribe', 'data': {'channels': ['channel1'], 'last': 0} @@ -147,7 +150,7 @@ class TestWebsocketCaryall(WebsocketCase): # receiving the message. subscribe_done_event.wait(timeout=5) self.env['bus.bus']._sendone('channel1', 'notif type', 'message') - dispatch._dispatch_notifications(next(iter(dispatch._ws_to_subscription.keys()))) + odoo_ws.trigger_notification_dispatching() self.assert_close_with_code(websocket, CloseCode.SESSION_EXPIRED) def test_channel_subscription_disconnect(self): @@ -198,3 +201,35 @@ class TestWebsocketCaryall(WebsocketCase): subscribe_done_event.wait(timeout=5) # channel is removed as expected when updating the subscription. self.assertNotIn((self.env.registry.db_name, 'my_channel'), dispatch._channels_to_ws) + + def test_trigger_notification(self): + original_subscribe = Websocket.subscribe + odoo_ws = None + + def patched_subscribe(self, *args): + nonlocal odoo_ws + odoo_ws = self + original_subscribe(self, *args) + + with patch.object(Websocket, 'subscribe', patched_subscribe): + websocket = self.websocket_connect() + self.env['bus.bus']._sendone('my_channel', 'notif_type', 'message') + websocket.send(json.dumps({ + 'event_name': 'subscribe', + 'data': {'channels': ['my_channel'], 'last': 0} + })) + + notifications = json.loads(websocket.recv()) + self.assertEqual(1, len(notifications)) + self.assertEqual(notifications[0]['message']['type'], 'notif_type') + self.assertEqual(notifications[0]['message']['payload'], 'message') + + self.env['bus.bus']._sendone('my_channel', 'notif_type', 'another_message') + odoo_ws.trigger_notification_dispatching() + + notifications = json.loads(websocket.recv()) + # First notification has been received, we should only receive + # the second one. + self.assertEqual(1, len(notifications)) + self.assertEqual(notifications[0]['message']['type'], 'notif_type') + self.assertEqual(notifications[0]['message']['payload'], 'another_message') diff --git a/addons/bus/websocket.py b/addons/bus/websocket.py index 950554b43a1..f1ab4e79a23 100644 --- a/addons/bus/websocket.py +++ b/addons/bus/websocket.py @@ -4,7 +4,7 @@ import hashlib import json import logging import psycopg2 -import queue +import random import socket import struct import selectors @@ -13,41 +13,36 @@ import time from collections import defaultdict, deque from contextlib import closing, suppress from enum import IntEnum -from itertools import count +from psycopg2.pool import PoolError from weakref import WeakSet from werkzeug.local import LocalStack from werkzeug.exceptions import BadRequest, HTTPException +import odoo from odoo import api +from .models.bus import dispatch from odoo.http import root, Request, Response, SessionExpiredException from odoo.modules.registry import Registry from odoo.service import model as service_model from odoo.service.server import CommonServer +from odoo.service.security import check_session 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() +MAX_TRY_ON_POOL_ERROR = 10 +DELAY_ON_POOL_ERROR = 0.03 - 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) +def acquire_cursor(db): + """ Try to acquire a cursor up to `MAX_TRY_ON_POOL_ERROR` """ + for tryno in range(1, MAX_TRY_ON_POOL_ERROR + 1): + with suppress(PoolError): + return odoo.registry(db).cursor() + time.sleep(random.uniform(DELAY_ON_POOL_ERROR, DELAY_ON_POOL_ERROR * tryno)) + raise PoolError('Failed to acquire cursor after %s retries' % MAX_TRY_ON_POOL_ERROR) # ------------------------------------------------------ @@ -184,10 +179,6 @@ _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, @@ -197,7 +188,6 @@ class Frame: rsv2=False, rsv3=False ): - self._send_order = next(self._frames_sent) self.opcode = opcode self.payload = payload self.fin = fin @@ -205,17 +195,6 @@ class Frame: 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): @@ -245,19 +224,24 @@ class Websocket: # How many seconds between each request. RL_DELAY = float(config['websocket_rate_limit_delay']) - def __init__(self, socket, session): + def __init__(self, sock, session): # Session linked to the current websocket connection. self._session = session - self._socket = socket + self._socket = sock 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() # Used for rate limiting. self._incoming_frame_timestamps = deque(maxlen=type(self).RL_BURST) + # Used to notify the websocket that bus notifications are + # available. + self._notif_sock_w, self._notif_sock_r = socket.socketpair() + self._channels = set() + self._last_notif_sent_id = 0 + # Websocket start up + self._selector = selectors.DefaultSelector() + self._selector.register(self._socket, selectors.EVENT_READ) + self._selector.register(self._notif_sock_r, selectors.EVENT_READ) self.state = ConnectionState.OPEN type(self)._instances.add(self) self._trigger_lifecycle_event(LifecycleEvent.OPEN) @@ -281,12 +265,10 @@ class Websocket: ) continue if not readables: - self._enqueue_ping_frame() + self._send_ping_frame() continue - if self._outgoing_frame_queue in readables: - self._send_next_frame() - if self.state is ConnectionState.CLOSED: - break + if self._notif_sock_r in readables: + self._dispatch_bus_notifications() if self._socket in readables: message = self._process_next_message() if message is not None: @@ -294,16 +276,6 @@ class Websocket: 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 @@ -314,7 +286,7 @@ class Websocket: acknowledgment if the connection was failed beforewards. """ if code is not CloseCode.ABNORMAL_CLOSURE: - self._enqueue_close_frame(code, reason) + self._send_close_frame(code, reason) else: self._terminate() @@ -328,6 +300,30 @@ class Websocket: cls._event_callbacks[LifecycleEvent.CLOSE].add(func) return func + def subscribe(self, channels, last): + """ Subscribe to bus channels. """ + self._channels = channels + if self._last_notif_sent_id < last: + self._last_notif_sent_id = last + # Dispatch past notifications if there are any. + self.trigger_notification_dispatching() + + def trigger_notification_dispatching(self): + """ + Warn the socket that notifications are available. Ignore if a + dispatch is already planned or if the socket is already in the + closing state. + """ + if self.state is not ConnectionState.OPEN: + return + readables = { + selector_key[0].fileobj for selector_key in + self._selector.select(0) + } + if self._notif_sock_r not in readables: + # Send a random bit to mark the socket as readable. + self._notif_sock_w.send(b'x') + # ------------------------------------------------------ # PRIVATE METHODS # ------------------------------------------------------ @@ -409,18 +405,6 @@ class Websocket: 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 - 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 @@ -465,6 +449,16 @@ class Websocket: if frame.fin: return bytes(message_fragments) + 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._send_frame(Frame(opcode, message)) + def _send_frame(self, frame): if frame.opcode in CTRL_OP and len(frame.payload) > 125: raise ProtocolError( @@ -499,20 +493,27 @@ class Websocket: 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. """ + if not isinstance(frame, CloseFrame): + return self.state = ConnectionState.CLOSING self._close_sent = True - self._outgoing_frame_queue.put(CloseFrame(code, reason)) + 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._notif_sock_r) - def _enqueue_ping_frame(self): - """ Put a ping frame in the outgoing frame queue. """ - self._outgoing_frame_queue.put(Frame(Opcode.PING)) + def _send_close_frame(self, code, reason=None): + """ Send a close frame. """ + self._send_frame(CloseFrame(code, reason)) - 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 _send_ping_frame(self): + """ Send a ping frame """ + self._send_frame(Frame(Opcode.PING)) + + def _send_pong_frame(self, payload): + """ Send a pong frame """ + self._send_frame(Frame(Opcode.PONG, payload)) def _terminate(self): """ Close the underlying TCP socket. """ @@ -529,11 +530,12 @@ class Websocket: self._selector.close() self._socket.close() self.state = ConnectionState.CLOSED + dispatch.unsubscribe(self) self._trigger_lifecycle_event(LifecycleEvent.CLOSE) def _handle_control_frame(self, frame): if frame.opcode is Opcode.PING: - self._enqueue_pong_frame(frame.payload) + self._send_pong_frame(frame.payload) elif frame.opcode is Opcode.CLOSE: self.state = ConnectionState.CLOSING self._close_received = True @@ -544,7 +546,7 @@ class Websocket: elif frame.payload: raise ProtocolError("Malformed closing frame") if not self._close_sent: - self._enqueue_close_frame(code, reason) + self._send_close_frame(code, reason) else: self._terminate() @@ -563,8 +565,10 @@ class Websocket: code = CloseCode.INCONSISTENT_DATA elif isinstance(exc, PayloadTooLargeException): code = CloseCode.MESSAGE_TOO_BIG - elif isinstance(exc, RateLimitExceededException): + elif isinstance(exc, (PoolError, RateLimitExceededException)): code = CloseCode.TRY_LATER + elif isinstance(exc, SessionExpiredException): + code = CloseCode.SESSION_EXPIRED if code is CloseCode.SERVER_ERROR: reason = None _logger.error(exc, exc_info=True) @@ -598,8 +602,7 @@ class Websocket: registered for this event type. Every callback is given both the environment and the related websocket. """ - registry = Registry(self._session.db) - with closing(registry.cursor()) as cr: + with closing(acquire_cursor(self._session.db)) as cr: env = api.Environment(cr, self._session.uid, self._session.context) for callback in type(self)._event_callbacks[event_type]: try: @@ -611,6 +614,28 @@ class Websocket: exc_info=True ) + def _dispatch_bus_notifications(self): + """ + Dispatch notifications related to the registered channels. If + the session is expired, close the connection with the + `SESSION_EXPIRED` close code. If no cursor can be acquired, + close the connection with the `TRY_LATER` close code. + """ + session = root.session_store.get(self._session.sid) + if not session: + raise SessionExpiredException() + with acquire_cursor(session.db) as cr: + env = api.Environment(cr, session.uid, session.context) + if session.uid is not None and not check_session(session, env): + raise SessionExpiredException() + # Mark the notification request as processed. + self._notif_sock_r.recv(1) + notifications = env['bus.bus']._poll(self._channels, self._last_notif_sent_id) + if not notifications: + return + self._last_notif_sent_id = notifications[-1]['id'] + self._send(notifications) + class TimeoutReason(IntEnum): KEEP_ALIVE = 0 @@ -721,7 +746,7 @@ class WebsocketRequest: ) as exc: raise InvalidDatabaseException() from exc - with closing(self.registry.cursor()) as cr: + with closing(acquire_cursor(self.db)) as cr: self.env = api.Environment(cr, self.session.uid, self.session.context) threading.current_thread().uid = self.env.uid service_model.retrying( @@ -854,6 +879,8 @@ class WebsocketConnectionHandler: req.serve_websocket_message(message) except SessionExpiredException: websocket.disconnect(CloseCode.SESSION_EXPIRED) + except PoolError: + websocket.disconnect(CloseCode.TRY_LATER) except Exception: _logger.exception("Exception occurred during websocket request handling")