diff --git a/mfr/__init__.py b/mfr/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/mfr/tasks/__init__.py b/mfr/tasks/__init__.py new file mode 100644 index 00000000..b28aa8c6 --- /dev/null +++ b/mfr/tasks/__init__.py @@ -0,0 +1,15 @@ +from mfr.tasks.app import app +from mfr.tasks.render import render +from mfr.tasks.core import celery_task +from mfr.tasks.core import backgrounded +from mfr.tasks.core import wait_on_celery +from mfr.tasks.exceptions import WaitTimeOutError + +__all__ = [ + 'app', + 'render', + 'celery_task', + 'backgrounded', + 'wait_on_celery', + 'WaitTimeOutError', +] diff --git a/mfr/tasks/app.py b/mfr/tasks/app.py new file mode 100644 index 00000000..4ad3321b --- /dev/null +++ b/mfr/tasks/app.py @@ -0,0 +1,44 @@ +import logging + +from celery import Celery +from celery.signals import task_failure + +import sentry_sdk +from sentry_sdk.integrations.celery import CeleryIntegration +from sentry_sdk.integrations.logging import LoggingIntegration + +from mfr.settings import config +from mfr.version import __version__ +from mfr.tasks import settings as tasks_settings + +logger = logging.getLogger(__name__) + +app = Celery() +app.config_from_object(tasks_settings) + + +def register_signal(): + """Adapted from `raven.contrib.celery.register_signal`. Remove args and + kwargs from logs so that keys aren't leaked to Sentry. + """ + def process_failure_signal(sender, task_id, *args, **kwargs): + scope = sentry_sdk.get_current_scope() + scope.set_tag('task_id', task_id) + scope.set_tag('task', sender) + sentry_sdk.capture_exception() + + task_failure.connect(process_failure_signal, weak=False) + + +sentry_dsn = config.get_nullable('SENTRY_DSN', None) +if sentry_dsn: + sentry_logging = LoggingIntegration( + level=logging.INFO, # Capture INFO level and above as breadcrumbs + event_level=None, # Do not send logs of any level as events + ) + sentry_sdk.init( + sentry_dsn, + release=__version__, + integrations=[CeleryIntegration(), sentry_logging] + ) + register_signal() diff --git a/mfr/tasks/core.py b/mfr/tasks/core.py new file mode 100644 index 00000000..59a6342a --- /dev/null +++ b/mfr/tasks/core.py @@ -0,0 +1,135 @@ +import os +import pickle +import asyncio +import logging +import functools + +from celery.backends.base import DisabledBackend + +from mfr.tasks.app import app +from mfr.tasks import settings +from mfr.tasks import exceptions + +logger = logging.getLogger(__name__) + + +def ensure_event_loop(): + """Ensure the existance of an eventloop + Useful for contexts where get_event_loop() may + raise an exception. + :returns: The new event loop + :rtype: BaseEventLoop + """ + try: + return asyncio.get_event_loop() + except (AssertionError, RuntimeError): + asyncio.set_event_loop(asyncio.new_event_loop()) + + # Note: No clever tricks are used here to dry up code + # This avoids an infinite loop if settings the event loop ever fails + return asyncio.get_event_loop() + + +def __coroutine_unwrapper(func): + @functools.wraps(func) + def wrapped(*args, **kwargs): + return ensure_event_loop().run_until_complete(func(*args, **kwargs)) + wrapped.as_async = func + return wrapped + + +async def backgrounded(func, *args, **kwargs): + """Runs the given function with the given arguments in + a background thread + """ + loop = asyncio.get_event_loop() + if asyncio.iscoroutinefunction(func): + func = __coroutine_unwrapper(func) + + return (await loop.run_in_executor( + None, # None uses the default executer, ThreadPoolExecuter + functools.partial(func, *args, **kwargs) + )) + + +def backgroundify(func): + @functools.wraps(func) + async def wrapped(*args, **kwargs): + return await backgrounded(func, *args, **kwargs) + return wrapped + + +def adhoc_file_backend(func, was_bound=False, basepath=None): + basepath = basepath or settings.ADHOC_BACKEND_PATH + + @functools.wraps(func) + def wrapped(task, *args, **kwargs): + if was_bound: + args = (task,) + args + + try: + result = func(*args, **kwargs) + except Exception as e: + result = e + + with open(os.path.join(basepath, task.request.id), 'wb') as result_file: + pickle.dump(result, result_file) + + if isinstance(result, Exception): + raise result + return result + return wrapped + + +def celery_task(func, *args, **kwargs): + """A wrapper around Celery.task. When the wrapped method is called it will be called using + Celery's Task.delay function and run in a background thread. + + If the celery backend is disabled, the task will be wrapped in a function that will write the + result to disk using the pickle serialization protocol. + """ + task_func = __coroutine_unwrapper(func) + + if isinstance(app.backend, DisabledBackend): + task_func = adhoc_file_backend( + task_func, + was_bound=kwargs.pop('bind', False) + ) + kwargs['bind'] = True + + logger.debug(f'celery_task: task_func:({task_func})') + + task = app.task(task_func, **kwargs) + task.adelay = backgroundify(task.delay) + + return task + + +@backgroundify +async def wait_on_celery(result, interval=None, timeout=None, basepath=None): + timeout = timeout or settings.WAIT_TIMEOUT + interval = interval or settings.WAIT_INTERVAL + basepath = basepath or settings.ADHOC_BACKEND_PATH + + waited = 0 + + while True: + if isinstance(app.backend, DisabledBackend): + try: + with open(os.path.join(basepath, result.id), 'rb') as result_file: + data = pickle.load(result_file) + if isinstance(data, Exception): + raise data + return data + except FileNotFoundError: + pass + else: + if result.ready(): + if result.failed(): + raise result.result + return result.result + + if waited > timeout: + raise exceptions.WaitTimeOutError + await asyncio.sleep(interval) + waited += interval diff --git a/mfr/tasks/exceptions.py b/mfr/tasks/exceptions.py new file mode 100644 index 00000000..46ddb8f2 --- /dev/null +++ b/mfr/tasks/exceptions.py @@ -0,0 +1,6 @@ +class MfrTaskError(Exception): + pass + + +class WaitTimeOutError(MfrTaskError): + pass diff --git a/mfr/tasks/render.py b/mfr/tasks/render.py new file mode 100644 index 00000000..0bcbd125 --- /dev/null +++ b/mfr/tasks/render.py @@ -0,0 +1,9 @@ +import logging + +from mfr.tasks import core +logger = logging.getLogger(__name__) + + +@core.celery_task +async def render(*args, **kwargs): + logger.critical(f'Received task with {args=} and {kwargs=}') diff --git a/mfr/tasks/settings.py b/mfr/tasks/settings.py new file mode 100644 index 00000000..dfdc0ad8 --- /dev/null +++ b/mfr/tasks/settings.py @@ -0,0 +1,39 @@ +import os + +from kombu import Queue, Exchange + +from mfr import settings + + +config = settings.child('TASKS_CONFIG') + +WAIT_TIMEOUT = int(config.get('WAIT_TIMEOUT', 20)) +WAIT_INTERVAL = float(config.get('WAIT_INTERVAL', 0.5)) +ADHOC_BACKEND_PATH = config.get('ADHOC_BACKEND_PATH', '/tmp') + +broker_url = config.get( + 'BROKER_URL', + 'amqp://{}:{}//'.format( + os.environ.get('RABBITMQ_PORT_5672_TCP_ADDR', ''), + os.environ.get('RABBITMQ_PORT_5672_TCP_PORT', ''), + ) +) + +task_default_queue = config.get('CELERY_DEFAULT_QUEUE', 'mfr') +task_queues = ( + Queue('mfr', Exchange('mfr'), routing_key='mfr'), +) + +task_always_eager = config.get_bool('CELERY_ALWAYS_EAGER', False) +result_backend = config.get_nullable('CELERY_RESULT_BACKEND', 'rpc://') +result_persistent = config.get_bool('CELERY_RESULT_PERSISTENT', True) +worker_disable_rate_limits = config.get_bool('CELERY_DISABLE_RATE_LIMITS', True) +result_expires = int(config.get('CELERY_TASK_RESULT_EXPIRES', 60)) +task_create_missing_queues = config.get_bool('CELERY_CREATE_MISSING_QUEUES', False) +task_acks_late = True +worker_hijack_root_logger = False +task_eager_propagates = True + +imports = [ + 'mfr.tasks.render', +] diff --git a/poetry.lock b/poetry.lock index f1a95e34..a61ec2c4 100644 --- a/poetry.lock +++ b/poetry.lock @@ -4129,7 +4129,7 @@ yarl = "1.17.0" type = "git" url = "https://github.com/CenterForOpenScience/waterbutler.git" reference = "feature/buff-worms" -resolved_reference = "85e5aebe9c72768820a0367101db5cf9fc361dfd" +resolved_reference = "75d049732ff161a3e3979b7aac67ce922e56d47f" [[package]] name = "wcwidth" @@ -4452,4 +4452,4 @@ propcache = ">=0.2.0" [metadata] lock-version = "2.1" python-versions = "^3.13" -content-hash = "a1bb6d0d549dccb62b77695b28a912d67e8a45ed48f0ff83d46cc5289d67b54a" +content-hash = "c72ab4a656dbde2aad8f98b0aa39d0ebb04b8b4f097c3060bbac98ae670c71dc" diff --git a/pyproject.toml b/pyproject.toml index d68b3b45..f8c158da 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,6 +45,7 @@ openpyxl = "^3.1" waterbutler = { git = "https://github.com/CenterForOpenScience/waterbutler.git", branch = "feature/buff-worms" } markupsafe = "2.0.1" +celery = "5.5.0" [tool.poetry.group.dev] optional = true diff --git a/tasks.py b/tasks.py index 4860de4d..f4f7c2af 100644 --- a/tasks.py +++ b/tasks.py @@ -58,3 +58,15 @@ def server(ctx): from mfr.server.app import serve serve() + +@task +def celery(ctx, loglevel='INFO', hostname='%h', concurrency=None): + from mfr.tasks.app import app + command = ['worker'] + if loglevel: + command.extend(['--loglevel', loglevel]) + if hostname: + command.extend(['--hostname', hostname]) + if concurrency: + command.extend(['--concurrency', concurrency]) + app.worker_main(command)