[FIX] bus: bus message dispatching

Before the websockets were introduced, longpolling coroutines were
sleeping until postgres notify. Each coroutine was then wake up and
notifications were fetched.

Since the websocket introduction, the main loop, responsible for listening
to postgres sends the notifications itself. This means, the postgres loop
is blocked during message fetch/dispatching and notifications are dispatched
in a sequential fashion resulting in a slow message dispatching.

In order to solve this issue, websocket coroutines are now responsible to
fetch/dispatch notifications, letting the main loop free to relay notifications
as they come and allowing notifications to be sent simultaneously.

When instructed to dispatch available notifications, the websocket coroutines
will try to acquire a cursor. Each coroutine will try up to `MAX_TRY_ON_POOL_ERROR`
times, sleeping between each try. If no cursor can be acquired, the connection is
closed with the TRY_LATER` close code.

closes odoo/odoo#98880

Signed-off-by: Antony Lesuisse <al@odoo.com>
This commit is contained in:
tsm-odoo
2022-09-05 14:50:17 +02:00
parent 24f33c0282
commit 58eade9f34
5 changed files with 178 additions and 151 deletions
+16 -59
View File
@@ -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)
@@ -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
@@ -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();
}
+40 -5
View File
@@ -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')
+112 -85
View File
@@ -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")