diff --git a/addons/bus/models/bus.py b/addons/bus/models/bus.py index 69bab16178b..a0c2cd549cc 100644 --- a/addons/bus/models/bus.py +++ b/addons/bus/models/bus.py @@ -151,9 +151,13 @@ class ImDispatch(threading.Thread): given channels will be sent through the websocket. If a subscription is already present, overwrite it. """ - channels = [channel_with_db(db, c) for c in channels] + 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(hashable(channel), set()).add(websocket) + self._channels_to_ws.setdefault(channel, set()).add(websocket) self._ws_to_subscription[websocket] = BusSubscription(channels, last) if not self.is_alive(): self.start() @@ -161,9 +165,17 @@ class ImDispatch(threading.Thread): self._dispatch_notifications(websocket) def unsubscribe(self, websocket): - self._ws_to_subscription.pop(websocket, None) - for websockets in self._channels_to_ws.values(): - websockets.discard(websocket) + websocket_subscription = self._ws_to_subscription.pop(websocket, None) + if not websocket_subscription: + return + self._clear_outdated_channels(websocket, websocket_subscription.channels) + + def _clear_outdated_channels(self, websocket, outdated_channels): + """ Remove channels from channel to websocket map. """ + for channel in outdated_channels: + self._channels_to_ws[channel].remove(websocket) + if not self._channels_to_ws[channel]: + self._channels_to_ws.pop(channel) def loop(self): """ Dispatch postgres notifications to the relevant websockets """ diff --git a/addons/bus/tests/test_websocket_caryall.py b/addons/bus/tests/test_websocket_caryall.py index f280c6944d3..19fb84b3b70 100644 --- a/addons/bus/tests/test_websocket_caryall.py +++ b/addons/bus/tests/test_websocket_caryall.py @@ -5,6 +5,7 @@ import json from collections import defaultdict from datetime import timedelta from freezegun import freeze_time +from threading import Event from unittest.mock import patch from odoo.api import Environment @@ -140,3 +141,52 @@ class TestWebsocketCaryall(WebsocketCase): self.env['bus.bus']._sendone('channel1', 'notif type', 'message') dispatch._dispatch_notifications(next(iter(dispatch._ws_to_subscription.keys()))) self.assert_close_with_code(websocket, CloseCode.SESSION_EXPIRED) + + def test_channel_subscription_disconnect(self): + subscribe_done_event = Event() + original_subscribe = dispatch.subscribe + + def patched_subscribe(*args): + original_subscribe(*args) + subscribe_done_event.set() + + with patch.object(dispatch, 'subscribe', patched_subscribe): + websocket = self.websocket_connect() + websocket.send(json.dumps({ + 'event_name': 'subscribe', + 'data': {'channels': ['my_channel'], 'last': 0} + })) + subscribe_done_event.wait(timeout=5) + # channel is added as expected to the channel to websocket map. + self.assertIn((self.env.registry.db_name, 'my_channel'), dispatch._channels_to_ws) + websocket.close(CloseCode.CLEAN) + self.wait_remaining_websocket_connections() + # channel is removed as expected when removing the last + # websocket that was listening to this channel. + self.assertNotIn((self.env.registry.db_name, 'my_channel'), dispatch._channels_to_ws) + + def test_channel_subscription_update(self): + subscribe_done_event = Event() + original_subscribe = dispatch.subscribe + + def patched_subscribe(*args): + original_subscribe(*args) + subscribe_done_event.set() + + with patch.object(dispatch, 'subscribe', patched_subscribe): + websocket = self.websocket_connect() + websocket.send(json.dumps({ + 'event_name': 'subscribe', + 'data': {'channels': ['my_channel'], 'last': 0} + })) + subscribe_done_event.wait(timeout=5) + subscribe_done_event.clear() + # channel is added as expected to the channel to websocket map. + self.assertIn((self.env.registry.db_name, 'my_channel'), dispatch._channels_to_ws) + websocket.send(json.dumps({ + 'event_name': 'subscribe', + 'data': {'channels': ['my_channel_2'], 'last': 0} + })) + 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)