mirror of
https://github.com/fastapi-practices/fastapi-best-architecture.git
synced 2026-09-21 21:15:13 +00:00
Update the ruff rules and format the code (#846)
* Update the ruff rules and format the code * Update the per-file-ignores * Update the ci * Update rules * Fix codes * Fix pagination * Update rules
This commit is contained in:
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
import sys
|
||||
|
||||
from backend.core.path_conf import BASE_PATH
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
from backend.app.task.celery import celery_app
|
||||
@@ -7,7 +5,7 @@ from backend.common.socketio.server import sio
|
||||
|
||||
|
||||
@sio.event
|
||||
async def task_worker_status(sid, data):
|
||||
async def task_worker_status(sid, data) -> None: # noqa: ANN001
|
||||
"""任务 Worker 状态事件"""
|
||||
worker = await run_in_threadpool(celery_app.control.ping)
|
||||
await sio.emit('task_worker_status', worker, sid)
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from fastapi import APIRouter
|
||||
|
||||
from backend.app.task.api.v1.control import router as task_control_router
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Path
|
||||
@@ -24,7 +22,7 @@ async def get_task_registered() -> ResponseSchemaModel[list[TaskRegisteredDetail
|
||||
raise errors.ServerError(msg='Celery Worker 暂不可用,请稍后重试')
|
||||
task_registered = []
|
||||
celery_app_tasks = celery_app.tasks
|
||||
for _, tasks in registered.items():
|
||||
for tasks in registered.values():
|
||||
for task in tasks:
|
||||
task_ins = celery_app_tasks.get(task)
|
||||
if task_ins:
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Path, Query
|
||||
|
||||
from backend.app.task.schema.scheduler import CreateTaskSchedulerParam, GetTaskSchedulerDetail, UpdateTaskSchedulerParam
|
||||
from backend.app.task.schema.scheduler import (
|
||||
CreateTaskSchedulerParam,
|
||||
GetTaskSchedulerDetail,
|
||||
UpdateTaskSchedulerParam,
|
||||
)
|
||||
from backend.app.task.service.scheduler_service import task_scheduler_service
|
||||
from backend.common.pagination import DependsPagination, PageData, paging_data
|
||||
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
|
||||
@@ -40,7 +42,7 @@ async def get_task_scheduler(
|
||||
)
|
||||
async def get_task_scheduler_paged(
|
||||
db: CurrentSession,
|
||||
name: Annotated[int, Path(description='任务调度名称')] = None,
|
||||
name: Annotated[int | None, Path(description='任务调度名称')] = None,
|
||||
type: Annotated[int | None, Query(description='任务调度类型')] = None,
|
||||
) -> ResponseSchemaModel[PageData[GetTaskSchedulerDetail]]:
|
||||
task_scheduler_select = await task_scheduler_service.get_select(name=name, type=type)
|
||||
@@ -70,7 +72,8 @@ async def create_task_scheduler(obj: CreateTaskSchedulerParam) -> ResponseModel:
|
||||
],
|
||||
)
|
||||
async def update_task_scheduler(
|
||||
pk: Annotated[int, Path(description='任务调度 ID')], obj: UpdateTaskSchedulerParam
|
||||
pk: Annotated[int, Path(description='任务调度 ID')],
|
||||
obj: UpdateTaskSchedulerParam,
|
||||
) -> ResponseModel:
|
||||
count = await task_scheduler_service.update(pk=pk, obj=obj)
|
||||
if count > 0:
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
import os
|
||||
|
||||
import celery
|
||||
@@ -10,10 +8,10 @@ from backend.core.conf import settings
|
||||
from backend.core.path_conf import BASE_PATH
|
||||
|
||||
|
||||
def find_task_packages():
|
||||
def find_task_packages() -> list[str]:
|
||||
packages = []
|
||||
task_dir = os.path.join(BASE_PATH, 'app', 'task', 'tasks')
|
||||
for root, dirs, files in os.walk(task_dir):
|
||||
task_dir = BASE_PATH / 'app' / 'task' / 'tasks'
|
||||
for root, _dirs, files in os.walk(task_dir):
|
||||
if 'tasks.py' in files:
|
||||
package = root.replace(str(BASE_PATH.parent) + os.path.sep, '').replace(os.path.sep, '.')
|
||||
packages.append(package)
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from sqlalchemy import Select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy_crud_plus import CRUDPlus
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from typing import Sequence
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy import Select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -86,7 +84,7 @@ class CRUDTaskScheduler(CRUDPlus[TaskScheduler]):
|
||||
TaskScheduler.no_changes = False
|
||||
return 1
|
||||
|
||||
async def set_status(self, db: AsyncSession, pk: int, status: bool) -> int:
|
||||
async def set_status(self, db: AsyncSession, pk: int, *, status: bool) -> int:
|
||||
"""
|
||||
设置任务调度状态
|
||||
|
||||
@@ -96,7 +94,7 @@ class CRUDTaskScheduler(CRUDPlus[TaskScheduler]):
|
||||
:return:
|
||||
"""
|
||||
task_scheduler = await self.get(db, pk)
|
||||
setattr(task_scheduler, 'enabled', status)
|
||||
task_scheduler.enabled = status
|
||||
TaskScheduler.no_changes = False
|
||||
return 1
|
||||
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
#!/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 sqlalchemy import PickleType
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from backend.app.task.model.result import Task, TaskExtended, TaskSet
|
||||
from backend.app.task.session import SessionManager
|
||||
@@ -24,7 +24,7 @@ class DatabaseBackend(BaseBackend):
|
||||
task_cls = Task
|
||||
taskset_cls = TaskSet
|
||||
|
||||
def __init__(self, dburi=None, engine_options=None, url=None, **kwargs):
|
||||
def __init__(self, dburi=None, engine_options=None, url=None, **kwargs) -> None: # noqa: ANN001
|
||||
# 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)
|
||||
@@ -44,7 +44,7 @@ class DatabaseBackend(BaseBackend):
|
||||
|
||||
if not self.url:
|
||||
raise ImproperlyConfigured(
|
||||
'Missing connection string! Do you have the database_url setting set to a real value?'
|
||||
'Missing connection string! Do you have the database_url setting set to a real value?',
|
||||
)
|
||||
|
||||
self.session_manager = SessionManager()
|
||||
@@ -54,24 +54,26 @@ class DatabaseBackend(BaseBackend):
|
||||
self._create_tables()
|
||||
|
||||
@property
|
||||
def extended_result(self):
|
||||
def extended_result(self): # noqa: ANN201
|
||||
return self.app.conf.find_value_for_key('extended', 'result')
|
||||
|
||||
def _create_tables(self):
|
||||
def _create_tables(self) -> None:
|
||||
"""Create the task and taskset tables."""
|
||||
self.ResultSession()
|
||||
self.result_session()
|
||||
|
||||
def ResultSession(self, session_manager=None):
|
||||
def result_session(self, session_manager=None) -> Session: # noqa: ANN001
|
||||
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
|
||||
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):
|
||||
def _store_result(self, task_id, result, state, traceback=None, request=None, **kwargs) -> None: # noqa: ANN001
|
||||
"""Store return value and state of an executed task."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
with session_cleanup(session):
|
||||
task = list(session.query(self.task_cls).filter(self.task_cls.task_id == task_id))
|
||||
task = task and task[0]
|
||||
@@ -84,9 +86,14 @@ class DatabaseBackend(BaseBackend):
|
||||
self._update_result(task, result, state, traceback=traceback, request=request)
|
||||
session.commit()
|
||||
|
||||
def _update_result(self, task, result, state, traceback=None, request=None):
|
||||
def _update_result(self, task, result, state, traceback=None, request=None) -> None: # noqa: ANN001
|
||||
meta = self._get_result_meta(
|
||||
result=result, state=state, traceback=traceback, request=request, format_date=False, encode=True
|
||||
result=result,
|
||||
state=state,
|
||||
traceback=traceback,
|
||||
request=request,
|
||||
format_date=False,
|
||||
encode=True,
|
||||
)
|
||||
|
||||
# Exclude the primary key id and task_id columns
|
||||
@@ -101,9 +108,9 @@ class DatabaseBackend(BaseBackend):
|
||||
setattr(task, column, value)
|
||||
|
||||
@retry
|
||||
def _get_task_meta_for(self, task_id):
|
||||
def _get_task_meta_for(self, task_id: str): # noqa: ANN202
|
||||
"""Get task meta-data for a task by id."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
with session_cleanup(session):
|
||||
task = list(session.query(self.task_cls).filter(self.task_cls.task_id == task_id))
|
||||
task = task and task[0]
|
||||
@@ -119,9 +126,9 @@ class DatabaseBackend(BaseBackend):
|
||||
return self.meta_from_decoded(data)
|
||||
|
||||
@retry
|
||||
def _save_group(self, group_id, result):
|
||||
def _save_group(self, group_id: str, result: PickleType): # noqa: ANN202
|
||||
"""Store the result of an executed group."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
with session_cleanup(session):
|
||||
group = self.taskset_cls(group_id, result)
|
||||
session.add(group)
|
||||
@@ -130,34 +137,34 @@ class DatabaseBackend(BaseBackend):
|
||||
return result
|
||||
|
||||
@retry
|
||||
def _restore_group(self, group_id):
|
||||
def _restore_group(self, group_id: str) -> dict | None:
|
||||
"""Get meta-data for group by id."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
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):
|
||||
def _delete_group(self, group_id: str) -> None:
|
||||
"""Delete meta-data for group by id."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
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):
|
||||
def _forget(self, task_id: str) -> None:
|
||||
"""Forget about result."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
with session_cleanup(session):
|
||||
session.query(self.task_cls).filter(self.task_cls.task_id == task_id).delete()
|
||||
session.commit()
|
||||
|
||||
def cleanup(self):
|
||||
def cleanup(self) -> None:
|
||||
"""Delete expired meta-data."""
|
||||
session = self.ResultSession()
|
||||
session = self.result_session()
|
||||
expires = self.expires
|
||||
now = self.app.now()
|
||||
with session_cleanup(session):
|
||||
@@ -165,7 +172,7 @@ class DatabaseBackend(BaseBackend):
|
||||
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
|
||||
def __reduce__(self, args=(), kwargs=None): # noqa: ANN001, ANN204
|
||||
kwargs = kwargs or {}
|
||||
kwargs.update({'dburi': self.url, 'expires': self.expires, 'engine_options': self.engine_options})
|
||||
return super().__reduce__(args, kwargs)
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from backend.common.enums import IntEnum, StrEnum
|
||||
|
||||
|
||||
|
||||
@@ -1,4 +1,2 @@
|
||||
#!/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
|
||||
from backend.app.task.model.result import TaskExtended as TaskResult # noqa: F401
|
||||
from backend.app.task.model.scheduler import TaskScheduler as TaskScheduler
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import sqlalchemy as sa
|
||||
@@ -25,14 +23,17 @@ class Task(MappedBase):
|
||||
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
|
||||
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):
|
||||
def __init__(self, task_id: str) -> None:
|
||||
self.task_id = task_id
|
||||
|
||||
def to_dict(self):
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
'task_id': self.task_id,
|
||||
'status': self.status,
|
||||
@@ -41,11 +42,11 @@ class Task(MappedBase):
|
||||
'date_done': self.date_done,
|
||||
}
|
||||
|
||||
def __repr__(self):
|
||||
return '<Task {0.task_id} state: {0.status}>'.format(self)
|
||||
def __repr__(self) -> str:
|
||||
return f'<Task {self.task_id} state: {self.status}>'
|
||||
|
||||
@classmethod
|
||||
def configure(cls, schema=None, name=None):
|
||||
def configure(cls, schema=None, name=None) -> None: # noqa: ANN001
|
||||
cls.__table__.schema = schema
|
||||
cls.id.default.schema = schema
|
||||
cls.__table__.name = name or cls.__tablename__
|
||||
@@ -64,7 +65,7 @@ class TaskExtended(Task):
|
||||
retries = sa.Column(sa.Integer, nullable=True)
|
||||
queue = sa.Column(sa.String(155), nullable=True)
|
||||
|
||||
def to_dict(self):
|
||||
def to_dict(self) -> dict:
|
||||
task_dict = super().to_dict()
|
||||
task_dict.update({
|
||||
'name': self.name,
|
||||
@@ -88,22 +89,22 @@ class TaskSet(MappedBase):
|
||||
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):
|
||||
def __init__(self, taskset_id, result) -> None: # noqa: ANN001
|
||||
self.taskset_id = taskset_id
|
||||
self.result = result
|
||||
|
||||
def to_dict(self):
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
'taskset_id': self.taskset_id,
|
||||
'result': self.result,
|
||||
'date_done': self.date_done,
|
||||
}
|
||||
|
||||
def __repr__(self):
|
||||
def __repr__(self) -> str:
|
||||
return f'<TaskSet: {self.taskset_id}>'
|
||||
|
||||
@classmethod
|
||||
def configure(cls, schema=None, name=None):
|
||||
def configure(cls, schema=None, name=None) -> None: # noqa: ANN001
|
||||
cls.__table__.schema = schema
|
||||
cls.id.default.schema = schema
|
||||
cls.__table__.name = name or cls.__tablename__
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
import asyncio
|
||||
|
||||
from datetime import datetime
|
||||
@@ -42,36 +40,42 @@ class TaskScheduler(Base):
|
||||
interval_period: Mapped[str | None] = mapped_column(String(255), comment='任务运行之间的周期类型')
|
||||
crontab: Mapped[str | None] = mapped_column(String(50), default='* * * * *', comment='任务运行的 Crontab 计划')
|
||||
one_off: Mapped[bool] = mapped_column(
|
||||
Boolean().with_variant(INTEGER, 'postgresql'), default=False, comment='是否仅运行一次'
|
||||
Boolean().with_variant(INTEGER, 'postgresql'),
|
||||
default=False,
|
||||
comment='是否仅运行一次',
|
||||
)
|
||||
enabled: Mapped[bool] = mapped_column(
|
||||
Boolean().with_variant(INTEGER, 'postgresql'), default=True, comment='是否启用任务'
|
||||
Boolean().with_variant(INTEGER, 'postgresql'),
|
||||
default=True,
|
||||
comment='是否启用任务',
|
||||
)
|
||||
total_run_count: Mapped[int] = mapped_column(default=0, comment='任务触发的总次数')
|
||||
last_run_time: Mapped[datetime | None] = mapped_column(TimeZone, default=None, comment='任务最后触发的时间')
|
||||
remark: Mapped[str | None] = mapped_column(
|
||||
LONGTEXT().with_variant(TEXT, 'postgresql'), default=None, comment='备注'
|
||||
LONGTEXT().with_variant(TEXT, 'postgresql'),
|
||||
default=None,
|
||||
comment='备注',
|
||||
)
|
||||
|
||||
no_changes: bool = False
|
||||
|
||||
@staticmethod
|
||||
def before_insert_or_update(mapper, connection, target):
|
||||
def before_insert_or_update(mapper, connection, target) -> None: # noqa: ANN001
|
||||
if target.expire_seconds is not None and target.expire_time:
|
||||
raise errors.ConflictError(msg='expires 和 expire_seconds 只能设置一个')
|
||||
|
||||
@classmethod
|
||||
def changed(cls, mapper, connection, target):
|
||||
def changed(cls, mapper, connection, target) -> None: # noqa: ANN001
|
||||
if not target.no_changes:
|
||||
cls.update_changed(mapper, connection, target)
|
||||
|
||||
@classmethod
|
||||
async def update_changed_async(cls):
|
||||
async def update_changed_async(cls) -> None:
|
||||
now = timezone.now()
|
||||
await redis_client.set(f'{settings.CELERY_REDIS_PREFIX}:last_update', timezone.to_str(now))
|
||||
|
||||
@classmethod
|
||||
def update_changed(cls, mapper, connection, target):
|
||||
def update_changed(cls, mapper, connection, target) -> None: # noqa: ANN001
|
||||
asyncio.create_task(cls.update_changed_async())
|
||||
|
||||
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from backend.common.schema import SchemaBase
|
||||
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
@@ -39,5 +37,5 @@ class GetTaskResultDetail(TaskResultSchemaBase):
|
||||
id: int = Field(description='任务结果 ID')
|
||||
|
||||
@field_serializer('args', 'kwargs', when_used='unless-none')
|
||||
def serialize_params(self, value: bytes | None, _info) -> Any:
|
||||
def serialize_params(self, value: bytes | None) -> Any:
|
||||
return celery_app.backend.decode(value)
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import ConfigDict, Field
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from sqlalchemy import Select
|
||||
|
||||
from backend.app.task.crud.crud_result import task_result_dao
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
import json
|
||||
|
||||
from typing import Sequence
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy import Select
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
@@ -21,7 +19,7 @@ class TaskSchedulerService:
|
||||
"""任务调度服务类"""
|
||||
|
||||
@staticmethod
|
||||
async def get(*, pk) -> TaskScheduler | None:
|
||||
async def get(*, pk: int) -> TaskScheduler | None:
|
||||
"""
|
||||
获取任务调度详情
|
||||
|
||||
@@ -81,9 +79,8 @@ class TaskSchedulerService:
|
||||
task_scheduler = await task_scheduler_dao.get(db, pk)
|
||||
if not task_scheduler:
|
||||
raise errors.NotFoundError(msg='任务调度不存在')
|
||||
if task_scheduler.name != obj.name:
|
||||
if await task_scheduler_dao.get_by_name(db, obj.name):
|
||||
raise errors.ConflictError(msg='任务调度已存在')
|
||||
if task_scheduler.name != obj.name and await task_scheduler_dao.get_by_name(db, obj.name):
|
||||
raise errors.ConflictError(msg='任务调度已存在')
|
||||
if task_scheduler.type == TaskSchedulerType.CRONTAB:
|
||||
crontab_verify(obj.crontab)
|
||||
count = await task_scheduler_dao.update(db, pk, obj)
|
||||
@@ -101,11 +98,11 @@ class TaskSchedulerService:
|
||||
task_scheduler = await task_scheduler_dao.get(db, pk)
|
||||
if not task_scheduler:
|
||||
raise errors.NotFoundError(msg='任务调度不存在')
|
||||
count = await task_scheduler_dao.set_status(db, pk, not task_scheduler.enabled)
|
||||
count = await task_scheduler_dao.set_status(db, pk, status=not task_scheduler.enabled)
|
||||
return count
|
||||
|
||||
@staticmethod
|
||||
async def delete(*, pk) -> int:
|
||||
async def delete(*, pk: int) -> int:
|
||||
"""
|
||||
删除任务调度
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from celery.backends.database.session import SessionManager as CelerySessionManager
|
||||
|
||||
|
||||
@@ -8,7 +6,7 @@ class SessionManager(CelerySessionManager):
|
||||
重写 celery SessionManager
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
# 禁止自动创建 celery 内部定义的任务结果表
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
import asyncio
|
||||
|
||||
from typing import Any
|
||||
@@ -17,7 +15,7 @@ class TaskBase(Task):
|
||||
autoretry_for = (SQLAlchemyError,)
|
||||
max_retries = settings.CELERY_TASK_MAX_RETRIES
|
||||
|
||||
async def before_start(self, task_id: str, args, kwargs) -> None:
|
||||
async def before_start(self, task_id: str, args, kwargs) -> None: # noqa: ANN001
|
||||
"""
|
||||
任务开始前执行钩子
|
||||
|
||||
@@ -26,7 +24,7 @@ class TaskBase(Task):
|
||||
"""
|
||||
await task_notification(msg=f'任务 {task_id} 开始执行')
|
||||
|
||||
async def on_success(self, retval: Any, task_id: str, args, kwargs) -> None:
|
||||
async def on_success(self, retval: Any, task_id: str, args, kwargs) -> None: # noqa: ANN001
|
||||
"""
|
||||
任务成功后执行钩子
|
||||
|
||||
@@ -36,7 +34,7 @@ class TaskBase(Task):
|
||||
"""
|
||||
await task_notification(msg=f'任务 {task_id} 执行成功')
|
||||
|
||||
def on_failure(self, exc: Exception, task_id: str, args, kwargs, einfo) -> None:
|
||||
def on_failure(self, exc: Exception, task_id: str, args, kwargs, einfo) -> None: # noqa: ANN001
|
||||
"""
|
||||
任务失败后执行钩子
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from celery.schedules import schedule
|
||||
|
||||
from backend.app.task.utils.tzcrontab import TzAwareCrontab
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from celery import shared_task
|
||||
|
||||
from backend.app.admin.service.login_log_service import login_log_service
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from time import sleep
|
||||
|
||||
from anyio import sleep as asleep
|
||||
@@ -24,4 +22,5 @@ async def task_demo_async() -> str:
|
||||
@celery_app.task(name='task_demo_params')
|
||||
async def task_demo_params(hello: str, world: str | None = None) -> str:
|
||||
"""参数示例任务,模拟传参操作"""
|
||||
await asleep(1)
|
||||
return hello + world
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
@@ -1,17 +1,17 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import math
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
from multiprocessing.util import Finalize
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from celery import current_app, schedules
|
||||
from celery.beat import ScheduleEntry, Scheduler
|
||||
from celery.signals import beat_init
|
||||
from celery.utils.log import get_logger
|
||||
from redis.asyncio.lock import Lock
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.exc import DatabaseError, InterfaceError
|
||||
|
||||
@@ -27,6 +27,9 @@ from backend.utils._await import run_await
|
||||
from backend.utils.serializers import select_as_dict
|
||||
from backend.utils.timezone import timezone
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from redis.asyncio.lock import Lock
|
||||
|
||||
# 此计划程序必须比常规的 5 分钟更频繁地唤醒,因为它需要考虑对计划的外部更改
|
||||
DEFAULT_MAX_INTERVAL = 5 # seconds
|
||||
|
||||
@@ -39,7 +42,7 @@ logger = get_logger('fba.schedulers')
|
||||
class ModelEntry(ScheduleEntry):
|
||||
"""任务调度实体"""
|
||||
|
||||
def __init__(self, model: TaskScheduler, app=None):
|
||||
def __init__(self, model: TaskScheduler, app=None) -> None: # noqa:ANN001,C901
|
||||
super().__init__(
|
||||
app=app or current_app._get_current_object(),
|
||||
name=model.name,
|
||||
@@ -72,7 +75,7 @@ class ModelEntry(ScheduleEntry):
|
||||
self.args = json.loads(model.args) if model.args else None
|
||||
self.kwargs = json.loads(model.kwargs) if model.kwargs else None
|
||||
except ValueError as exc:
|
||||
logger.error(f'禁用参数错误的任务:{self.name};error: {str(exc)}')
|
||||
logger.error(f'禁用参数错误的任务:{self.name};error: {exc!s}')
|
||||
asyncio.create_task(self._disable(model))
|
||||
|
||||
self.options = {}
|
||||
@@ -103,7 +106,7 @@ class ModelEntry(ScheduleEntry):
|
||||
model.no_changes = True
|
||||
self.model.enabled = self.enabled = model.enabled = False
|
||||
async with async_db_session.begin():
|
||||
setattr(model, 'enabled', False)
|
||||
model.enabled = False
|
||||
|
||||
def is_due(self) -> tuple[bool, int | float]:
|
||||
"""任务到期状态"""
|
||||
@@ -130,7 +133,7 @@ class ModelEntry(ScheduleEntry):
|
||||
|
||||
return self.schedule.is_due(self.last_run_at)
|
||||
|
||||
def __next__(self):
|
||||
def __next__(self): # noqa: ANN204
|
||||
self.model.last_run_time = timezone.now()
|
||||
self.model.total_run_count += 1
|
||||
self.model.no_changes = True
|
||||
@@ -138,7 +141,7 @@ class ModelEntry(ScheduleEntry):
|
||||
|
||||
next = __next__
|
||||
|
||||
async def save(self, fields: tuple = ()):
|
||||
async def save(self, fields: tuple = ()) -> None:
|
||||
"""
|
||||
保存任务状态字段
|
||||
|
||||
@@ -158,7 +161,7 @@ class ModelEntry(ScheduleEntry):
|
||||
logger.warning(f'任务 {self.model.name} 不存在,跳过更新')
|
||||
|
||||
@classmethod
|
||||
async def from_entry(cls, name, app=None, **entry):
|
||||
async def from_entry(cls, name, app=None, **entry) -> ModelEntry: # noqa: ANN001
|
||||
"""保存或更新本地任务调度"""
|
||||
async with async_db_session.begin() as db:
|
||||
stmt = select(TaskScheduler).where(TaskScheduler.name == name)
|
||||
@@ -175,7 +178,7 @@ class ModelEntry(ScheduleEntry):
|
||||
return res
|
||||
|
||||
@staticmethod
|
||||
async def to_model_schedule(name: str, task: str, schedule: schedules.schedule | TzAwareCrontab):
|
||||
async def to_model_schedule(name: str, task: str, schedule: schedules.schedule | TzAwareCrontab) -> TaskScheduler:
|
||||
schedule = schedules.maybe_schedule(schedule)
|
||||
|
||||
async with async_db_session() as db:
|
||||
@@ -218,7 +221,7 @@ class ModelEntry(ScheduleEntry):
|
||||
schedule: schedules.schedule | TzAwareCrontab,
|
||||
args: tuple | None = None,
|
||||
kwargs: dict | None = None,
|
||||
options: dict = None,
|
||||
options: dict | None = None,
|
||||
**entry,
|
||||
) -> dict:
|
||||
model_schedule = await cls.to_model_schedule(name, task, schedule)
|
||||
@@ -226,7 +229,7 @@ class ModelEntry(ScheduleEntry):
|
||||
for k in ['id', 'created_time', 'updated_time']:
|
||||
try:
|
||||
del model_dict[k]
|
||||
except KeyError:
|
||||
except KeyError: # noqa:PERF203
|
||||
continue
|
||||
model_dict.update(
|
||||
args=json.dumps(args, ensure_ascii=False) if args else None,
|
||||
@@ -239,12 +242,13 @@ class ModelEntry(ScheduleEntry):
|
||||
@classmethod
|
||||
def _unpack_options(
|
||||
cls,
|
||||
queue: str = None,
|
||||
exchange: str = None,
|
||||
routing_key: str = None,
|
||||
start_time: datetime = None,
|
||||
expires: datetime = None,
|
||||
expire_seconds: int = None,
|
||||
queue: str | None = None,
|
||||
exchange: str | None = None,
|
||||
routing_key: str | None = None,
|
||||
start_time: datetime | None = None,
|
||||
expires: datetime | None = None,
|
||||
expire_seconds: int | None = None,
|
||||
*,
|
||||
one_off: bool = False,
|
||||
) -> dict:
|
||||
data = {
|
||||
@@ -277,14 +281,14 @@ class DatabaseScheduler(Scheduler):
|
||||
lock: Lock | None = None
|
||||
lock_key = f'{settings.CELERY_REDIS_PREFIX}:beat_lock'
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
self.app = kwargs['app']
|
||||
self._dirty = set()
|
||||
super().__init__(*args, **kwargs)
|
||||
self._finalize = Finalize(self, self.sync, exitpriority=5)
|
||||
self.max_interval = kwargs.get('max_interval') or self.app.conf.beat_max_loop_interval or DEFAULT_MAX_INTERVAL
|
||||
|
||||
def install_default_entries(self, data):
|
||||
def install_default_entries(self, data) -> None: # noqa: ANN001
|
||||
"""重写父函数"""
|
||||
entries = {}
|
||||
if self.app.conf.result_expires:
|
||||
@@ -298,31 +302,31 @@ class DatabaseScheduler(Scheduler):
|
||||
)
|
||||
self.update_from_dict(entries)
|
||||
|
||||
def schedules_equal(self, *args, **kwargs):
|
||||
def schedules_equal(self, *args, **kwargs) -> bool:
|
||||
"""重写父函数"""
|
||||
if self._heap_invalidated:
|
||||
self._heap_invalidated = False
|
||||
return False
|
||||
return super().schedules_equal(*args, **kwargs)
|
||||
|
||||
def reserve(self, entry):
|
||||
def reserve(self, entry): # noqa: ANN001, ANN201
|
||||
"""重写父函数"""
|
||||
new_entry = next(entry)
|
||||
# 需要按名称存储条目,因为条目可能会发生变化
|
||||
self._dirty.add(new_entry.name)
|
||||
return new_entry
|
||||
|
||||
def setup_schedule(self):
|
||||
def setup_schedule(self) -> None:
|
||||
"""重写父函数"""
|
||||
logger.info('setup_schedule')
|
||||
tasks = self.schedule
|
||||
self.install_default_entries(tasks)
|
||||
self.update_from_dict(self.app.conf.beat_schedule)
|
||||
|
||||
def sync(self):
|
||||
def sync(self) -> None:
|
||||
"""重写父函数"""
|
||||
_tried = set()
|
||||
_failed = set()
|
||||
tried = set()
|
||||
failed = set()
|
||||
try:
|
||||
while self._dirty:
|
||||
name = self._dirty.pop()
|
||||
@@ -330,28 +334,27 @@ class DatabaseScheduler(Scheduler):
|
||||
tasks = self.schedule
|
||||
run_await(tasks[name].save)()
|
||||
logger.debug(f'保存任务 {name} 最新状态到数据库')
|
||||
_tried.add(name)
|
||||
tried.add(name)
|
||||
except KeyError as e:
|
||||
logger.error(f'保存任务 {name} 最新状态失败:{e} ')
|
||||
_failed.add(name)
|
||||
except DatabaseError as e:
|
||||
logger.exception('同步时出现数据库错误: %r', e)
|
||||
failed.add(name)
|
||||
except DatabaseError:
|
||||
logger.exception('同步时出现数据库错误')
|
||||
except InterfaceError as e:
|
||||
logger.warning(f'DatabaseScheduler InterfaceError:{str(e)},等待下次调用时重试...')
|
||||
logger.warning(f'DatabaseScheduler InterfaceError:{e!s},等待下次调用时重试...')
|
||||
finally:
|
||||
# 请稍后重试(仅针对失败的)
|
||||
self._dirty |= _failed
|
||||
self._dirty |= failed
|
||||
|
||||
def tick(self, **kwargs):
|
||||
def tick(self, **kwargs) -> float:
|
||||
"""重写父函数"""
|
||||
if self.lock:
|
||||
logger.debug('beat: Extending lock...')
|
||||
run_await(self.lock.extend)(DEFAULT_MAX_LOCK_TIMEOUT, replace_ttl=True)
|
||||
|
||||
result = super().tick(**kwargs)
|
||||
return result
|
||||
return super().tick(**kwargs)
|
||||
|
||||
def close(self):
|
||||
def close(self) -> None:
|
||||
"""重写父函数"""
|
||||
if self.lock:
|
||||
logger.info('beat: Releasing lock')
|
||||
@@ -361,22 +364,23 @@ class DatabaseScheduler(Scheduler):
|
||||
|
||||
super().close()
|
||||
|
||||
def update_from_dict(self, beat_dict: dict):
|
||||
def update_from_dict(self, beat_dict: dict) -> None:
|
||||
"""重写父函数"""
|
||||
s = {}
|
||||
for name, entry_fields in beat_dict.items():
|
||||
try:
|
||||
|
||||
try:
|
||||
for name, entry_fields in beat_dict.items():
|
||||
entry = run_await(self.Entry.from_entry)(name, app=self.app, **entry_fields)
|
||||
if entry.model.enabled:
|
||||
s[name] = entry
|
||||
except Exception as e:
|
||||
logger.error(f'添加任务 {name} 到数据库失败')
|
||||
raise e
|
||||
except Exception:
|
||||
logger.error(f'添加任务 {name} 到数据库失败')
|
||||
raise
|
||||
|
||||
tasks = self.schedule
|
||||
tasks.update(s)
|
||||
|
||||
def schedule_changed(self) -> bool:
|
||||
def schedule_changed(self) -> bool | None:
|
||||
"""任务调度变更状态"""
|
||||
now = timezone.now()
|
||||
last_update = run_await(redis_client.get)(f'{settings.CELERY_REDIS_PREFIX}:last_update')
|
||||
@@ -386,21 +390,21 @@ class DatabaseScheduler(Scheduler):
|
||||
|
||||
last, ts = self._last_update, timezone.from_str(last_update)
|
||||
try:
|
||||
if ts and ts > (last if last else ts):
|
||||
if ts and ts > (last or ts):
|
||||
return True
|
||||
finally:
|
||||
self._last_update = now
|
||||
|
||||
async def get_all_task_schedulers(self):
|
||||
async def get_all_task_schedulers(self) -> dict:
|
||||
"""获取所有任务调度"""
|
||||
async with async_db_session() as db:
|
||||
logger.debug('DatabaseScheduler: Fetching database schedule')
|
||||
stmt = select(TaskScheduler).where(TaskScheduler.enabled == 1)
|
||||
query = await db.execute(stmt)
|
||||
tasks = query.scalars().all()
|
||||
schedulers = query.scalars().all()
|
||||
s = {}
|
||||
for task in tasks:
|
||||
s[task.name] = self.Entry(task, app=self.app)
|
||||
for scheduler in schedulers:
|
||||
s[scheduler.name] = self.Entry(scheduler, app=self.app)
|
||||
return s
|
||||
|
||||
@property
|
||||
@@ -433,7 +437,7 @@ class DatabaseScheduler(Scheduler):
|
||||
|
||||
|
||||
@beat_init.connect
|
||||
def acquire_distributed_beat_lock(sender=None, **kwargs):
|
||||
def acquire_distributed_beat_lock(sender=None, **kwargs) -> None: # noqa: ANN001
|
||||
"""
|
||||
尝试在启动时获取锁
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
from datetime import datetime
|
||||
|
||||
from celery import schedules
|
||||
@@ -12,7 +10,7 @@ from backend.utils.timezone import timezone
|
||||
class TzAwareCrontab(schedules.crontab):
|
||||
"""时区感知 Crontab"""
|
||||
|
||||
def __init__(self, minute='*', hour='*', day_of_week='*', day_of_month='*', month_of_year='*', app=None):
|
||||
def __init__(self, minute='*', hour='*', day_of_week='*', day_of_month='*', month_of_year='*', app=None) -> None: # noqa: ANN001
|
||||
super().__init__(
|
||||
minute=minute,
|
||||
hour=hour,
|
||||
|
||||
Reference in New Issue
Block a user