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:
Wu Clan
2025-10-10 19:02:49 +08:00
committed by GitHub
parent 51354593d0
commit 5762834744
269 changed files with 975 additions and 1198 deletions
+1 -1
View File
@@ -12,7 +12,7 @@ jobs:
name: lint ${{ matrix.python-version }}
strategy:
matrix:
python-version: [ '3.10', '3.11', '3.12', '3.13' ]
python-version: [ '3.10', '3.11', '3.12', '3.13', '3.14' ]
fail-fast: false
steps:
- uses: actions/checkout@v4
+174
View File
@@ -0,0 +1,174 @@
line-length = 120
preview = true
fix = true
unsafe-fixes = true
show-fixes = true
required-version = ">=0.13.0"
[lint]
select = [
"FAST",
"ANN001",
"ANN201",
"ANN202",
"ANN204",
"ANN205",
"ANN206",
"ASYNC110",
"ASYNC116",
"ASYNC210",
"ASYNC212",
"ASYNC230",
"ASYNC240",
"ASYNC250",
"ASYNC251",
"S310",
"FBT001",
"FBT002",
"B002",
"B005",
"B006",
"B007",
"B008",
"B009",
"B010",
"B013",
"B014",
"B019",
"B020",
"B021",
"B024",
"B025",
"B026",
"B027",
"B039",
"COM",
"C402",
"C403",
"C404",
"C408",
"C410",
"C411",
"C414",
"C416",
"C417",
"C418",
"C419",
"C420",
"DTZ",
"EXE",
"ISC001",
"ISC002",
"ISC003",
"PIE",
"PYI009",
"PYI010",
"PYI011",
"PYI012",
"PYI013",
"PYI016",
"PYI017",
"PYI019",
"PYI020",
"PYI021",
"PYI024",
"PYI026",
"PYI030",
"PYI033",
"PYI034",
"PYI036",
"PYI041",
"PYI042",
"PYI055",
"PYI061",
"PYI062",
"PYI063",
"Q001",
"Q002",
"RSE102",
"RET501",
"RET505",
"RET506",
"RET507",
"RET508",
"SIM101",
"SIM102",
"SIM103",
"SIM107",
"SIM108",
"SIM109",
"SIM110",
"SIM114",
"SIM115",
"SIM201",
"SIM202",
"SIM210",
"SIM211",
"SIM212",
"SIM300",
"SIM401",
"SIM910",
"TID252",
"TC",
"FLY",
"I",
"C901",
"N",
"PERF",
"E",
"W",
"D404",
"D417",
"D419",
"F",
"PGH",
"PLC1901",
"UP",
"FURB",
"RUF",
"TRY",
]
ignore = [
"COM812",
"PGH003",
"RUF001",
"RUF002",
"RUF003",
"RUF006",
"RUF012",
"TRY400",
"TRY003",
"TRY301"
]
[lint.per-file-ignores]
"**/model/*.py" = ["TC003"]
"backend/common/socketio/server.py" = ["ANN001"]
"backend/common/exception/exception_handler.py" = ["ANN202","RUF029"]
[lint.flake8-pytest-style]
parametrize-names-type = "list"
parametrize-values-row-type = "list"
parametrize-values-type = "list"
[lint.flake8-quotes]
inline-quotes = "single"
[lint.flake8-type-checking]
runtime-evaluated-base-classes = ["pydantic.BaseModel", "sqlalchemy.orm.DeclarativeBase"]
[lint.flake8-unused-arguments]
ignore-variadic-names = true
[lint.isort]
case-sensitive = true
lines-between-types = 1
order-by-type = true
[lint.pylint]
allow-dunder-method-names = ["__tablename__", "__table_args__"]
[format]
docstring-code-format = true
preview = true
quote-style = "single"
-55
View File
@@ -1,55 +0,0 @@
line-length = 120
cache-dir = ".ruff_cache"
target-version = "py310"
unsafe-fixes = true
show-fixes = true
[lint]
select = [
"E",
"F",
"I",
"TC",
# W
"W505",
# PT
"PT018",
# SIM
"SIM101",
"SIM114",
# PGH
"PGH004",
# PL
"PLE1142",
# RUF
"RUF100",
# UP
"UP007"
]
preview = true
ignore = ["FURB101"]
[lint.flake8-pytest-style]
mark-parentheses = false
parametrize-names-type = "list"
parametrize-values-row-type = "list"
parametrize-values-type = "tuple"
[lint.flake8-unused-arguments]
ignore-variadic-names = true
[lint.isort]
lines-between-types = 1
order-by-type = true
[lint.per-file-ignores]
"**/api/v1/*.py" = ["TC"]
"**/model/*.py" = ["TC003"]
"**/model/__init__.py" = ["F401"]
"**/tests/*.py" = ["E402"]
[format]
preview = true
quote-style = "single"
docstring-code-format = true
skip-magic-trailing-comma = false
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.common.i18n import i18n
__version__ = '1.8.2'
+5 -6
View File
@@ -1,8 +1,6 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
# ruff: noqa: F403, F401, I001, RUF100
import asyncio
import os
from logging.config import fileConfig
from alembic import context
@@ -40,11 +38,12 @@ target_metadata = MappedBase.metadata
# other values from the config, defined by the needs of env.py,
alembic_config.set_main_option(
'sqlalchemy.url', SQLALCHEMY_DATABASE_URL.render_as_string(hide_password=False).replace('%', '%%')
'sqlalchemy.url',
SQLALCHEMY_DATABASE_URL.render_as_string(hide_password=False).replace('%', '%%'),
)
def run_migrations_offline():
def run_migrations_offline() -> None:
"""Run migrations in 'offline' mode.
This configures the context with just a URL
@@ -73,7 +72,7 @@ def run_migrations_offline():
def do_run_migrations(connection: Connection) -> None:
# 当迁移无变化时,不生成迁移记录
def process_revision_directives(context, revision, directives):
def process_revision_directives(context, revision, directives) -> None: # noqa: ANN001
if alembic_config.cmd_opts.autogenerate:
script = directives[0]
if script.upgrade_ops.is_empty():
+2 -8
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import os.path
from backend.core.path_conf import BASE_PATH
@@ -8,14 +6,10 @@ from backend.utils.import_parse import get_model_objects
def get_app_models() -> list[type]:
"""获取 app 所有模型类"""
app_path = os.path.join(BASE_PATH, 'app')
app_path = BASE_PATH / 'app'
list_dirs = os.listdir(app_path)
apps = []
for d in list_dirs:
if os.path.isdir(os.path.join(app_path, d)) and d != '__pycache__':
apps.append(d)
apps = [d for d in list_dirs if os.path.isdir(os.path.join(app_path, d)) and d != '__pycache__']
objs = []
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from backend.app.admin.api.v1.auth import router as auth_router
-2
View File
@@ -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.admin.api.v1.auth.auth import router as auth_router
+4 -3
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Request, Response
@@ -29,7 +27,10 @@ async def login_swagger(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetS
dependencies=[Depends(RateLimiter(times=5, minutes=1))],
)
async def login(
request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks
request: Request,
response: Response,
obj: AuthLoginParam,
background_tasks: BackgroundTasks,
) -> ResponseSchemaModel[GetLoginToken]:
data = await auth_service.login(request=request, response=response, obj=obj, background_tasks=background_tasks)
return response_base.success(data=data)
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from uuid import uuid4
from fast_captcha import img_captcha
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from backend.app.admin.api.v1.log.login_log import router as login_log
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Query
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Query
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from backend.app.admin.api.v1.monitor.online import router as token_router
+2 -4
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
from typing import Annotated
@@ -37,8 +35,8 @@ async def get_sessions(
'browser': extra_info.get('browser', '未知'),
'device': extra_info.get('device', '未知'),
'last_login_time': extra_info.get('last_login_time', '未知'),
}
)
},
),
)
for key in token_keys:
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from backend.common.response.response_schema import ResponseModel, response_base
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from starlette.concurrency import run_in_threadpool
-2
View File
@@ -1,5 +1,3 @@
# !/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from backend.app.admin.api.v1.sys.data_rule import router as data_rule_router
+4 -4
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query
@@ -59,7 +57,8 @@ async def get_data_rule(
],
)
async def get_data_rules_paged(
db: CurrentSession, name: Annotated[str | None, Query(description='规则名称')] = None
db: CurrentSession,
name: Annotated[str | None, Query(description='规则名称')] = None,
) -> ResponseSchemaModel[PageData[GetDataRuleDetail]]:
data_rule_select = await data_rule_service.get_select(name=name)
page_data = await paging_data(db, data_rule_select)
@@ -88,7 +87,8 @@ async def create_data_rule(obj: CreateDataRuleParam) -> ResponseModel:
],
)
async def update_data_rule(
pk: Annotated[int, Path(description='数据规则 ID')], obj: UpdateDataRuleParam
pk: Annotated[int, Path(description='数据规则 ID')],
obj: UpdateDataRuleParam,
) -> ResponseModel:
count = await data_rule_service.update(pk=pk, obj=obj)
if count > 0:
+5 -5
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query
@@ -85,7 +83,8 @@ async def create_data_scope(obj: CreateDataScopeParam) -> ResponseModel:
],
)
async def update_data_scope(
pk: Annotated[int, Path(description='数据范围 ID')], obj: UpdateDataScopeParam
pk: Annotated[int, Path(description='数据范围 ID')],
obj: UpdateDataScopeParam,
) -> ResponseModel:
count = await data_scope_service.update(pk=pk, obj=obj)
if count > 0:
@@ -102,8 +101,9 @@ async def update_data_scope(
],
)
async def update_data_scope_rules(
pk: Annotated[int, Path(description='数据范围 ID')], rule_ids: UpdateDataScopeRuleParam
):
pk: Annotated[int, Path(description='数据范围 ID')],
rule_ids: UpdateDataScopeRuleParam,
) -> ResponseModel:
count = await data_scope_service.update_data_scope_rule(pk=pk, rule_ids=rule_ids)
if count > 0:
return response_base.success()
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query, Request
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, File, UploadFile
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated, Any
from fastapi import APIRouter, Depends, Path, Query, Request
+4 -5
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated, Any
from fastapi import APIRouter, Depends, File, Path, UploadFile
@@ -46,8 +44,9 @@ async def install_plugin(
plugin_name = await plugin_service.install(type=type, file=file, repo_url=repo_url)
return response_base.success(
res=CustomResponse(
code=200, msg=f'插件 {plugin_name} 安装成功,请根据插件说明(README.md)进行相关配置并重启服务'
)
code=200,
msg=f'插件 {plugin_name} 安装成功,请根据插件说明(README.md)进行相关配置并重启服务',
),
)
@@ -63,7 +62,7 @@ async def install_plugin(
async def uninstall_plugin(plugin: Annotated[str, Path(description='插件名称')]) -> ResponseModel:
await plugin_service.uninstall(plugin=plugin)
return response_base.success(
res=CustomResponse(code=200, msg=f'插件 {plugin} 卸载成功,请根据插件说明(README.md)移除相关配置并重启服务')
res=CustomResponse(code=200, msg=f'插件 {plugin} 卸载成功,请根据插件说明(README.md)移除相关配置并重启服务'),
)
+4 -4
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query
@@ -106,7 +104,8 @@ async def update_role(pk: Annotated[int, Path(description='角色 ID')], obj: Up
],
)
async def update_role_menus(
pk: Annotated[int, Path(description='角色 ID')], menu_ids: UpdateRoleMenuParam
pk: Annotated[int, Path(description='角色 ID')],
menu_ids: UpdateRoleMenuParam,
) -> ResponseModel:
count = await role_service.update_role_menu(pk=pk, menu_ids=menu_ids)
if count > 0:
@@ -123,7 +122,8 @@ async def update_role_menus(
],
)
async def update_role_scopes(
pk: Annotated[int, Path(description='角色 ID')], scope_ids: UpdateRoleScopeParam
pk: Annotated[int, Path(description='角色 ID')],
scope_ids: UpdateRoleScopeParam,
) -> ResponseModel:
count = await role_service.update_role_scope(pk=pk, scope_ids=scope_ids)
if count > 0:
+7 -5
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Body, Depends, Path, Query, Request
@@ -73,7 +71,9 @@ async def create_user(request: Request, obj: AddUserParam) -> ResponseSchemaMode
@router.put('/{pk}', summary='更新用户信息', dependencies=[DependsRBAC])
async def update_user(
request: Request, pk: Annotated[int, Path(description='用户 ID')], obj: UpdateUserParam
request: Request,
pk: Annotated[int, Path(description='用户 ID')],
obj: UpdateUserParam,
) -> ResponseModel:
count = await user_service.update(request=request, pk=pk, obj=obj)
if count > 0:
@@ -115,7 +115,8 @@ async def reset_user_password(
@router.put('/me/nickname', summary='更新当前用户昵称', dependencies=[DependsJwtAuth])
async def update_user_nickname(
request: Request, nickname: Annotated[str, Body(embed=True, description='用户昵称')]
request: Request,
nickname: Annotated[str, Body(embed=True, description='用户昵称')],
) -> ResponseModel:
count = await user_service.update_nickname(request=request, nickname=nickname)
if count > 0:
@@ -125,7 +126,8 @@ async def update_user_nickname(
@router.put('/me/avatar', summary='更新当前用户头像', dependencies=[DependsJwtAuth])
async def update_user_avatar(
request: Request, avatar: Annotated[str, Body(embed=True, description='用户头像地址')]
request: Request,
avatar: Annotated[str, Body(embed=True, description='用户头像地址')],
) -> ResponseModel:
count = await user_service.update_avatar(request=request, avatar=avatar)
if count > 0:
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
+1 -3
View File
@@ -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
+1 -3
View File
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from collections.abc import Sequence
from sqlalchemy import Select, select
from sqlalchemy.ext.asyncio import AsyncSession
+1 -3
View File
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from collections.abc import Sequence
from fastapi import Request
from sqlalchemy.ext.asyncio import AsyncSession
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from sqlalchemy import Select
from sqlalchemy import delete as sa_delete
from sqlalchemy.ext.asyncio import AsyncSession
+1 -3
View File
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from collections.abc import Sequence
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from sqlalchemy import Select
from sqlalchemy import delete as sa_delete
from sqlalchemy.ext.asyncio import AsyncSession
+1 -3
View File
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from collections.abc import Sequence
from sqlalchemy import Select, select
from sqlalchemy.ext.asyncio import AsyncSession
+8 -6
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import bcrypt
from sqlalchemy import select
@@ -214,7 +212,7 @@ class CRUDUser(CRUDPlus[User]):
**filters,
)
async def set_super(self, db: AsyncSession, user_id: int, is_super: bool) -> int:
async def set_super(self, db: AsyncSession, user_id: int, *, is_super: bool) -> int:
"""
设置用户超级管理员状态
@@ -225,7 +223,7 @@ class CRUDUser(CRUDPlus[User]):
"""
return await self.update_model(db, user_id, {'is_superuser': is_super})
async def set_staff(self, db: AsyncSession, user_id: int, is_staff: bool) -> int:
async def set_staff(self, db: AsyncSession, user_id: int, *, is_staff: bool) -> int:
"""
设置用户后台登录状态
@@ -247,7 +245,7 @@ class CRUDUser(CRUDPlus[User]):
"""
return await self.update_model(db, user_id, {'status': status})
async def set_multi_login(self, db: AsyncSession, user_id: int, multi_login: bool) -> int:
async def set_multi_login(self, db: AsyncSession, user_id: int, *, multi_login: bool) -> int:
"""
设置用户多端登录状态
@@ -259,7 +257,11 @@ class CRUDUser(CRUDPlus[User]):
return await self.update_model(db, user_id, {'is_multi_login': multi_login})
async def get_with_relation(
self, db: AsyncSession, *, user_id: int | None = None, username: str | None = None
self,
db: AsyncSession,
*,
user_id: int | None = None,
username: str | None = None,
) -> User | None:
"""
获取用户关联信息
+8 -10
View File
@@ -1,10 +1,8 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.app.admin.model.data_rule import DataRule
from backend.app.admin.model.data_scope import DataScope
from backend.app.admin.model.dept import Dept
from backend.app.admin.model.login_log import LoginLog
from backend.app.admin.model.menu import Menu
from backend.app.admin.model.opera_log import OperaLog
from backend.app.admin.model.role import Role
from backend.app.admin.model.user import User
from backend.app.admin.model.data_rule import DataRule as DataRule
from backend.app.admin.model.data_scope import DataScope as DataScope
from backend.app.admin.model.dept import Dept as Dept
from backend.app.admin.model.login_log import LoginLog as LoginLog
from backend.app.admin.model.menu import Menu as Menu
from backend.app.admin.model.opera_log import OperaLog as OperaLog
from backend.app.admin.model.role import Role as Role
from backend.app.admin.model.user import User as User
+1 -3
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from __future__ import annotations
from typing import TYPE_CHECKING
@@ -25,7 +23,7 @@ class DataRule(Base):
column: Mapped[str] = mapped_column(String(20), comment='模型字段名')
operator: Mapped[int] = mapped_column(comment='运算符(0and、1or')
expression: Mapped[int] = mapped_column(
comment='表达式(0==、1!=、2>、3>=、4<、5<=、6in、7not_in'
comment='表达式(0==、1!=、2>、3>=、4<、5<=、6in、7not_in',
)
value: Mapped[str] = mapped_column(String(255), comment='规则值')
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from __future__ import annotations
from typing import TYPE_CHECKING
+11 -7
View File
@@ -1,8 +1,6 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING
from sqlalchemy import BigInteger, Boolean, ForeignKey, String
from sqlalchemy.dialects.postgresql import INTEGER
@@ -27,15 +25,21 @@ class Dept(Base):
email: Mapped[str | None] = mapped_column(String(50), default=None, comment='邮箱')
status: Mapped[int] = mapped_column(default=1, comment='部门状态(0停用 1正常)')
del_flag: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=False, comment='删除标志(0删除 1存在)'
Boolean().with_variant(INTEGER, 'postgresql'),
default=False,
comment='删除标志(0删除 1存在)',
)
# 父级部门一对多
parent_id: Mapped[int | None] = mapped_column(
BigInteger, ForeignKey('sys_dept.id', ondelete='SET NULL'), default=None, index=True, comment='父部门ID'
BigInteger,
ForeignKey('sys_dept.id', ondelete='SET NULL'),
default=None,
index=True,
comment='父部门ID',
)
parent: Mapped[Optional['Dept']] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[Optional[list['Dept']]] = relationship(init=False, back_populates='parent')
parent: Mapped[Dept | None] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[list[Dept] | None] = relationship(init=False, back_populates='parent')
# 部门用户一对多
users: Mapped[list[User]] = relationship(init=False, back_populates='dept')
+4 -3
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from sqlalchemy import String
@@ -31,5 +29,8 @@ class LoginLog(DataClassBase):
msg: Mapped[str] = mapped_column(LONGTEXT().with_variant(TEXT, 'postgresql'), comment='提示消息')
login_time: Mapped[datetime] = mapped_column(TimeZone, comment='登录时间')
created_time: Mapped[datetime] = mapped_column(
TimeZone, init=False, default_factory=timezone.now, comment='创建时间'
TimeZone,
init=False,
default_factory=timezone.now,
comment='创建时间',
)
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from sqlalchemy import BigInteger, Column, ForeignKey, Table
from backend.common.model import MappedBase
+14 -8
View File
@@ -1,8 +1,6 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING
from sqlalchemy import BigInteger, ForeignKey, String
from sqlalchemy.dialects.mysql import LONGTEXT
@@ -34,18 +32,26 @@ class Menu(Base):
display: Mapped[int] = mapped_column(default=1, comment='是否显示(0否 1是)')
cache: Mapped[int] = mapped_column(default=1, comment='是否缓存(0否 1是)')
link: Mapped[str | None] = mapped_column(
LONGTEXT().with_variant(TEXT, 'postgresql'), default=None, comment='外链地址'
LONGTEXT().with_variant(TEXT, 'postgresql'),
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='备注',
)
# 父级菜单一对多
parent_id: Mapped[int | None] = mapped_column(
BigInteger, ForeignKey('sys_menu.id', ondelete='SET NULL'), default=None, index=True, comment='父菜单ID'
BigInteger,
ForeignKey('sys_menu.id', ondelete='SET NULL'),
default=None,
index=True,
comment='父菜单ID',
)
parent: Mapped[Optional['Menu']] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[Optional[list['Menu']]] = relationship(init=False, back_populates='parent')
parent: Mapped[Menu | None] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[list[Menu] | None] = relationship(init=False, back_populates='parent')
# 菜单角色多对多
roles: Mapped[list[Role]] = relationship(init=False, secondary=sys_role_menu, back_populates='menus')
+4 -3
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from sqlalchemy import String
@@ -37,5 +35,8 @@ class OperaLog(DataClassBase):
cost_time: Mapped[float] = mapped_column(insert_default=0.0, comment='请求耗时(ms')
opera_time: Mapped[datetime] = mapped_column(TimeZone, comment='操作时间')
created_time: Mapped[datetime] = mapped_column(
TimeZone, init=False, default_factory=timezone.now, comment='创建时间'
TimeZone,
init=False,
default_factory=timezone.now,
comment='创建时间',
)
+6 -4
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from __future__ import annotations
from typing import TYPE_CHECKING
@@ -25,10 +23,14 @@ class Role(Base):
name: Mapped[str] = mapped_column(String(20), unique=True, comment='角色名称')
status: Mapped[int] = mapped_column(default=1, comment='角色状态(0停用 1正常)')
is_filter_scopes: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=True, comment='过滤数据权限(0否 1是)'
Boolean().with_variant(INTEGER, 'postgresql'),
default=True,
comment='过滤数据权限(0否 1是)',
)
remark: Mapped[str | None] = mapped_column(
LONGTEXT().with_variant(TEXT, 'postgresql'), default=None, comment='备注'
LONGTEXT().with_variant(TEXT, 'postgresql'),
default=None,
comment='备注',
)
# 角色用户多对多
+16 -7
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from __future__ import annotations
from datetime import datetime
@@ -34,22 +32,33 @@ class User(Base):
avatar: Mapped[str | None] = mapped_column(String(255), default=None, comment='头像')
status: Mapped[int] = mapped_column(default=1, index=True, comment='用户账号状态(0停用 1正常)')
is_superuser: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=False, comment='超级权限(0否 1是)'
Boolean().with_variant(INTEGER, 'postgresql'),
default=False,
comment='超级权限(0否 1是)',
)
is_staff: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=False, comment='后台管理登陆(0否 1是)'
Boolean().with_variant(INTEGER, 'postgresql'),
default=False,
comment='后台管理登陆(0否 1是)',
)
is_multi_login: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=False, comment='是否重复登陆(0否 1是)'
Boolean().with_variant(INTEGER, 'postgresql'),
default=False,
comment='是否重复登陆(0否 1是)',
)
join_time: Mapped[datetime] = mapped_column(TimeZone, init=False, default_factory=timezone.now, comment='注册时间')
last_login_time: Mapped[datetime | None] = mapped_column(
TimeZone, init=False, onupdate=timezone.now, comment='上次登录'
TimeZone,
init=False,
onupdate=timezone.now,
comment='上次登录',
)
# 部门用户一对多
dept_id: Mapped[int | None] = mapped_column(
ForeignKey('sys_dept.id', ondelete='SET NULL'), default=None, comment='部门关联ID'
ForeignKey('sys_dept.id', ondelete='SET NULL'),
default=None,
comment='部门关联ID',
)
dept: Mapped[Dept | None] = relationship(init=False, back_populates='users')
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from pydantic import Field
from backend.common.schema import SchemaBase
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from typing import Any
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import Field
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from typing import Any
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
+17 -16
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import Request, Response
from fastapi.security import HTTPBasicCredentials
from sqlalchemy.ext.asyncio import AsyncSession
@@ -49,7 +47,6 @@ class AuthService:
if user.password is None:
raise errors.AuthorizationError(msg='用户名或密码有误')
else:
if not password_verify(password, user.password):
raise errors.AuthorizationError(msg='用户名或密码有误')
@@ -70,14 +67,19 @@ class AuthService:
await user_dao.update_login_time(db, obj.username)
access_token = await create_access_token(
user.id,
user.is_multi_login,
multi_login=user.is_multi_login,
# extra info
swagger=True,
)
return access_token.access_token, user
async def login(
self, *, request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks
self,
*,
request: Request,
response: Response,
obj: AuthLoginParam,
background_tasks: BackgroundTasks,
) -> GetLoginToken:
"""
用户登录
@@ -102,7 +104,7 @@ class AuthService:
await db.refresh(user)
access_token = await create_access_token(
user.id,
user.is_multi_login,
multi_login=user.is_multi_login,
# extra info
username=user.username,
nickname=user.nickname,
@@ -112,7 +114,11 @@ class AuthService:
browser=request.state.browser,
device=request.state.device,
)
refresh_token = await create_refresh_token(access_token.session_uuid, user.id, user.is_multi_login)
refresh_token = await create_refresh_token(
access_token.session_uuid,
user.id,
multi_login=user.is_multi_login,
)
response.set_cookie(
key=settings.COOKIE_REFRESH_TOKEN_KEY,
value=refresh_token.refresh_token,
@@ -128,7 +134,6 @@ class AuthService:
log.error('登陆错误: 用户密码有误')
task = BackgroundTask(
login_log_service.create,
**dict(
db=db,
request=request,
user_uuid=user.uuid if user else uuid4_str(),
@@ -136,16 +141,14 @@ class AuthService:
login_time=timezone.now(),
status=LoginLogStatusType.fail.value,
msg=e.msg,
),
)
raise errors.RequestError(code=e.code, msg=e.msg, background=task)
except Exception as e:
log.error(f'登陆错误: {e}')
raise e
raise
else:
background_tasks.add_task(
login_log_service.create,
**dict(
db=db,
request=request,
user_uuid=user.uuid,
@@ -153,7 +156,6 @@ class AuthService:
login_time=timezone.now(),
status=LoginLogStatusType.success.value,
msg=t('success.login.success'),
),
)
data = GetLoginToken(
access_token=access_token.access_token,
@@ -204,16 +206,15 @@ class AuthService:
user = await user_dao.get(db, token_payload.id)
if not user:
raise errors.NotFoundError(msg='用户不存在')
elif not user.status:
if not user.status:
raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员')
if not user.is_multi_login:
if await redis_client.keys(match=f'{settings.TOKEN_REDIS_PREFIX}:{user.id}:*'):
if not user.is_multi_login and await redis_client.keys(match=f'{settings.TOKEN_REDIS_PREFIX}:{user.id}:*'):
raise errors.ForbiddenError(msg='此用户已在异地登录,请重新登录并及时修改密码')
new_token = await create_new_token(
refresh_token,
token_payload.session_uuid,
user.id,
user.is_multi_login,
multi_login=user.is_multi_login,
# extra info
username=user.username,
nickname=user.nickname,
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from collections.abc import Sequence
from sqlalchemy import Select
@@ -103,8 +101,7 @@ class DataRuleService:
data_rule = await data_rule_dao.get(db, pk)
if not data_rule:
raise errors.NotFoundError(msg='数据规则不存在')
if data_rule.name != obj.name:
if await data_rule_dao.get_by_name(db, obj.name):
if data_rule.name != obj.name and await data_rule_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='数据规则已存在')
count = await data_rule_dao.update(db, pk, obj)
return count
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from collections.abc import Sequence
from sqlalchemy import Select
@@ -94,8 +92,7 @@ class DataScopeService:
data_scope = await data_scope_dao.get(db, pk)
if not data_scope:
raise errors.NotFoundError(msg='数据范围不存在')
if data_scope.name != obj.name:
if await data_scope_dao.get_by_name(db, obj.name):
if data_scope.name != obj.name and await data_scope_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='数据范围已存在')
count = await data_scope_dao.update(db, pk, obj)
for role in await data_scope.awaitable_attrs.roles:
+7 -5
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Any
from fastapi import Request
@@ -33,7 +31,12 @@ class DeptService:
@staticmethod
async def get_tree(
*, request: Request, name: str | None, leader: str | None, phone: str | None, status: int | None
*,
request: Request,
name: str | None,
leader: str | None,
phone: str | None,
status: int | None,
) -> list[dict[str, Any]]:
"""
获取部门树形结构
@@ -81,8 +84,7 @@ class DeptService:
dept = await dept_dao.get(db, pk)
if not dept:
raise errors.NotFoundError(msg='部门不存在')
if dept.name != obj.name:
if await dept_dao.get_by_name(db, obj.name):
if dept.name != obj.name and await dept_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='部门名称已存在')
if obj.parent_id:
parent_dept = await dept_dao.get(db, obj.parent_id)
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from fastapi import Request
+2 -6
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Any
from fastapi import Request
@@ -61,8 +59,7 @@ class MenuService:
menu_ids = set()
if roles:
for role in roles:
for menu in role.menus:
menu_ids.add(menu.id)
menu_ids.update(menu.id for menu in role.menus)
menu_data = await menu_dao.get_sidebar(db, list(menu_ids))
menu_tree = get_vben5_tree_data(menu_data)
return menu_tree
@@ -98,8 +95,7 @@ class MenuService:
menu = await menu_dao.get(db, pk)
if not menu:
raise errors.NotFoundError(msg='菜单不存在')
if menu.title != obj.title:
if await menu_dao.get_by_title(db, obj.title):
if menu.title != obj.title and await menu_dao.get_by_title(db, obj.title):
raise errors.ConflictError(msg='菜单标题已存在')
if obj.parent_id:
parent_menu = await menu_dao.get(db, obj.parent_id)
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from sqlalchemy import Select
from backend.app.admin.crud.crud_opera_log import opera_log_dao
+12 -16
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import io
import json
import os
@@ -8,6 +6,8 @@ import zipfile
from typing import Any
import anyio
from fastapi import UploadFile
from backend.common.enums import PluginType, StatusType
@@ -26,14 +26,10 @@ class PluginService:
@staticmethod
async def get_all() -> list[dict[str, Any]]:
"""获取所有插件"""
keys = []
result = []
async for key in redis_client.scan_iter(f'{settings.PLUGIN_REDIS_PREFIX}:*'):
keys.append(key)
keys = [key async for key in redis_client.scan_iter(f'{settings.PLUGIN_REDIS_PREFIX}:*')]
for info in await redis_client.mget(*keys):
result.append(json.loads(info))
result = [json.loads(info) for info in await redis_client.mget(*keys)]
return result
@@ -61,24 +57,24 @@ class PluginService:
return await install_git_plugin(repo_url)
@staticmethod
async def uninstall(*, plugin: str):
async def uninstall(*, plugin: str) -> None:
"""
卸载插件
:param plugin: 插件名称
:return:
"""
plugin_dir = os.path.join(PLUGIN_DIR, plugin)
if not os.path.exists(plugin_dir):
plugin_dir = anyio.Path(PLUGIN_DIR / plugin)
if not await plugin_dir.exists():
raise errors.NotFoundError(msg='插件不存在')
await uninstall_requirements_async(plugin)
bacup_dir = os.path.join(PLUGIN_DIR, f'{plugin}.{timezone.now().strftime("%Y%m%d%H%M%S")}.backup')
bacup_dir = PLUGIN_DIR / f'{plugin}.{timezone.now().strftime("%Y%m%d%H%M%S")}.backup'
shutil.move(plugin_dir, bacup_dir)
await redis_client.delete(f'{settings.PLUGIN_REDIS_PREFIX}:{plugin}')
await redis_client.set(f'{settings.PLUGIN_REDIS_PREFIX}:changed', 'ture')
@staticmethod
async def update_status(*, plugin: str):
async def update_status(*, plugin: str) -> None:
"""
更新插件状态
@@ -107,8 +103,8 @@ class PluginService:
:param plugin: 插件名称
:return:
"""
plugin_dir = os.path.join(PLUGIN_DIR, plugin)
if not os.path.exists(plugin_dir):
plugin_dir = anyio.Path(PLUGIN_DIR / plugin)
if not await plugin_dir.exists():
raise errors.NotFoundError(msg='插件不存在')
bio = io.BytesIO()
@@ -117,7 +113,7 @@ class PluginService:
dirs[:] = [d for d in dirs if d != '__pycache__']
for file in files:
file_path = os.path.join(root, file)
arcname = os.path.relpath(file_path, start=plugin_dir)
arcname = os.path.relpath(file_path, start=plugin_dir) # noqa: ASYNC240
zf.write(file_path, os.path.join(plugin, arcname))
bio.seek(0)
+3 -5
View File
@@ -1,6 +1,5 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Any, Sequence
from collections.abc import Sequence
from typing import Any
from sqlalchemy import Select
@@ -114,8 +113,7 @@ class RoleService:
role = await role_dao.get(db, pk)
if not role:
raise errors.NotFoundError(msg='角色不存在')
if role.name != obj.name:
if await role_dao.get_by_name(db, obj.name):
if role.name != obj.name and await role_dao.get_by_name(db, obj.name):
raise errors.ConflictError(msg='角色已存在')
count = await role_dao.update(db, pk, obj)
for user in await role.awaitable_attrs.users:
+9 -11
View File
@@ -1,8 +1,6 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import random
from typing import Sequence
from collections.abc import Sequence
from fastapi import Request
from sqlalchemy import Select
@@ -83,7 +81,7 @@ class UserService:
superuser_verify(request)
if await user_dao.get_by_username(db, obj.username):
raise errors.ConflictError(msg='用户名已注册')
obj.nickname = obj.nickname if obj.nickname else f'#{random.randrange(88888, 99999)}'
obj.nickname = obj.nickname or f'#{random.randrange(88888, 99999)}'
if not obj.password:
raise errors.RequestError(msg='密码不允许为空')
if not await dept_dao.get(db, obj.dept_id):
@@ -108,8 +106,7 @@ class UserService:
user = await user_dao.get_with_relation(db, user_id=pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
if obj.username != user.username:
if await user_dao.get_by_username(db, obj.username):
if obj.username != user.username and await user_dao.get_by_username(db, obj.username):
raise errors.ConflictError(msg='用户名已注册')
for role_id in obj.roles:
if not await role_dao.get(db, role_id):
@@ -119,7 +116,7 @@ class UserService:
return count
@staticmethod
async def update_permission(*, request: Request, pk: int, type: UserPermissionType) -> int:
async def update_permission(*, request: Request, pk: int, type: UserPermissionType) -> int: # noqa: C901
"""
更新用户权限
@@ -137,14 +134,14 @@ class UserService:
raise errors.NotFoundError(msg='用户不存在')
if pk == request.user.id:
raise errors.ForbiddenError(msg='禁止修改自身权限')
count = await user_dao.set_super(db, pk, not user.status)
count = await user_dao.set_super(db, pk, is_super=not user.status)
case UserPermissionType.staff:
user = await user_dao.get(db, pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
if pk == request.user.id:
raise errors.ForbiddenError(msg='禁止修改自身权限')
count = await user_dao.set_staff(db, pk, not user.is_staff)
count = await user_dao.set_staff(db, pk, is_staff=not user.is_staff)
case UserPermissionType.status:
user = await user_dao.get(db, pk)
if not user:
@@ -158,7 +155,7 @@ class UserService:
raise errors.NotFoundError(msg='用户不存在')
multi_login = user.is_multi_login if pk != user.id else request.user.is_multi_login
new_multi_login = not multi_login
count = await user_dao.set_multi_login(db, pk, new_multi_login)
count = await user_dao.set_multi_login(db, pk, multi_login=new_multi_login)
token = get_token(request)
token_payload = jwt_decode(token)
if pk == user.id:
@@ -166,7 +163,8 @@ class UserService:
if not new_multi_login:
key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{user.id}'
await redis_client.delete_prefix(
key_prefix, exclude=f'{key_prefix}:{token_payload.session_uuid}'
key_prefix,
exclude=f'{key_prefix}:{token_payload.session_uuid}',
)
else:
# 系统管理员修改他人时,他人 token 全部失效
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from starlette.testclient import TestClient
+1 -3
View File
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Generator
from collections.abc import Generator
import pytest
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
+1 -3
View File
@@ -1,6 +1,4 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import AsyncGenerator
from collections.abc import AsyncGenerator
from sqlalchemy.ext.asyncio.session import AsyncSession
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from fastapi import APIRouter
from backend.app.admin.api.router import v1 as admin_v1
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import sys
from backend.core.path_conf import BASE_PATH
+1 -3
View File
@@ -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)
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
-2
View File
@@ -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
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
+1 -3
View File
@@ -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:
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query
+8 -5
View File
@@ -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:
+3 -5
View File
@@ -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)
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
-2
View File
@@ -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
+3 -5
View File
@@ -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
+34 -27
View File
@@ -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)
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.common.enums import IntEnum, StrEnum
+2 -4
View File
@@ -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
+14 -13
View File
@@ -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__
+13 -9
View File
@@ -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())
-2
View File
@@ -1,2 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.common.schema import SchemaBase
+1 -3
View File
@@ -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)
-2
View File
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
-2
View File
@@ -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

Some files were not shown because too many files have changed in this diff Show More