diff --git a/dash/background_callback/managers/celery_manager.py b/dash/background_callback/managers/celery_manager.py index eb2864caa0..a7bdb65683 100644 --- a/dash/background_callback/managers/celery_manager.py +++ b/dash/background_callback/managers/celery_manager.py @@ -39,14 +39,13 @@ def __init__(self, celery_app, cache_by=None, expire=None): import celery # type: ignore[import-not-found] # pylint: disable=import-outside-toplevel,import-error from celery.backends.base import ( # type: ignore[import-not-found] # pylint: disable=import-outside-toplevel,import-error DisabledBackend, + BaseKeyValueStoreBackend, ) except ImportError as missing_imports: - raise ImportError( - """\ + raise ImportError("""\ CeleryManager requires extra dependencies which can be installed doing - $ pip install "dash[celery]"\n""" - ) from missing_imports + $ pip install "dash[celery]"\n""") from missing_imports if not isinstance(celery_app, celery.Celery): raise ValueError("First argument must be a celery.Celery object") @@ -54,6 +53,11 @@ def __init__(self, celery_app, cache_by=None, expire=None): if isinstance(celery_app.backend, DisabledBackend): raise ValueError("Celery instance must be configured with a result backend") + if not isinstance(celery_app.backend, BaseKeyValueStoreBackend): + raise ValueError( + "Celery must be configured with a key-value store backend (e.g. Redis or Filesystem)" + ) + self.handle = celery_app self.expire = expire super().__init__(cache_by) @@ -89,8 +93,17 @@ def get_task(self, job): return None + @staticmethod + def _ensure_bytes(o) -> bytes: + if isinstance(o, bytes): + return o + return str(o).encode() + def clear_cache_entry(self, key): - self.handle.backend.delete(key) + # delete should not be called when the entry is not present + value = self.handle.backend.get(self._ensure_bytes(key)) + if value is not None: + self.handle.backend.delete(self._ensure_bytes(key)) def get_or_create_signing_secret(self, generate): backend = self.handle.backend @@ -103,24 +116,30 @@ def get_or_create_signing_secret(self, generate): return backend.get(self.SIGNING_SECRET_KEY) or secret def call_job_fn(self, key, job_fn, args, context): - task = job_fn.delay(key, self._make_progress_key(key), args, context) + result_key = self._ensure_bytes(key) + progress_key = self._ensure_bytes(self._make_progress_key(key)) + set_props_key = self._ensure_bytes(self._make_set_props_key(key)) + task = job_fn.delay(result_key, progress_key, set_props_key, args, context) return task.task_id def get_progress(self, key): - progress_key = self._make_progress_key(key) + progress_key = self._ensure_bytes(self._make_progress_key(key)) progress_data = self.handle.backend.get(progress_key) if progress_data: - self.handle.backend.delete(progress_key) + self.clear_cache_entry(progress_key) return json.loads(progress_data) return None def result_ready(self, key): - return self.handle.backend.get(key) is not None + result_key = self._ensure_bytes(key) + return self.handle.backend.get(result_key) is not None def get_result(self, key, job): + result_key = self._ensure_bytes(key) + progress_key = self._ensure_bytes(self._make_progress_key(key)) # Get result value - result = self.handle.backend.get(key) + result = self.handle.backend.get(result_key) if result is None: return self.UNDEFINED @@ -128,22 +147,23 @@ def get_result(self, key, job): # Clear result if not caching if self.cache_by is None: - self.clear_cache_entry(key) + self.clear_cache_entry(result_key) else: if self.expire: # Set/update expiration time - self.handle.backend.expire(key, self.expire) - self.clear_cache_entry(self._make_progress_key(key)) + self.handle.backend.expire(result_key, self.expire) + self.clear_cache_entry(progress_key) self.terminate_job(job) return result def get_updated_props(self, key): - updated_props = self.handle.backend.get(self._make_set_props_key(key)) + set_props_key = self._ensure_bytes(self._make_set_props_key(key)) + updated_props = self.handle.backend.get(set_props_key) if updated_props is None: return {} - self.clear_cache_entry(key) + self.clear_cache_entry(set_props_key) return json.loads(updated_props) @@ -153,7 +173,7 @@ def _make_job_fn(fn, celery_app, progress, key): # pylint: disable=too-many-sta @celery_app.task(name=f"background_callback_{key}") def job_fn( - result_key, progress_key, user_callback_args, context=None + result_key, progress_key, set_props_key, user_callback_args, context=None ): # pylint: disable=too-many-statements def _set_progress(progress_value): if not isinstance(progress_value, (list, tuple)): @@ -165,7 +185,7 @@ def _set_progress(progress_value): def _set_props(_id, props): cache.set( - f"{result_key}-set_props", + set_props_key, json.dumps({_id: props}, cls=PlotlyJSONEncoder), ) diff --git a/requirements/ci.txt b/requirements/ci.txt index 8e18280d04..2581571520 100644 --- a/requirements/ci.txt +++ b/requirements/ci.txt @@ -2,6 +2,7 @@ black==22.3.0 flake8==7.0.0 flaky==3.8.1 +filelock>=3.0 flask-talisman==1.0.0 ipython<9.0.0 mimesis<=11.1.0 diff --git a/tests/async_tests/conftest.py b/tests/async_tests/conftest.py index b701eea91a..8ac581bd59 100644 --- a/tests/async_tests/conftest.py +++ b/tests/async_tests/conftest.py @@ -2,12 +2,11 @@ import pytest - if "REDIS_URL" in os.environ: - managers = ["celery", "diskcache"] + managers = ["celery-filesystem", "celery-redis", "diskcache"] else: - print("Skipping celery tests because REDIS_URL is not defined") - managers = ["diskcache"] + print("Skipping celery tests on Redis because REDIS_URL is not defined") + managers = ["celery-filesystem", "diskcache"] @pytest.fixture(params=managers) diff --git a/tests/async_tests/utils.py b/tests/async_tests/utils.py index b7074b0735..2b88f8a3ca 100644 --- a/tests/async_tests/utils.py +++ b/tests/async_tests/utils.py @@ -36,7 +36,7 @@ def get_background_callback_manager(): """ Get the long callback mangaer configured by environment variables """ - if os.environ.get("LONG_CALLBACK_MANAGER", None) == "celery": + if os.environ.get("LONG_CALLBACK_MANAGER", None) == "celery-redis": from dash.background_callback import CeleryManager from celery import Celery @@ -77,8 +77,8 @@ def kill(proc_pid): def setup_background_callback_app(manager_name, app_name): from dash.testing.application_runners import import_app - if manager_name == "celery": - os.environ["LONG_CALLBACK_MANAGER"] = "celery" + if manager_name == "celery-redis": + os.environ["LONG_CALLBACK_MANAGER"] = "celery-redis" redis_url = os.environ["REDIS_URL"].rstrip("/") os.environ["CELERY_BROKER"] = f"{redis_url}/0" os.environ["CELERY_BACKEND"] = f"{redis_url}/1" diff --git a/tests/background_callback/conftest.py b/tests/background_callback/conftest.py index b701eea91a..8ac581bd59 100644 --- a/tests/background_callback/conftest.py +++ b/tests/background_callback/conftest.py @@ -2,12 +2,11 @@ import pytest - if "REDIS_URL" in os.environ: - managers = ["celery", "diskcache"] + managers = ["celery-filesystem", "celery-redis", "diskcache"] else: - print("Skipping celery tests because REDIS_URL is not defined") - managers = ["diskcache"] + print("Skipping celery tests on Redis because REDIS_URL is not defined") + managers = ["celery-filesystem", "diskcache"] @pytest.fixture(params=managers) diff --git a/tests/background_callback/utils.py b/tests/background_callback/utils.py index 1cefd4ecc3..9beb37f000 100644 --- a/tests/background_callback/utils.py +++ b/tests/background_callback/utils.py @@ -35,7 +35,7 @@ def get_background_callback_manager(): """ Get the long callback mangaer configured by environment variables """ - if os.environ.get("LONG_CALLBACK_MANAGER", None) == "celery": + if os.environ.get("LONG_CALLBACK_MANAGER", None) == "celery-redis": from dash.background_callback import CeleryManager from celery import Celery import redis @@ -44,10 +44,36 @@ def get_background_callback_manager(): __name__, broker=os.environ.get("CELERY_BROKER"), backend=os.environ.get("CELERY_BACKEND"), + broker_connection_retry_on_startup=True, ) background_callback_manager = CeleryManager(celery_app) redis_conn = redis.Redis(host="localhost", port=6379, db=1) background_callback_manager.test_lock = redis_conn.lock("test-lock") + elif os.environ.get("LONG_CALLBACK_MANAGER", None) == "celery-filesystem": + from dash.background_callback import CeleryManager + from celery import Celery + from filelock import FileLock + + celery_broker_path = os.environ.get("CELERY_BROKER_FILESYSTEM_DIRECTORY") + assert ( + celery_broker_path is not None + ), "CELERY_BROKER_FILESYSTEM_DIRECTORY must be set" + + celery_app = Celery( + __name__, + broker=os.environ.get("CELERY_BROKER"), + backend=os.environ.get("CELERY_BACKEND"), + broker_transport_options={ + "data_folder_in": celery_broker_path, + "data_folder_out": celery_broker_path, + "control_folder": celery_broker_path, + }, + ) + background_callback_manager = CeleryManager(celery_app) + + background_callback_manager.test_lock = FileLock( + os.path.join(celery_broker_path, "test-lock") + ) elif os.environ.get("LONG_CALLBACK_MANAGER", None) == "diskcache": import diskcache @@ -77,17 +103,27 @@ def kill(proc_pid): def setup_background_callback_app(manager_name, app_name): from dash.testing.application_runners import import_app - if manager_name == "celery": - os.environ["LONG_CALLBACK_MANAGER"] = "celery" - redis_url = os.environ["REDIS_URL"].rstrip("/") - os.environ["CELERY_BROKER"] = f"{redis_url}/0" - os.environ["CELERY_BACKEND"] = f"{redis_url}/1" - - # Clear redis of cached values - redis_conn = redis.Redis(host="localhost", port=6379, db=1) - cache_keys = redis_conn.keys() - if cache_keys: - redis_conn.delete(*cache_keys) + if manager_name in ["celery-redis", "celery-filesystem"]: + os.environ["LONG_CALLBACK_MANAGER"] = manager_name + + if manager_name == "celery-redis": + redis_url = os.environ["REDIS_URL"].rstrip("/") + os.environ["CELERY_BROKER"] = f"{redis_url}/0" + os.environ["CELERY_BACKEND"] = f"{redis_url}/1" + + # Clear redis of cached values + redis_conn = redis.Redis(host="localhost", port=6379, db=1) + cache_keys = redis_conn.keys() + if cache_keys: + redis_conn.delete(*cache_keys) + elif manager_name == "celery-filesystem": + celery_filesystem_directory = tempfile.mkdtemp(prefix="lc-celery-") + os.environ["CELERY_BROKER"] = "filesystem://" + os.environ[ + "CELERY_BROKER_FILESYSTEM_DIRECTORY" + ] = celery_filesystem_directory + print(f"{celery_filesystem_directory=}") + os.environ["CELERY_BACKEND"] = f"file://{celery_filesystem_directory}" worker = subprocess.Popen( [