diff --git a/backend/app/admin/api/v1/auth/auth.py b/backend/app/admin/api/v1/auth/auth.py index 7aa1cfc0..cd5cac88 100644 --- a/backend/app/admin/api/v1/auth/auth.py +++ b/backend/app/admin/api/v1/auth/auth.py @@ -10,13 +10,16 @@ from backend.app.admin.schema.user import AuthLoginParam from backend.app.admin.service.auth_service import auth_service from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.post('/login/swagger', summary='swagger 调试专用', description='用于快捷获取 token 进行 swagger 认证') -async def login_swagger(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetSwaggerToken: - token, user = await auth_service.swagger_login(obj=obj) +async def login_swagger( + db: CurrentSessionTransaction, obj: Annotated[HTTPBasicCredentials, Depends()] +) -> GetSwaggerToken: + token, user = await auth_service.swagger_login(db=db, obj=obj) return GetSwaggerToken(access_token=token, user=user) @@ -27,24 +30,24 @@ async def login_swagger(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetS dependencies=[Depends(RateLimiter(times=5, minutes=1))], ) async def login( - request: Request, + db: CurrentSessionTransaction, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks, ) -> ResponseSchemaModel[GetLoginToken]: - data = await auth_service.login(request=request, response=response, obj=obj, background_tasks=background_tasks) + data = await auth_service.login(db=db, response=response, obj=obj, background_tasks=background_tasks) return response_base.success(data=data) @router.get('/codes', summary='获取所有授权码', description='适配 vben admin v5', dependencies=[DependsJwtAuth]) -async def get_codes(request: Request) -> ResponseSchemaModel[list[str]]: - codes = await auth_service.get_codes(request=request) +async def get_codes(db: CurrentSession, request: Request) -> ResponseSchemaModel[list[str]]: + codes = await auth_service.get_codes(db=db, request=request) return response_base.success(data=codes) @router.post('/refresh', summary='刷新 token') -async def refresh_token(request: Request) -> ResponseSchemaModel[GetNewToken]: - data = await auth_service.refresh_token(request=request) +async def refresh_token(db: CurrentSession, request: Request) -> ResponseSchemaModel[GetNewToken]: + data = await auth_service.refresh_token(db=db, request=request) return response_base.success(data=data) diff --git a/backend/app/admin/api/v1/auth/captcha.py b/backend/app/admin/api/v1/auth/captcha.py index 58d4ef48..dee8a41c 100644 --- a/backend/app/admin/api/v1/auth/captcha.py +++ b/backend/app/admin/api/v1/auth/captcha.py @@ -1,7 +1,7 @@ from uuid import uuid4 from fast_captcha import img_captcha -from fastapi import APIRouter, Depends, Request +from fastapi import APIRouter, Depends from fastapi_limiter.depends import RateLimiter from starlette.concurrency import run_in_threadpool @@ -18,7 +18,7 @@ router = APIRouter() summary='获取登录验证码', dependencies=[Depends(RateLimiter(times=5, seconds=10))], ) -async def get_captcha(request: Request) -> ResponseSchemaModel[GetCaptchaDetail]: +async def get_captcha() -> ResponseSchemaModel[GetCaptchaDetail]: """ 此接口可能存在性能损耗,尽管是异步接口,但是验证码生成是IO密集型任务,使用线程池尽量减少性能损耗 """ diff --git a/backend/app/admin/api/v1/log/login_log.py b/backend/app/admin/api/v1/log/login_log.py index f308c578..0e9e64d7 100644 --- a/backend/app/admin/api/v1/log/login_log.py +++ b/backend/app/admin/api/v1/log/login_log.py @@ -4,12 +4,12 @@ from fastapi import APIRouter, Depends, Query from backend.app.admin.schema.login_log import DeleteLoginLogParam, GetLoginLogDetail from backend.app.admin.service.login_log_service import login_log_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @@ -22,14 +22,14 @@ router = APIRouter() DependsPagination, ], ) -async def get_login_logs_paged( +async def get_login_logs_paginated( db: CurrentSession, username: Annotated[str | None, Query(description='用户名')] = None, status: Annotated[int | None, Query(description='状态')] = None, ip: Annotated[str | None, Query(description='IP 地址')] = None, ) -> ResponseSchemaModel[PageData[GetLoginLogDetail]]: - log_select = await login_log_service.get_select(username=username, status=status, ip=ip) - page_data = await paging_data(db, log_select) + page_data = await login_log_service.get_list(db=db, username=username, status=status, ip=ip) + return response_base.success(data=page_data) @@ -41,8 +41,8 @@ async def get_login_logs_paged( DependsRBAC, ], ) -async def delete_login_logs(obj: DeleteLoginLogParam) -> ResponseModel: - count = await login_log_service.delete(obj=obj) +async def delete_login_logs(db: CurrentSessionTransaction, obj: DeleteLoginLogParam) -> ResponseModel: + count = await login_log_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -56,6 +56,6 @@ async def delete_login_logs(obj: DeleteLoginLogParam) -> ResponseModel: DependsRBAC, ], ) -async def delete_all_login_logs() -> ResponseModel: - await login_log_service.delete_all() +async def delete_all_login_logs(db: CurrentSessionTransaction) -> ResponseModel: + await login_log_service.delete_all(db=db) return response_base.success() diff --git a/backend/app/admin/api/v1/log/opera_log.py b/backend/app/admin/api/v1/log/opera_log.py index 64c0951b..7c937d89 100644 --- a/backend/app/admin/api/v1/log/opera_log.py +++ b/backend/app/admin/api/v1/log/opera_log.py @@ -4,12 +4,12 @@ from fastapi import APIRouter, Depends, Query from backend.app.admin.schema.opera_log import DeleteOperaLogParam, GetOperaLogDetail from backend.app.admin.service.opera_log_service import opera_log_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @@ -22,14 +22,14 @@ router = APIRouter() DependsPagination, ], ) -async def get_opera_logs_paged( +async def get_opera_logs_paginated( db: CurrentSession, username: Annotated[str | None, Query(description='用户名')] = None, status: Annotated[int | None, Query(description='状态')] = None, ip: Annotated[str | None, Query(description='IP 地址')] = None, ) -> ResponseSchemaModel[PageData[GetOperaLogDetail]]: - log_select = await opera_log_service.get_select(username=username, status=status, ip=ip) - page_data = await paging_data(db, log_select) + page_data = await opera_log_service.get_list(db=db, username=username, status=status, ip=ip) + return response_base.success(data=page_data) @@ -41,8 +41,8 @@ async def get_opera_logs_paged( DependsRBAC, ], ) -async def delete_opera_logs(obj: DeleteOperaLogParam) -> ResponseModel: - count = await opera_log_service.delete(obj=obj) +async def delete_opera_logs(db: CurrentSessionTransaction, obj: DeleteOperaLogParam) -> ResponseModel: + count = await opera_log_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -56,6 +56,6 @@ async def delete_opera_logs(obj: DeleteOperaLogParam) -> ResponseModel: DependsRBAC, ], ) -async def delete_all_opera_logs() -> ResponseModel: - await opera_log_service.delete_all() +async def delete_all_opera_logs(db: CurrentSessionTransaction) -> ResponseModel: + await opera_log_service.delete_all(db=db) return response_base.success() diff --git a/backend/app/admin/api/v1/monitor/online.py b/backend/app/admin/api/v1/monitor/online.py index e22e078c..999cf693 100644 --- a/backend/app/admin/api/v1/monitor/online.py +++ b/backend/app/admin/api/v1/monitor/online.py @@ -2,7 +2,7 @@ import json from typing import Annotated -from fastapi import APIRouter, Path, Query, Request +from fastapi import APIRouter, Path, Query from backend.app.admin.schema.token import GetTokenDetail from backend.common.enums import StatusType @@ -76,7 +76,6 @@ async def get_sessions( dependencies=[DependsSuperUser], ) async def delete_session( - request: Request, pk: Annotated[int, Path(description='用户 ID')], session_uuid: Annotated[str, Query(description='会话 UUID')], ) -> ResponseModel: diff --git a/backend/app/admin/api/v1/sys/data_rule.py b/backend/app/admin/api/v1/sys/data_rule.py index f018ca09..5757b351 100644 --- a/backend/app/admin/api/v1/sys/data_rule.py +++ b/backend/app/admin/api/v1/sys/data_rule.py @@ -10,12 +10,12 @@ from backend.app.admin.schema.data_rule import ( UpdateDataRuleParam, ) from backend.app.admin.service.data_rule_service import data_rule_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @@ -35,16 +35,17 @@ async def get_data_rule_model_columns( @router.get('/all', summary='获取所有数据规则', dependencies=[DependsJwtAuth]) -async def get_all_data_rules() -> ResponseSchemaModel[list[GetDataRuleDetail]]: - data = await data_rule_service.get_all() +async def get_all_data_rules(db: CurrentSession) -> ResponseSchemaModel[list[GetDataRuleDetail]]: + data = await data_rule_service.get_all(db=db) return response_base.success(data=data) @router.get('/{pk}', summary='获取数据规则详情', dependencies=[DependsJwtAuth]) async def get_data_rule( + db: CurrentSession, pk: Annotated[int, Path(description='数据规则 ID')], ) -> ResponseSchemaModel[GetDataRuleDetail]: - data = await data_rule_service.get(pk=pk) + data = await data_rule_service.get(db=db, pk=pk) return response_base.success(data=data) @@ -56,12 +57,11 @@ async def get_data_rule( DependsPagination, ], ) -async def get_data_rules_paged( +async def get_data_rules_paginated( 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) + page_data = await data_rule_service.get_list(db=db, name=name) return response_base.success(data=page_data) @@ -73,8 +73,8 @@ async def get_data_rules_paged( DependsRBAC, ], ) -async def create_data_rule(obj: CreateDataRuleParam) -> ResponseModel: - await data_rule_service.create(obj=obj) +async def create_data_rule(db: CurrentSessionTransaction, obj: CreateDataRuleParam) -> ResponseModel: + await data_rule_service.create(db=db, obj=obj) return response_base.success() @@ -87,10 +87,11 @@ async def create_data_rule(obj: CreateDataRuleParam) -> ResponseModel: ], ) async def update_data_rule( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='数据规则 ID')], obj: UpdateDataRuleParam, ) -> ResponseModel: - count = await data_rule_service.update(pk=pk, obj=obj) + count = await data_rule_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -104,8 +105,8 @@ async def update_data_rule( DependsRBAC, ], ) -async def delete_data_rules(obj: DeleteDataRuleParam) -> ResponseModel: - count = await data_rule_service.delete(obj=obj) +async def delete_data_rules(db: CurrentSessionTransaction, obj: DeleteDataRuleParam) -> ResponseModel: + count = await data_rule_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/admin/api/v1/sys/data_scope.py b/backend/app/admin/api/v1/sys/data_scope.py index eb9db7b3..460ae60c 100644 --- a/backend/app/admin/api/v1/sys/data_scope.py +++ b/backend/app/admin/api/v1/sys/data_scope.py @@ -11,35 +11,37 @@ from backend.app.admin.schema.data_scope import ( UpdateDataScopeRuleParam, ) from backend.app.admin.service.data_scope_service import data_scope_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.get('/all', summary='获取所有数据范围', dependencies=[DependsJwtAuth]) -async def get_all_data_scope() -> ResponseSchemaModel[list[GetDataScopeDetail]]: - data = await data_scope_service.get_all() +async def get_all_data_scope(db: CurrentSession) -> ResponseSchemaModel[list[GetDataScopeDetail]]: + data = await data_scope_service.get_all(db=db) return response_base.success(data=data) @router.get('/{pk}', summary='获取数据范围详情', dependencies=[DependsJwtAuth]) async def get_data_scope( + db: CurrentSession, pk: Annotated[int, Path(description='数据范围 ID')], ) -> ResponseSchemaModel[GetDataScopeDetail]: - data = await data_scope_service.get(pk=pk) + data = await data_scope_service.get(db=db, pk=pk) return response_base.success(data=data) @router.get('/{pk}/rules', summary='获取数据范围所有规则', dependencies=[DependsJwtAuth]) async def get_data_scope_rules( + db: CurrentSession, pk: Annotated[int, Path(description='数据范围 ID')], ) -> ResponseSchemaModel[GetDataScopeWithRelationDetail]: - data = await data_scope_service.get_rules(pk=pk) + data = await data_scope_service.get_rules(db=db, pk=pk) return response_base.success(data=data) @@ -51,13 +53,12 @@ async def get_data_scope_rules( DependsPagination, ], ) -async def get_data_scopes_paged( +async def get_data_scopes_paginated( db: CurrentSession, name: Annotated[str | None, Query(description='范围名称')] = None, status: Annotated[int | None, Query(description='状态')] = None, ) -> ResponseSchemaModel[PageData[GetDataScopeDetail]]: - data_scope_select = await data_scope_service.get_select(name=name, status=status) - page_data = await paging_data(db, data_scope_select) + page_data = await data_scope_service.get_list(db=db, name=name, status=status) return response_base.success(data=page_data) @@ -69,8 +70,8 @@ async def get_data_scopes_paged( DependsRBAC, ], ) -async def create_data_scope(obj: CreateDataScopeParam) -> ResponseModel: - await data_scope_service.create(obj=obj) +async def create_data_scope(db: CurrentSessionTransaction, obj: CreateDataScopeParam) -> ResponseModel: + await data_scope_service.create(db=db, obj=obj) return response_base.success() @@ -83,10 +84,11 @@ async def create_data_scope(obj: CreateDataScopeParam) -> ResponseModel: ], ) async def update_data_scope( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='数据范围 ID')], obj: UpdateDataScopeParam, ) -> ResponseModel: - count = await data_scope_service.update(pk=pk, obj=obj) + count = await data_scope_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -101,10 +103,11 @@ async def update_data_scope( ], ) async def update_data_scope_rules( + db: CurrentSessionTransaction, 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) + count = await data_scope_service.update_data_scope_rule(db=db, pk=pk, rule_ids=rule_ids) if count > 0: return response_base.success() return response_base.fail() @@ -118,8 +121,8 @@ async def update_data_scope_rules( DependsRBAC, ], ) -async def delete_data_scopes(obj: DeleteDataScopeParam) -> ResponseModel: - count = await data_scope_service.delete(obj=obj) +async def delete_data_scopes(db: CurrentSessionTransaction, obj: DeleteDataScopeParam) -> ResponseModel: + count = await data_scope_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/admin/api/v1/sys/dept.py b/backend/app/admin/api/v1/sys/dept.py index dbc8678c..86ef0efb 100644 --- a/backend/app/admin/api/v1/sys/dept.py +++ b/backend/app/admin/api/v1/sys/dept.py @@ -8,25 +8,29 @@ from backend.common.response.response_schema import ResponseModel, ResponseSchem from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.get('/{pk}', summary='获取部门详情', dependencies=[DependsJwtAuth]) -async def get_dept(pk: Annotated[int, Path(description='部门 ID')]) -> ResponseSchemaModel[GetDeptDetail]: - data = await dept_service.get(pk=pk) +async def get_dept( + db: CurrentSession, pk: Annotated[int, Path(description='部门 ID')] +) -> ResponseSchemaModel[GetDeptDetail]: + data = await dept_service.get(db=db, pk=pk) return response_base.success(data=data) @router.get('', summary='获取部门树', dependencies=[DependsJwtAuth]) async def get_dept_tree( + db: CurrentSession, request: Request, name: Annotated[str | None, Query(description='部门名称')] = None, leader: Annotated[str | None, Query(description='部门负责人')] = None, phone: Annotated[str | None, Query(description='联系电话')] = None, status: Annotated[int | None, Query(description='状态')] = None, ) -> ResponseSchemaModel[list[GetDeptTree]]: - dept = await dept_service.get_tree(request=request, name=name, leader=leader, phone=phone, status=status) + dept = await dept_service.get_tree(db=db, request=request, name=name, leader=leader, phone=phone, status=status) return response_base.success(data=dept) @@ -38,8 +42,8 @@ async def get_dept_tree( DependsRBAC, ], ) -async def create_dept(obj: CreateDeptParam) -> ResponseModel: - await dept_service.create(obj=obj) +async def create_dept(db: CurrentSessionTransaction, obj: CreateDeptParam) -> ResponseModel: + await dept_service.create(db=db, obj=obj) return response_base.success() @@ -51,8 +55,10 @@ async def create_dept(obj: CreateDeptParam) -> ResponseModel: DependsRBAC, ], ) -async def update_dept(pk: Annotated[int, Path(description='部门 ID')], obj: UpdateDeptParam) -> ResponseModel: - count = await dept_service.update(pk=pk, obj=obj) +async def update_dept( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='部门 ID')], obj: UpdateDeptParam +) -> ResponseModel: + count = await dept_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -66,8 +72,8 @@ async def update_dept(pk: Annotated[int, Path(description='部门 ID')], obj: Up DependsRBAC, ], ) -async def delete_dept(pk: Annotated[int, Path(description='部门 ID')]) -> ResponseModel: - count = await dept_service.delete(pk=pk) +async def delete_dept(db: CurrentSessionTransaction, pk: Annotated[int, Path(description='部门 ID')]) -> ResponseModel: + count = await dept_service.delete(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/admin/api/v1/sys/menu.py b/backend/app/admin/api/v1/sys/menu.py index e44034d4..58a7ff6d 100644 --- a/backend/app/admin/api/v1/sys/menu.py +++ b/backend/app/admin/api/v1/sys/menu.py @@ -8,28 +8,32 @@ from backend.common.response.response_schema import ResponseModel, ResponseSchem from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.get('/sidebar', summary='获取用户菜单侧边栏', description='已适配 vben admin v5', dependencies=[DependsJwtAuth]) -async def get_user_sidebar(request: Request) -> ResponseSchemaModel[list[dict[str, Any] | None]]: - menu = await menu_service.get_sidebar(request=request) +async def get_user_sidebar(db: CurrentSession, request: Request) -> ResponseSchemaModel[list[dict[str, Any] | None]]: + menu = await menu_service.get_sidebar(db=db, request=request) return response_base.success(data=menu) @router.get('/{pk}', summary='获取菜单详情', dependencies=[DependsJwtAuth]) -async def get_menu(pk: Annotated[int, Path(description='菜单 ID')]) -> ResponseSchemaModel[GetMenuDetail]: - data = await menu_service.get(pk=pk) +async def get_menu( + db: CurrentSession, pk: Annotated[int, Path(description='菜单 ID')] +) -> ResponseSchemaModel[GetMenuDetail]: + data = await menu_service.get(db=db, pk=pk) return response_base.success(data=data) @router.get('', summary='获取菜单树', dependencies=[DependsJwtAuth]) async def get_menu_tree( + db: CurrentSession, title: Annotated[str | None, Query(description='菜单标题')] = None, status: Annotated[int | None, Query(description='状体')] = None, ) -> ResponseSchemaModel[list[GetMenuTree]]: - menu = await menu_service.get_tree(title=title, status=status) + menu = await menu_service.get_tree(db=db, title=title, status=status) return response_base.success(data=menu) @@ -41,8 +45,8 @@ async def get_menu_tree( DependsRBAC, ], ) -async def create_menu(obj: CreateMenuParam) -> ResponseModel: - await menu_service.create(obj=obj) +async def create_menu(db: CurrentSessionTransaction, obj: CreateMenuParam) -> ResponseModel: + await menu_service.create(db=db, obj=obj) return response_base.success() @@ -54,8 +58,10 @@ async def create_menu(obj: CreateMenuParam) -> ResponseModel: DependsRBAC, ], ) -async def update_menu(pk: Annotated[int, Path(description='菜单 ID')], obj: UpdateMenuParam) -> ResponseModel: - count = await menu_service.update(pk=pk, obj=obj) +async def update_menu( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='菜单 ID')], obj: UpdateMenuParam +) -> ResponseModel: + count = await menu_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -69,8 +75,8 @@ async def update_menu(pk: Annotated[int, Path(description='菜单 ID')], obj: Up DependsRBAC, ], ) -async def delete_menu(pk: Annotated[int, Path(description='菜单 ID')]) -> ResponseModel: - count = await menu_service.delete(pk=pk) +async def delete_menu(db: CurrentSessionTransaction, pk: Annotated[int, Path(description='菜单 ID')]) -> ResponseModel: + count = await menu_service.delete(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/admin/api/v1/sys/role.py b/backend/app/admin/api/v1/sys/role.py index 74972308..6b6dc5ab 100644 --- a/backend/app/admin/api/v1/sys/role.py +++ b/backend/app/admin/api/v1/sys/role.py @@ -13,39 +13,44 @@ from backend.app.admin.schema.role import ( UpdateRoleScopeParam, ) from backend.app.admin.service.role_service import role_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.get('/all', summary='获取所有角色', dependencies=[DependsJwtAuth]) -async def get_all_roles() -> ResponseSchemaModel[list[GetRoleDetail]]: - data = await role_service.get_all() +async def get_all_roles(db: CurrentSession) -> ResponseSchemaModel[list[GetRoleDetail]]: + data = await role_service.get_all(db=db) return response_base.success(data=data) @router.get('/{pk}/menus', summary='获取角色菜单树', dependencies=[DependsJwtAuth]) async def get_role_menu_tree( + db: CurrentSession, pk: Annotated[int, Path(description='角色 ID')], ) -> ResponseSchemaModel[list[GetMenuTree] | None]: - menu = await role_service.get_menu_tree(pk=pk) + menu = await role_service.get_menu_tree(db=db, pk=pk) return response_base.success(data=menu) @router.get('/{pk}/scopes', summary='获取角色所有数据范围', dependencies=[DependsJwtAuth]) -async def get_role_scopes(pk: Annotated[int, Path(description='角色 ID')]) -> ResponseSchemaModel[list[int]]: - rule = await role_service.get_scopes(pk=pk) +async def get_role_scopes( + db: CurrentSession, pk: Annotated[int, Path(description='角色 ID')] +) -> ResponseSchemaModel[list[int]]: + rule = await role_service.get_scopes(db=db, pk=pk) return response_base.success(data=rule) @router.get('/{pk}', summary='获取角色详情', dependencies=[DependsJwtAuth]) -async def get_role(pk: Annotated[int, Path(description='角色 ID')]) -> ResponseSchemaModel[GetRoleWithRelationDetail]: - data = await role_service.get(pk=pk) +async def get_role( + db: CurrentSession, pk: Annotated[int, Path(description='角色 ID')] +) -> ResponseSchemaModel[GetRoleWithRelationDetail]: + data = await role_service.get(db=db, pk=pk) return response_base.success(data=data) @@ -57,13 +62,12 @@ async def get_role(pk: Annotated[int, Path(description='角色 ID')]) -> Respons DependsPagination, ], ) -async def get_roles_paged( +async def get_roles_paginated( db: CurrentSession, name: Annotated[str | None, Query(description='角色名称')] = None, status: Annotated[int | None, Query(description='状态')] = None, ) -> ResponseSchemaModel[PageData[GetRoleDetail]]: - role_select = await role_service.get_select(name=name, status=status) - page_data = await paging_data(db, role_select) + page_data = await role_service.get_list(db=db, name=name, status=status) return response_base.success(data=page_data) @@ -75,8 +79,8 @@ async def get_roles_paged( DependsRBAC, ], ) -async def create_role(obj: CreateRoleParam) -> ResponseModel: - await role_service.create(obj=obj) +async def create_role(db: CurrentSessionTransaction, obj: CreateRoleParam) -> ResponseModel: + await role_service.create(db=db, obj=obj) return response_base.success() @@ -88,8 +92,10 @@ async def create_role(obj: CreateRoleParam) -> ResponseModel: DependsRBAC, ], ) -async def update_role(pk: Annotated[int, Path(description='角色 ID')], obj: UpdateRoleParam) -> ResponseModel: - count = await role_service.update(pk=pk, obj=obj) +async def update_role( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='角色 ID')], obj: UpdateRoleParam +) -> ResponseModel: + count = await role_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -104,10 +110,11 @@ async def update_role(pk: Annotated[int, Path(description='角色 ID')], obj: Up ], ) async def update_role_menus( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='角色 ID')], menu_ids: UpdateRoleMenuParam, ) -> ResponseModel: - count = await role_service.update_role_menu(pk=pk, menu_ids=menu_ids) + count = await role_service.update_role_menu(db=db, pk=pk, menu_ids=menu_ids) if count > 0: return response_base.success() return response_base.fail() @@ -122,10 +129,11 @@ async def update_role_menus( ], ) async def update_role_scopes( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='角色 ID')], scope_ids: UpdateRoleScopeParam, ) -> ResponseModel: - count = await role_service.update_role_scope(pk=pk, scope_ids=scope_ids) + count = await role_service.update_role_scope(db=db, pk=pk, scope_ids=scope_ids) if count > 0: return response_base.success() return response_base.fail() @@ -139,8 +147,8 @@ async def update_role_scopes( DependsRBAC, ], ) -async def delete_roles(obj: DeleteRoleParam) -> ResponseModel: - count = await role_service.delete(obj=obj) +async def delete_roles(db: CurrentSessionTransaction, obj: DeleteRoleParam) -> ResponseModel: + count = await role_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/admin/api/v1/sys/user.py b/backend/app/admin/api/v1/sys/user.py index 6e43ef1a..6912f2b6 100644 --- a/backend/app/admin/api/v1/sys/user.py +++ b/backend/app/admin/api/v1/sys/user.py @@ -12,12 +12,12 @@ from backend.app.admin.schema.user import ( ) from backend.app.admin.service.user_service import user_service from backend.common.enums import UserPermissionType -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth, DependsSuperUser from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @@ -30,15 +30,18 @@ async def get_current_user(request: Request) -> ResponseSchemaModel[GetCurrentUs @router.get('/{pk}', summary='获取用户信息', dependencies=[DependsJwtAuth]) async def get_userinfo( + db: CurrentSession, pk: Annotated[int, Path(description='用户 ID')], ) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]: - data = await user_service.get_userinfo(pk=pk) + data = await user_service.get_userinfo(db=db, pk=pk) return response_base.success(data=data) @router.get('/{pk}/roles', summary='获取用户所有角色', dependencies=[DependsJwtAuth]) -async def get_user_roles(pk: Annotated[int, Path(description='用户 ID')]) -> ResponseSchemaModel[list[GetRoleDetail]]: - data = await user_service.get_roles(pk=pk) +async def get_user_roles( + db: CurrentSession, pk: Annotated[int, Path(description='用户 ID')] +) -> ResponseSchemaModel[list[GetRoleDetail]]: + data = await user_service.get_roles(db=db, pk=pk) return response_base.success(data=data) @@ -50,31 +53,33 @@ async def get_user_roles(pk: Annotated[int, Path(description='用户 ID')]) -> R DependsPagination, ], ) -async def get_users_paged( +async def get_users_paginated( db: CurrentSession, dept: Annotated[int | None, Query(description='部门 ID')] = None, username: Annotated[str | None, Query(description='用户名')] = None, phone: Annotated[str | None, Query(description='手机号')] = None, status: Annotated[int | None, Query(description='状态')] = None, ) -> ResponseSchemaModel[PageData[GetUserInfoWithRelationDetail]]: - user_select = await user_service.get_select(dept=dept, username=username, phone=phone, status=status) - page_data = await paging_data(db, user_select) + page_data = await user_service.get_list(db=db, dept=dept, username=username, phone=phone, status=status) return response_base.success(data=page_data) @router.post('', summary='创建用户', dependencies=[DependsSuperUser]) -async def create_user(obj: AddUserParam) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]: - await user_service.create(obj=obj) - data = await user_service.get_userinfo(username=obj.username) +async def create_user( + db: CurrentSessionTransaction, obj: AddUserParam +) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]: + await user_service.create(db=db, obj=obj) + data = await user_service.get_userinfo(db=db, username=obj.username) return response_base.success(data=data) @router.put('/{pk}', summary='更新用户信息', dependencies=[DependsSuperUser]) async def update_user( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='用户 ID')], obj: UpdateUserParam, ) -> ResponseModel: - count = await user_service.update(pk=pk, obj=obj) + count = await user_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -82,19 +87,22 @@ async def update_user( @router.put('/{pk}/permissions', summary='更新用户权限', dependencies=[DependsSuperUser]) async def update_user_permission( + db: CurrentSessionTransaction, request: Request, pk: Annotated[int, Path(description='用户 ID')], type: Annotated[UserPermissionType, Query(description='权限类型')], ) -> ResponseModel: - count = await user_service.update_permission(request=request, pk=pk, type=type) + count = await user_service.update_permission(db=db, request=request, pk=pk, type=type) if count > 0: return response_base.success() return response_base.fail() @router.put('/me/password', summary='更新当前用户密码', dependencies=[DependsJwtAuth]) -async def update_user_password(request: Request, obj: ResetPasswordParam) -> ResponseModel: - count = await user_service.update_password(request=request, obj=obj) +async def update_user_password( + db: CurrentSessionTransaction, request: Request, obj: ResetPasswordParam +) -> ResponseModel: + count = await user_service.update_password(db=db, request=request, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -102,10 +110,11 @@ async def update_user_password(request: Request, obj: ResetPasswordParam) -> Res @router.put('/{pk}/password', summary='重置用户密码', dependencies=[DependsSuperUser]) async def reset_user_password( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='用户 ID')], password: Annotated[str, Body(embed=True, description='新密码')], ) -> ResponseModel: - count = await user_service.reset_password(pk=pk, password=password) + count = await user_service.reset_password(db=db, pk=pk, password=password) if count > 0: return response_base.success() return response_base.fail() @@ -113,10 +122,11 @@ async def reset_user_password( @router.put('/me/nickname', summary='更新当前用户昵称', dependencies=[DependsJwtAuth]) async def update_user_nickname( + db: CurrentSessionTransaction, request: Request, nickname: Annotated[str, Body(embed=True, description='用户昵称')], ) -> ResponseModel: - count = await user_service.update_nickname(request=request, nickname=nickname) + count = await user_service.update_nickname(db=db, request=request, nickname=nickname) if count > 0: return response_base.success() return response_base.fail() @@ -124,10 +134,11 @@ async def update_user_nickname( @router.put('/me/avatar', summary='更新当前用户头像', dependencies=[DependsJwtAuth]) async def update_user_avatar( + db: CurrentSessionTransaction, request: Request, avatar: Annotated[str, Body(embed=True, description='用户头像地址')], ) -> ResponseModel: - count = await user_service.update_avatar(request=request, avatar=avatar) + count = await user_service.update_avatar(db=db, request=request, avatar=avatar) if count > 0: return response_base.success() return response_base.fail() @@ -135,11 +146,12 @@ async def update_user_avatar( @router.put('/me/email', summary='更新当前用户邮箱', dependencies=[DependsJwtAuth]) async def update_user_email( + db: CurrentSessionTransaction, request: Request, captcha: Annotated[str, Body(embed=True, description='邮箱验证码')], email: Annotated[str, Body(embed=True, description='用户邮箱')], ) -> ResponseModel: - count = await user_service.update_email(request=request, captcha=captcha, email=email) + count = await user_service.update_email(db=db, request=request, captcha=captcha, email=email) if count > 0: return response_base.success() return response_base.fail() @@ -153,8 +165,8 @@ async def update_user_email( DependsRBAC, ], ) -async def delete_user(pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel: - count = await user_service.delete(pk=pk) +async def delete_user(db: CurrentSessionTransaction, pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel: + count = await user_service.delete(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/admin/crud/crud_data_rule.py b/backend/app/admin/crud/crud_data_rule.py index 378e63ff..48dc6d16 100644 --- a/backend/app/admin/crud/crud_data_rule.py +++ b/backend/app/admin/crud/crud_data_rule.py @@ -21,9 +21,9 @@ class CRUDDataRule(CRUDPlus[DataRule]): """ return await self.select_model(db, pk) - async def get_list(self, name: str | None) -> Select: + async def get_select(self, name: str | None) -> Select: """ - 获取规则列表 + 获取规则列表查询表达式 :param name: 规则名称 :return: diff --git a/backend/app/admin/crud/crud_data_scope.py b/backend/app/admin/crud/crud_data_scope.py index de33f205..2383db82 100644 --- a/backend/app/admin/crud/crud_data_scope.py +++ b/backend/app/admin/crud/crud_data_scope.py @@ -50,9 +50,9 @@ class CRUDDataScope(CRUDPlus[DataScope]): """ return await self.select_models(db) - async def get_list(self, name: str | None, status: int | None) -> Select: + async def get_select(self, name: str | None, status: int | None) -> Select: """ - 获取数据范围列表 + 获取数据范围列表查询表达式 :param name: 范围名称 :param status: 范围状态 diff --git a/backend/app/admin/crud/crud_login_log.py b/backend/app/admin/crud/crud_login_log.py index 7d62aead..b658a76c 100644 --- a/backend/app/admin/crud/crud_login_log.py +++ b/backend/app/admin/crud/crud_login_log.py @@ -10,9 +10,9 @@ from backend.app.admin.schema.login_log import CreateLoginLogParam class CRUDLoginLog(CRUDPlus[LoginLog]): """登录日志数据库操作类""" - async def get_list(self, username: str | None, status: int | None, ip: str | None) -> Select: + async def get_select(self, username: str | None, status: int | None, ip: str | None) -> Select: """ - 获取登录日志列表 + 获取登录日志列表查询表达式 :param username: 用户名 :param status: 登录状态 diff --git a/backend/app/admin/crud/crud_opera_log.py b/backend/app/admin/crud/crud_opera_log.py index 222f52b7..900cc8cc 100644 --- a/backend/app/admin/crud/crud_opera_log.py +++ b/backend/app/admin/crud/crud_opera_log.py @@ -10,9 +10,9 @@ from backend.app.admin.schema.opera_log import CreateOperaLogParam class CRUDOperaLogDao(CRUDPlus[OperaLog]): """操作日志数据库操作类""" - async def get_list(self, username: str | None, status: int | None, ip: str | None) -> Select: + async def get_select(self, username: str | None, status: int | None, ip: str | None) -> Select: """ - 获取操作日志列表 + 获取操作日志列表查询表达式 :param username: 用户名 :param status: 操作状态 diff --git a/backend/app/admin/crud/crud_role.py b/backend/app/admin/crud/crud_role.py index 56377927..dd859262 100644 --- a/backend/app/admin/crud/crud_role.py +++ b/backend/app/admin/crud/crud_role.py @@ -45,9 +45,9 @@ class CRUDRole(CRUDPlus[Role]): """ return await self.select_models(db) - async def get_list(self, name: str | None, status: int | None) -> Select: + async def get_select(self, name: str | None, status: int | None) -> Select: """ - 获取角色列表 + 获取角色列表查询表达式 :param name: 角色名称 :param status: 角色状态 diff --git a/backend/app/admin/crud/crud_user.py b/backend/app/admin/crud/crud_user.py index 3ddff010..b75c8d04 100644 --- a/backend/app/admin/crud/crud_user.py +++ b/backend/app/admin/crud/crud_user.py @@ -181,9 +181,9 @@ class CRUDUser(CRUDPlus[User]): new_pwd = get_hash_password(password, salt) return await self.update_model(db, pk, {'password': new_pwd, 'salt': salt}) - async def get_list(self, dept: int | None, username: str | None, phone: str | None, status: int | None) -> Select: + async def get_select(self, dept: int | None, username: str | None, phone: str | None, status: int | None) -> Select: """ - 获取用户列表 + 获取用户列表查询表达式 :param dept: 部门 ID :param username: 用户名 diff --git a/backend/app/admin/service/auth_service.py b/backend/app/admin/service/auth_service.py index ccd2d345..7fba7f0b 100644 --- a/backend/app/admin/service/auth_service.py +++ b/backend/app/admin/service/auth_service.py @@ -9,6 +9,7 @@ from backend.app.admin.model import User from backend.app.admin.schema.token import GetLoginToken, GetNewToken from backend.app.admin.schema.user import AuthLoginParam from backend.app.admin.service.login_log_service import login_log_service +from backend.common.context import ctx from backend.common.enums import LoginLogStatusType from backend.common.exception import errors from backend.common.i18n import t @@ -23,7 +24,7 @@ from backend.common.security.jwt import ( password_verify, ) from backend.core.conf import settings -from backend.database.db import async_db_session, uuid4_str +from backend.database.db import uuid4_str from backend.database.redis import redis_client from backend.utils.timezone import timezone @@ -55,28 +56,28 @@ class AuthService: return user - async def swagger_login(self, *, obj: HTTPBasicCredentials) -> tuple[str, User]: + async def swagger_login(self, *, db: AsyncSession, obj: HTTPBasicCredentials) -> tuple[str, User]: """ Swagger 文档登录 + :param db: 数据库会话 :param obj: 登录凭证 :return: """ - async with async_db_session.begin() as db: - user = await self.user_verify(db, obj.username, obj.password) - await user_dao.update_login_time(db, obj.username) - access_token = await create_access_token( - user.id, - multi_login=user.is_multi_login, - # extra info - swagger=True, - ) - return access_token.access_token, user + user = await self.user_verify(db, obj.username, obj.password) + await user_dao.update_login_time(db, obj.username) + access_token = await create_access_token( + user.id, + multi_login=user.is_multi_login, + # extra info + swagger=True, + ) + return access_token.access_token, user async def login( self, *, - request: Request, + db: AsyncSession, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks, @@ -84,102 +85,100 @@ class AuthService: """ 用户登录 + :param db: 数据库会话 :param request: 请求对象 :param response: 响应对象 :param obj: 登录参数 :param background_tasks: 后台任务 :return: """ - async with async_db_session.begin() as db: - user = None - try: - user = await self.user_verify(db, obj.username, obj.password) - captcha_code = await redis_client.get(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{obj.uuid}') - if not captcha_code: - raise errors.RequestError(msg=t('error.captcha.expired')) - if captcha_code.lower() != obj.captcha.lower(): - raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) - await redis_client.delete(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{obj.uuid}') - await user_dao.update_login_time(db, obj.username) - await db.refresh(user) - access_token = await create_access_token( - user.id, - multi_login=user.is_multi_login, - # extra info - username=user.username, - nickname=user.nickname, - last_login_time=timezone.to_str(user.last_login_time), - ip=request.state.ip, - os=request.state.os, - browser=request.state.browser, - device=request.state.device, - ) - 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, - max_age=settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, - expires=timezone.to_utc(refresh_token.refresh_token_expire_time), - httponly=True, - ) - except errors.NotFoundError as e: - log.error('登陆错误: 用户名不存在') - raise errors.NotFoundError(msg=e.msg) - except (errors.RequestError, errors.CustomError) as e: - if not user: - log.error('登陆错误: 用户密码有误') - task = BackgroundTask( - login_log_service.create, - db=db, - request=request, - user_uuid=user.uuid if user else uuid4_str(), - username=obj.username, - 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 - else: - background_tasks.add_task( - login_log_service.create, - db=db, - request=request, - user_uuid=user.uuid, - username=obj.username, - login_time=timezone.now(), - status=LoginLogStatusType.success.value, - msg=t('success.login.success'), - ) - data = GetLoginToken( - access_token=access_token.access_token, - access_token_expire_time=access_token.access_token_expire_time, - session_uuid=access_token.session_uuid, - user=user, # type: ignore - ) - return data + user = None + try: + user = await self.user_verify(db, obj.username, obj.password) + captcha_code = await redis_client.get(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{obj.uuid}') + if not captcha_code: + raise errors.RequestError(msg=t('error.captcha.expired')) + if captcha_code.lower() != obj.captcha.lower(): + raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) + await redis_client.delete(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{obj.uuid}') + await user_dao.update_login_time(db, obj.username) + await db.refresh(user) + access_token = await create_access_token( + user.id, + multi_login=user.is_multi_login, + # extra info + username=user.username, + nickname=user.nickname, + last_login_time=timezone.to_str(user.last_login_time), + ip=ctx.ip, + os=ctx.os, + browser=ctx.browser, + device=ctx.device, + ) + 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, + max_age=settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, + expires=timezone.to_utc(refresh_token.refresh_token_expire_time), + httponly=True, + ) + except errors.NotFoundError as e: + log.error('登陆错误: 用户名不存在') + raise errors.NotFoundError(msg=e.msg) + except (errors.RequestError, errors.CustomError) as e: + if not user: + log.error('登陆错误: 用户密码有误') + task = BackgroundTask( + login_log_service.create, + db=db, + user_uuid=user.uuid if user else uuid4_str(), + username=obj.username, + 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 + else: + background_tasks.add_task( + login_log_service.create, + db=db, + user_uuid=user.uuid, + username=obj.username, + login_time=timezone.now(), + status=LoginLogStatusType.success.value, + msg=t('success.login.success'), + ) + data = GetLoginToken( + access_token=access_token.access_token, + access_token_expire_time=access_token.access_token_expire_time, + session_uuid=access_token.session_uuid, + user=user, # type: ignore + ) + return data @staticmethod - async def get_codes(*, request: Request) -> list[str]: + async def get_codes(*, db: AsyncSession, request: Request) -> list[str]: """ 获取用户权限码 + :param db: 数据库会话 :param request: FastAPI 请求对象 :return: """ codes = set() if request.user.is_superuser: - async with async_db_session.begin() as db: - menus = await menu_dao.get_all(db, None, None) - for menu in menus: - if menu.perms: - codes.add(*menu.perms.split(',')) + menus = await menu_dao.get_all(db, None, None) + for menu in menus: + if menu.perms: + codes.add(*menu.perms.split(',')) else: roles = request.user.roles if roles: @@ -191,10 +190,11 @@ class AuthService: return list(codes) @staticmethod - async def refresh_token(*, request: Request) -> GetNewToken: + async def refresh_token(*, db: AsyncSession, request: Request) -> GetNewToken: """ 刷新令牌 + :param db: 数据库会话 :param request: FastAPI 请求对象 :return: """ @@ -202,34 +202,34 @@ class AuthService: if not refresh_token: raise errors.RequestError(msg='Refresh Token 已过期,请重新登录') token_payload = jwt_decode(refresh_token) - async with async_db_session() as db: - user = await user_dao.get(db, token_payload.id) - if not user: - raise errors.NotFoundError(msg='用户不存在') - if not user.status: - raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员') - 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, - multi_login=user.is_multi_login, - # extra info - username=user.username, - nickname=user.nickname, - last_login_time=timezone.to_str(user.last_login_time), - ip=request.state.ip, - os=request.state.os, - browser=request.state.browser, - device_type=request.state.device, - ) - data = GetNewToken( - access_token=new_token.new_access_token, - access_token_expire_time=new_token.new_access_token_expire_time, - session_uuid=new_token.session_uuid, - ) - return data + + user = await user_dao.get(db, token_payload.id) + if not user: + raise errors.NotFoundError(msg='用户不存在') + if not user.status: + raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员') + 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, + multi_login=user.is_multi_login, + # extra info + username=user.username, + nickname=user.nickname, + last_login_time=timezone.to_str(user.last_login_time), + ip=ctx.ip, + os=ctx.os, + browser=ctx.browser, + device_type=ctx.device, + ) + data = GetNewToken( + access_token=new_token.new_access_token, + access_token_expire_time=new_token.new_access_token_expire_time, + session_uuid=new_token.session_uuid, + ) + return data @staticmethod async def logout(*, request: Request, response: Response) -> None: diff --git a/backend/app/admin/service/data_rule_service.py b/backend/app/admin/service/data_rule_service.py index f0fb8828..0c6aebc1 100644 --- a/backend/app/admin/service/data_rule_service.py +++ b/backend/app/admin/service/data_rule_service.py @@ -1,6 +1,7 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_data_rule import data_rule_dao from backend.app.admin.model import DataRule @@ -11,8 +12,8 @@ from backend.app.admin.schema.data_rule import ( UpdateDataRuleParam, ) from backend.common.exception import errors +from backend.common.pagination import paging_data from backend.core.conf import settings -from backend.database.db import async_db_session from backend.utils.import_parse import dynamic_import_data_model @@ -20,18 +21,19 @@ class DataRuleService: """数据规则服务类""" @staticmethod - async def get(*, pk: int) -> DataRule: + async def get(*, db: AsyncSession, pk: int) -> DataRule: """ 获取数据规则详情 + :param db: 数据库会话 :param pk: 规则 ID :return: """ - async with async_db_session() as db: - data_rule = await data_rule_dao.get(db, pk) - if not data_rule: - raise errors.NotFoundError(msg='数据规则不存在') - return data_rule + + data_rule = await data_rule_dao.get(db, pk) + if not data_rule: + raise errors.NotFoundError(msg='数据规则不存在') + return data_rule @staticmethod async def get_models() -> list[str]: @@ -58,65 +60,72 @@ class DataRuleService: return model_columns @staticmethod - async def get_select(*, name: str | None) -> Select: + async def get_list(*, db: AsyncSession, name: str | None) -> dict[str, Any]: """ - 获取数据规则列表查询条件 + 获取数据规则列表 + :param db: 数据库会话 :param name: 规则名称 :return: """ - return await data_rule_dao.get_list(name=name) + data_rule_select = await data_rule_dao.get_select(name=name) + return await paging_data(db, data_rule_select) @staticmethod - async def get_all() -> Sequence[DataRule]: - """获取所有数据规则""" - async with async_db_session() as db: - data_rules = await data_rule_dao.get_all(db) - return data_rules + async def get_all(*, db: AsyncSession) -> Sequence[DataRule]: + """ + 获取所有数据规则 + + :param db: 数据库会话 + :return: + """ + + data_rules = await data_rule_dao.get_all(db) + return data_rules @staticmethod - async def create(*, obj: CreateDataRuleParam) -> None: + async def create(*, db: AsyncSession, obj: CreateDataRuleParam) -> None: """ 创建数据规则 + :param db: 数据库会话 :param obj: 规则创建参数 :return: """ - async with async_db_session.begin() as db: - data_rule = await data_rule_dao.get_by_name(db, obj.name) - if data_rule: - raise errors.ConflictError(msg='数据规则已存在') - await data_rule_dao.create(db, obj) + data_rule = await data_rule_dao.get_by_name(db, obj.name) + if data_rule: + raise errors.ConflictError(msg='数据规则已存在') + await data_rule_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateDataRuleParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateDataRuleParam) -> int: """ 更新数据规则 + :param db: 数据库会话 :param pk: 规则 ID :param obj: 规则更新参数 :return: """ - async with async_db_session.begin() as db: - data_rule = await data_rule_dao.get(db, pk) - if not data_rule: - raise errors.NotFoundError(msg='数据规则不存在') - 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 + data_rule = await data_rule_dao.get(db, pk) + if not data_rule: + raise errors.NotFoundError(msg='数据规则不存在') + 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 @staticmethod - async def delete(*, obj: DeleteDataRuleParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteDataRuleParam) -> int: """ 批量删除数据规则 + :param db: 数据库会话 :param obj: 规则 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await data_rule_dao.delete(db, obj.pks) - return count + count = await data_rule_dao.delete(db, obj.pks) + return count data_rule_service: DataRuleService = DataRuleService() diff --git a/backend/app/admin/service/data_scope_service.py b/backend/app/admin/service/data_scope_service.py index 6a2382ba..54ad2ad7 100644 --- a/backend/app/admin/service/data_scope_service.py +++ b/backend/app/admin/service/data_scope_service.py @@ -1,6 +1,7 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_data_scope import data_scope_dao from backend.app.admin.model import DataScope @@ -11,8 +12,8 @@ from backend.app.admin.schema.data_scope import ( UpdateDataScopeRuleParam, ) from backend.common.exception import errors +from backend.common.pagination import paging_data from backend.core.conf import settings -from backend.database.db import async_db_session from backend.database.redis import redis_client @@ -20,88 +21,97 @@ class DataScopeService: """数据范围服务类""" @staticmethod - async def get(*, pk: int) -> DataScope: + async def get(*, db: AsyncSession, pk: int) -> DataScope: """ 获取数据范围详情 + :param db: 数据库会话 :param pk: 范围 ID :return: """ - async with async_db_session() as db: - data_scope = await data_scope_dao.get(db, pk) - if not data_scope: - raise errors.NotFoundError(msg='数据范围不存在') - return data_scope + + data_scope = await data_scope_dao.get(db, pk) + if not data_scope: + raise errors.NotFoundError(msg='数据范围不存在') + return data_scope @staticmethod - async def get_all() -> Sequence[DataScope]: - """获取所有数据范围""" - async with async_db_session() as db: - data_scopes = await data_scope_dao.get_all(db) - return data_scopes + async def get_all(*, db: AsyncSession) -> Sequence[DataScope]: + """ + 获取所有数据范围 + + :param db: 数据库会话 + :return: + """ + + data_scopes = await data_scope_dao.get_all(db) + return data_scopes @staticmethod - async def get_rules(*, pk: int) -> DataScope: + async def get_rules(*, db: AsyncSession, pk: int) -> DataScope: """ 获取数据范围规则 + :param db: 数据库会话 :param pk: 范围 ID :return: """ - async with async_db_session() as db: - data_scope = await data_scope_dao.get_with_relation(db, pk) - if not data_scope: - raise errors.NotFoundError(msg='数据范围不存在') - return data_scope + + data_scope = await data_scope_dao.get_with_relation(db, pk) + if not data_scope: + raise errors.NotFoundError(msg='数据范围不存在') + return data_scope @staticmethod - async def get_select(*, name: str | None, status: int | None) -> Select: + async def get_list(*, db: AsyncSession, name: str | None, status: int | None) -> dict[str, Any]: """ - 获取数据范围列表查询条件 + 获取数据范围列表 + :param db: 数据库会话 :param name: 范围名称 :param status: 范围状态 :return: """ - return await data_scope_dao.get_list(name, status) + data_scope_select = await data_scope_dao.get_select(name, status) + return await paging_data(db, data_scope_select) @staticmethod - async def create(*, obj: CreateDataScopeParam) -> None: + async def create(*, db: AsyncSession, obj: CreateDataScopeParam) -> None: """ 创建数据范围 + :param db: 数据库会话 :param obj: 数据范围参数 :return: """ - async with async_db_session.begin() as db: - data_scope = await data_scope_dao.get_by_name(db, obj.name) - if data_scope: - raise errors.ConflictError(msg='数据范围已存在') - await data_scope_dao.create(db, obj) + data_scope = await data_scope_dao.get_by_name(db, obj.name) + if data_scope: + raise errors.ConflictError(msg='数据范围已存在') + await data_scope_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateDataScopeParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateDataScopeParam) -> int: """ 更新数据范围 + :param db: 数据库会话 :param pk: 范围 ID :param obj: 数据范围更新参数 :return: """ - async with async_db_session.begin() as db: - data_scope = await data_scope_dao.get(db, pk) - if not data_scope: - raise errors.NotFoundError(msg='数据范围不存在') - 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: - for user in await role.awaitable_attrs.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + data_scope = await data_scope_dao.get(db, pk) + if not data_scope: + raise errors.NotFoundError(msg='数据范围不存在') + 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: + for user in await role.awaitable_attrs.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_data_scope_rule(*, pk: int, rule_ids: UpdateDataScopeRuleParam) -> int: + async def update_data_scope_rule(*, db: AsyncSession, pk: int, rule_ids: UpdateDataScopeRuleParam) -> int: """ 更新数据范围规则 @@ -109,27 +119,26 @@ class DataScopeService: :param rule_ids: 规则 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await data_scope_dao.update_rules(db, pk, rule_ids) - return count + count = await data_scope_dao.update_rules(db, pk, rule_ids) + return count @staticmethod - async def delete(*, obj: DeleteDataScopeParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteDataScopeParam) -> int: """ 批量删除数据范围 + :param db: 数据库会话 :param obj: 范围 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await data_scope_dao.delete(db, obj.pks) - for pk in obj.pks: - data_rule = await data_scope_dao.get(db, pk) - if data_rule: - for role in await data_rule.awaitable_attrs.roles: - for user in await role.awaitable_attrs.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + count = await data_scope_dao.delete(db, obj.pks) + for pk in obj.pks: + data_rule = await data_scope_dao.get(db, pk) + if data_rule: + for role in await data_rule.awaitable_attrs.roles: + for user in await role.awaitable_attrs.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count data_scope_service: DataScopeService = DataScopeService() diff --git a/backend/app/admin/service/dept_service.py b/backend/app/admin/service/dept_service.py index 32c79629..d75cd687 100644 --- a/backend/app/admin/service/dept_service.py +++ b/backend/app/admin/service/dept_service.py @@ -1,13 +1,13 @@ from typing import Any from fastapi import Request +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_dept import dept_dao from backend.app.admin.model import Dept from backend.app.admin.schema.dept import CreateDeptParam, UpdateDeptParam from backend.common.exception import errors from backend.core.conf import settings -from backend.database.db import async_db_session from backend.database.redis import redis_client from backend.utils.build_tree import get_tree_data @@ -16,22 +16,24 @@ class DeptService: """部门服务类""" @staticmethod - async def get(*, pk: int) -> Dept: + async def get(*, db: AsyncSession, pk: int) -> Dept: """ 获取部门详情 + :param db: 数据库会话 :param pk: 部门 ID :return: """ - async with async_db_session() as db: - dept = await dept_dao.get(db, pk) - if not dept: - raise errors.NotFoundError(msg='部门不存在') - return dept + + dept = await dept_dao.get(db, pk) + if not dept: + raise errors.NotFoundError(msg='部门不存在') + return dept @staticmethod async def get_tree( *, + db: AsyncSession, request: Request, name: str | None, leader: str | None, @@ -41,6 +43,7 @@ class DeptService: """ 获取部门树形结构 + :param db: 数据库会话 :param request: FastAPI 请求对象 :param name: 部门名称 :param leader: 部门负责人 @@ -48,72 +51,72 @@ class DeptService: :param status: 状态 :return: """ - async with async_db_session() as db: - dept_select = await dept_dao.get_all(request, db, name, leader, phone, status) - tree_data = get_tree_data(dept_select) - return tree_data + + dept_select = await dept_dao.get_all(request, db, name, leader, phone, status) + tree_data = get_tree_data(dept_select) + return tree_data @staticmethod - async def create(*, obj: CreateDeptParam) -> None: + async def create(*, db: AsyncSession, obj: CreateDeptParam) -> None: """ 创建部门 + :param db: 数据库会话 :param obj: 部门创建参数 :return: """ - async with async_db_session.begin() as db: - dept = await dept_dao.get_by_name(db, obj.name) - if dept: - raise errors.ConflictError(msg='部门名称已存在') - if obj.parent_id: - parent_dept = await dept_dao.get(db, obj.parent_id) - if not parent_dept: - raise errors.NotFoundError(msg='父级部门不存在') - await dept_dao.create(db, obj) + dept = await dept_dao.get_by_name(db, obj.name) + if dept: + raise errors.ConflictError(msg='部门名称已存在') + if obj.parent_id: + parent_dept = await dept_dao.get(db, obj.parent_id) + if not parent_dept: + raise errors.NotFoundError(msg='父级部门不存在') + await dept_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateDeptParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateDeptParam) -> int: """ 更新部门 + :param db: 数据库会话 :param pk: 部门 ID :param obj: 部门更新参数 :return: """ - async with async_db_session.begin() as db: - dept = await dept_dao.get(db, pk) - if not dept: - raise errors.NotFoundError(msg='部门不存在') - 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) - if not parent_dept: - raise errors.NotFoundError(msg='父级部门不存在') - if obj.parent_id == dept.id: - raise errors.ForbiddenError(msg='禁止关联自身为父级') - count = await dept_dao.update(db, pk, obj) - return count + dept = await dept_dao.get(db, pk) + if not dept: + raise errors.NotFoundError(msg='部门不存在') + 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) + if not parent_dept: + raise errors.NotFoundError(msg='父级部门不存在') + if obj.parent_id == dept.id: + raise errors.ForbiddenError(msg='禁止关联自身为父级') + count = await dept_dao.update(db, pk, obj) + return count @staticmethod - async def delete(*, pk: int) -> int: + async def delete(*, db: AsyncSession, pk: int) -> int: """ 删除部门 + :param db: 数据库会话 :param pk: 部门 ID :return: """ - async with async_db_session.begin() as db: - dept = await dept_dao.get_with_relation(db, pk) - if dept.users: - raise errors.ConflictError(msg='部门下存在用户,无法删除') - children = await dept_dao.get_children(db, pk) - if children: - raise errors.ConflictError(msg='部门下存在子部门,无法删除') - count = await dept_dao.delete(db, pk) - for user in dept.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + dept = await dept_dao.get_with_relation(db, pk) + if dept.users: + raise errors.ConflictError(msg='部门下存在用户,无法删除') + children = await dept_dao.get_children(db, pk) + if children: + raise errors.ConflictError(msg='部门下存在子部门,无法删除') + count = await dept_dao.delete(db, pk) + for user in dept.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count dept_service: DeptService = DeptService() diff --git a/backend/app/admin/service/login_log_service.py b/backend/app/admin/service/login_log_service.py index 811b397e..21db54d3 100644 --- a/backend/app/admin/service/login_log_service.py +++ b/backend/app/admin/service/login_log_service.py @@ -1,36 +1,36 @@ from datetime import datetime +from typing import Any -from fastapi import Request -from sqlalchemy import Select from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_login_log import login_log_dao from backend.app.admin.schema.login_log import CreateLoginLogParam, DeleteLoginLogParam from backend.common.context import ctx from backend.common.log import log -from backend.database.db import async_db_session +from backend.common.pagination import paging_data class LoginLogService: """登录日志服务类""" @staticmethod - async def get_select(*, username: str | None, status: int | None, ip: str | None) -> Select: + async def get_list(*, db: AsyncSession, username: str | None, status: int | None, ip: str | None) -> dict[str, Any]: """ - 获取登录日志列表查询条件 + 获取登录日志列表 + :param db: 数据库会话 :param username: 用户名 :param status: 状态 :param ip: IP 地址 :return: """ - return await login_log_dao.get_list(username=username, status=status, ip=ip) + log_select = await login_log_dao.get_select(username=username, status=status, ip=ip) + return await paging_data(db, log_select) @staticmethod async def create( *, db: AsyncSession, - request: Request, user_uuid: str, username: str, login_time: datetime, @@ -41,7 +41,6 @@ class LoginLogService: 创建登录日志 :param db: 数据库会话 - :param request: FastAPI 请求对象 :param user_uuid: 用户 UUID :param username: 用户名 :param login_time: 登录时间 @@ -70,22 +69,21 @@ class LoginLogService: log.error(f'登录日志创建失败: {e}') @staticmethod - async def delete(*, obj: DeleteLoginLogParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteLoginLogParam) -> int: """ 批量删除登录日志 + :param db: 数据库会话 :param obj: 日志 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await login_log_dao.delete(db, obj.pks) - return count + count = await login_log_dao.delete(db, obj.pks) + return count @staticmethod - async def delete_all() -> None: + async def delete_all(*, db: AsyncSession) -> None: """清空所有登录日志""" - async with async_db_session.begin() as db: - await login_log_dao.delete_all(db) + await login_log_dao.delete_all(db) login_log_service: LoginLogService = LoginLogService() diff --git a/backend/app/admin/service/menu_service.py b/backend/app/admin/service/menu_service.py index b61e4238..9ce705a8 100644 --- a/backend/app/admin/service/menu_service.py +++ b/backend/app/admin/service/menu_service.py @@ -1,13 +1,13 @@ from typing import Any from fastapi import Request +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_menu import menu_dao from backend.app.admin.model import Menu from backend.app.admin.schema.menu import CreateMenuParam, UpdateMenuParam from backend.common.exception import errors from backend.core.conf import settings -from backend.database.db import async_db_session from backend.database.redis import redis_client from backend.utils.build_tree import get_tree_data, get_vben5_tree_data @@ -16,118 +16,124 @@ class MenuService: """菜单服务类""" @staticmethod - async def get(*, pk: int) -> Menu: + async def get(*, db: AsyncSession, pk: int) -> Menu: """ 获取菜单详情 + :param db: 数据库会话 :param pk: 菜单 ID :return: """ - async with async_db_session() as db: - menu = await menu_dao.get(db, menu_id=pk) - if not menu: - raise errors.NotFoundError(msg='菜单不存在') - return menu + + menu = await menu_dao.get(db, menu_id=pk) + if not menu: + raise errors.NotFoundError(msg='菜单不存在') + return menu @staticmethod - async def get_tree(*, title: str | None, status: int | None) -> list[dict[str, Any]]: + async def get_tree(*, db: AsyncSession, title: str | None, status: int | None) -> list[dict[str, Any]]: """ 获取菜单树形结构 + :param db: 数据库会话 :param title: 菜单标题 :param status: 状态 :return: """ - async with async_db_session() as db: - menu_data = await menu_dao.get_all(db, title=title, status=status) - menu_tree = get_tree_data(menu_data) - return menu_tree + + menu_data = await menu_dao.get_all(db, title=title, status=status) + menu_tree = get_tree_data(menu_data) + return menu_tree @staticmethod - async def get_sidebar(*, request: Request) -> list[dict[str, Any] | None]: + async def get_sidebar(*, db: AsyncSession, request: Request) -> list[dict[str, Any] | None]: """ 获取用户的菜单侧边栏 + :param db: 数据库会话 :param request: FastAPI 请求对象 :return: """ - async with async_db_session() as db: - if request.user.is_superuser: - menu_data = await menu_dao.get_sidebar(db, None) - else: - roles = request.user.roles - menu_ids = set() - if roles: - for role in roles: - 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 + + if request.user.is_superuser: + menu_data = await menu_dao.get_sidebar(db, None) + else: + roles = request.user.roles + menu_ids = set() + if roles: + for role in roles: + 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 @staticmethod - async def create(*, obj: CreateMenuParam) -> None: + async def create(*, db: AsyncSession, obj: CreateMenuParam) -> None: """ 创建菜单 + :param db: 数据库会话 :param obj: 菜单创建参数 :return: """ - async with async_db_session.begin() as db: - title = await menu_dao.get_by_title(db, obj.title) - if title: - raise errors.ConflictError(msg='菜单标题已存在') - if obj.parent_id: - parent_menu = await menu_dao.get(db, obj.parent_id) - if not parent_menu: - raise errors.NotFoundError(msg='父级菜单不存在') - await menu_dao.create(db, obj) + + title = await menu_dao.get_by_title(db, obj.title) + if title: + raise errors.ConflictError(msg='菜单标题已存在') + if obj.parent_id: + parent_menu = await menu_dao.get(db, obj.parent_id) + if not parent_menu: + raise errors.NotFoundError(msg='父级菜单不存在') + await menu_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateMenuParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateMenuParam) -> int: """ 更新菜单 + :param db: 数据库会话 :param pk: 菜单 ID :param obj: 菜单更新参数 :return: """ - async with async_db_session.begin() as db: - menu = await menu_dao.get(db, pk) - if not menu: - raise errors.NotFoundError(msg='菜单不存在') - 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) - if not parent_menu: - raise errors.NotFoundError(msg='父级菜单不存在') - if obj.parent_id == menu.id: - raise errors.ForbiddenError(msg='禁止关联自身为父级') - count = await menu_dao.update(db, pk, obj) - for role in await menu.awaitable_attrs.roles: - for user in await role.awaitable_attrs.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + + menu = await menu_dao.get(db, pk) + if not menu: + raise errors.NotFoundError(msg='菜单不存在') + 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) + if not parent_menu: + raise errors.NotFoundError(msg='父级菜单不存在') + if obj.parent_id == menu.id: + raise errors.ForbiddenError(msg='禁止关联自身为父级') + count = await menu_dao.update(db, pk, obj) + for role in await menu.awaitable_attrs.roles: + for user in await role.awaitable_attrs.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def delete(*, pk: int) -> int: + async def delete(*, db: AsyncSession, pk: int) -> int: """ 删除菜单 + :param db: 数据库会话 :param pk: 菜单 ID :return: """ - async with async_db_session.begin() as db: - children = await menu_dao.get_children(db, pk) - if children: - raise errors.ConflictError(msg='菜单下存在子菜单,无法删除') - menu = await menu_dao.get(db, pk) - count = await menu_dao.delete(db, pk) - if menu: - for role in await menu.awaitable_attrs.roles: - for user in await role.awaitable_attrs.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + + children = await menu_dao.get_children(db, pk) + if children: + raise errors.ConflictError(msg='菜单下存在子菜单,无法删除') + menu = await menu_dao.get(db, pk) + count = await menu_dao.delete(db, pk) + if menu: + for role in await menu.awaitable_attrs.roles: + for user in await role.awaitable_attrs.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count menu_service: MenuService = MenuService() diff --git a/backend/app/admin/service/opera_log_service.py b/backend/app/admin/service/opera_log_service.py index d8cfcb0a..af92be37 100644 --- a/backend/app/admin/service/opera_log_service.py +++ b/backend/app/admin/service/opera_log_service.py @@ -1,64 +1,72 @@ -from sqlalchemy import Select +from typing import Any + +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_opera_log import opera_log_dao from backend.app.admin.schema.opera_log import CreateOperaLogParam, DeleteOperaLogParam -from backend.database.db import async_db_session +from backend.common.pagination import paging_data class OperaLogService: """操作日志服务类""" @staticmethod - async def get_select(*, username: str | None, status: int | None, ip: str | None) -> Select: + async def get_list(*, db: AsyncSession, username: str | None, status: int | None, ip: str | None) -> dict[str, Any]: """ - 获取操作日志列表查询条件 + 获取操作日志列表 + :param db: 数据库会话 :param username: 用户名 :param status: 状态 :param ip: IP 地址 :return: """ - return await opera_log_dao.get_list(username=username, status=status, ip=ip) + log_select = await opera_log_dao.get_select(username=username, status=status, ip=ip) + return await paging_data(db, log_select) @staticmethod - async def create(*, obj: CreateOperaLogParam) -> None: + async def create(*, db: AsyncSession, obj: CreateOperaLogParam) -> None: """ 创建操作日志 + :param db: 数据库会话 :param obj: 操作日志创建参数 :return: """ - async with async_db_session.begin() as db: - await opera_log_dao.create(db, obj) + await opera_log_dao.create(db, obj) @staticmethod - async def bulk_create(*, objs: list[CreateOperaLogParam]) -> None: + async def bulk_create(*, db: AsyncSession, objs: list[CreateOperaLogParam]) -> None: """ 批量创建操作日志 + :param db: 数据库会话 :param objs: 操作日志创建参数列表 :return: """ - async with async_db_session.begin() as db: - await opera_log_dao.bulk_create(db, objs) + await opera_log_dao.bulk_create(db, objs) @staticmethod - async def delete(*, obj: DeleteOperaLogParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteOperaLogParam) -> int: """ 批量删除操作日志 + :param db: 数据库会话 :param obj: 日志 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await opera_log_dao.delete(db, obj.pks) - return count + count = await opera_log_dao.delete(db, obj.pks) + return count @staticmethod - async def delete_all() -> None: - """清空所有操作日志""" - async with async_db_session.begin() as db: - await opera_log_dao.delete_all(db) + async def delete_all(*, db: AsyncSession) -> None: + """ + 清空所有操作日志 + + :param db: 数据库会话 + :return: + """ + await opera_log_dao.delete_all(db) opera_log_service: OperaLogService = OperaLogService() diff --git a/backend/app/admin/service/role_service.py b/backend/app/admin/service/role_service.py index 1cab676f..ab750c79 100644 --- a/backend/app/admin/service/role_service.py +++ b/backend/app/admin/service/role_service.py @@ -1,7 +1,7 @@ from collections.abc import Sequence from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_data_scope import data_scope_dao from backend.app.admin.crud.crud_menu import menu_dao @@ -15,8 +15,8 @@ from backend.app.admin.schema.role import ( UpdateRoleScopeParam, ) from backend.common.exception import errors +from backend.common.pagination import paging_data from backend.core.conf import settings -from backend.database.db import async_db_session from backend.database.redis import redis_client from backend.utils.build_tree import get_tree_data @@ -25,161 +25,176 @@ class RoleService: """角色服务类""" @staticmethod - async def get(*, pk: int) -> Role: + async def get(*, db: AsyncSession, pk: int) -> Role: """ 获取角色详情 + :param db: 数据库会话 :param pk: 角色 ID :return: """ - async with async_db_session() as db: - role = await role_dao.get_with_relation(db, pk) - if not role: - raise errors.NotFoundError(msg='角色不存在') - return role + + role = await role_dao.get_with_relation(db, pk) + if not role: + raise errors.NotFoundError(msg='角色不存在') + return role @staticmethod - async def get_all() -> Sequence[Role]: - """获取所有角色""" - async with async_db_session() as db: - roles = await role_dao.get_all(db) - return roles - - @staticmethod - async def get_select(*, name: str | None, status: int | None) -> Select: + async def get_all(*, db: AsyncSession) -> Sequence[Role]: """ - 获取角色列表查询条件 + 获取所有角色 + :param db: 数据库会话 + :return: + """ + + roles = await role_dao.get_all(db) + return roles + + @staticmethod + async def get_list(*, db: AsyncSession, name: str | None, status: int | None) -> dict[str, Any]: + """ + 获取角色列表 + + :param db: 数据库会话 :param name: 角色名称 :param status: 状态 :return: """ - return await role_dao.get_list(name=name, status=status) + role_select = await role_dao.get_select(name=name, status=status) + return await paging_data(db, role_select) @staticmethod - async def get_menu_tree(*, pk: int) -> list[dict[str, Any] | None]: + async def get_menu_tree(*, db: AsyncSession, pk: int) -> list[dict[str, Any] | None]: """ 获取角色的菜单树形结构 + :param db: 数据库会话 :param pk: 角色 ID :return: """ - async with async_db_session() as db: - role = await role_dao.get_with_relation(db, pk) - if not role: - raise errors.NotFoundError(msg='角色不存在') - menu_tree = get_tree_data(role.menus) if role.menus else [] - return menu_tree + + role = await role_dao.get_with_relation(db, pk) + if not role: + raise errors.NotFoundError(msg='角色不存在') + menu_tree = get_tree_data(role.menus) if role.menus else [] + return menu_tree @staticmethod - async def get_scopes(*, pk: int) -> list[int]: + async def get_scopes(*, db: AsyncSession, pk: int) -> list[int]: """ 获取角色数据范围列表 + :param db: 数据库会话 :param pk: :return: """ - async with async_db_session() as db: - role = await role_dao.get_with_relation(db, pk) - if not role: - raise errors.NotFoundError(msg='角色不存在') - scope_ids = [scope.id for scope in role.scopes] - return scope_ids + + role = await role_dao.get_with_relation(db, pk) + if not role: + raise errors.NotFoundError(msg='角色不存在') + scope_ids = [scope.id for scope in role.scopes] + return scope_ids @staticmethod - async def create(*, obj: CreateRoleParam) -> None: + async def create(*, db: AsyncSession, obj: CreateRoleParam) -> None: """ 创建角色 + :param db: 数据库会话 :param obj: 角色创建参数 :return: """ - async with async_db_session.begin() as db: - role = await role_dao.get_by_name(db, obj.name) - if role: - raise errors.ConflictError(msg='角色已存在') - await role_dao.create(db, obj) + + role = await role_dao.get_by_name(db, obj.name) + if role: + raise errors.ConflictError(msg='角色已存在') + await role_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateRoleParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateRoleParam) -> int: """ 更新角色 + :param db: 数据库会话 :param pk: 角色 ID :param obj: 角色更新参数 :return: """ - async with async_db_session.begin() as db: - role = await role_dao.get(db, pk) - if not role: - raise errors.NotFoundError(msg='角色不存在') - 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: - await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + + role = await role_dao.get(db, pk) + if not role: + raise errors.NotFoundError(msg='角色不存在') + 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: + await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_role_menu(*, pk: int, menu_ids: UpdateRoleMenuParam) -> int: + async def update_role_menu(*, db: AsyncSession, pk: int, menu_ids: UpdateRoleMenuParam) -> int: """ 更新角色菜单 + :param db: 数据库会话 :param pk: 角色 ID :param menu_ids: 菜单 ID 列表 :return: """ - async with async_db_session.begin() as db: - role = await role_dao.get(db, pk) - if not role: - raise errors.NotFoundError(msg='角色不存在') - for menu_id in menu_ids.menus: - menu = await menu_dao.get(db, menu_id) - if not menu: - raise errors.NotFoundError(msg='菜单不存在') - count = await role_dao.update_menus(db, pk, menu_ids) - for user in await role.awaitable_attrs.users: - await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + + role = await role_dao.get(db, pk) + if not role: + raise errors.NotFoundError(msg='角色不存在') + for menu_id in menu_ids.menus: + menu = await menu_dao.get(db, menu_id) + if not menu: + raise errors.NotFoundError(msg='菜单不存在') + count = await role_dao.update_menus(db, pk, menu_ids) + for user in await role.awaitable_attrs.users: + await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_role_scope(*, pk: int, scope_ids: UpdateRoleScopeParam) -> int: + async def update_role_scope(*, db: AsyncSession, pk: int, scope_ids: UpdateRoleScopeParam) -> int: """ 更新角色数据范围 + :param db: 数据库会话 :param pk: 角色 ID :param scope_ids: 权限规则 ID 列表 :return: """ - async with async_db_session.begin() as db: - role = await role_dao.get(db, pk) - if not role: - raise errors.NotFoundError(msg='角色不存在') - for scope_id in scope_ids.scopes: - scope = await data_scope_dao.get(db, scope_id) - if not scope: - raise errors.NotFoundError(msg='数据范围不存在') - count = await role_dao.update_scopes(db, pk, scope_ids) - for user in await role.awaitable_attrs.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + + role = await role_dao.get(db, pk) + if not role: + raise errors.NotFoundError(msg='角色不存在') + for scope_id in scope_ids.scopes: + scope = await data_scope_dao.get(db, scope_id) + if not scope: + raise errors.NotFoundError(msg='数据范围不存在') + count = await role_dao.update_scopes(db, pk, scope_ids) + for user in await role.awaitable_attrs.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def delete(*, obj: DeleteRoleParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteRoleParam) -> int: """ 批量删除角色 + :param db: 数据库会话 :param obj: 角色 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await role_dao.delete(db, obj.pks) - for pk in obj.pks: - role = await role_dao.get(db, pk) - if role: - for user in await role.awaitable_attrs.users: - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + + count = await role_dao.delete(db, obj.pks) + for pk in obj.pks: + role = await role_dao.get(db, pk) + if role: + for user in await role.awaitable_attrs.users: + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count role_service: RoleService = RoleService() diff --git a/backend/app/admin/service/user_service.py b/backend/app/admin/service/user_service.py index af02c310..f214220e 100644 --- a/backend/app/admin/service/user_service.py +++ b/backend/app/admin/service/user_service.py @@ -1,9 +1,10 @@ import random from collections.abc import Sequence +from typing import Any from fastapi import Request -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_dept import dept_dao from backend.app.admin.crud.crud_role import role_dao @@ -17,10 +18,10 @@ from backend.app.admin.schema.user import ( from backend.common.context import ctx from backend.common.enums import UserPermissionType from backend.common.exception import errors +from backend.common.pagination import paging_data from backend.common.response.response_code import CustomErrorCode from backend.common.security.jwt import get_token, jwt_decode, password_verify from backend.core.conf import settings -from backend.database.db import async_db_session from backend.database.redis import redis_client @@ -28,287 +29,289 @@ class UserService: """用户服务类""" @staticmethod - async def get_userinfo(*, pk: int | None = None, username: str | None = None) -> User: + async def get_userinfo(*, db: AsyncSession, pk: int | None = None, username: str | None = None) -> User: """ 获取用户信息 + :param db: 数据库会话 :param pk: 用户 ID :param username: 用户名 :return: """ - async with async_db_session() as db: - user = await user_dao.get_with_relation(db, user_id=pk, username=username) - if not user: - raise errors.NotFoundError(msg='用户不存在') - return user + user = await user_dao.get_with_relation(db, user_id=pk, username=username) + if not user: + raise errors.NotFoundError(msg='用户不存在') + return user @staticmethod - async def get_roles(*, pk: int) -> Sequence[Role]: + async def get_roles(*, db: AsyncSession, pk: int) -> Sequence[Role]: """ 获取用户所有角色 + :param db: 数据库会话 :param pk: 用户 ID :return: """ - async with async_db_session() as db: - user = await user_dao.get_with_relation(db, user_id=pk) - if not user: - raise errors.NotFoundError(msg='用户不存在') - return user.roles + user = await user_dao.get_with_relation(db, user_id=pk) + if not user: + raise errors.NotFoundError(msg='用户不存在') + return user.roles @staticmethod - async def get_select(*, dept: int, username: str, phone: str, status: int) -> Select: + async def get_list(*, db: AsyncSession, dept: int, username: str, phone: str, status: int) -> dict[str, Any]: """ - 获取用户列表查询条件 + 获取用户列表 + :param db: 数据库会话 :param dept: 部门 ID :param username: 用户名 :param phone: 手机号 :param status: 状态 :return: """ - return await user_dao.get_list(dept=dept, username=username, phone=phone, status=status) + user_select = await user_dao.get_select(dept=dept, username=username, phone=phone, status=status) + return await paging_data(db, user_select) @staticmethod - async def create(*, obj: AddUserParam) -> None: + async def create(*, db: AsyncSession, obj: AddUserParam) -> None: """ 创建用户 + :param db: 数据库会话 :param obj: 用户添加参数 :return: """ - async with async_db_session.begin() as db: - if await user_dao.get_by_username(db, obj.username): - raise errors.ConflictError(msg='用户名已注册') - 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): - raise errors.NotFoundError(msg='部门不存在') - for role_id in obj.roles: - if not await role_dao.get(db, role_id): - raise errors.NotFoundError(msg='角色不存在') - await user_dao.add(db, obj) + if await user_dao.get_by_username(db, obj.username): + raise errors.ConflictError(msg='用户名已注册') + 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): + raise errors.NotFoundError(msg='部门不存在') + for role_id in obj.roles: + if not await role_dao.get(db, role_id): + raise errors.NotFoundError(msg='角色不存在') + await user_dao.add(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateUserParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateUserParam) -> int: """ 更新用户信息 + :param db: 数据库会话 :param pk: 用户 ID :param obj: 用户更新参数 :return: """ - async with async_db_session.begin() as db: - user = await user_dao.get_with_relation(db, user_id=pk) - if not user: - raise errors.NotFoundError(msg='用户不存在') - 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): - raise errors.NotFoundError(msg='角色不存在') - count = await user_dao.update(db, user, obj) - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + user = await user_dao.get_with_relation(db, user_id=pk) + if not user: + raise errors.NotFoundError(msg='用户不存在') + 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): + raise errors.NotFoundError(msg='角色不存在') + count = await user_dao.update(db, user, obj) + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_permission(*, request: Request, pk: int, type: UserPermissionType) -> int: # noqa: C901 + async def update_permission(*, db: AsyncSession, request: Request, pk: int, type: UserPermissionType) -> int: # noqa: C901 """ 更新用户权限 + :param db: 数据库会话 :param request: FastAPI 请求对象 :param pk: 用户 ID :param type: 权限类型 :return: """ - async with async_db_session.begin() as db: - match type: - case UserPermissionType.superuser: - 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_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, is_staff=not user.is_staff) - case UserPermissionType.status: - 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_status(db, pk, 0 if user.status == 1 else 1) - case UserPermissionType.multi_login: - user = await user_dao.get(db, pk) - if not user: - 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, multi_login=new_multi_login) - token = get_token(request) - token_payload = jwt_decode(token) - if pk == user.id: - # 系统管理员修改自身时,除当前 token 外,其他 token 失效 - 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}', - ) - else: - # 系统管理员修改他人时,他人 token 全部失效 - if not new_multi_login: - key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{user.id}' - await redis_client.delete_prefix(key_prefix) - case _: - raise errors.RequestError(msg='权限类型不存在') + match type: + case UserPermissionType.superuser: + 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_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, is_staff=not user.is_staff) + case UserPermissionType.status: + 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_status(db, pk, 0 if user.status == 1 else 1) + case UserPermissionType.multi_login: + user = await user_dao.get(db, pk) + if not user: + 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, multi_login=new_multi_login) + token = get_token(request) + token_payload = jwt_decode(token) + if pk == user.id: + # 系统管理员修改自身时,除当前 token 外,其他 token 失效 + 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}', + ) + else: + # 系统管理员修改他人时,他人 token 全部失效 + if not new_multi_login: + key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{user.id}' + await redis_client.delete_prefix(key_prefix) + case _: + raise errors.RequestError(msg='权限类型不存在') await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') return count @staticmethod - async def reset_password(*, pk: int, password: str) -> int: + async def reset_password(*, db: AsyncSession, pk: int, password: str) -> int: """ 重置用户密码 + :param db: 数据库会话 :param pk: 用户 ID :param password: 新密码 :return: """ - async with async_db_session.begin() as db: - user = await user_dao.get(db, pk) - if not user: - raise errors.NotFoundError(msg='用户不存在') - count = await user_dao.reset_password(db, user.id, password) - key_prefix = [ - f'{settings.TOKEN_REDIS_PREFIX}:{user.id}', - f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}', - f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}', - ] - for prefix in key_prefix: - await redis_client.delete(prefix) - return count + user = await user_dao.get(db, pk) + if not user: + raise errors.NotFoundError(msg='用户不存在') + count = await user_dao.reset_password(db, user.id, password) + key_prefix = [ + f'{settings.TOKEN_REDIS_PREFIX}:{user.id}', + f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}', + f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}', + ] + for prefix in key_prefix: + await redis_client.delete(prefix) + return count @staticmethod - async def update_nickname(*, request: Request, nickname: str) -> int: + async def update_nickname(*, db: AsyncSession, request: Request, nickname: str) -> int: """ 更新当前用户昵称 + :param db: 数据库会话 :param request: FastAPI 请求对象 :param nickname: 用户昵称 :return: """ - async with async_db_session.begin() as db: - token = get_token(request) - token_payload = jwt_decode(token) - user = await user_dao.get(db, token_payload.id) - if not user: - raise errors.NotFoundError(msg='用户不存在') - count = await user_dao.update_nickname(db, token_payload.id, nickname) - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + token = get_token(request) + token_payload = jwt_decode(token) + user = await user_dao.get(db, token_payload.id) + if not user: + raise errors.NotFoundError(msg='用户不存在') + count = await user_dao.update_nickname(db, token_payload.id, nickname) + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_avatar(*, request: Request, avatar: str) -> int: + async def update_avatar(*, db: AsyncSession, request: Request, avatar: str) -> int: """ 更新当前用户头像 + :param db: 数据库会话 :param request: FastAPI 请求对象 :param avatar: 头像地址 :return: """ - async with async_db_session.begin() as db: - token = get_token(request) - token_payload = jwt_decode(token) - user = await user_dao.get(db, token_payload.id) - if not user: - raise errors.NotFoundError(msg='用户不存在') - count = await user_dao.update_avatar(db, token_payload.id, avatar) - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + token = get_token(request) + token_payload = jwt_decode(token) + user = await user_dao.get(db, token_payload.id) + if not user: + raise errors.NotFoundError(msg='用户不存在') + count = await user_dao.update_avatar(db, token_payload.id, avatar) + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_email(*, request: Request, captcha: str, email: str) -> int: + async def update_email(*, db: AsyncSession, request: Request, captcha: str, email: str) -> int: """ 更新当前用户邮箱 + :param db: 数据库会话 :param request: FastAPI 请求对象 :param captcha: 邮箱验证码 :param email: 邮箱 :return: """ - async with async_db_session.begin() as db: - token = get_token(request) - token_payload = jwt_decode(token) - user = await user_dao.get(db, token_payload.id) - if not user: - raise errors.NotFoundError(msg='用户不存在') - captcha_code = await redis_client.get(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ctx.ip}') - if not captcha_code: - raise errors.RequestError(msg='验证码已失效,请重新获取') - if captcha != captcha_code: - raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) - await redis_client.delete(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ctx.ip}') - count = await user_dao.update_email(db, token_payload.id, email) - await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') - return count + token = get_token(request) + token_payload = jwt_decode(token) + user = await user_dao.get(db, token_payload.id) + if not user: + raise errors.NotFoundError(msg='用户不存在') + captcha_code = await redis_client.get(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ctx.ip}') + if not captcha_code: + raise errors.RequestError(msg='验证码已失效,请重新获取') + if captcha != captcha_code: + raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) + await redis_client.delete(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ctx.ip}') + count = await user_dao.update_email(db, token_payload.id, email) + await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') + return count @staticmethod - async def update_password(*, request: Request, obj: ResetPasswordParam) -> int: + async def update_password(*, db: AsyncSession, request: Request, obj: ResetPasswordParam) -> int: """ 更新当前用户密码 + :param db: 数据库会话 :param request: FastAPI 请求对象 :param obj: 密码重置参数 :return: """ - async with async_db_session.begin() as db: - token = get_token(request) - token_payload = jwt_decode(token) - user = await user_dao.get(db, token_payload.id) - if not user: - raise errors.NotFoundError(msg='用户不存在') - if not password_verify(obj.old_password, user.password): - raise errors.RequestError(msg='原密码错误') - if obj.new_password != obj.confirm_password: - raise errors.RequestError(msg='密码输入不一致') - count = await user_dao.reset_password(db, user.id, obj.new_password) - key_prefix = [ - f'{settings.TOKEN_REDIS_PREFIX}:{user.id}', - f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}', - f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}', - ] - for prefix in key_prefix: - await redis_client.delete_prefix(prefix) - return count + token = get_token(request) + token_payload = jwt_decode(token) + user = await user_dao.get(db, token_payload.id) + if not user: + raise errors.NotFoundError(msg='用户不存在') + if not password_verify(obj.old_password, user.password): + raise errors.RequestError(msg='原密码错误') + if obj.new_password != obj.confirm_password: + raise errors.RequestError(msg='密码输入不一致') + count = await user_dao.reset_password(db, user.id, obj.new_password) + key_prefix = [ + f'{settings.TOKEN_REDIS_PREFIX}:{user.id}', + f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}', + f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}', + ] + for prefix in key_prefix: + await redis_client.delete_prefix(prefix) + return count @staticmethod - async def delete(*, pk: int) -> int: + async def delete(*, db: AsyncSession, pk: int) -> int: """ 删除用户 + :param db: 数据库会话 :param pk: 用户 ID :return: """ - async with async_db_session.begin() as db: - user = await user_dao.get(db, pk) - if not user: - raise errors.NotFoundError(msg='用户不存在') - count = await user_dao.delete(db, user.id) - key_prefix = [ - f'{settings.TOKEN_REDIS_PREFIX}:{user.id}', - f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}', - ] - for key in key_prefix: - await redis_client.delete_prefix(key) - return count + user = await user_dao.get(db, pk) + if not user: + raise errors.NotFoundError(msg='用户不存在') + count = await user_dao.delete(db, user.id) + key_prefix = [ + f'{settings.TOKEN_REDIS_PREFIX}:{user.id}', + f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}', + ] + for key in key_prefix: + await redis_client.delete_prefix(key) + return count user_service: UserService = UserService() diff --git a/backend/app/task/api/v1/result.py b/backend/app/task/api/v1/result.py index 458e6ea6..08f07038 100644 --- a/backend/app/task/api/v1/result.py +++ b/backend/app/task/api/v1/result.py @@ -4,21 +4,22 @@ from fastapi import APIRouter, Depends, Path, Query from backend.app.task.schema.result import DeleteTaskResultParam, GetTaskResultDetail from backend.app.task.service.result_service import task_result_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.get('/{pk}', summary='获取任务结果详情', dependencies=[DependsJwtAuth]) async def get_task_result( + db: CurrentSession, pk: Annotated[int, Path(description='任务结果 ID')], ) -> ResponseSchemaModel[GetTaskResultDetail]: - result = await task_result_service.get(pk=pk) + result = await task_result_service.get(db=db, pk=pk) return response_base.success(data=result) @@ -30,13 +31,12 @@ async def get_task_result( DependsPagination, ], ) -async def get_task_results_paged( +async def get_task_results_paginated( db: CurrentSession, name: Annotated[str | None, Query(description='任务名称')] = None, task_id: Annotated[str | None, Query(description='任务 ID')] = None, ) -> ResponseSchemaModel[PageData[GetTaskResultDetail]]: - result_select = await task_result_service.get_select(name=name, task_id=task_id) - page_data = await paging_data(db, result_select) + page_data = await task_result_service.get_list(db=db, name=name, task_id=task_id) return response_base.success(data=page_data) @@ -48,8 +48,8 @@ async def get_task_results_paged( DependsRBAC, ], ) -async def delete_task_result(obj: DeleteTaskResultParam) -> ResponseModel: - count = await task_result_service.delete(obj=obj) +async def delete_task_result(db: CurrentSessionTransaction, obj: DeleteTaskResultParam) -> ResponseModel: + count = await task_result_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/app/task/api/v1/scheduler.py b/backend/app/task/api/v1/scheduler.py index efa3c5d3..14b7e10b 100644 --- a/backend/app/task/api/v1/scheduler.py +++ b/backend/app/task/api/v1/scheduler.py @@ -8,27 +8,28 @@ from backend.app.task.schema.scheduler import ( UpdateTaskSchedulerParam, ) from backend.app.task.service.scheduler_service import task_scheduler_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction router = APIRouter() @router.get('/all', summary='获取所有任务调度', dependencies=[DependsJwtAuth]) -async def get_all_task_schedulers() -> ResponseSchemaModel[list[GetTaskSchedulerDetail]]: - schedulers = await task_scheduler_service.get_all() +async def get_all_task_schedulers(db: CurrentSession) -> ResponseSchemaModel[list[GetTaskSchedulerDetail]]: + schedulers = await task_scheduler_service.get_all(db=db) return response_base.success(data=schedulers) @router.get('/{pk}', summary='获取任务调度详情', dependencies=[DependsJwtAuth]) async def get_task_scheduler( + db: CurrentSession, pk: Annotated[int, Path(description='任务调度 ID')], ) -> ResponseSchemaModel[GetTaskSchedulerDetail]: - task_scheduler = await task_scheduler_service.get(pk=pk) + task_scheduler = await task_scheduler_service.get(db=db, pk=pk) return response_base.success(data=task_scheduler) @@ -40,13 +41,12 @@ async def get_task_scheduler( DependsPagination, ], ) -async def get_task_scheduler_paged( +async def get_task_scheduler_paginated( db: CurrentSession, 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) - page_data = await paging_data(db, task_scheduler_select) + page_data = await task_scheduler_service.get_list(db=db, name=name, type=type) return response_base.success(data=page_data) @@ -58,8 +58,8 @@ async def get_task_scheduler_paged( DependsRBAC, ], ) -async def create_task_scheduler(obj: CreateTaskSchedulerParam) -> ResponseModel: - await task_scheduler_service.create(obj=obj) +async def create_task_scheduler(db: CurrentSessionTransaction, obj: CreateTaskSchedulerParam) -> ResponseModel: + await task_scheduler_service.create(db=db, obj=obj) return response_base.success() @@ -72,10 +72,11 @@ async def create_task_scheduler(obj: CreateTaskSchedulerParam) -> ResponseModel: ], ) async def update_task_scheduler( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='任务调度 ID')], obj: UpdateTaskSchedulerParam, ) -> ResponseModel: - count = await task_scheduler_service.update(pk=pk, obj=obj) + count = await task_scheduler_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -89,8 +90,10 @@ async def update_task_scheduler( DependsRBAC, ], ) -async def update_task_scheduler_status(pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel: - count = await task_scheduler_service.update_status(pk=pk) +async def update_task_scheduler_status( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='任务调度 ID')] +) -> ResponseModel: + count = await task_scheduler_service.update_status(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() @@ -104,8 +107,10 @@ async def update_task_scheduler_status(pk: Annotated[int, Path(description='任 DependsRBAC, ], ) -async def delete_task_scheduler(pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel: - count = await task_scheduler_service.delete(pk=pk) +async def delete_task_scheduler( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='任务调度 ID')] +) -> ResponseModel: + count = await task_scheduler_service.delete(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() @@ -119,6 +124,6 @@ async def delete_task_scheduler(pk: Annotated[int, Path(description='任务调 DependsRBAC, ], ) -async def execute_task(pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel: - await task_scheduler_service.execute(pk=pk) +async def execute_task(db: CurrentSession, pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel: + await task_scheduler_service.execute(db=db, pk=pk) return response_base.success() diff --git a/backend/app/task/crud/crud_result.py b/backend/app/task/crud/crud_result.py index cc442a29..e716f875 100644 --- a/backend/app/task/crud/crud_result.py +++ b/backend/app/task/crud/crud_result.py @@ -18,9 +18,9 @@ class CRUDTaskResult(CRUDPlus[TaskResult]): """ return await self.select_model(db, pk) - async def get_list(self, name: str | None, task_id: str | None) -> Select: + async def get_select(self, name: str | None, task_id: str | None) -> Select: """ - 获取任务结果列表 + 获取任务结果列表查询表达式 :param name: 任务名称 :param task_id: 任务 ID diff --git a/backend/app/task/crud/crud_scheduler.py b/backend/app/task/crud/crud_scheduler.py index e8d01989..1b7c636b 100644 --- a/backend/app/task/crud/crud_scheduler.py +++ b/backend/app/task/crud/crud_scheduler.py @@ -31,9 +31,9 @@ class CRUDTaskScheduler(CRUDPlus[TaskScheduler]): """ return await self.select_models(db) - async def get_list(self, name: str | None, type: int | None) -> Select: + async def get_select(self, name: str | None, type: int | None) -> Select: """ - 获取任务调度列表 + 获取任务调度列表查询表达式 :param name: 任务调度名称 :param type: 任务调度类型 diff --git a/backend/app/task/service/result_service.py b/backend/app/task/service/result_service.py index 01592fb6..d3e02da8 100644 --- a/backend/app/task/service/result_service.py +++ b/backend/app/task/service/result_service.py @@ -1,49 +1,55 @@ -from sqlalchemy import Select +from typing import Any + +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.task.crud.crud_result import task_result_dao 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 +from backend.common.pagination import paging_data class TaskResultService: @staticmethod - async def get(*, pk: int) -> TaskResult: + async def get(*, db: AsyncSession, pk: int) -> TaskResult: """ 获取任务结果详情 + :param db: 数据库会话 :param pk: 任务 ID :return: """ - async with async_db_session() as db: - result = await task_result_dao.get(db, pk) - if not result: - raise errors.NotFoundError(msg='任务结果不存在') - return result + + result = await task_result_dao.get(db, pk) + if not result: + raise errors.NotFoundError(msg='任务结果不存在') + return result @staticmethod - async def get_select(*, name: str | None, task_id: str | None) -> Select: + async def get_list(*, db: AsyncSession, name: str | None, task_id: str | None) -> dict[str, Any]: """ - 获取任务结果列表查询条件 + 获取任务结果列表 + :param db: 数据库会话 :param name: 任务名称 :param task_id: 任务 ID :return: """ - return await task_result_dao.get_list(name, task_id) + result_select = await task_result_dao.get_select(name, task_id) + return await paging_data(db, result_select) @staticmethod - async def delete(*, obj: DeleteTaskResultParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteTaskResultParam) -> int: """ 批量删除任务结果 + :param db: 数据库会话 :param obj: 任务结果 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await task_result_dao.delete(db, obj.pks) - return count + + count = await task_result_dao.delete(db, obj.pks) + return count task_result_service: TaskResultService = TaskResultService() diff --git a/backend/app/task/service/scheduler_service.py b/backend/app/task/service/scheduler_service.py index 60106179..85276a3e 100644 --- a/backend/app/task/service/scheduler_service.py +++ b/backend/app/task/service/scheduler_service.py @@ -1,8 +1,9 @@ import json from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from starlette.concurrency import run_in_threadpool from backend.app.task.celery import celery_app @@ -12,132 +13,145 @@ from backend.app.task.model import TaskScheduler from backend.app.task.schema.scheduler import CreateTaskSchedulerParam, UpdateTaskSchedulerParam from backend.app.task.utils.tzcrontab import crontab_verify from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data class TaskSchedulerService: """任务调度服务类""" @staticmethod - async def get(*, pk: int) -> TaskScheduler | None: + async def get(*, db: AsyncSession, pk: int) -> TaskScheduler | None: """ 获取任务调度详情 + :param db: 数据库会话 :param pk: 任务调度 ID :return: """ - async with async_db_session() as db: - task_scheduler = await task_scheduler_dao.get(db, pk) - if not task_scheduler: - raise errors.NotFoundError(msg='任务调度不存在') - return task_scheduler + + task_scheduler = await task_scheduler_dao.get(db, pk) + if not task_scheduler: + raise errors.NotFoundError(msg='任务调度不存在') + return task_scheduler @staticmethod - async def get_all() -> Sequence[TaskScheduler]: - """获取所有任务调度""" - async with async_db_session() as db: - task_schedulers = await task_scheduler_dao.get_all(db) - return task_schedulers - - @staticmethod - async def get_select(*, name: str | None, type: int | None) -> Select: + async def get_all(*, db: AsyncSession) -> Sequence[TaskScheduler]: """ - 获取任务调度列表查询条件 + 获取所有任务调度 + :param db: 数据库会话 + :return: + """ + + task_schedulers = await task_scheduler_dao.get_all(db) + return task_schedulers + + @staticmethod + async def get_list(*, db: AsyncSession, name: str | None, type: int | None) -> dict[str, Any]: + """ + 获取任务调度列表 + + :param db: 数据库会话 :param name: 任务调度名称 :param type: 任务调度类型 :return: """ - return await task_scheduler_dao.get_list(name=name, type=type) + task_scheduler_select = await task_scheduler_dao.get_select(name=name, type=type) + return await paging_data(db, task_scheduler_select) @staticmethod - async def create(*, obj: CreateTaskSchedulerParam) -> None: + async def create(*, db: AsyncSession, obj: CreateTaskSchedulerParam) -> None: """ 创建任务调度 + :param db: 数据库会话 :param obj: 任务调度创建参数 :return: """ - async with async_db_session.begin() as db: - task_scheduler = await task_scheduler_dao.get_by_name(db, obj.name) - if task_scheduler: - raise errors.ConflictError(msg='任务调度已存在') - if obj.type == TaskSchedulerType.CRONTAB: - crontab_verify(obj.crontab) - await task_scheduler_dao.create(db, obj) + + task_scheduler = await task_scheduler_dao.get_by_name(db, obj.name) + if task_scheduler: + raise errors.ConflictError(msg='任务调度已存在') + if obj.type == TaskSchedulerType.CRONTAB: + crontab_verify(obj.crontab) + await task_scheduler_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateTaskSchedulerParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateTaskSchedulerParam) -> int: """ 更新任务调度 + :param db: 数据库会话 :param pk: 任务调度 ID :param obj: 任务调度更新参数 :return: """ - async with async_db_session.begin() as db: - task_scheduler = await task_scheduler_dao.get(db, pk) - if not task_scheduler: - raise errors.NotFoundError(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) - return count + + task_scheduler = await task_scheduler_dao.get(db, pk) + if not task_scheduler: + raise errors.NotFoundError(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) + return count @staticmethod - async def update_status(*, pk: int) -> int: + async def update_status(*, db: AsyncSession, pk: int) -> int: """ 更新任务调度状态 + :param db: 数据库会话 :param pk: 任务调度 ID :return: """ - async with async_db_session.begin() as db: - 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, status=not task_scheduler.enabled) - return count + + 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, status=not task_scheduler.enabled) + return count @staticmethod - async def delete(*, pk: int) -> int: + async def delete(*, db: AsyncSession, pk: int) -> int: """ 删除任务调度 + :param db: 数据库会话 :param pk: 用户 ID :return: """ - async with async_db_session.begin() as db: - task_scheduler = await task_scheduler_dao.get(db, pk) - if not task_scheduler: - raise errors.NotFoundError(msg='任务调度不存在') - count = await task_scheduler_dao.delete(db, pk) - return count + + task_scheduler = await task_scheduler_dao.get(db, pk) + if not task_scheduler: + raise errors.NotFoundError(msg='任务调度不存在') + count = await task_scheduler_dao.delete(db, pk) + return count @staticmethod - async def execute(*, pk: int) -> None: + async def execute(*, db: AsyncSession, pk: int) -> None: """ 执行任务 + :param db: 数据库会话 :param pk: 任务调度 ID :return: """ - async with async_db_session() as db: - workers = await run_in_threadpool(celery_app.control.ping, timeout=0.5) - if not workers: - raise errors.ServerError(msg='Celery Worker 暂不可用,请稍后重试') - task_scheduler = await task_scheduler_dao.get(db, pk) - if not task_scheduler: - raise errors.NotFoundError(msg='任务调度不存在') - try: - args = json.loads(task_scheduler.args) if task_scheduler.args else None - kwargs = json.loads(task_scheduler.kwargs) if task_scheduler.kwargs else None - except (TypeError, json.JSONDecodeError): - raise errors.RequestError(msg='执行失败,任务参数非法') - else: - celery_app.send_task(name=task_scheduler.task, args=args, kwargs=kwargs) + + workers = await run_in_threadpool(celery_app.control.ping, timeout=0.5) + if not workers: + raise errors.ServerError(msg='Celery Worker 暂不可用,请稍后重试') + task_scheduler = await task_scheduler_dao.get(db, pk) + if not task_scheduler: + raise errors.NotFoundError(msg='任务调度不存在') + try: + args = json.loads(task_scheduler.args) if task_scheduler.args else None + kwargs = json.loads(task_scheduler.kwargs) if task_scheduler.kwargs else None + except (TypeError, json.JSONDecodeError): + raise errors.RequestError(msg='执行失败,任务参数非法') + else: + celery_app.send_task(name=task_scheduler.task, args=args, kwargs=kwargs) task_scheduler_service: TaskSchedulerService = TaskSchedulerService() diff --git a/backend/common/exception/exception_handler.py b/backend/common/exception/exception_handler.py index 872a8d72..72288cf0 100644 --- a/backend/common/exception/exception_handler.py +++ b/backend/common/exception/exception_handler.py @@ -33,11 +33,10 @@ def _get_exception_code(status_code: int) -> int: return status_code -async def _validation_exception_handler(request: Request, exc: RequestValidationError | ValidationError): +async def _validation_exception_handler(exc: RequestValidationError | ValidationError): """ 数据验证异常处理 - :param request: 请求对象 :param exc: 验证异常 :return: """ @@ -114,7 +113,7 @@ def register_exception(app: FastAPI) -> None: :param exc: 验证异常 :return: """ - return await _validation_exception_handler(request, exc) + return await _validation_exception_handler(exc) @app.exception_handler(ValidationError) async def pydantic_validation_exception_handler(request: Request, exc: ValidationError): @@ -125,7 +124,7 @@ def register_exception(app: FastAPI) -> None: :param exc: 验证异常 :return: """ - return await _validation_exception_handler(request, exc) + return await _validation_exception_handler(exc) @app.exception_handler(AssertionError) async def assertion_error_handler(request: Request, exc: AssertionError): diff --git a/backend/database/db.py b/backend/database/db.py index bcedbbd4..60ee9ef9 100644 --- a/backend/database/db.py +++ b/backend/database/db.py @@ -6,7 +6,12 @@ from uuid import uuid4 from fastapi import Depends from sqlalchemy import URL -from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.ext.asyncio import ( + AsyncEngine, + AsyncSession, + async_sessionmaker, + create_async_engine, +) from backend.common.log import log from backend.common.model import MappedBase @@ -74,6 +79,12 @@ async def get_db() -> AsyncGenerator[AsyncSession, None]: yield session +async def get_db_transaction() -> AsyncGenerator[AsyncSession, None]: + """获取带有事务的数据库会话""" + async with async_db_session.begin() as session: + yield session + + async def create_tables() -> None: """创建数据库表""" async with async_engine.begin() as coon: @@ -93,3 +104,4 @@ async_engine, async_db_session = create_async_engine_and_session(SQLALCHEMY_DATA # Session Annotated CurrentSession = Annotated[AsyncSession, Depends(get_db)] +CurrentSessionTransaction = Annotated[AsyncSession, Depends(get_db_transaction)] diff --git a/backend/middleware/opera_log_middleware.py b/backend/middleware/opera_log_middleware.py index 690e9429..816d97a9 100644 --- a/backend/middleware/opera_log_middleware.py +++ b/backend/middleware/opera_log_middleware.py @@ -17,6 +17,7 @@ from backend.common.log import log from backend.common.queue import batch_dequeue from backend.common.response.response_code import StandardResponseCode from backend.core.conf import settings +from backend.database.db import async_db_session from backend.utils.encrypt import AESCipher, ItsDCipher, Md5Cipher from backend.utils.trace_id import get_request_trace_id @@ -202,7 +203,8 @@ class OperaLogMiddleware(BaseHTTPMiddleware): try: if settings.DATABASE_ECHO: log.info('自动执行【操作日志批量创建】任务...') - await opera_log_service.bulk_create(objs=logs) + async with async_db_session.begin() as db: + await opera_log_service.bulk_create(db=db, objs=logs) finally: if not cls.opera_log_queue.empty(): cls.opera_log_queue.task_done() diff --git a/backend/plugin/code_generator/api/v1/business.py b/backend/plugin/code_generator/api/v1/business.py index aab86308..cc51b532 100644 --- a/backend/plugin/code_generator/api/v1/business.py +++ b/backend/plugin/code_generator/api/v1/business.py @@ -2,12 +2,12 @@ from typing import Annotated from fastapi import APIRouter, Depends, Path, Query -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.code_generator.schema.business import ( CreateGenBusinessParam, GetGenBusinessDetail, @@ -21,16 +21,17 @@ router = APIRouter() @router.get('/all', summary='获取所有代码生成业务', dependencies=[DependsJwtAuth]) -async def get_all_businesses() -> ResponseSchemaModel[list[GetGenBusinessDetail]]: - data = await gen_business_service.get_all() +async def get_all_businesses(db: CurrentSession) -> ResponseSchemaModel[list[GetGenBusinessDetail]]: + data = await gen_business_service.get_all(db=db) return response_base.success(data=data) @router.get('/{pk}', summary='获取代码生成业务详情', dependencies=[DependsJwtAuth]) async def get_business( + db: CurrentSession, pk: Annotated[int, Path(description='业务 ID')], ) -> ResponseSchemaModel[GetGenBusinessDetail]: - data = await gen_business_service.get(pk=pk) + data = await gen_business_service.get(db=db, pk=pk) return response_base.success(data=data) @@ -42,20 +43,20 @@ async def get_business( DependsPagination, ], ) -async def get_businesses_paged( +async def get_businesses_paginated( db: CurrentSession, table_name: Annotated[str | None, Query(description='代码生成业务表名称')] = None, ) -> ResponseSchemaModel[PageData[GetGenBusinessDetail]]: - business_select = await gen_business_service.get_select(table_name=table_name) - page_data = await paging_data(db, business_select) + page_data = await gen_business_service.get_list(db=db, table_name=table_name) return response_base.success(data=page_data) @router.get('/{pk}/columns', summary='获取代码生成业务所有模型列', dependencies=[DependsJwtAuth]) async def get_business_all_columns( + db: CurrentSession, pk: Annotated[int, Path(description='业务 ID')], ) -> ResponseSchemaModel[list[GetGenColumnDetail]]: - data = await gen_column_service.get_columns(business_id=pk) + data = await gen_column_service.get_columns(db=db, business_id=pk) return response_base.success(data=data) @@ -68,8 +69,8 @@ async def get_business_all_columns( DependsRBAC, ], ) -async def create_business(obj: CreateGenBusinessParam) -> ResponseModel: - await gen_business_service.create(obj=obj) +async def create_business(db: CurrentSessionTransaction, obj: CreateGenBusinessParam) -> ResponseModel: + await gen_business_service.create(db=db, obj=obj) return response_base.success() @@ -82,10 +83,11 @@ async def create_business(obj: CreateGenBusinessParam) -> ResponseModel: ], ) async def update_business( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='业务 ID')], obj: UpdateGenBusinessParam, ) -> ResponseModel: - count = await gen_business_service.update(pk=pk, obj=obj) + count = await gen_business_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -99,8 +101,10 @@ async def update_business( DependsRBAC, ], ) -async def delete_business(pk: Annotated[int, Path(description='业务 ID')]) -> ResponseModel: - count = await gen_business_service.delete(pk=pk) +async def delete_business( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='业务 ID')] +) -> ResponseModel: + count = await gen_business_service.delete(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/code_generator/api/v1/code.py b/backend/plugin/code_generator/api/v1/code.py index 144c0f25..ddb278d6 100644 --- a/backend/plugin/code_generator/api/v1/code.py +++ b/backend/plugin/code_generator/api/v1/code.py @@ -8,6 +8,7 @@ from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC from backend.core.conf import settings +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.code_generator.schema.code import ImportParam from backend.plugin.code_generator.service.code_service import gen_service @@ -16,9 +17,10 @@ router = APIRouter() @router.get('/tables', summary='获取数据库表') async def get_all_tables( + db: CurrentSession, table_schema: Annotated[str, Query(description='数据库名')] = 'fba', ) -> ResponseSchemaModel[list[dict[str, str | None]]]: - data = await gen_service.get_tables(table_schema=table_schema) + data = await gen_service.get_tables(db=db, table_schema=table_schema) return response_base.success(data=data) @@ -30,20 +32,24 @@ async def get_all_tables( DependsRBAC, ], ) -async def import_table(obj: ImportParam) -> ResponseModel: - await gen_service.import_business_and_model(obj=obj) +async def import_table(db: CurrentSessionTransaction, obj: ImportParam) -> ResponseModel: + await gen_service.import_business_and_model(db=db, obj=obj) return response_base.success() @router.get('/{pk}/previews', summary='代码生成预览', dependencies=[DependsJwtAuth]) -async def preview_code(pk: Annotated[int, Path(description='业务 ID')]) -> ResponseSchemaModel[dict[str, bytes]]: - data = await gen_service.preview(pk=pk) +async def preview_code( + db: CurrentSession, pk: Annotated[int, Path(description='业务 ID')] +) -> ResponseSchemaModel[dict[str, bytes]]: + data = await gen_service.preview(db=db, pk=pk) return response_base.success(data=data) @router.get('/{pk}/paths', summary='获取代码生成路径', dependencies=[DependsJwtAuth]) -async def get_generate_paths(pk: Annotated[int, Path(description='业务 ID')]) -> ResponseSchemaModel[list[str]]: - data = await gen_service.get_generate_path(pk=pk) +async def get_generate_paths( + db: CurrentSession, pk: Annotated[int, Path(description='业务 ID')] +) -> ResponseSchemaModel[list[str]]: + data = await gen_service.get_generate_path(db=db, pk=pk) return response_base.success(data=data) @@ -56,14 +62,14 @@ async def get_generate_paths(pk: Annotated[int, Path(description='业务 ID')]) DependsRBAC, ], ) -async def generate_code(pk: Annotated[int, Path(description='业务 ID')]) -> ResponseModel: - await gen_service.generate(pk=pk) +async def generate_code(db: CurrentSession, pk: Annotated[int, Path(description='业务 ID')]) -> ResponseModel: + await gen_service.generate(db=db, pk=pk) return response_base.success() @router.get('/{pk}', summary='下载代码', dependencies=[DependsJwtAuth]) -async def download_code(pk: Annotated[int, Path(description='业务 ID')]): # noqa: ANN201 - bio = await gen_service.download(pk=pk) +async def download_code(db: CurrentSession, pk: Annotated[int, Path(description='业务 ID')]): # noqa: ANN201 + bio = await gen_service.download(db=db, pk=pk) return StreamingResponse( bio, media_type='application/x-zip-compressed', diff --git a/backend/plugin/code_generator/api/v1/column.py b/backend/plugin/code_generator/api/v1/column.py index 5d933c97..5828d531 100644 --- a/backend/plugin/code_generator/api/v1/column.py +++ b/backend/plugin/code_generator/api/v1/column.py @@ -6,6 +6,7 @@ from backend.common.response.response_schema import ResponseModel, ResponseSchem from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.code_generator.schema.column import ( CreateGenColumnParam, GetGenColumnDetail, @@ -23,8 +24,10 @@ async def get_column_types() -> ResponseSchemaModel[list[str]]: @router.get('/{pk}', summary='获取代码生成模型列详情', dependencies=[DependsJwtAuth]) -async def get_column(pk: Annotated[int, Path(description='模型列 ID')]) -> ResponseSchemaModel[GetGenColumnDetail]: - data = await gen_column_service.get(pk=pk) +async def get_column( + db: CurrentSession, pk: Annotated[int, Path(description='模型列 ID')] +) -> ResponseSchemaModel[GetGenColumnDetail]: + data = await gen_column_service.get(db=db, pk=pk) return response_base.success(data=data) @@ -36,8 +39,8 @@ async def get_column(pk: Annotated[int, Path(description='模型列 ID')]) -> Re DependsRBAC, ], ) -async def create_column(obj: CreateGenColumnParam) -> ResponseModel: - await gen_column_service.create(obj=obj) +async def create_column(db: CurrentSessionTransaction, obj: CreateGenColumnParam) -> ResponseModel: + await gen_column_service.create(db=db, obj=obj) return response_base.success() @@ -49,8 +52,10 @@ async def create_column(obj: CreateGenColumnParam) -> ResponseModel: DependsRBAC, ], ) -async def update_column(pk: Annotated[int, Path(description='模型列 ID')], obj: UpdateGenColumnParam) -> ResponseModel: - count = await gen_column_service.update(pk=pk, obj=obj) +async def update_column( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='模型列 ID')], obj: UpdateGenColumnParam +) -> ResponseModel: + count = await gen_column_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -64,8 +69,10 @@ async def update_column(pk: Annotated[int, Path(description='模型列 ID')], ob DependsRBAC, ], ) -async def delete_column(pk: Annotated[int, Path(description='模型列 ID')]) -> ResponseModel: - count = await gen_column_service.delete(pk=pk) +async def delete_column( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='模型列 ID')] +) -> ResponseModel: + count = await gen_column_service.delete(db=db, pk=pk) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/code_generator/crud/crud_business.py b/backend/plugin/code_generator/crud/crud_business.py index e356a8e6..45fe8cea 100644 --- a/backend/plugin/code_generator/crud/crud_business.py +++ b/backend/plugin/code_generator/crud/crud_business.py @@ -40,9 +40,9 @@ class CRUDGenBusiness(CRUDPlus[GenBusiness]): """ return await self.select_models(db) - async def get_list(self, table_name: str | None) -> Select: + async def get_select(self, table_name: str | None) -> Select: """ - 获取所有代码生成业务 + 获取所有代码生成业务查询表达式 :param table_name: 业务表名 :return: diff --git a/backend/plugin/code_generator/service/business_service.py b/backend/plugin/code_generator/service/business_service.py index c1b63584..9e966f19 100644 --- a/backend/plugin/code_generator/service/business_service.py +++ b/backend/plugin/code_generator/service/business_service.py @@ -1,9 +1,10 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data from backend.plugin.code_generator.crud.crud_business import gen_business_dao from backend.plugin.code_generator.model import GenBusiness from backend.plugin.code_generator.schema.business import CreateGenBusinessParam, UpdateGenBusinessParam @@ -13,71 +14,82 @@ class GenBusinessService: """代码生成业务服务类""" @staticmethod - async def get(*, pk: int) -> GenBusiness: + async def get(*, db: AsyncSession, pk: int) -> GenBusiness: """ 获取指定 ID 的业务 + :param db: 数据库会话 :param pk: 业务 ID :return: """ - async with async_db_session() as db: - business = await gen_business_dao.get(db, pk) - if not business: - raise errors.NotFoundError(msg='代码生成业务不存在') - return business + + business = await gen_business_dao.get(db, pk) + if not business: + raise errors.NotFoundError(msg='代码生成业务不存在') + return business @staticmethod - async def get_all() -> Sequence[GenBusiness]: - """获取所有业务""" - async with async_db_session() as db: - return await gen_business_dao.get_all(db) - - @staticmethod - async def get_select(*, table_name: str) -> Select: + async def get_all(*, db: AsyncSession) -> Sequence[GenBusiness]: """ - 获取代码生成业务列表查询条件 + 获取所有业务 + :param db: 数据库会话 + :return: + """ + + return await gen_business_dao.get_all(db) + + @staticmethod + async def get_list(*, db: AsyncSession, table_name: str) -> dict[str, Any]: + """ + 获取代码生成业务列表 + + :param db: 数据库会话 :param table_name: 业务表名 :return: """ - return await gen_business_dao.get_list(table_name=table_name) + business_select = await gen_business_dao.get_select(table_name=table_name) + return await paging_data(db, business_select) @staticmethod - async def create(*, obj: CreateGenBusinessParam) -> None: + async def create(*, db: AsyncSession, obj: CreateGenBusinessParam) -> None: """ 创建业务 + :param db: 数据库会话 :param obj: 创建业务参数 :return: """ - async with async_db_session.begin() as db: - business = await gen_business_dao.get_by_name(db, obj.table_name) - if business: - raise errors.ConflictError(msg='代码生成业务已存在') - await gen_business_dao.create(db, obj) + + business = await gen_business_dao.get_by_name(db, obj.table_name) + if business: + raise errors.ConflictError(msg='代码生成业务已存在') + await gen_business_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateGenBusinessParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateGenBusinessParam) -> int: """ 更新业务 + :param db: 数据库会话 :param pk: 业务 ID :param obj: 更新业务参数 :return: """ - async with async_db_session.begin() as db: - return await gen_business_dao.update(db, pk, obj) + + return await gen_business_dao.update(db, pk, obj) @staticmethod - async def delete(*, pk: int) -> int: + async def delete(*, db: AsyncSession, pk: int) -> int: """ 删除业务 + :param db: 数据库会话 :param pk: 业务 ID :return: """ - async with async_db_session.begin() as db: - return await gen_business_dao.delete(db, pk) + + return await gen_business_dao.delete(db, pk) gen_business_service: GenBusinessService = GenBusinessService() diff --git a/backend/plugin/code_generator/service/code_service.py b/backend/plugin/code_generator/service/code_service.py index d54b2c65..9ab8a8d4 100644 --- a/backend/plugin/code_generator/service/code_service.py +++ b/backend/plugin/code_generator/service/code_service.py @@ -9,10 +9,10 @@ import anyio from anyio import open_file from pydantic.alias_generators import to_pascal from sqlalchemy import RowMapping +from sqlalchemy.ext.asyncio import AsyncSession from backend.common.exception import errors from backend.core.path_conf import BASE_PATH -from backend.database.db import async_db_session from backend.plugin.code_generator.crud.crud_business import gen_business_dao from backend.plugin.code_generator.crud.crud_code import gen_dao from backend.plugin.code_generator.crud.crud_column import gen_column_dao @@ -29,76 +29,79 @@ class GenService: """代码生成服务类""" @staticmethod - async def get_tables(*, table_schema: str) -> Sequence[RowMapping]: + async def get_tables(*, db: AsyncSession, table_schema: str) -> Sequence[RowMapping]: """ 获取指定 schema 下的所有表名 + :param db: 数据库会话 :param table_schema: 数据库 schema 名称 :return: """ - async with async_db_session() as db: - return await gen_dao.get_all_tables(db, table_schema) + + return await gen_dao.get_all_tables(db, table_schema) @staticmethod - async def import_business_and_model(*, obj: ImportParam) -> None: + async def import_business_and_model(*, db: AsyncSession, obj: ImportParam) -> None: """ 导入业务和模型列数据 + :param db: 数据库会话 :param obj: 导入参数对象 :return: """ - async with async_db_session.begin() as db: - table_info = await gen_dao.get_table(db, obj.table_name) - if not table_info: - raise errors.NotFoundError(msg='数据库表不存在') - business_info = await gen_business_dao.get_by_name(db, obj.table_name) - if business_info: - raise errors.ConflictError(msg='已存在相同数据库表业务') + table_info = await gen_dao.get_table(db, obj.table_name) + if not table_info: + raise errors.NotFoundError(msg='数据库表不存在') - table_name = table_info[0] - new_business = GenBusiness( - **CreateGenBusinessParam( - app_name=obj.app, - table_name=table_name, - doc_comment=table_info[1] or table_name.split('_')[-1], - table_comment=table_info[1], - class_name=to_pascal(table_name), - schema_name=to_pascal(table_name), - filename=table_name, - ).model_dump(), + business_info = await gen_business_dao.get_by_name(db, obj.table_name) + if business_info: + raise errors.ConflictError(msg='已存在相同数据库表业务') + + table_name = table_info[0] + new_business = GenBusiness( + **CreateGenBusinessParam( + app_name=obj.app, + table_name=table_name, + doc_comment=table_info[1] or table_name.split('_')[-1], + table_comment=table_info[1], + class_name=to_pascal(table_name), + schema_name=to_pascal(table_name), + filename=table_name, + ).model_dump(), + ) + db.add(new_business) + await db.flush() + + column_info = await gen_dao.get_all_columns(db, obj.table_schema, table_name) + for column in column_info: + column_type = column[-1].split('(')[0].upper() + pd_type = sql_type_to_pydantic(column_type) + await gen_column_dao.create( + db, + CreateGenColumnParam( + name=column[0], + comment=column[-2], + type=column_type, + sort=column[-3], + length=column[-1].split('(')[1][:-1] if pd_type == 'str' and '(' in column[-1] else 0, + is_pk=column[1], + is_nullable=column[2], + gen_business_id=new_business.id, + ), + pd_type=pd_type, ) - db.add(new_business) - await db.flush() - - column_info = await gen_dao.get_all_columns(db, obj.table_schema, table_name) - for column in column_info: - column_type = column[-1].split('(')[0].upper() - pd_type = sql_type_to_pydantic(column_type) - await gen_column_dao.create( - db, - CreateGenColumnParam( - name=column[0], - comment=column[-2], - type=column_type, - sort=column[-3], - length=column[-1].split('(')[1][:-1] if pd_type == 'str' and '(' in column[-1] else 0, - is_pk=column[1], - is_nullable=column[2], - gen_business_id=new_business.id, - ), - pd_type=pd_type, - ) @staticmethod - async def render_tpl_code(*, business: GenBusiness) -> dict[str, str]: + async def render_tpl_code(*, db: AsyncSession, business: GenBusiness) -> dict[str, str]: """ 渲染模板代码 + :param db: 数据库会话 :param business: 业务对象 :return: """ - gen_models = await gen_column_service.get_columns(business_id=business.id) + gen_models = await gen_column_service.get_columns(db=db, business_id=business.id) if not gen_models: raise errors.NotFoundError(msg='代码生成模型表为空') @@ -108,160 +111,167 @@ class GenService: for tpl_path in gen_template.get_template_files() } - async def preview(self, *, pk: int) -> dict[str, bytes]: + async def preview(self, *, db: AsyncSession, pk: int) -> dict[str, bytes]: """ 预览生成的代码 + :param db: 数据库会话 :param pk: 业务 ID :return: """ - async with async_db_session() as db: - business = await gen_business_dao.get(db, pk) - if not business: - raise errors.NotFoundError(msg='业务不存在') - tpl_code_map = await self.render_tpl_code(business=business) + business = await gen_business_dao.get(db, pk) + if not business: + raise errors.NotFoundError(msg='业务不存在') - codes = {} - for tpl_path, code in tpl_code_map.items(): - if tpl_path.startswith('python'): - rootpath = f'fastapi_best_architecture/backend/app/{business.app_name}' - template_name = tpl_path.split('/')[-1] - match template_name: - case 'api.jinja': - filepath = f'{rootpath}/api/{business.api_version}/{business.filename}.py' - case 'crud.jinja': - filepath = f'{rootpath}/crud/crud_{business.filename}.py' - case 'model.jinja': - filepath = f'{rootpath}/model/{business.filename}.py' - case 'schema.jinja': - filepath = f'{rootpath}/schema/{business.filename}.py' - case 'service.jinja': - filepath = f'{rootpath}/service/{business.filename}_service.py' + tpl_code_map = await self.render_tpl_code(db=db, business=business) + + codes = {} + for tpl_path, code in tpl_code_map.items(): + if tpl_path.startswith('python'): + rootpath = f'fastapi_best_architecture/backend/app/{business.app_name}' + template_name = tpl_path.split('/')[-1] + filepath = None + match template_name: + case 'api.jinja': + filepath = f'{rootpath}/api/{business.api_version}/{business.filename}.py' + case 'crud.jinja': + filepath = f'{rootpath}/crud/crud_{business.filename}.py' + case 'model.jinja': + filepath = f'{rootpath}/model/{business.filename}.py' + case 'schema.jinja': + filepath = f'{rootpath}/schema/{business.filename}.py' + case 'service.jinja': + filepath = f'{rootpath}/service/{business.filename}_service.py' + + if filepath: codes[filepath] = code.encode('utf-8') - return codes + return codes @staticmethod - async def get_generate_path(*, pk: int) -> list[str]: + async def get_generate_path(*, db: AsyncSession, pk: int) -> list[str]: """ 获取代码生成路径 + :param db: 数据库会话 :param pk: 业务 ID :return: """ - async with async_db_session() as db: - business = await gen_business_dao.get(db, pk) - if not business: - raise errors.NotFoundError(msg='业务不存在') - gen_path = business.gen_path or '.../backend/app/' - target_files = gen_template.get_code_gen_paths(business) + business = await gen_business_dao.get(db, pk) + if not business: + raise errors.NotFoundError(msg='业务不存在') - return [os.path.join(gen_path, *target_file.split('/')) for target_file in target_files] + gen_path = business.gen_path or '.../backend/app/' + target_files = gen_template.get_code_gen_paths(business) - async def generate(self, *, pk: int) -> str: + return [os.path.join(gen_path, *target_file.split('/')) for target_file in target_files] + + async def generate(self, *, db: AsyncSession, pk: int) -> str: """ 生成代码文件 + :param db: 数据库会话 :param pk: 业务 ID :return: """ - async with async_db_session() as db: - business = await gen_business_dao.get(db, pk) - if not business: - raise errors.NotFoundError(msg='业务不存在') - tpl_code_map = await self.render_tpl_code(business=business) - gen_path = business.gen_path or BASE_PATH / 'app' + business = await gen_business_dao.get(db, pk) + if not business: + raise errors.NotFoundError(msg='业务不存在') - for tpl_path, code in tpl_code_map.items(): - code_filepath = os.path.join( - gen_path, - *gen_template.get_code_gen_path(tpl_path, business).split('/'), - ) + tpl_code_map = await self.render_tpl_code(db=db, business=business) + gen_path = business.gen_path or BASE_PATH / 'app' - # 写入 init 文件 - code_folder = anyio.Path(code_filepath).parent - await code_folder.mkdir(parents=True, exist_ok=True) + for tpl_path, code in tpl_code_map.items(): + code_filepath = os.path.join( + gen_path, + *gen_template.get_code_gen_path(tpl_path, business).split('/'), + ) - init_filepath = code_folder.joinpath('__init__.py') - if not await init_filepath.exists(): - async with await open_file(init_filepath, 'w', encoding='utf-8') as f: - await f.write(gen_template.init_content) + # 写入 init 文件 + code_folder = anyio.Path(code_filepath).parent + await code_folder.mkdir(parents=True, exist_ok=True) - # api __init__.py - if 'api' in code_filepath: - api_init_filepath = code_folder.parent.joinpath('__init__.py') - async with await open_file(api_init_filepath, 'w', encoding='utf-8') as f: - await f.write(gen_template.init_content) + init_filepath = code_folder.joinpath('__init__.py') + if not await init_filepath.exists(): + async with await open_file(init_filepath, 'w', encoding='utf-8') as f: + await f.write(gen_template.init_content) - # app __init__.py - if 'service' in code_filepath: - app_init_filepath = code_folder.parent.joinpath('__init__.py') - async with await open_file(app_init_filepath, 'w', encoding='utf-8') as f: - await f.write(gen_template.init_content) + # api __init__.py + if 'api' in code_filepath: + api_init_filepath = code_folder.parent.joinpath('__init__.py') + async with await open_file(api_init_filepath, 'w', encoding='utf-8') as f: + await f.write(gen_template.init_content) - # model init 文件补充 - if code_folder.name == 'model': - async with await open_file(init_filepath, 'a', encoding='utf-8') as f: - await f.write( - f'from backend.app.{business.app_name}.model.{business.table_name} ' - f'import {to_pascal(business.table_name)}\n', - ) + # app __init__.py + if 'service' in code_filepath: + app_init_filepath = code_folder.parent.joinpath('__init__.py') + async with await open_file(app_init_filepath, 'w', encoding='utf-8') as f: + await f.write(gen_template.init_content) - # 写入代码文件 - async with await open_file(code_filepath, 'w', encoding='utf-8') as f: - await f.write(code) + # model init 文件补充 + if code_folder.name == 'model': + async with await open_file(init_filepath, 'a', encoding='utf-8') as f: + await f.write( + f'from backend.app.{business.app_name}.model.{business.table_name} ' + f'import {to_pascal(business.table_name)}\n', + ) + + # 写入代码文件 + async with await open_file(code_filepath, 'w', encoding='utf-8') as f: + await f.write(code) return gen_path - async def download(self, *, pk: int) -> io.BytesIO: + async def download(self, *, db: AsyncSession, pk: int) -> io.BytesIO: """ 下载生成的代码 + :param db: 数据库会话 :param pk: 业务 ID :return: """ - async with async_db_session() as db: - business = await gen_business_dao.get(db, pk) - if not business: - raise errors.NotFoundError(msg='业务不存在') - bio = io.BytesIO() - with zipfile.ZipFile(bio, 'w') as zf: - tpl_code_map = await self.render_tpl_code(business=business) - for tpl_path, code in tpl_code_map.items(): - code_filepath = gen_template.get_code_gen_path(tpl_path, business) + business = await gen_business_dao.get(db, pk) + if not business: + raise errors.NotFoundError(msg='业务不存在') - # 写入 init 文件 - code_dir = os.path.dirname(code_filepath) - init_filepath = os.path.join(code_dir, '__init__.py') - if 'model' not in code_filepath.split('/'): - zf.writestr(init_filepath, gen_template.init_content) - else: - zf.writestr( - init_filepath, - f'{gen_template.init_content}' - f'from backend.app.{business.app_name}.model.{business.table_name} ' - f'import {to_pascal(business.table_name)}\n', - ) + bio = io.BytesIO() + with zipfile.ZipFile(bio, 'w') as zf: + tpl_code_map = await self.render_tpl_code(db=db, business=business) + for tpl_path, code in tpl_code_map.items(): + code_filepath = gen_template.get_code_gen_path(tpl_path, business) - # api __init__.py - if 'api' in code_dir: - api_init_filepath = os.path.join(os.path.dirname(code_dir), '__init__.py') - zf.writestr(api_init_filepath, gen_template.init_content) + # 写入 init 文件 + code_dir = os.path.dirname(code_filepath) + init_filepath = os.path.join(code_dir, '__init__.py') + if 'model' not in code_filepath.split('/'): + zf.writestr(init_filepath, gen_template.init_content) + else: + zf.writestr( + init_filepath, + f'{gen_template.init_content}' + f'from backend.app.{business.app_name}.model.{business.table_name} ' + f'import {to_pascal(business.table_name)}\n', + ) - # app __init__.py - if 'service' in code_dir: - app_init_filepath = os.path.join(os.path.dirname(code_dir), '__init__.py') - zf.writestr(app_init_filepath, gen_template.init_content) + # api __init__.py + if 'api' in code_dir: + api_init_filepath = os.path.join(os.path.dirname(code_dir), '__init__.py') + zf.writestr(api_init_filepath, gen_template.init_content) - # 写入代码文件 - zf.writestr(code_filepath, code) + # app __init__.py + if 'service' in code_dir: + app_init_filepath = os.path.join(os.path.dirname(code_dir), '__init__.py') + zf.writestr(app_init_filepath, gen_template.init_content) - bio.seek(0) - return bio + # 写入代码文件 + zf.writestr(code_filepath, code) + + bio.seek(0) + return bio gen_service: GenService = GenService() diff --git a/backend/plugin/code_generator/service/column_service.py b/backend/plugin/code_generator/service/column_service.py index 82c6b3f2..8b9b4df8 100644 --- a/backend/plugin/code_generator/service/column_service.py +++ b/backend/plugin/code_generator/service/column_service.py @@ -1,7 +1,8 @@ from collections.abc import Sequence +from sqlalchemy.ext.asyncio import AsyncSession + from backend.common.exception import errors -from backend.database.db import async_db_session from backend.plugin.code_generator.crud.crud_column import gen_column_dao from backend.plugin.code_generator.enums import GenMySQLColumnType from backend.plugin.code_generator.model import GenColumn @@ -13,18 +14,19 @@ class GenColumnService: """代码生成模型列服务类""" @staticmethod - async def get(*, pk: int) -> GenColumn: + async def get(*, db: AsyncSession, pk: int) -> GenColumn: """ 获取指定 ID 的模型列 + :param db: 数据库会话 :param pk: 模型列 ID :return: """ - async with async_db_session() as db: - column = await gen_column_dao.get(db, pk) - if not column: - raise errors.NotFoundError(msg='代码生成模型列不存在') - return column + + column = await gen_column_dao.get(db, pk) + if not column: + raise errors.NotFoundError(msg='代码生成模型列不存在') + return column @staticmethod async def get_types() -> list[str]: @@ -34,61 +36,65 @@ class GenColumnService: return types @staticmethod - async def get_columns(*, business_id: int) -> Sequence[GenColumn]: + async def get_columns(*, db: AsyncSession, business_id: int) -> Sequence[GenColumn]: """ 获取指定业务的所有模型列 + :param db: 数据库会话 :param business_id: 业务 ID :return: """ - async with async_db_session() as db: - return await gen_column_dao.get_all_by_business(db, business_id) + + return await gen_column_dao.get_all_by_business(db, business_id) @staticmethod - async def create(*, obj: CreateGenColumnParam) -> None: + async def create(*, db: AsyncSession, obj: CreateGenColumnParam) -> None: """ 创建模型列 + :param db: 数据库会话 :param obj: 创建模型列参数 :return: """ - async with async_db_session.begin() as db: - gen_columns = await gen_column_dao.get_all_by_business(db, obj.gen_business_id) - if obj.name in [gen_column.name for gen_column in gen_columns]: - raise errors.ForbiddenError(msg='模型列已存在') - pd_type = sql_type_to_pydantic(obj.type) - await gen_column_dao.create(db, obj, pd_type=pd_type) + gen_columns = await gen_column_dao.get_all_by_business(db, obj.gen_business_id) + if obj.name in [gen_column.name for gen_column in gen_columns]: + raise errors.ForbiddenError(msg='模型列已存在') + + pd_type = sql_type_to_pydantic(obj.type) + await gen_column_dao.create(db, obj, pd_type=pd_type) @staticmethod - async def update(*, pk: int, obj: UpdateGenColumnParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateGenColumnParam) -> int: """ 更新模型列 + :param db: 数据库会话 :param pk: 模型列 ID :param obj: 更新模型列参数 :return: """ - async with async_db_session.begin() as db: - column = await gen_column_dao.get(db, pk) - if obj.name != column.name: - gen_columns = await gen_column_dao.get_all_by_business(db, obj.gen_business_id) - if obj.name in [gen_column.name for gen_column in gen_columns]: - raise errors.ConflictError(msg='模型列名已存在') - pd_type = sql_type_to_pydantic(obj.type) - return await gen_column_dao.update(db, pk, obj, pd_type=pd_type) + column = await gen_column_dao.get(db, pk) + if obj.name != column.name: + gen_columns = await gen_column_dao.get_all_by_business(db, obj.gen_business_id) + if obj.name in [gen_column.name for gen_column in gen_columns]: + raise errors.ConflictError(msg='模型列名已存在') + + pd_type = sql_type_to_pydantic(obj.type) + return await gen_column_dao.update(db, pk, obj, pd_type=pd_type) @staticmethod - async def delete(*, pk: int) -> int: + async def delete(*, db: AsyncSession, pk: int) -> int: """ 删除模型列 + :param db: 数据库会话 :param pk: 模型列 ID :return: """ - async with async_db_session.begin() as db: - return await gen_column_dao.delete(db, pk) + + return await gen_column_dao.delete(db, pk) gen_column_service: GenColumnService = GenColumnService() diff --git a/backend/plugin/code_generator/templates/python/api.jinja b/backend/plugin/code_generator/templates/python/api.jinja index 7c7bb688..3704b4aa 100644 --- a/backend/plugin/code_generator/templates/python/api.jinja +++ b/backend/plugin/code_generator/templates/python/api.jinja @@ -9,19 +9,22 @@ from backend.app.{{ app_name }}.schema.{{ table_name }} import ( Update{{ schema_name }}Param, ) from backend.app.{{ app_name }}.service.{{ table_name }}_service import {{ table_name }}_service -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.database.db import CurrentSession router = APIRouter() @router.get('/{pk}', summary='获取{{ doc_comment }}详情', dependencies=[DependsJwtAuth]) -async def get_{{ table_name }}(pk: Annotated[int, Path(description='{{ doc_comment }} ID')]) -> ResponseSchemaModel[Get{{ schema_name }}Detail]: - {{ table_name }} = await {{ table_name }}_service.get(pk=pk) +async def get_{{ table_name }}( + db: CurrentSession, pk: Annotated[int, Path(description='{{ doc_comment }} ID')] +) -> ResponseSchemaModel[Get{{ schema_name }}Detail]: + {{ table_name }} = await {{ table_name }}_service.get(db=db, pk=pk) return response_base.success(data={{ table_name }}) @@ -33,9 +36,8 @@ async def get_{{ table_name }}(pk: Annotated[int, Path(description='{{ doc_comme DependsPagination, ], ) -async def get_{{ table_name }}s_paged(db: CurrentSession) -> ResponseSchemaModel[PageData[Get{{ schema_name }}Detail]]: - {{ table_name }}_select = await {{ table_name }}_service.get_select() - page_data = await paging_data(db, {{ table_name }}_select) +async def get_{{ table_name }}s_paginated(db: CurrentSession) -> ResponseSchemaModel[PageData[Get{{ schema_name }}Detail]]: + page_data = = await {{ table_name }}_service.get_list(db=db) return response_base.success(data=page_data) @@ -47,8 +49,8 @@ async def get_{{ table_name }}s_paged(db: CurrentSession) -> ResponseSchemaModel DependsRBAC, ], ) -async def create_{{ table_name }}(obj: Create{{ schema_name }}Param) -> ResponseModel: - await {{ table_name }}_service.create(obj=obj) +async def create_{{ table_name }}(db: CurrentSessionTransaction, obj: Create{{ schema_name }}Param) -> ResponseModel: + await {{ table_name }}_service.create(db=db, obj=obj) return response_base.success() @@ -60,8 +62,10 @@ async def create_{{ table_name }}(obj: Create{{ schema_name }}Param) -> Response DependsRBAC, ], ) -async def update_{{ table_name }}(pk: Annotated[int, Path(description='{{ doc_comment }} ID')], obj: Update{{ schema_name }}Param) -> ResponseModel: - count = await {{ table_name }}_service.update(pk=pk, obj=obj) +async def update_{{ table_name }}( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='{{ doc_comment }} ID')], obj: Update{{ schema_name }}Param +) -> ResponseModel: + count = await {{ table_name }}_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -75,8 +79,8 @@ async def update_{{ table_name }}(pk: Annotated[int, Path(description='{{ doc_co DependsRBAC, ], ) -async def delete_{{ table_name }}s(obj: Delete{{ schema_name }}Param) -> ResponseModel: - count = await {{ table_name }}_service.delete(obj=obj) +async def delete_{{ table_name }}s(db: CurrentSessionTransaction, obj: Delete{{ schema_name }}Param) -> ResponseModel: + count = await {{ table_name }}_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/code_generator/templates/python/crud.jinja b/backend/plugin/code_generator/templates/python/crud.jinja index 62bf5438..da54c098 100644 --- a/backend/plugin/code_generator/templates/python/crud.jinja +++ b/backend/plugin/code_generator/templates/python/crud.jinja @@ -19,8 +19,8 @@ class CRUD{{ class_name }}(CRUDPlus[{{ schema_name }}]): """ return await self.select_model(db, pk) - async def get_list(self) -> Select: - """获取{{ doc_comment }}列表""" + async def get_select(self) -> Select: + """获取{{ doc_comment }}列表查询表达式""" return await self.select_order('id', 'desc') async def get_all(self, db: AsyncSession) -> Sequence[{{ class_name }}]: diff --git a/backend/plugin/code_generator/templates/python/service.jinja b/backend/plugin/code_generator/templates/python/service.jinja index 611f3d5a..920b7499 100644 --- a/backend/plugin/code_generator/templates/python/service.jinja +++ b/backend/plugin/code_generator/templates/python/service.jinja @@ -1,76 +1,86 @@ -from typing import Sequence +from typing import Any, Sequence -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.{{ app_name }}.crud.crud_{{ table_name }} import {{ table_name }}_dao from backend.app.{{ app_name }}.model import {{ class_name }} from backend.app.{{ app_name }}.schema.{{ table_name }} import Create{{ schema_name }}Param, Delete{{ schema_name }}Param, Update{{ schema_name }}Param from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data class {{ class_name }}Service: @staticmethod - async def get(*, pk: int) -> {{ class_name }}: + async def get(*, db: AsyncSession, pk: int) -> {{ class_name }}: """ 获取{{ doc_comment }} + :param db: 数据库会话 :param pk: {{ doc_comment }} ID :return: """ - async with async_db_session() as db: - {{ table_name }} = await {{ table_name }}_dao.get(db, pk) - if not {{ table_name }}: - raise errors.NotFoundError(msg='{{ doc_comment }}不存在') - return {{ table_name }} + {{ table_name }} = await {{ table_name }}_dao.get(db, pk) + if not {{ table_name }}: + raise errors.NotFoundError(msg='{{ doc_comment }}不存在') + return {{ table_name }} @staticmethod - async def get_select() -> Select: - """获取{{ doc_comment }}查询对象""" - return await {{ table_name }}_dao.get_list() + async def get_list(db: AsyncSession) -> dict[str, Any]: + """ + 获取{{ doc_comment }}列表 + + :param db: 数据库会话 + :return: + """ + {{ table_name }}_select = await {{ table_name }}_dao.get_select() + return await paging_data(db, {{ table_name }}_select) @staticmethod - async def get_all() -> Sequence[{{ class_name }}]: - """获取所有{{ doc_comment }}""" - async with async_db_session() as db: - {{ table_name }}s = await {{ table_name }}_dao.get_all(db) - return {{ table_name }}s + async def get_all(*, db: AsyncSession) -> Sequence[{{ class_name }}]: + """ + 获取所有{{ doc_comment }} + + :param db: 数据库会话 + :return: + """ + {{ table_name }}s = await {{ table_name }}_dao.get_all(db) + return {{ table_name }}s @staticmethod - async def create(*, obj: Create{{ schema_name }}Param) -> None: + async def create(*, db: AsyncSession, obj: Create{{ schema_name }}Param) -> None: """ 创建{{ doc_comment }} + :param db: 数据库会话 :param obj: 创建{{ doc_comment }}参数 :return: """ - async with async_db_session.begin() as db: - await {{ table_name }}_dao.create(db, obj) + await {{ table_name }}_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: Update{{ schema_name }}Param) -> int: + async def update(*, db: AsyncSession, pk: int, obj: Update{{ schema_name }}Param) -> int: """ 更新{{ doc_comment }} + :param db: 数据库会话 :param pk: {{ doc_comment }} ID :param obj: 更新{{ doc_comment }}参数 :return: """ - async with async_db_session.begin() as db: - count = await {{ table_name }}_dao.update(db, pk, obj) - return count + count = await {{ table_name }}_dao.update(db, pk, obj) + return count @staticmethod - async def delete(*, obj: Delete{{ schema_name }}Param) -> int: + async def delete(*, db: AsyncSession, obj: Delete{{ schema_name }}Param) -> int: """ 删除{{ doc_comment }} + :param db: 数据库会话 :param obj: {{ doc_comment }} ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await {{ table_name }}_dao.delete(db, obj.pks) - return count + count = await {{ table_name }}_dao.delete(db, obj.pks) + return count {{ table_name }}_service: {{ class_name }}Service = {{ class_name }}Service() diff --git a/backend/plugin/config/api/v1/sys/config.py b/backend/plugin/config/api/v1/sys/config.py index 8a729960..efdcdd91 100644 --- a/backend/plugin/config/api/v1/sys/config.py +++ b/backend/plugin/config/api/v1/sys/config.py @@ -2,12 +2,12 @@ from typing import Annotated from fastapi import APIRouter, Body, Depends, Path, Query -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.config.schema.config import ( CreateConfigParam, GetConfigDetail, @@ -21,15 +21,18 @@ router = APIRouter() @router.get('/all', summary='获取所有参数配置', dependencies=[DependsJwtAuth]) async def get_all_configs( + db: CurrentSession, type: Annotated[str | None, Query(description='参数配置类型')] = None, ) -> ResponseSchemaModel[list[GetConfigDetail]]: - configs = await config_service.get_all(type=type) + configs = await config_service.get_all(db=db, type=type) return response_base.success(data=configs) @router.get('/{pk}', summary='获取参数配置详情', dependencies=[DependsJwtAuth]) -async def get_config(pk: Annotated[int, Path(description='参数配置 ID')]) -> ResponseSchemaModel[GetConfigDetail]: - config = await config_service.get(pk=pk) +async def get_config( + db: CurrentSession, pk: Annotated[int, Path(description='参数配置 ID')] +) -> ResponseSchemaModel[GetConfigDetail]: + config = await config_service.get(db=db, pk=pk) return response_base.success(data=config) @@ -41,13 +44,12 @@ async def get_config(pk: Annotated[int, Path(description='参数配置 ID')]) -> DependsPagination, ], ) -async def get_configs_paged( +async def get_configs_paginated( db: CurrentSession, name: Annotated[str | None, Query(description='参数配置名称')] = None, type: Annotated[str | None, Query(description='参数配置类型')] = None, ) -> ResponseSchemaModel[PageData[GetConfigDetail]]: - config_select = await config_service.get_select(name=name, type=type) - page_data = await paging_data(db, config_select) + page_data = await config_service.get_list(db=db, name=name, type=type) return response_base.success(data=page_data) @@ -59,14 +61,14 @@ async def get_configs_paged( DependsRBAC, ], ) -async def create_config(obj: CreateConfigParam) -> ResponseModel: - await config_service.create(obj=obj) +async def create_config(db: CurrentSessionTransaction, obj: CreateConfigParam) -> ResponseModel: + await config_service.create(db=db, obj=obj) return response_base.success() @router.put('', summary='批量更新参数配置', dependencies=[Depends(RequestPermission('sys.config.edits')), DependsRBAC]) -async def bulk_update_config(objs: list[UpdateConfigsParam]) -> ResponseModel: - count = await config_service.bulk_update(objs=objs) +async def bulk_update_config(db: CurrentSessionTransaction, objs: list[UpdateConfigsParam]) -> ResponseModel: + count = await config_service.bulk_update(db=db, objs=objs) if count > 0: return response_base.success() return response_base.fail() @@ -80,8 +82,10 @@ async def bulk_update_config(objs: list[UpdateConfigsParam]) -> ResponseModel: DependsRBAC, ], ) -async def update_config(pk: Annotated[int, Path(description='参数配置 ID')], obj: UpdateConfigParam) -> ResponseModel: - count = await config_service.update(pk=pk, obj=obj) +async def update_config( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='参数配置 ID')], obj: UpdateConfigParam +) -> ResponseModel: + count = await config_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -95,8 +99,10 @@ async def update_config(pk: Annotated[int, Path(description='参数配置 ID')], DependsRBAC, ], ) -async def delete_configs(pks: Annotated[list[int], Body(description='参数配置 ID 列表')]) -> ResponseModel: - count = await config_service.delete(pks=pks) +async def delete_configs( + db: CurrentSessionTransaction, pks: Annotated[list[int], Body(description='参数配置 ID 列表')] +) -> ResponseModel: + count = await config_service.delete(db=db, pks=pks) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/config/crud/crud_config.py b/backend/plugin/config/crud/crud_config.py index 6b451a0c..15651d2b 100644 --- a/backend/plugin/config/crud/crud_config.py +++ b/backend/plugin/config/crud/crud_config.py @@ -41,9 +41,9 @@ class CRUDConfig(CRUDPlus[Config]): """ return await self.select_model_by_column(db, key=key) - async def get_list(self, name: str | None, type: str | None) -> Select: + async def get_select(self, name: str | None, type: str | None) -> Select: """ - 获取参数配置列表 + 获取参数配置列表查询表达式 :param name: 参数配置名称 :param type: 参数配置类型 diff --git a/backend/plugin/config/service/config_service.py b/backend/plugin/config/service/config_service.py index 4565cdad..23c3e034 100644 --- a/backend/plugin/config/service/config_service.py +++ b/backend/plugin/config/service/config_service.py @@ -1,9 +1,10 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data from backend.plugin.config.crud.crud_config import config_dao from backend.plugin.config.model import Config from backend.plugin.config.schema.config import ( @@ -17,107 +18,115 @@ class ConfigService: """参数配置服务类""" @staticmethod - async def get(*, pk: int) -> Config: + async def get(*, db: AsyncSession, pk: int) -> Config: """ 获取参数配置详情 + :param db: 数据库会话 :param pk: 参数配置 ID :return: """ - async with async_db_session() as db: - config = await config_dao.get(db, pk) - if not config: - raise errors.NotFoundError(msg='参数配置不存在') - return config + + config = await config_dao.get(db, pk) + if not config: + raise errors.NotFoundError(msg='参数配置不存在') + return config @staticmethod - async def get_all(*, type: str | None) -> Sequence[Config | None]: + async def get_all(*, db: AsyncSession, type: str | None) -> Sequence[Config | None]: """ 获取所有参数配置 + :param db: 数据库会话 :param type: 参数配置类型 :return: """ - async with async_db_session() as db: - return await config_dao.get_all(db, type) + + return await config_dao.get_all(db, type) @staticmethod - async def get_select(*, name: str | None, type: str | None) -> Select: + async def get_list(*, db: AsyncSession, name: str | None, type: str | None) -> dict[str, Any]: """ - 获取参数配置列表查询条件 + 获取参数配置列表 + :param db: 数据库会话 :param name: 参数配置名称 :param type: 参数配置类型 :return: """ - return await config_dao.get_list(name=name, type=type) + config_select = await config_dao.get_select(name=name, type=type) + return await paging_data(db, config_select) @staticmethod - async def create(*, obj: CreateConfigParam) -> None: + async def create(*, db: AsyncSession, obj: CreateConfigParam) -> None: """ 创建参数配置 + :param db: 数据库会话 :param obj: 参数配置创建参数 :return: """ - async with async_db_session.begin() as db: - config = await config_dao.get_by_key(db, obj.key) - if config: - raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在') - await config_dao.create(db, obj) + + config = await config_dao.get_by_key(db, obj.key) + if config: + raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在') + await config_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateConfigParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateConfigParam) -> int: """ 更新参数配置 + :param db: 数据库会话 :param pk: 参数配置 ID :param obj: 参数配置更新参数 :return: """ - async with async_db_session.begin() as db: - config = await config_dao.get(db, pk) - if not config: - raise errors.NotFoundError(msg='参数配置不存在') - if config.key != obj.key: - config = await config_dao.get_by_key(db, obj.key) - if config: - raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在') - count = await config_dao.update(db, pk, obj) - return count + + config = await config_dao.get(db, pk) + if not config: + raise errors.NotFoundError(msg='参数配置不存在') + if config.key != obj.key: + config = await config_dao.get_by_key(db, obj.key) + if config: + raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在') + count = await config_dao.update(db, pk, obj) + return count @staticmethod - async def bulk_update(*, objs: list[UpdateConfigsParam]) -> int: + async def bulk_update(*, db: AsyncSession, objs: list[UpdateConfigsParam]) -> int: """ 批量更新参数配置 + :param db: 数据库会话 :param objs: 参数配置批量更新参数 :return: """ - async with async_db_session.begin() as db: - for _batch in range(0, len(objs), 1000): - for obj in objs: - config = await config_dao.get(db, obj.id) - if not config: - raise errors.NotFoundError(msg='参数配置不存在') - if config.key != obj.key: - config = await config_dao.get_by_key(db, obj.key) - if config: - raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在') - count = await config_dao.bulk_update(db, objs) - return count + + for _batch in range(0, len(objs), 1000): + for obj in objs: + config = await config_dao.get(db, obj.id) + if not config: + raise errors.NotFoundError(msg='参数配置不存在') + if config.key != obj.key: + config = await config_dao.get_by_key(db, obj.key) + if config: + raise errors.ConflictError(msg=f'参数配置 {obj.key} 已存在') + count = await config_dao.bulk_update(db, objs) + return count @staticmethod - async def delete(*, pks: list[int]) -> int: + async def delete(*, db: AsyncSession, pks: list[int]) -> int: """ 批量删除参数配置 + :param db: 数据库会话 :param pks: 参数配置 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await config_dao.delete(db, pks) - return count + + count = await config_dao.delete(db, pks) + return count config_service: ConfigService = ConfigService() diff --git a/backend/plugin/dict/api/v1/sys/dict_data.py b/backend/plugin/dict/api/v1/sys/dict_data.py index 319badb8..e7342aa8 100644 --- a/backend/plugin/dict/api/v1/sys/dict_data.py +++ b/backend/plugin/dict/api/v1/sys/dict_data.py @@ -2,12 +2,12 @@ from typing import Annotated from fastapi import APIRouter, Depends, Path, Query -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.dict.schema.dict_data import ( CreateDictDataParam, DeleteDictDataParam, @@ -20,24 +20,26 @@ router = APIRouter() @router.get('/all', summary='获取所有字典数据', dependencies=[DependsJwtAuth]) -async def get_all_dict_datas() -> ResponseSchemaModel[list[GetDictDataDetail]]: - data = await dict_data_service.get_all() +async def get_all_dict_datas(db: CurrentSession) -> ResponseSchemaModel[list[GetDictDataDetail]]: + data = await dict_data_service.get_all(db=db) return response_base.success(data=data) @router.get('/{pk}', summary='获取字典数据详情', dependencies=[DependsJwtAuth]) async def get_dict_data( + db: CurrentSession, pk: Annotated[int, Path(description='字典数据 ID')], ) -> ResponseSchemaModel[GetDictDataDetail]: - data = await dict_data_service.get(pk=pk) + data = await dict_data_service.get(db=db, pk=pk) return response_base.success(data=data) @router.get('/type-codes/{code}', summary='获取字典数据列表', dependencies=[DependsJwtAuth]) async def get_dict_data_by_type_code( + db: CurrentSession, code: Annotated[str, Path(description='字典类型编码')], ) -> ResponseSchemaModel[list[GetDictDataDetail]]: - data = await dict_data_service.get_by_type_code(code=code) + data = await dict_data_service.get_by_type_code(db=db, code=code) return response_base.success(data=data) @@ -49,7 +51,7 @@ async def get_dict_data_by_type_code( DependsPagination, ], ) -async def get_dict_datas_paged( +async def get_dict_datas_paginated( db: CurrentSession, type_code: Annotated[str | None, Query(description='字典类型编码')] = None, label: Annotated[str | None, Query(description='字典数据标签')] = None, @@ -57,14 +59,14 @@ async def get_dict_datas_paged( status: Annotated[int | None, Query(description='状态')] = None, type_id: Annotated[int | None, Query(description='字典类型 ID')] = None, ) -> ResponseSchemaModel[PageData[GetDictDataDetail]]: - dict_data_select = await dict_data_service.get_select( + page_data = await dict_data_service.get_list( + db=db, type_code=type_code, label=label, value=value, status=status, type_id=type_id, ) - page_data = await paging_data(db, dict_data_select) return response_base.success(data=page_data) @@ -76,8 +78,8 @@ async def get_dict_datas_paged( DependsRBAC, ], ) -async def create_dict_data(obj: CreateDictDataParam) -> ResponseModel: - await dict_data_service.create(obj=obj) +async def create_dict_data(db: CurrentSessionTransaction, obj: CreateDictDataParam) -> ResponseModel: + await dict_data_service.create(db=db, obj=obj) return response_base.success() @@ -90,10 +92,11 @@ async def create_dict_data(obj: CreateDictDataParam) -> ResponseModel: ], ) async def update_dict_data( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='字典数据 ID')], obj: UpdateDictDataParam, ) -> ResponseModel: - count = await dict_data_service.update(pk=pk, obj=obj) + count = await dict_data_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -107,8 +110,8 @@ async def update_dict_data( DependsRBAC, ], ) -async def delete_dict_datas(obj: DeleteDictDataParam) -> ResponseModel: - count = await dict_data_service.delete(obj=obj) +async def delete_dict_datas(db: CurrentSessionTransaction, obj: DeleteDictDataParam) -> ResponseModel: + count = await dict_data_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/dict/api/v1/sys/dict_type.py b/backend/plugin/dict/api/v1/sys/dict_type.py index 6c9d8fc7..5f0df875 100644 --- a/backend/plugin/dict/api/v1/sys/dict_type.py +++ b/backend/plugin/dict/api/v1/sys/dict_type.py @@ -2,12 +2,12 @@ from typing import Annotated from fastapi import APIRouter, Depends, Path, Query -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.dict.schema.dict_type import ( CreateDictTypeParam, DeleteDictTypeParam, @@ -20,16 +20,17 @@ router = APIRouter() @router.get('/all', summary='获取所有字典数据', dependencies=[DependsJwtAuth]) -async def get_all_dict_types() -> ResponseSchemaModel[list[GetDictTypeDetail]]: - data = await dict_type_service.get_all() +async def get_all_dict_types(db: CurrentSession) -> ResponseSchemaModel[list[GetDictTypeDetail]]: + data = await dict_type_service.get_all(db=db) return response_base.success(data=data) @router.get('/{pk}', summary='获取字典类型详情', dependencies=[DependsJwtAuth]) async def get_dict_type( + db: CurrentSession, pk: Annotated[int, Path(description='字典类型 ID')], ) -> ResponseSchemaModel[GetDictTypeDetail]: - data = await dict_type_service.get(pk=pk) + data = await dict_type_service.get(db=db, pk=pk) return response_base.success(data=data) @@ -41,13 +42,12 @@ async def get_dict_type( DependsPagination, ], ) -async def get_dict_types_paged( +async def get_dict_types_paginated( db: CurrentSession, name: Annotated[str | None, Query(description='字典类型名称')] = None, code: Annotated[str | None, Query(description='字典类型编码')] = None, ) -> ResponseSchemaModel[PageData[GetDictTypeDetail]]: - dict_type_select = await dict_type_service.get_select(name=name, code=code) - page_data = await paging_data(db, dict_type_select) + page_data = await dict_type_service.get_list(db=db, name=name, code=code) return response_base.success(data=page_data) @@ -59,8 +59,8 @@ async def get_dict_types_paged( DependsRBAC, ], ) -async def create_dict_type(obj: CreateDictTypeParam) -> ResponseModel: - await dict_type_service.create(obj=obj) +async def create_dict_type(db: CurrentSessionTransaction, obj: CreateDictTypeParam) -> ResponseModel: + await dict_type_service.create(db=db, obj=obj) return response_base.success() @@ -73,10 +73,11 @@ async def create_dict_type(obj: CreateDictTypeParam) -> ResponseModel: ], ) async def update_dict_type( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='字典类型 ID')], obj: UpdateDictTypeParam, ) -> ResponseModel: - count = await dict_type_service.update(pk=pk, obj=obj) + count = await dict_type_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -90,8 +91,8 @@ async def update_dict_type( DependsRBAC, ], ) -async def delete_dict_types(obj: DeleteDictTypeParam) -> ResponseModel: - count = await dict_type_service.delete(obj=obj) +async def delete_dict_types(db: CurrentSessionTransaction, obj: DeleteDictTypeParam) -> ResponseModel: + count = await dict_type_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/dict/crud/crud_dict_data.py b/backend/plugin/dict/crud/crud_dict_data.py index 123dda13..13ee0950 100644 --- a/backend/plugin/dict/crud/crud_dict_data.py +++ b/backend/plugin/dict/crud/crud_dict_data.py @@ -46,7 +46,7 @@ class CRUDDictData(CRUDPlus[DictData]): """ return await self.select_models(db, load_strategies={'type': 'noload'}) - async def get_list( + async def get_select( self, type_code: str | None, label: str | None, @@ -55,7 +55,7 @@ class CRUDDictData(CRUDPlus[DictData]): type_id: int | None, ) -> Select: """ - 获取字典数据列表 + 获取字典数据列表查询表达式 :param type_code: 字典类型编码 :param label: 字典数据标签 diff --git a/backend/plugin/dict/crud/crud_dict_type.py b/backend/plugin/dict/crud/crud_dict_type.py index d37d7da8..592cf61d 100644 --- a/backend/plugin/dict/crud/crud_dict_type.py +++ b/backend/plugin/dict/crud/crud_dict_type.py @@ -30,9 +30,9 @@ class CRUDDictType(CRUDPlus[DictType]): """ return await self.select_models(db, load_strategies={'datas': 'noload'}) - async def get_list(self, *, name: str | None, code: str | None) -> Select: + async def get_select(self, *, name: str | None, code: str | None) -> Select: """ - 获取字典类型列表 + 获取字典类型列表查询表达式 :param name: 字典类型名称 :param code: 字典类型编码 diff --git a/backend/plugin/dict/service/dict_data_service.py b/backend/plugin/dict/service/dict_data_service.py index bacf8808..8ec9f6f8 100644 --- a/backend/plugin/dict/service/dict_data_service.py +++ b/backend/plugin/dict/service/dict_data_service.py @@ -1,9 +1,10 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data from backend.plugin.dict.crud.crud_dict_data import dict_data_dao from backend.plugin.dict.crud.crud_dict_type import dict_type_dao from backend.plugin.dict.model import DictData @@ -14,51 +15,60 @@ class DictDataService: """字典数据服务类""" @staticmethod - async def get(*, pk: int) -> DictData: + async def get(*, db: AsyncSession, pk: int) -> DictData: """ 获取字典数据详情 + :param db: 数据库会话 :param pk: 字典数据 ID :return: """ - async with async_db_session() as db: - dict_data = await dict_data_dao.get(db, pk) - if not dict_data: - raise errors.NotFoundError(msg='字典数据不存在') - return dict_data + + dict_data = await dict_data_dao.get(db, pk) + if not dict_data: + raise errors.NotFoundError(msg='字典数据不存在') + return dict_data @staticmethod - async def get_by_type_code(*, code: str) -> Sequence[DictData]: + async def get_by_type_code(*, db: AsyncSession, code: str) -> Sequence[DictData]: """ 获取字典数据详情 + :param db: 数据库会话 :param code: 字典类型编码 :return: """ - async with async_db_session() as db: - dict_datas = await dict_data_dao.get_by_type_code(db, code) - if not dict_datas: - raise errors.NotFoundError(msg='字典数据不存在') - return dict_datas + + dict_datas = await dict_data_dao.get_by_type_code(db, code) + if not dict_datas: + raise errors.NotFoundError(msg='字典数据不存在') + return dict_datas @staticmethod - async def get_all() -> Sequence[DictData]: - async with async_db_session() as db: - dict_datas = await dict_data_dao.get_all(db) - return dict_datas + async def get_all(*, db: AsyncSession) -> Sequence[DictData]: + """ + 获取所有字典数据 + + :param db: 数据库会话 + :return: + """ + dict_datas = await dict_data_dao.get_all(db) + return dict_datas @staticmethod - async def get_select( + async def get_list( *, + db: AsyncSession, type_code: str | None, label: str | None, value: str | None, status: int | None, type_id: int | None, - ) -> Select: + ) -> dict[str, Any]: """ - 获取字典数据列表查询条件 + 获取字典数据列表 + :param db: 数据库会话 :param type_code: 字典类型编码 :param label: 字典数据标签 :param value: 字典数据键值 @@ -66,63 +76,67 @@ class DictDataService: :param type_id: 字典类型 ID :return: """ - return await dict_data_dao.get_list( + dict_data_select = await dict_data_dao.get_select( type_code=type_code, label=label, value=value, status=status, type_id=type_id, ) + return await paging_data(db, dict_data_select) @staticmethod - async def create(*, obj: CreateDictDataParam) -> None: + async def create(*, db: AsyncSession, obj: CreateDictDataParam) -> None: """ 创建字典数据 + :param db: 数据库会话 :param obj: 字典数据创建参数 :return: """ - async with async_db_session.begin() as db: - dict_data = await dict_data_dao.get_by_label(db, obj.label) - if dict_data: - raise errors.ConflictError(msg='字典数据已存在') - dict_type = await dict_type_dao.get(db, obj.type_id) - if not dict_type: - raise errors.NotFoundError(msg='字典类型不存在') - await dict_data_dao.create(db, obj, dict_type.code) + + dict_data = await dict_data_dao.get_by_label(db, obj.label) + if dict_data: + raise errors.ConflictError(msg='字典数据已存在') + dict_type = await dict_type_dao.get(db, obj.type_id) + if not dict_type: + raise errors.NotFoundError(msg='字典类型不存在') + await dict_data_dao.create(db, obj, dict_type.code) @staticmethod - async def update(*, pk: int, obj: UpdateDictDataParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateDictDataParam) -> int: """ 更新字典数据 + :param db: 数据库会话 :param pk: 字典数据 ID :param obj: 字典数据更新参数 :return: """ - async with async_db_session.begin() as db: - dict_data = await dict_data_dao.get(db, pk) - if not dict_data: - raise errors.NotFoundError(msg='字典数据不存在') - if dict_data.label != obj.label and await dict_data_dao.get_by_label(db, obj.label): - raise errors.ConflictError(msg='字典数据已存在') - dict_type = await dict_type_dao.get(db, obj.type_id) - if not dict_type: - raise errors.NotFoundError(msg='字典类型不存在') - count = await dict_data_dao.update(db, pk, obj, dict_type.code) - return count + + dict_data = await dict_data_dao.get(db, pk) + if not dict_data: + raise errors.NotFoundError(msg='字典数据不存在') + if dict_data.label != obj.label and await dict_data_dao.get_by_label(db, obj.label): + raise errors.ConflictError(msg='字典数据已存在') + dict_type = await dict_type_dao.get(db, obj.type_id) + if not dict_type: + raise errors.NotFoundError(msg='字典类型不存在') + count = await dict_data_dao.update(db, pk, obj, dict_type.code) + return count @staticmethod - async def delete(*, obj: DeleteDictDataParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteDictDataParam) -> int: """ 批量删除字典数据 + :param db: 数据库会话 :param obj: 字典数据 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await dict_data_dao.delete(db, obj.pks) - return count + + count = await dict_data_dao.delete(db, obj.pks) + return count dict_data_service: DictDataService = DictDataService() diff --git a/backend/plugin/dict/service/dict_type_service.py b/backend/plugin/dict/service/dict_type_service.py index fb0fc713..584e3e2a 100644 --- a/backend/plugin/dict/service/dict_type_service.py +++ b/backend/plugin/dict/service/dict_type_service.py @@ -1,9 +1,10 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data from backend.plugin.dict.crud.crud_dict_type import dict_type_dao from backend.plugin.dict.model import DictType from backend.plugin.dict.schema.dict_type import CreateDictTypeParam, DeleteDictTypeParam, UpdateDictTypeParam @@ -13,79 +14,90 @@ class DictTypeService: """字典类型服务类""" @staticmethod - async def get(*, pk: int) -> DictType: + async def get(*, db: AsyncSession, pk: int) -> DictType: """ 获取字典类型详情 + :param db: 数据库会话 :param pk: 字典类型 ID :return: """ - async with async_db_session() as db: - dict_type = await dict_type_dao.get(db, pk) - if not dict_type: - raise errors.NotFoundError(msg='字典类型不存在') - return dict_type + + dict_type = await dict_type_dao.get(db, pk) + if not dict_type: + raise errors.NotFoundError(msg='字典类型不存在') + return dict_type @staticmethod - async def get_all() -> Sequence[DictType]: - async with async_db_session() as db: - dict_datas = await dict_type_dao.get_all(db) - return dict_datas - - @staticmethod - async def get_select(*, name: str | None, code: str | None) -> Select: + async def get_all(*, db: AsyncSession) -> Sequence[DictType]: """ - 获取字典类型列表查询条件 + 获取所有字典类型 + :param db: 数据库会话 + :return: + """ + dict_datas = await dict_type_dao.get_all(db) + return dict_datas + + @staticmethod + async def get_list(*, db: AsyncSession, name: str | None, code: str | None) -> dict[str, Any]: + """ + 获取字典类型列表 + + :param db: 数据库会话 :param name: 字典类型名称 :param code: 字典类型编码 :return: """ - return await dict_type_dao.get_list(name=name, code=code) + dict_type_select = await dict_type_dao.get_select(name=name, code=code) + return await paging_data(db, dict_type_select) @staticmethod - async def create(*, obj: CreateDictTypeParam) -> None: + async def create(*, db: AsyncSession, obj: CreateDictTypeParam) -> None: """ 创建字典类型 + :param db: 数据库会话 :param obj: 字典类型创建参数 :return: """ - async with async_db_session.begin() as db: - dict_type = await dict_type_dao.get_by_code(db, obj.code) - if dict_type: - raise errors.ConflictError(msg='字典类型已存在') - await dict_type_dao.create(db, obj) + + dict_type = await dict_type_dao.get_by_code(db, obj.code) + if dict_type: + raise errors.ConflictError(msg='字典类型已存在') + await dict_type_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateDictTypeParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateDictTypeParam) -> int: """ 更新字典类型 + :param db: 数据库会话 :param pk: 字典类型 ID :param obj: 字典类型更新参数 :return: """ - async with async_db_session.begin() as db: - dict_type = await dict_type_dao.get(db, pk) - if not dict_type: - raise errors.NotFoundError(msg='字典类型不存在') - if dict_type.code != obj.code and await dict_type_dao.get_by_code(db, obj.code): - raise errors.ConflictError(msg='字典类型已存在') - count = await dict_type_dao.update(db, pk, obj) - return count + + dict_type = await dict_type_dao.get(db, pk) + if not dict_type: + raise errors.NotFoundError(msg='字典类型不存在') + if dict_type.code != obj.code and await dict_type_dao.get_by_code(db, obj.code): + raise errors.ConflictError(msg='字典类型已存在') + count = await dict_type_dao.update(db, pk, obj) + return count @staticmethod - async def delete(*, obj: DeleteDictTypeParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteDictTypeParam) -> int: """ 批量删除字典类型 + :param db: 数据库会话 :param obj: 字典类型 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await dict_type_dao.delete(db, obj.pks) - return count + + count = await dict_type_dao.delete(db, obj.pks) + return count dict_type_service: DictTypeService = DictTypeService() diff --git a/backend/plugin/email/api/v1/email.py b/backend/plugin/email/api/v1/email.py index e4ba2578..5652c72d 100644 --- a/backend/plugin/email/api/v1/email.py +++ b/backend/plugin/email/api/v1/email.py @@ -2,7 +2,7 @@ import random from typing import Annotated -from fastapi import APIRouter, Body, Request +from fastapi import APIRouter, Body from backend.common.context import ctx from backend.common.response.response_schema import ResponseModel, response_base @@ -17,7 +17,6 @@ router = APIRouter() @router.post('/captcha', summary='发送电子邮件验证码', dependencies=[DependsJwtAuth]) async def send_email_captcha( - request: Request, db: CurrentSession, recipients: Annotated[str | list[str], Body(embed=True, description='邮件接收者')], ) -> ResponseModel: diff --git a/backend/plugin/notice/api/v1/sys/notice.py b/backend/plugin/notice/api/v1/sys/notice.py index aeb9fa5c..f33a99ac 100644 --- a/backend/plugin/notice/api/v1/sys/notice.py +++ b/backend/plugin/notice/api/v1/sys/notice.py @@ -2,12 +2,12 @@ from typing import Annotated from fastapi import APIRouter, Depends, Path, Query -from backend.common.pagination import DependsPagination, PageData, paging_data +from backend.common.pagination import DependsPagination, PageData from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.database.db import CurrentSession +from backend.database.db import CurrentSession, CurrentSessionTransaction from backend.plugin.notice.schema.notice import ( CreateNoticeParam, DeleteNoticeParam, @@ -20,8 +20,10 @@ router = APIRouter() @router.get('/{pk}', summary='获取通知公告详情', dependencies=[DependsJwtAuth]) -async def get_notice(pk: Annotated[int, Path(description='通知公告 ID')]) -> ResponseSchemaModel[GetNoticeDetail]: - notice = await notice_service.get(pk=pk) +async def get_notice( + db: CurrentSession, pk: Annotated[int, Path(description='通知公告 ID')] +) -> ResponseSchemaModel[GetNoticeDetail]: + notice = await notice_service.get(db=db, pk=pk) return response_base.success(data=notice) @@ -33,14 +35,13 @@ async def get_notice(pk: Annotated[int, Path(description='通知公告 ID')]) -> DependsPagination, ], ) -async def get_notices_paged( +async def get_notices_paginated( db: CurrentSession, title: Annotated[str | None, Query(description='标题')] = None, type: Annotated[int | None, Query(description='类型')] = None, status: Annotated[int | None, Query(description='状态')] = None, ) -> ResponseSchemaModel[PageData[GetNoticeDetail]]: - notice_select = await notice_service.get_select(title=title, type=type, status=status) - page_data = await paging_data(db, notice_select) + page_data = await notice_service.get_list(db=db, title=title, type=type, status=status) return response_base.success(data=page_data) @@ -52,8 +53,8 @@ async def get_notices_paged( DependsRBAC, ], ) -async def create_notice(obj: CreateNoticeParam) -> ResponseModel: - await notice_service.create(obj=obj) +async def create_notice(db: CurrentSessionTransaction, obj: CreateNoticeParam) -> ResponseModel: + await notice_service.create(db=db, obj=obj) return response_base.success() @@ -65,8 +66,10 @@ async def create_notice(obj: CreateNoticeParam) -> ResponseModel: DependsRBAC, ], ) -async def update_notice(pk: Annotated[int, Path(description='通知公告 ID')], obj: UpdateNoticeParam) -> ResponseModel: - count = await notice_service.update(pk=pk, obj=obj) +async def update_notice( + db: CurrentSessionTransaction, pk: Annotated[int, Path(description='通知公告 ID')], obj: UpdateNoticeParam +) -> ResponseModel: + count = await notice_service.update(db=db, pk=pk, obj=obj) if count > 0: return response_base.success() return response_base.fail() @@ -80,8 +83,8 @@ async def update_notice(pk: Annotated[int, Path(description='通知公告 ID')], DependsRBAC, ], ) -async def delete_notices(obj: DeleteNoticeParam) -> ResponseModel: - count = await notice_service.delete(obj=obj) +async def delete_notices(db: CurrentSessionTransaction, obj: DeleteNoticeParam) -> ResponseModel: + count = await notice_service.delete(db=db, obj=obj) if count > 0: return response_base.success() return response_base.fail() diff --git a/backend/plugin/notice/crud/crud_notice.py b/backend/plugin/notice/crud/crud_notice.py index 391caef5..19d54433 100644 --- a/backend/plugin/notice/crud/crud_notice.py +++ b/backend/plugin/notice/crud/crud_notice.py @@ -21,9 +21,9 @@ class CRUDNotice(CRUDPlus[Notice]): """ return await self.select_model(db, pk) - async def get_list(self, title: str, type: int | None, status: int | None) -> Select: + async def get_select(self, title: str, type: int | None, status: int | None) -> Select: """ - 获取通知公告列表 + 获取通知公告列表查询表达式 :param title: 通知公告标题 :param type: 通知公告类型 diff --git a/backend/plugin/notice/service/notice_service.py b/backend/plugin/notice/service/notice_service.py index 058ad7a7..c51818da 100644 --- a/backend/plugin/notice/service/notice_service.py +++ b/backend/plugin/notice/service/notice_service.py @@ -1,9 +1,10 @@ from collections.abc import Sequence +from typing import Any -from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession from backend.common.exception import errors -from backend.database.db import async_db_session +from backend.common.pagination import paging_data from backend.plugin.notice.crud.crud_notice import notice_dao from backend.plugin.notice.model import Notice from backend.plugin.notice.schema.notice import CreateNoticeParam, DeleteNoticeParam, UpdateNoticeParam @@ -13,69 +14,87 @@ class NoticeService: """通知公告服务类""" @staticmethod - async def get(*, pk: int) -> Notice: + async def get(*, db: AsyncSession, pk: int) -> Notice: """ 获取通知公告 + :param db: 数据库会话 :param pk: 通知公告 ID :return: """ - async with async_db_session() as db: - notice = await notice_dao.get(db, pk) - if not notice: - raise errors.NotFoundError(msg='通知公告不存在') - return notice + + notice = await notice_dao.get(db, pk) + if not notice: + raise errors.NotFoundError(msg='通知公告不存在') + return notice @staticmethod - async def get_select(title: str | None, type: int | None, status: int | None) -> Select: - """获取通知公告查询对象""" - return await notice_dao.get_list(title, type, status) + async def get_list(db: AsyncSession, title: str | None, type: int | None, status: int | None) -> dict[str, Any]: + """ + 获取通知公告列表 + + :param db: 数据库会话 + :param title: 通知公告标题 + :param type: 通知公告类型 + :param status: 通知公告状态 + :return: + """ + notice_select = await notice_dao.get_select(title, type, status) + return await paging_data(db, notice_select) @staticmethod - async def get_all() -> Sequence[Notice]: - """获取所有通知公告""" - async with async_db_session() as db: - notices = await notice_dao.get_all(db) - return notices + async def get_all(*, db: AsyncSession) -> Sequence[Notice]: + """ + 获取所有通知公告 + + :param db: 数据库会话 + :return: + """ + + notices = await notice_dao.get_all(db) + return notices @staticmethod - async def create(*, obj: CreateNoticeParam) -> None: + async def create(*, db: AsyncSession, obj: CreateNoticeParam) -> None: """ 创建通知公告 + :param db: 数据库会话 :param obj: 创建通知公告参数 :return: """ - async with async_db_session.begin() as db: - await notice_dao.create(db, obj) + + await notice_dao.create(db, obj) @staticmethod - async def update(*, pk: int, obj: UpdateNoticeParam) -> int: + async def update(*, db: AsyncSession, pk: int, obj: UpdateNoticeParam) -> int: """ 更新通知公告 + :param db: 数据库会话 :param pk: 通知公告 ID :param obj: 更新通知公告参数 :return: """ - async with async_db_session.begin() as db: - notice = await notice_dao.get(db, pk) - if not notice: - raise errors.NotFoundError(msg='通知公告不存在') - count = await notice_dao.update(db, pk, obj) - return count + + notice = await notice_dao.get(db, pk) + if not notice: + raise errors.NotFoundError(msg='通知公告不存在') + count = await notice_dao.update(db, pk, obj) + return count @staticmethod - async def delete(*, obj: DeleteNoticeParam) -> int: + async def delete(*, db: AsyncSession, obj: DeleteNoticeParam) -> int: """ 批量删除通知公告 + :param db: 数据库会话 :param obj: 通知公告 ID 列表 :return: """ - async with async_db_session.begin() as db: - count = await notice_dao.delete(db, obj.pks) - return count + + count = await notice_dao.delete(db, obj.pks) + return count notice_service: NoticeService = NoticeService() diff --git a/backend/plugin/oauth2/api/v1/github.py b/backend/plugin/oauth2/api/v1/github.py index bca49317..06f3466b 100644 --- a/backend/plugin/oauth2/api/v1/github.py +++ b/backend/plugin/oauth2/api/v1/github.py @@ -8,6 +8,7 @@ from starlette.responses import RedirectResponse from backend.common.enums import UserSocialType from backend.common.response.response_schema import ResponseSchemaModel, response_base from backend.core.conf import settings +from backend.database.db import CurrentSessionTransaction from backend.plugin.oauth2.service.oauth2_service import oauth2_service router = APIRouter() @@ -28,7 +29,7 @@ async def get_github_oauth2_url(request: Request) -> ResponseSchemaModel[str]: dependencies=[Depends(RateLimiter(times=5, minutes=1))], ) async def github_oauth2_callback( # noqa: ANN201 - request: Request, + db: CurrentSessionTransaction, response: Response, background_tasks: BackgroundTasks, oauth2: Annotated[ @@ -40,7 +41,7 @@ async def github_oauth2_callback( # noqa: ANN201 access_token = token['access_token'] user = await github_client.get_userinfo(access_token) data = await oauth2_service.create_with_login( - request=request, + db=db, response=response, background_tasks=background_tasks, user=user, diff --git a/backend/plugin/oauth2/api/v1/google.py b/backend/plugin/oauth2/api/v1/google.py index 8871751b..2be4f268 100644 --- a/backend/plugin/oauth2/api/v1/google.py +++ b/backend/plugin/oauth2/api/v1/google.py @@ -8,6 +8,7 @@ from starlette.responses import RedirectResponse from backend.common.enums import UserSocialType from backend.common.response.response_schema import ResponseSchemaModel, response_base from backend.core.conf import settings +from backend.database.db import CurrentSessionTransaction from backend.plugin.oauth2.service.oauth2_service import oauth2_service router = APIRouter() @@ -28,7 +29,7 @@ async def get_google_oauth2_url(request: Request) -> ResponseSchemaModel[str]: dependencies=[Depends(RateLimiter(times=5, minutes=1))], ) async def google_oauth2_callback( # noqa: ANN201 - request: Request, + db: CurrentSessionTransaction, response: Response, background_tasks: BackgroundTasks, oauth2: Annotated[ @@ -40,7 +41,7 @@ async def google_oauth2_callback( # noqa: ANN201 access_token = token['access_token'] user = await google_client.get_userinfo(access_token) data = await oauth2_service.create_with_login( - request=request, + db=db, response=response, background_tasks=background_tasks, user=user, diff --git a/backend/plugin/oauth2/api/v1/linux_do.py b/backend/plugin/oauth2/api/v1/linux_do.py index 2549b3b5..627a5ec2 100644 --- a/backend/plugin/oauth2/api/v1/linux_do.py +++ b/backend/plugin/oauth2/api/v1/linux_do.py @@ -8,6 +8,7 @@ from starlette.responses import RedirectResponse from backend.common.enums import UserSocialType from backend.common.response.response_schema import ResponseSchemaModel, response_base from backend.core.conf import settings +from backend.database.db import CurrentSessionTransaction from backend.plugin.oauth2.service.oauth2_service import oauth2_service router = APIRouter() @@ -28,7 +29,7 @@ async def get_linux_do_oauth2_url(request: Request) -> ResponseSchemaModel[str]: dependencies=[Depends(RateLimiter(times=5, minutes=1))], ) async def linux_do_oauth2_callback( # noqa: ANN201 - request: Request, + db: CurrentSessionTransaction, response: Response, background_tasks: BackgroundTasks, oauth2: Annotated[ @@ -40,7 +41,7 @@ async def linux_do_oauth2_callback( # noqa: ANN201 access_token = token['access_token'] user = await linux_do_client.get_userinfo(access_token) data = await oauth2_service.create_with_login( - request=request, + db=db, response=response, background_tasks=background_tasks, user=user, diff --git a/backend/plugin/oauth2/service/oauth2_service.py b/backend/plugin/oauth2/service/oauth2_service.py index 166f2ac0..a7fffc8d 100644 --- a/backend/plugin/oauth2/service/oauth2_service.py +++ b/backend/plugin/oauth2/service/oauth2_service.py @@ -1,17 +1,18 @@ from typing import Any from fast_captcha import text_captcha -from fastapi import BackgroundTasks, Request, Response +from fastapi import BackgroundTasks, Response +from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_user import user_dao from backend.app.admin.schema.token import GetLoginToken from backend.app.admin.schema.user import AddOAuth2UserParam from backend.app.admin.service.login_log_service import login_log_service +from backend.common.context import ctx from backend.common.enums import LoginLogStatusType, UserSocialType from backend.common.i18n import t from backend.common.security import jwt from backend.core.conf import settings -from backend.database.db import async_db_session from backend.database.redis import redis_client from backend.plugin.oauth2.crud.crud_user_social import user_social_dao from backend.plugin.oauth2.schema.user_social import CreateUserSocialParam @@ -24,7 +25,7 @@ class OAuth2Service: @staticmethod async def create_with_login( *, - request: Request, + db: AsyncSession, response: Response, background_tasks: BackgroundTasks, user: dict[str, Any], @@ -33,111 +34,110 @@ class OAuth2Service: """ 创建 OAuth2 用户并登录 - :param request: FastAPI 请求对象 + :param db: 数据库会话 :param response: FastAPI 响应对象 :param background_tasks: FastAPI 后台任务 :param user: OAuth2 用户信息 :param social: 社交平台类型 :return: """ - async with async_db_session.begin() as db: - sid = user.get('uuid') - username = user.get('username') - nickname = user.get('nickname') - email = user.get('email') - avatar = user.get('avatar_url') - if social == UserSocialType.github: - sid = user.get('id') - username = user.get('login') - nickname = user.get('name') + sid = user.get('uuid') + username = user.get('username') + nickname = user.get('nickname') + email = user.get('email') + avatar = user.get('avatar_url') - if social == UserSocialType.google: - sid = user.get('id') - username = user.get('name') - nickname = user.get('given_name') - avatar = user.get('picture') + if social == UserSocialType.github: + sid = user.get('id') + username = user.get('login') + nickname = user.get('name') - if social == UserSocialType.linux_do: - sid = user.get('id') - nickname = user.get('name') + if social == UserSocialType.google: + sid = user.get('id') + username = user.get('name') + nickname = user.get('given_name') + avatar = user.get('picture') - user_social = await user_social_dao.get_by_sid(db, str(sid), str(social.value)) - if user_social: - sys_user = await user_dao.get(db, user_social.user_id) - # 更新用户头像 - if not sys_user.avatar and avatar is not None: - await user_dao.update_avatar(db, sys_user.id, avatar) - else: - sys_user = None - # 检测系统用户是否已存在 - if email: - sys_user = await user_dao.check_email(db, email) # 通过邮箱验证绑定保证邮箱真实性 + if social == UserSocialType.linux_do: + sid = user.get('id') + nickname = user.get('name') - # 创建系统用户 - if not sys_user: - while await user_dao.get_by_username(db, username): - username = f'{username}_{text_captcha(5)}' - new_sys_user = AddOAuth2UserParam( - username=username, - password=None, - nickname=nickname, - email=email, - avatar=avatar, - ) - await user_dao.add_by_oauth2(db, new_sys_user) - await db.flush() - sys_user = await user_dao.get_by_username(db, username) + user_social = await user_social_dao.get_by_sid(db, str(sid), str(social.value)) + if user_social: + sys_user = await user_dao.get(db, user_social.user_id) + # 更新用户头像 + if not sys_user.avatar and avatar is not None: + await user_dao.update_avatar(db, sys_user.id, avatar) + else: + sys_user = None + # 检测系统用户是否已存在 + if email: + sys_user = await user_dao.check_email(db, email) # 通过邮箱验证绑定保证邮箱真实性 - # 绑定社交账号 - new_user_social = CreateUserSocialParam(sid=str(sid), source=social.value, user_id=sys_user.id) - await user_social_dao.create(db, new_user_social) + # 创建系统用户 + if not sys_user: + while await user_dao.get_by_username(db, username): + username = f'{username}_{text_captcha(5)}' + new_sys_user = AddOAuth2UserParam( + username=username, + password=None, + nickname=nickname, + email=email, + avatar=avatar, + ) + await user_dao.add_by_oauth2(db, new_sys_user) + await db.flush() + sys_user = await user_dao.get_by_username(db, username) - # 创建 token - access_token = await jwt.create_access_token( - sys_user.id, - multi_login=sys_user.is_multi_login, - # extra info - username=sys_user.username, - nickname=sys_user.nickname or f'#{text_captcha(5)}', - last_login_time=timezone.to_str(timezone.now()), - ip=request.state.ip, - os=request.state.os, - browser=request.state.browser, - device=request.state.device, - ) - refresh_token = await jwt.create_refresh_token( - access_token.session_uuid, - sys_user.id, - multi_login=sys_user.is_multi_login, - ) - await user_dao.update_login_time(db, sys_user.username) - await db.refresh(sys_user) - background_tasks.add_task( - login_log_service.create, - db=db, - request=request, - user_uuid=sys_user.uuid, - username=sys_user.username, - login_time=timezone.now(), - status=LoginLogStatusType.success.value, - msg=t('success.login.oauth2_success'), - ) - await redis_client.delete(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}') - response.set_cookie( - key=settings.COOKIE_REFRESH_TOKEN_KEY, - value=refresh_token.refresh_token, - max_age=settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, - expires=timezone.to_utc(refresh_token.refresh_token_expire_time), - httponly=True, - ) - data = GetLoginToken( - access_token=access_token.access_token, - access_token_expire_time=access_token.access_token_expire_time, - session_uuid=access_token.session_uuid, - user=sys_user, # type: ignore - ) - return data + # 绑定社交账号 + new_user_social = CreateUserSocialParam(sid=str(sid), source=social.value, user_id=sys_user.id) + await user_social_dao.create(db, new_user_social) + + # 创建 token + access_token = await jwt.create_access_token( + sys_user.id, + multi_login=sys_user.is_multi_login, + # extra info + username=sys_user.username, + nickname=sys_user.nickname or f'#{text_captcha(5)}', + last_login_time=timezone.to_str(timezone.now()), + ip=ctx.ip, + os=ctx.os, + browser=ctx.browser, + device=ctx.device, + ) + refresh_token = await jwt.create_refresh_token( + access_token.session_uuid, + sys_user.id, + multi_login=sys_user.is_multi_login, + ) + await user_dao.update_login_time(db, sys_user.username) + await db.refresh(sys_user) + background_tasks.add_task( + login_log_service.create, + db=db, + user_uuid=sys_user.uuid, + username=sys_user.username, + login_time=timezone.now(), + status=LoginLogStatusType.success.value, + msg=t('success.login.oauth2_success'), + ) + await redis_client.delete(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{ctx.ip}') + response.set_cookie( + key=settings.COOKIE_REFRESH_TOKEN_KEY, + value=refresh_token.refresh_token, + max_age=settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, + expires=timezone.to_utc(refresh_token.refresh_token_expire_time), + httponly=True, + ) + data = GetLoginToken( + access_token=access_token.access_token, + access_token_expire_time=access_token.access_token_expire_time, + session_uuid=access_token.session_uuid, + user=sys_user, # type: ignore + ) + return data oauth2_service: OAuth2Service = OAuth2Service()