diff --git a/addons/bus/models/bus.py b/addons/bus/models/bus.py index 17e68d1c8f6..d3a3d8b2402 100644 --- a/addons/bus/models/bus.py +++ b/addons/bus/models/bus.py @@ -3,6 +3,7 @@ import contextlib import datetime import json import logging +import math import os import random import selectors @@ -24,6 +25,21 @@ TIMEOUT = 50 # custom function to call instead of default PostgreSQL's `pg_notify` ODOO_NOTIFY_FUNCTION = os.getenv('ODOO_NOTIFY_FUNCTION', 'pg_notify') + +def get_notify_payload_max_length(default=8000): + try: + length = int(os.environ.get('ODOO_NOTIFY_PAYLOAD_MAX_LENGTH', default)) + except ValueError: + _logger.warning("ODOO_NOTIFY_PAYLOAD_MAX_LENGTH has to be an integer, " + "defaulting to %d bytes", default) + length = default + return length + + +# max length in bytes for the NOTIFY query payload +NOTIFY_PAYLOAD_MAX_LENGTH = get_notify_payload_max_length() + + #---------------------------------------------------------- # Bus #---------------------------------------------------------- @@ -46,6 +62,26 @@ def channel_with_db(dbname, channel): return channel +def get_notify_payloads(channels): + """ + Generates the json payloads for the imbus NOTIFY. + Splits recursively payloads that are too large. + + :param list channels: + :return: list of payloads of json dumps + :rtype: list[str] + """ + if not channels: + return [] + payload = json_dump(channels) + if len(channels) == 1 or len(payload.encode()) < NOTIFY_PAYLOAD_MAX_LENGTH: + return [payload] + else: + pivot = math.ceil(len(channels) / 2) + return (get_notify_payloads(channels[:pivot]) + + get_notify_payloads(channels[pivot:])) + + class ImBus(models.Model): _name = 'bus.bus' @@ -84,7 +120,12 @@ class ImBus(models.Model): def notify(): with odoo.sql_db.db_connect('postgres').cursor() as cr: query = sql.SQL("SELECT {}('imbus', %s)").format(sql.Identifier(ODOO_NOTIFY_FUNCTION)) - cr.execute(query, (json_dump(list(channels)), )) + payloads = get_notify_payloads(list(channels)) + if len(payloads) > 1: + _logger.info("The imbus notification payload was too large, " + "it's been split into %d payloads.", len(payloads)) + for payload in payloads: + cr.execute(query, (payload,)) @api.model def _sendone(self, channel, notification_type, message): diff --git a/addons/bus/tests/__init__.py b/addons/bus/tests/__init__.py index f8f91a9b5ef..ad34feffec9 100644 --- a/addons/bus/tests/__init__.py +++ b/addons/bus/tests/__init__.py @@ -3,6 +3,7 @@ from . import test_assetsbundle from . import test_health from . import test_ir_model from . import test_ir_websocket +from . import test_notify from . import test_websocket_caryall from . import test_websocket_controller from . import test_websocket_rate_limiting diff --git a/addons/bus/tests/test_notify.py b/addons/bus/tests/test_notify.py new file mode 100644 index 00000000000..3fd64376bb1 --- /dev/null +++ b/addons/bus/tests/test_notify.py @@ -0,0 +1,49 @@ +# Part of Odoo. See LICENSE file for full copyright and licensing details. + +from odoo.tests import BaseCase + +from ..models.bus import json_dump, get_notify_payloads, NOTIFY_PAYLOAD_MAX_LENGTH + + +class NotifyTests(BaseCase): + + def test_get_notify_payloads(self): + """ + Asserts that the implementation of `get_notify_payloads` + actually splits correctly large payloads + """ + def check_payloads_size(payloads): + for payload in payloads: + self.assertLess(len(payload.encode()), NOTIFY_PAYLOAD_MAX_LENGTH) + + channel = ('dummy_db', 'dummy_model', 12345) + channels = [channel] + self.assertLess(len(json_dump(channels).encode()), NOTIFY_PAYLOAD_MAX_LENGTH) + payloads = get_notify_payloads(channels) + self.assertEqual(len(payloads), 1, + "The payload is less then the threshold, " + "there should be 1 payload only, as it shouldn't be split") + channels = [channel] * 100 + self.assertLess(len(json_dump(channels).encode()), NOTIFY_PAYLOAD_MAX_LENGTH) + payloads = get_notify_payloads(channels) + self.assertEqual(len(payloads), 1, + "The payload is less then the threshold, " + "there should be 1 payload only, as it shouldn't be split") + check_payloads_size(payloads) + channels = [channel] * 1000 + self.assertGreaterEqual(len(json_dump(channels).encode()), NOTIFY_PAYLOAD_MAX_LENGTH) + payloads = get_notify_payloads(channels) + self.assertGreater(len(payloads), 1, + "Payload was larger than the threshold, it should've been split") + check_payloads_size(payloads) + + fat_channel = tuple(item * 1000 for item in channel) + channels = [fat_channel] + self.assertEqual(len(channels), 1, "There should be only 1 channel") + self.assertGreaterEqual(len(json_dump(channels).encode()), NOTIFY_PAYLOAD_MAX_LENGTH) + payloads = get_notify_payloads(channels) + self.assertEqual(len(payloads), 1, + "Payload was larger than the threshold, but shouldn't be split, " + "as it contains only 1 channel") + with self.assertRaises(AssertionError): + check_payloads_size(payloads)