From 1dfcd7ae3cb59ec820384b0edacb1a472c7a9740 Mon Sep 17 00:00:00 2001 From: Wu Clan Date: Thu, 21 Aug 2025 23:46:52 +0800 Subject: [PATCH] Update the celery task result table creation logic (#783) * Update the celery task result table creation logic * Disable beat_sync_every config * Update the prepared comment --- backend/app/task/celery.py | 20 +-- backend/app/task/crud/crud_result.py | 2 +- backend/app/task/database.py | 171 +++++++++++++++++++++ backend/app/task/model/__init__.py | 1 + backend/app/task/model/result.py | 110 ++++++++++++- backend/app/task/service/result_service.py | 2 +- backend/app/task/session.py | 15 ++ 7 files changed, 304 insertions(+), 17 deletions(-) create mode 100644 backend/app/task/database.py create mode 100644 backend/app/task/session.py diff --git a/backend/app/task/celery.py b/backend/app/task/celery.py index 609bef44..3cb510f0 100644 --- a/backend/app/task/celery.py +++ b/backend/app/task/celery.py @@ -5,7 +5,6 @@ import os import celery import celery_aio_pool -from backend.app.task.model.result import OVERWRITE_CELERY_RESULT_GROUP_TABLE_NAME, OVERWRITE_CELERY_RESULT_TABLE_NAME from backend.app.task.tasks.beat import LOCAL_BEAT_SCHEDULE from backend.core.conf import settings from backend.core.path_conf import BASE_PATH @@ -30,22 +29,19 @@ def init_celery() -> celery.Celery: celery.app.trace.build_tracer = celery_aio_pool.build_async_tracer celery.app.trace.reset_worker_optimizations() + # https://docs.celeryq.dev/en/stable/userguide/configuration.html app = celery.Celery( 'fba_celery', - broker=f'redis://:{settings.REDIS_PASSWORD}@{settings.REDIS_HOST}:{settings.REDIS_PORT}/{settings.CELERY_BROKER_REDIS_DATABASE}' + broker_url=f'redis://:{settings.REDIS_PASSWORD}@{settings.REDIS_HOST}:{settings.REDIS_PORT}/{settings.CELERY_BROKER_REDIS_DATABASE}' if settings.CELERY_BROKER == 'redis' else f'amqp://{settings.CELERY_RABBITMQ_USERNAME}:{settings.CELERY_RABBITMQ_PASSWORD}@{settings.CELERY_RABBITMQ_HOST}:{settings.CELERY_RABBITMQ_PORT}', broker_connection_retry_on_startup=True, - backend=f'db+{settings.DATABASE_TYPE}+{"pymysql" if settings.DATABASE_TYPE == "mysql" else "psycopg"}' + result_backend=f'db+{settings.DATABASE_TYPE}+{"pymysql" if settings.DATABASE_TYPE == "mysql" else "psycopg"}' f'://{settings.DATABASE_USER}:{settings.DATABASE_PASSWORD}@{settings.DATABASE_HOST}:{settings.DATABASE_PORT}/{settings.DATABASE_SCHEMA}', - database_engine_options={'echo': settings.DATABASE_ECHO}, - database_table_names={ - 'task': OVERWRITE_CELERY_RESULT_TABLE_NAME, - 'group': OVERWRITE_CELERY_RESULT_GROUP_TABLE_NAME, - }, result_extended=True, - # result_expires=0, # 清理任务结果,默认每天凌晨 4 点,0 或 None 表示不清理 - # beat_sync_every=1, # 保存任务状态周期,默认 3 * 60 秒 + database_engine_options={'echo': settings.DATABASE_ECHO}, + # result_expires=0, + # beat_sync_every=1, beat_schedule=LOCAL_BEAT_SCHEDULE, beat_scheduler='backend.app.task.utils.schedulers:DatabaseScheduler', task_cls='backend.app.task.tasks.base:TaskBase', @@ -54,6 +50,10 @@ def init_celery() -> celery.Celery: timezone=settings.DATETIME_TIMEZONE, ) + # 在 Celery 中设置此参数无效 + # 参数:https://github.com/celery/celery/issues/7270 + app.loader.override_backends = {'db': 'backend.app.task.database:DatabaseBackend'} + # 自动发现任务 packages = find_task_packages() app.autodiscover_tasks(packages) diff --git a/backend/app/task/crud/crud_result.py b/backend/app/task/crud/crud_result.py index 3b169222..224295f5 100644 --- a/backend/app/task/crud/crud_result.py +++ b/backend/app/task/crud/crud_result.py @@ -4,7 +4,7 @@ from sqlalchemy import Select from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy_crud_plus import CRUDPlus -from backend.app.task.model.result import TaskResult +from backend.app.task.model import TaskResult class CRUDTaskResult(CRUDPlus[TaskResult]): diff --git a/backend/app/task/database.py b/backend/app/task/database.py new file mode 100644 index 00000000..3f589f84 --- /dev/null +++ b/backend/app/task/database.py @@ -0,0 +1,171 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from celery import states +from celery.backends.base import BaseBackend +from celery.backends.database import retry, session_cleanup +from celery.exceptions import ImproperlyConfigured +from celery.utils.time import maybe_timedelta + +from backend.app.task.model.result import Task, TaskExtended, TaskSet +from backend.app.task.session import SessionManager + +""" +重写 from celery.backends.database 内部 DatabaseBackend 类,此类实现与模型配合不佳,导致 fba 创建表和 alembic 迁移困难 +""" + + +class DatabaseBackend(BaseBackend): + """The database result backend.""" + + # ResultSet.iterate should sleep this much between each pool, + # to not bombard the database with queries. + subpolling_interval = 0.5 + + task_cls = Task + taskset_cls = TaskSet + + def __init__(self, dburi=None, engine_options=None, url=None, **kwargs): + # The `url` argument was added later and is used by + # the app to set backend by url (celery.app.backends.by_url) + super().__init__(expires_type=maybe_timedelta, url=url, **kwargs) + conf = self.app.conf + + if self.extended_result: + self.task_cls = TaskExtended + + self.url = url or dburi or conf.database_url + self.engine_options = dict(engine_options or {}, **conf.database_engine_options or {}) + self.short_lived_sessions = kwargs.get('short_lived_sessions', conf.database_short_lived_sessions) + + schemas = conf.database_table_schemas or {} + tablenames = conf.database_table_names or {} + self.task_cls.configure(schema=schemas.get('task'), name=tablenames.get('task')) + self.taskset_cls.configure(schema=schemas.get('group'), name=tablenames.get('group')) + + if not self.url: + raise ImproperlyConfigured( + 'Missing connection string! Do you have the database_url setting set to a real value?' + ) + + self.session_manager = SessionManager() + + create_tables_at_setup = conf.database_create_tables_at_setup + if create_tables_at_setup is True: + self._create_tables() + + @property + def extended_result(self): + return self.app.conf.find_value_for_key('extended', 'result') + + def _create_tables(self): + """Create the task and taskset tables.""" + self.ResultSession() + + def ResultSession(self, session_manager=None): + if session_manager is None: + session_manager = self.session_manager + return session_manager.session_factory( + dburi=self.url, short_lived_sessions=self.short_lived_sessions, **self.engine_options + ) + + @retry + def _store_result(self, task_id, result, state, traceback=None, request=None, **kwargs): + """Store return value and state of an executed task.""" + session = self.ResultSession() + with session_cleanup(session): + task = list(session.query(self.task_cls).filter(self.task_cls.task_id == task_id)) + task = task and task[0] + if not task: + task = self.task_cls(task_id) + task.task_id = task_id + session.add(task) + session.flush() + + self._update_result(task, result, state, traceback=traceback, request=request) + session.commit() + + def _update_result(self, task, result, state, traceback=None, request=None): + meta = self._get_result_meta( + result=result, state=state, traceback=traceback, request=request, format_date=False, encode=True + ) + + # Exclude the primary key id and task_id columns + # as we should not set it None + columns = [column.name for column in self.task_cls.__table__.columns if column.name not in {'id', 'task_id'}] + + # Iterate through the columns name of the table + # to set the value from meta. + # If the value is not present in meta, set None + for column in columns: + value = meta.get(column) + setattr(task, column, value) + + @retry + def _get_task_meta_for(self, task_id): + """Get task meta-data for a task by id.""" + session = self.ResultSession() + with session_cleanup(session): + task = list(session.query(self.task_cls).filter(self.task_cls.task_id == task_id)) + task = task and task[0] + if not task: + task = self.task_cls(task_id) + task.status = states.PENDING + task.result = None + data = task.to_dict() + if data.get('args', None) is not None: + data['args'] = self.decode(data['args']) + if data.get('kwargs', None) is not None: + data['kwargs'] = self.decode(data['kwargs']) + return self.meta_from_decoded(data) + + @retry + def _save_group(self, group_id, result): + """Store the result of an executed group.""" + session = self.ResultSession() + with session_cleanup(session): + group = self.taskset_cls(group_id, result) + session.add(group) + session.flush() + session.commit() + return result + + @retry + def _restore_group(self, group_id): + """Get meta-data for group by id.""" + session = self.ResultSession() + with session_cleanup(session): + group = session.query(self.taskset_cls).filter(self.taskset_cls.taskset_id == group_id).first() + if group: + return group.to_dict() + + @retry + def _delete_group(self, group_id): + """Delete meta-data for group by id.""" + session = self.ResultSession() + with session_cleanup(session): + session.query(self.taskset_cls).filter(self.taskset_cls.taskset_id == group_id).delete() + session.flush() + session.commit() + + @retry + def _forget(self, task_id): + """Forget about result.""" + session = self.ResultSession() + with session_cleanup(session): + session.query(self.task_cls).filter(self.task_cls.task_id == task_id).delete() + session.commit() + + def cleanup(self): + """Delete expired meta-data.""" + session = self.ResultSession() + expires = self.expires + now = self.app.now() + with session_cleanup(session): + session.query(self.task_cls).filter(self.task_cls.date_done < (now - expires)).delete() + session.query(self.taskset_cls).filter(self.taskset_cls.date_done < (now - expires)).delete() + session.commit() + + def __reduce__(self, args=(), kwargs=None): + kwargs = {} if not kwargs else kwargs + kwargs.update({'dburi': self.url, 'expires': self.expires, 'engine_options': self.engine_options}) + return super().__reduce__(args, kwargs) diff --git a/backend/app/task/model/__init__.py b/backend/app/task/model/__init__.py index 3f980cc8..60f2aaac 100644 --- a/backend/app/task/model/__init__.py +++ b/backend/app/task/model/__init__.py @@ -1,3 +1,4 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- +from backend.app.task.model.result import TaskExtended as TaskResult from backend.app.task.model.scheduler import TaskScheduler diff --git a/backend/app/task/model/result.py b/backend/app/task/model/result.py index a2b2576a..b846ddab 100644 --- a/backend/app/task/model/result.py +++ b/backend/app/task/model/result.py @@ -1,9 +1,109 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- -from celery.backends.database.models import TaskExtended as TaskResult +from datetime import datetime, timezone -OVERWRITE_CELERY_RESULT_TABLE_NAME = 'task_result' -OVERWRITE_CELERY_RESULT_GROUP_TABLE_NAME = 'task_group_result' +import sqlalchemy as sa -# 重写表名配置 -TaskResult.configure(name=OVERWRITE_CELERY_RESULT_TABLE_NAME) +from celery import states +from sqlalchemy.types import PickleType + +from backend.common.model import MappedBase + +""" +重写 celery.backends.database.models 内部所有模型,适配 fba 创建表和 alembic 迁移 +""" + + +class Task(MappedBase): + """Task result/status.""" + + __tablename__ = 'task_result' + __table_args__ = {'comment': '任务结果表'} + + id = sa.Column(sa.Integer, sa.Sequence('task_id_sequence'), primary_key=True, autoincrement=True) + task_id = sa.Column(sa.String(155), unique=True) + status = sa.Column(sa.String(50), default=states.PENDING) + result = sa.Column(PickleType, nullable=True) + date_done = sa.Column( + sa.DateTime, default=datetime.now(timezone.utc), onupdate=datetime.now(timezone.utc), nullable=True + ) + traceback = sa.Column(sa.Text, nullable=True) + + def __init__(self, task_id): + self.task_id = task_id + + def to_dict(self): + return { + 'task_id': self.task_id, + 'status': self.status, + 'result': self.result, + 'traceback': self.traceback, + 'date_done': self.date_done, + } + + def __repr__(self): + return ''.format(self) + + @classmethod + def configure(cls, schema=None, name=None): + cls.__table__.schema = schema + cls.id.default.schema = schema + cls.__table__.name = name or cls.__tablename__ + + +class TaskExtended(Task): + """For the extend result.""" + + __tablename__ = 'task_result' + __table_args__ = {'extend_existing': True, 'comment': '任务结果表'} + + name = sa.Column(sa.String(155), nullable=True) + args = sa.Column(sa.LargeBinary, nullable=True) + kwargs = sa.Column(sa.LargeBinary, nullable=True) + worker = sa.Column(sa.String(155), nullable=True) + retries = sa.Column(sa.Integer, nullable=True) + queue = sa.Column(sa.String(155), nullable=True) + + def to_dict(self): + task_dict = super().to_dict() + task_dict.update({ + 'name': self.name, + 'args': self.args, + 'kwargs': self.kwargs, + 'worker': self.worker, + 'retries': self.retries, + 'queue': self.queue, + }) + return task_dict + + +class TaskSet(MappedBase): + """TaskSet result.""" + + __tablename__ = 'task_set_result' + __table_args__ = {'comment': '任务集结果表'} + + id = sa.Column(sa.Integer, sa.Sequence('taskset_id_sequence'), autoincrement=True, primary_key=True) + taskset_id = sa.Column(sa.String(155), unique=True) + result = sa.Column(PickleType, nullable=True) + date_done = sa.Column(sa.DateTime, default=datetime.now(timezone.utc), nullable=True) + + def __init__(self, taskset_id, result): + self.taskset_id = taskset_id + self.result = result + + def to_dict(self): + return { + 'taskset_id': self.taskset_id, + 'result': self.result, + 'date_done': self.date_done, + } + + def __repr__(self): + return f'' + + @classmethod + def configure(cls, schema=None, name=None): + cls.__table__.schema = schema + cls.id.default.schema = schema + cls.__table__.name = name or cls.__tablename__ diff --git a/backend/app/task/service/result_service.py b/backend/app/task/service/result_service.py index bb739137..dccaf084 100644 --- a/backend/app/task/service/result_service.py +++ b/backend/app/task/service/result_service.py @@ -3,7 +3,7 @@ from sqlalchemy import Select from backend.app.task.crud.crud_result import task_result_dao -from backend.app.task.model.result import TaskResult +from backend.app.task.model import TaskResult from backend.app.task.schema.result import DeleteTaskResultParam from backend.common.exception import errors from backend.database.db import async_db_session diff --git a/backend/app/task/session.py b/backend/app/task/session.py new file mode 100644 index 00000000..4ee41d50 --- /dev/null +++ b/backend/app/task/session.py @@ -0,0 +1,15 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from celery.backends.database.session import SessionManager as CelerySessionManager + + +class SessionManager(CelerySessionManager): + """ + 重写 celery SessionManager + """ + + def __init__(self): + super().__init__() + + # 禁止自动创建 celery 内部定义的任务结果表 + self.prepared = True