diff --git a/HISTORY.rst b/HISTORY.rst index 1fb4b72..5857531 100644 --- a/HISTORY.rst +++ b/HISTORY.rst @@ -6,6 +6,7 @@ not released yet ---------------- * Fix documentation +* Allow overriding of Celery task kwargs (see ``_get_celery_task_kwargs()`` method) 0.2.2 (2020-02-11) ------------------ diff --git a/example/tests/test_receivers.py b/example/tests/test_receivers.py index d2f89ba..ebd5356 100644 --- a/example/tests/test_receivers.py +++ b/example/tests/test_receivers.py @@ -47,8 +47,16 @@ def test_receiver_should_pass_serialized_kwargs_to_celery_task(self): receiver.receive(self.signal_kwargs) commit() - receiver.celery_task.delay.assert_called_once_with( - 'unittest.mock.MagicMock', - 'pynotify.serializers.ModelSerializer', - receiver.serializer_class().serialize(self.signal_kwargs), + receiver.celery_task.delay.assert_called_with( + handler_class='unittest.mock.MagicMock', + serializer_class='pynotify.serializers.ModelSerializer', + signal_kwargs=receiver.serializer_class().serialize(self.signal_kwargs), ) + + @override_settings(PYNOTIFY_CELERY_TASK='tests.test_receivers.mock_task') + def test_receiver_should_allow_overriding_of_celery_task_kwargs(self): + receiver = AsynchronousReceiver(MagicMock) + receiver._get_celery_task_kwargs = MagicMock(return_value={'abc': 1}) + receiver.receive(self.signal_kwargs) + commit() + receiver.celery_task.delay.assert_called_with(abc=1) diff --git a/pynotify/receivers.py b/pynotify/receivers.py index 264cf76..e08f921 100644 --- a/pynotify/receivers.py +++ b/pynotify/receivers.py @@ -51,11 +51,15 @@ def _get_celery_task(self): ) return locate(celery_task) + def _get_celery_task_kwargs(self): + return { + 'handler_class': get_import_path(self.handler_class), + 'serializer_class': get_import_path(self.serializer_class), + 'signal_kwargs': self.serializer_class().serialize(self.signal_kwargs), + } + def receive(self, signal_kwargs): + self.signal_kwargs = signal_kwargs # Call of the Celery task should be performed after current DB transaction is commited to avoid race condition, # e.g. accessing referenced object in the task before it has finished saving into DB. - on_commit(lambda: self.celery_task.delay( - get_import_path(self.handler_class), - get_import_path(self.serializer_class), - self.serializer_class().serialize(signal_kwargs), - )) + on_commit(lambda: self.celery_task.delay(**self._get_celery_task_kwargs()))