diff --git a/lib/galaxy/app_unittest_utils/celery_helper.py b/lib/galaxy/app_unittest_utils/celery_helper.py deleted file mode 100644 index d70e0eafb94..00000000000 --- a/lib/galaxy/app_unittest_utils/celery_helper.py +++ /dev/null @@ -1,21 +0,0 @@ -from functools import wraps - - -def rebind_container_to_task(app): - import galaxy.app - - galaxy.app.app = app - from galaxy.celery import ( - CELERY_TASKS, - tasks, - ) - - def magic_bind_dynamic(func): - return wraps(func)(app.magic_partial(func, shared=None)) - - for task in CELERY_TASKS: - task_fn = getattr(tasks, task, None) - if task_fn: - task_fn = getattr(task_fn, "__wrapped__", task_fn) - container_bound_task = magic_bind_dynamic(task_fn) - setattr(tasks, task, container_bound_task) diff --git a/lib/galaxy/app_unittest_utils/galaxy_mock.py b/lib/galaxy/app_unittest_utils/galaxy_mock.py index 53ece3ff6fa..4760f0ba5a0 100644 --- a/lib/galaxy/app_unittest_utils/galaxy_mock.py +++ b/lib/galaxy/app_unittest_utils/galaxy_mock.py @@ -11,6 +11,7 @@ from galaxy import ( quota, ) from galaxy.auth import AuthManager +from galaxy.celery import set_thread_app from galaxy.config import CommonConfigurationMixin from galaxy.jobs.manager import NoopManager from galaxy.managers.users import UserManager @@ -34,7 +35,6 @@ from galaxy.util import StructuredExecutionTimer from galaxy.util.bunch import Bunch from galaxy.util.dbkeys import GenomeBuilds from galaxy.web_stack import ApplicationStack -from .celery_helper import rebind_container_to_task # ============================================================================= @@ -103,7 +103,7 @@ class MockApp(di.Container, GalaxyDataTestApp): self.interactivetool_manager = Bunch(create_interactivetool=lambda *args, **kwargs: None) self.is_job_handler = False self.biotools_metadata_source = None - rebind_container_to_task(self) + set_thread_app(self) def url_for(*args, **kwds): return "/mock/url" diff --git a/lib/galaxy/celery/__init__.py b/lib/galaxy/celery/__init__.py index 6aa5a4c8535..24fdfe3b952 100644 --- a/lib/galaxy/celery/__init__.py +++ b/lib/galaxy/celery/__init__.py @@ -3,16 +3,21 @@ from functools import ( lru_cache, wraps, ) +from threading import local +from typing import ( + Any, + Dict, +) from celery import ( Celery, shared_task, ) from kombu import serialization -from lagom import magic_bind_to_container from galaxy.config import Configuration from galaxy.main_config import find_config +from galaxy.util import ExecutionTimer from galaxy.util.custom_logging import get_logger from galaxy.util.properties import load_app_properties from ._serialization import ( @@ -22,16 +27,39 @@ from ._serialization import ( log = get_logger(__name__) +MAIN_TASK_MODULE = "galaxy.celery.tasks" +TASKS_MODULES = [MAIN_TASK_MODULE] +PYDANTIC_AWARE_SERIALIZER_NAME = "pydantic-aware-json" + +APP_LOCAL = local() + +serialization.register( + PYDANTIC_AWARE_SERIALIZER_NAME, encoder=schema_dumps, decoder=schema_loads, content_type="application/json" +) + + +def set_thread_app(app): + APP_LOCAL.app = app + + +def get_galaxy_app(): + try: + return APP_LOCAL.app + except AttributeError: + import galaxy.app + + if galaxy.app.app: + return galaxy.app.app + return build_app() + @lru_cache(maxsize=1) -def get_galaxy_app(): - import galaxy.app - - if galaxy.app.app: - return galaxy.app.app +def build_app(): kwargs = get_app_properties() if kwargs: kwargs["check_migrate_databases"] = False + import galaxy.app + galaxy_app = galaxy.app.GalaxyManagerApplication(configure_logging=False, **kwargs) return galaxy_app @@ -63,7 +91,13 @@ def get_config(): def get_broker(): config = get_config() if config: - return config.amqp_internal_connection + return config.celery_broker or config.amqp_internal_connection + + +def get_backend(): + config = get_config() + if config: + return config.celery_backend def get_history_audit_table_prune_interval(): @@ -75,7 +109,15 @@ def get_history_audit_table_prune_interval(): broker = get_broker() -celery_app = Celery("galaxy", broker=broker, include=["galaxy.celery.tasks"]) +backend = get_backend() +celery_app_kwd: Dict[str, Any] = { + "broker": broker, + "include": TASKS_MODULES, +} +if backend: + celery_app_kwd["backend"] = backend + +celery_app = Celery("galaxy", **celery_app_kwd) prune_interval = get_history_audit_table_prune_interval() if prune_interval > 0: celery_app.conf.beat_schedule = { @@ -87,28 +129,33 @@ if prune_interval > 0: celery_app.conf.timezone = "UTC" -CELERY_TASKS = [] -PYDANTIC_AWARE_SERIALIER_NAME = "pydantic-aware-json" - - -serialization.register( - PYDANTIC_AWARE_SERIALIER_NAME, encoder=schema_dumps, decoder=schema_loads, content_type="application/json" -) - - -def galaxy_task(*args, **celery_task_kwd): +def galaxy_task(*args, action=None, **celery_task_kwd): if "serializer" not in celery_task_kwd: - celery_task_kwd["serializer"] = PYDANTIC_AWARE_SERIALIER_NAME + celery_task_kwd["serializer"] = PYDANTIC_AWARE_SERIALIZER_NAME def decorate(func): - CELERY_TASKS.append(func.__name__) - @shared_task(**celery_task_kwd) @wraps(func) def wrapper(*args, **kwds): app = get_galaxy_app() assert app - return magic_bind_to_container(app)(func)(*args, **kwds) + desc = func.__name__ + if action is not None: + desc += f" to {action}" + + try: + timer = app.execution_timer_factory.get_timer("internals.tasks.{func.__name__}", desc) + except AttributeError: + timer = ExecutionTimer() + + try: + rval = app.magic_partial(func)(*args, **kwds) + message = f"Successfully executed Celery task {desc} {timer}" + log.info(message) + return rval + except Exception: + log.warning(f"Celery task execution failed for {desc} {timer}") + raise return wrapper diff --git a/lib/galaxy/celery/tasks.py b/lib/galaxy/celery/tasks.py index 8914cfaf193..08f20deadcf 100644 --- a/lib/galaxy/celery/tasks.py +++ b/lib/galaxy/celery/tasks.py @@ -5,13 +5,12 @@ from galaxy.managers.hdas import HDAManager from galaxy.managers.lddas import LDDAManager from galaxy.model.scoped_session import galaxy_scoped_session from galaxy.structured_app import MinimalManagerApp -from galaxy.util import ExecutionTimer from galaxy.util.custom_logging import get_logger log = get_logger(__name__) -@galaxy_task(ignore_result=True) +@galaxy_task(ignore_result=True, action="recalcuate a user's disk usage") def recalculate_user_disk_usage(session: galaxy_scoped_session, user_id=None): if user_id: user = session.query(model.User).get(user_id) @@ -24,13 +23,13 @@ def recalculate_user_disk_usage(session: galaxy_scoped_session, user_id=None): log.error("Recalculate user disk usage task received without user_id.") -@galaxy_task(ignore_result=True) +@galaxy_task(ignore_result=True, action="purge a history dataset") def purge_hda(hda_manager: HDAManager, hda_id): hda = hda_manager.by_id(hda_id) hda_manager._purge(hda) -@galaxy_task +@galaxy_task(action="set dataset association metadata") def set_metadata( hda_manager: HDAManager, ldda_manager: LDDAManager, dataset_id, model_class="HistoryDatasetAssociation" ): @@ -41,7 +40,7 @@ def set_metadata( dataset.datatype.set_meta(dataset) -@galaxy_task(ignore_result=True) +@galaxy_task(ignore_result=True, action="setting up export history job") def export_history( app: MinimalManagerApp, sa_session: galaxy_scoped_session, @@ -61,9 +60,7 @@ def export_history( job_manager.enqueue(job) -@galaxy_task +@galaxy_task(action="pruning history audit table") def prune_history_audit_table(sa_session: galaxy_scoped_session): """Prune ever growing history_audit table.""" - timer = ExecutionTimer() model.HistoryAudit.prune(sa_session) - log.debug(f"Successfully pruned history_audit table {timer}") diff --git a/lib/galaxy/config/schemas/config_schema.yml b/lib/galaxy/config/schemas/config_schema.yml index 2cc4869a34c..01da519db83 100644 --- a/lib/galaxy/config/schemas/config_schema.yml +++ b/lib/galaxy/config/schemas/config_schema.yml @@ -3405,6 +3405,18 @@ mapping: Activate this only if you have setup a Celery worker for Galaxy. For details, see https://docs.galaxyproject.org/en/master/admin/production.html + celery_broker: + type: str + required: false + desc: | + Celery broker (if unset falls back to amqp_internal_connection). + + celery_backend: + type: str + required: false + desc: | + If set, it will be the results backend for Celery. + use_pbkdf2: type: bool default: true diff --git a/lib/galaxy_test/driver/driver_util.py b/lib/galaxy_test/driver/driver_util.py index a20a769c6ba..d7fe6e462b4 100644 --- a/lib/galaxy_test/driver/driver_util.py +++ b/lib/galaxy_test/driver/driver_util.py @@ -26,7 +26,6 @@ import yaml from paste import httpserver from galaxy.app import UniverseApplication as GalaxyUniverseApplication -from galaxy.app_unittest_utils.celery_helper import rebind_container_to_task from galaxy.config import LOGGING_CONFIG_DEFAULT from galaxy.model import mapping from galaxy.model.database_utils import ( @@ -663,9 +662,6 @@ def build_galaxy_app(simple_kwargs) -> GalaxyUniverseApplication: simple_kwargs = load_app_properties(kwds=simple_kwargs) # Build the Universe Application app = GalaxyUniverseApplication(**simple_kwargs) - if not simple_kwargs.get("enable_celery_tasks"): - rebind_container_to_task(app) - log.info("Embedded Galaxy application started") global galaxy_context diff --git a/lib/galaxy_test/driver/integration_util.py b/lib/galaxy_test/driver/integration_util.py index ec1adf7aaed..a01808a4fc9 100644 --- a/lib/galaxy_test/driver/integration_util.py +++ b/lib/galaxy_test/driver/integration_util.py @@ -181,8 +181,21 @@ class IntegrationInstance(UsesApiTestCaseMixin): return os.path.realpath(os.path.join(cls._test_driver.galaxy_test_tmp_dir, name)) +def setup_celery_includes(): + from galaxy.celery import TASKS_MODULES + + def celery_includes(): + return TASKS_MODULES + + return pytest.fixture(scope="session")(celery_includes) + + class UsesCeleryTasks: - enable_celery_tasks = True + @classmethod + def setup_celery_config(cls, config): + config["enable_celery_tasks"] = True + config["celery_broker"] = "memory://" + config["celery_backend"] = "cache+memory://" @pytest.fixture(autouse=True) def _request_celery_app(self, celery_app):