Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Empty file added mfr/__init__.py
Empty file.
15 changes: 15 additions & 0 deletions mfr/tasks/__init__.py
Original file line number Diff line number Diff line change
@@ -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',
]
44 changes: 44 additions & 0 deletions mfr/tasks/app.py
Original file line number Diff line number Diff line change
@@ -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()
135 changes: 135 additions & 0 deletions mfr/tasks/core.py
Original file line number Diff line number Diff line change
@@ -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
6 changes: 6 additions & 0 deletions mfr/tasks/exceptions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
class MfrTaskError(Exception):
pass


class WaitTimeOutError(MfrTaskError):
pass
9 changes: 9 additions & 0 deletions mfr/tasks/render.py
Original file line number Diff line number Diff line change
@@ -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=}')
39 changes: 39 additions & 0 deletions mfr/tasks/settings.py
Original file line number Diff line number Diff line change
@@ -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',
]
4 changes: 2 additions & 2 deletions poetry.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
12 changes: 12 additions & 0 deletions tasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)