Skip to content
Draft
54 changes: 37 additions & 17 deletions dash/background_callback/managers/celery_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,21 +39,25 @@ 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")

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)
Expand Down Expand Up @@ -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
Expand All @@ -103,47 +116,54 @@ 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

result = json.loads(result)

# 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)

Expand All @@ -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)):
Expand All @@ -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),
)

Expand Down
1 change: 1 addition & 0 deletions requirements/ci.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 3 additions & 4 deletions tests/async_tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
6 changes: 3 additions & 3 deletions tests/async_tests/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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"
Expand Down
7 changes: 3 additions & 4 deletions tests/background_callback/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
60 changes: 48 additions & 12 deletions tests/background_callback/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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(
[
Expand Down