Compare commits

...
87 Commits
Author SHA1 Message Date
Wu Clan eeb709c6aa Update the version number to 1.8.0 (#771) 2025-08-15 20:06:21 +08:00
Wu Clan a6bbf2971d Update the menu title in SQL scripts (#770) 2025-08-15 19:58:46 +08:00
Dylan cd48bb4210 Add i18n support for response message (#753)
* feat: i18n support

* Optimize i18n

* Update the locale in the code

* Update the zh-CN file

* Update the en-US file

* Update the reload filter

* Update locale success plugin value

* Fix lint

* Update pydantic error message translation

* Update to minimal implementation

* Fix minimal missing code
2025-08-15 19:57:10 +08:00
Wu Clan 4500dd0128 Add a standalone email sending plugin (#769) 2025-08-13 18:09:42 +08:00
Wu Clan 2b6d8222ad Update the content layout of the config file (#768) 2025-08-13 18:07:02 +08:00
Wu Clan 1b47ab7e83 Optimize the timezone datetime return encoder (#767)
* Optimize the timezone datetime return encoder

* Update the default datetime conversion
2025-08-13 11:16:58 +08:00
Wu Clan bd804e0a38 Update the description for the run file (#766) 2025-08-12 23:28:43 +08:00
Wu Clan 8e8af2032a Optimize naming and preview in code generation (#764) 2025-08-12 16:42:25 +08:00
IAseven e09062eb39 Optimize the opera log storage logic through queue (#750)
*  feat: 操作日志中间件添加批量插入功能

* Delete GEMINI.md

* 🌈 style: 修复格式化错误

* 🐞 fix: 通过asyncio.wait_for兼容py3.10中asyncio.timeout不存在

* 🦄 refactor: 重新组织操作日志批量插入代码逻辑

* 优化代码实现

* 恢复默认配置

* 恢复默认 .gitignore 文件

* 更新队列批处理逻辑
2025-08-07 17:36:10 +08:00
Wu Clan 8c00492e44 Update the naming of table creation function (#760) 2025-08-07 16:57:57 +08:00
Wu Clan 0237d4c7b1 Update log output config and format (#759) 2025-08-07 16:46:56 +08:00
Wu Clan fe3a3b4e86 Optimize the data sort logic of tree nodes (#758) 2025-08-06 18:23:09 +08:00
Wu Clan 65ec721a1c Add business pagination in the code generator (#757) 2025-08-05 21:04:27 +08:00
Wu Clan 4eb76ad6ea Update the opera log desensitization method (#756)
* Update the opera log desensitization method

* Update request args function return
2025-08-05 18:38:49 +08:00
Wu Clan 1f8687155a Fix message format in validation exception handler (#755) 2025-08-05 18:11:02 +08:00
Wu Clan 8591d4e592 Refactor task routes and add control routes (#749) 2025-08-04 13:18:07 +08:00
Wu Clan dedf4e7bae Refactor code generation files and routes (#748)
* Refactor code generation files and routes

* Fix lint
2025-08-04 13:16:41 +08:00
Wu Clan 1e4aa88487 Fix the kwargs params of schedule task (#747) 2025-08-04 13:15:48 +08:00
Wu Clan 0dd745b7a7 Add schedule task demo that contains params (#746) 2025-07-31 18:29:54 +08:00
Wu Clan 6b2402f212 Add some interfaces for user profiles (#745) 2025-07-31 16:46:02 +08:00
Wu Clan 24a487eeea Simplify the plugin status update logic (#744) 2025-07-30 11:31:37 +08:00
Wu Clan 4f574189c7 Fix the error trigger when model auto import (#743) 2025-07-29 22:52:34 +08:00
Wu Clan b559a74cea Add update support for user email and phone (#742) 2025-07-29 22:52:08 +08:00
Wu Clan 83dcdbe59d Update the OAuth2 login password policy (#741)
* Update the OAuth2 login password policy

* Update the crud pwd

* Update the reset pwd service
2025-07-29 22:51:31 +08:00
Wu Clan d64f7c2911 Fix the opera log field encryption (#739) 2025-07-25 19:23:36 +08:00
Wu Clan 53e64bce37 Add auth whitelist regular expression config (#738) 2025-07-24 21:34:24 +08:00
Wu Clan 00a781357b Fix celery CLI option to required (#737) 2025-07-24 21:32:57 +08:00
Wu Clan d7f87ed0ad Update the default cache period for userinfo (#734) 2025-07-21 21:31:55 +08:00
Wu Clan bda9b1d463 Add distributed lock for scheduled task (#732)
* Add distributed lock for scheduled task

* Add the task to extend lock

* Fix the close
2025-07-21 12:29:12 +08:00
Wu Clan e0a106ec51 Simplify task crontab expression validation (#733) 2025-07-18 21:11:54 +08:00
Wu Clan 016361bd68 Update the changelog for v1.7.0 (#729) 2025-07-16 13:40:55 +08:00
Wu Clan f2d3c39425 Fix login and operation log clearing (#728) 2025-07-16 13:34:37 +08:00
Wu Clan e45d2d6fe7 Add CLI support for startup celery services (#724)
* Add CLI support for startup celery services

* update the cli file
2025-07-16 12:28:57 +08:00
Wu Clan 4ddf84fa22 Bump granian from 2.4.0 to 2.4.2 (#727) 2025-07-16 12:28:02 +08:00
Wu Clan 326a1883e8 Fix the parsing of execution task params (#725)
* Fix the parsing of execution task params

* Enhanced checksums
2025-07-16 12:10:33 +08:00
Wu Clan 9ff36d4498 Delete the default value of schema enum data (#723) 2025-07-15 10:43:58 +08:00
Wu Clan 98ef07ad32 Simplify celery task crontab config (#722) 2025-07-15 00:28:10 +08:00
Wu Clan 802b0d456a Optimize celery integrations and events (#721) 2025-07-14 19:17:24 +08:00
Wu Clan 6767f0e2e6 Update the celery task comment and name (#720) 2025-07-11 21:08:56 +08:00
Wu Clan d72a05c965 Fix the celery task scheduler query (#719) 2025-07-11 21:08:44 +08:00
Wu Clan ce3be1db8e Add support for celery dynamic tasks (#715)
* Add support for celery dynamic tasks

* Update the celery conf

* Update the celery task tables name

* Refactor the celery task-related interfaces

* Optimize auto-discovery tasks

* Remove redundant config

* Refine the business codes

* Optimize crontab validation returns

* Update dependencies in pyproject toml

* Fix some bugs

* Update dependencies

* Update the version to 1.7.0

* Fix update and delete event
2025-07-11 07:54:33 +08:00
Wu Clan e84ef04f15 Update the CLI startup service mode (#718) 2025-07-10 20:33:50 +08:00
Wu Clan adee3a2177 Simplify user permission database queries (#717) 2025-07-09 20:44:32 +08:00
Wu Clan 526e0aab9a Optimize the analysis of get plugins (#716)
* Optimize the analysis of get plugins

* update var
2025-07-08 12:20:13 +08:00
Wu Clan 2bbbbe764a Update the log output default style (#714)
* Update the log output default style

* Update the log file compression

* fix line

* Add log summary
2025-07-06 14:32:36 +08:00
Wu Clan ef5e921c08 Update the middleware logging accuracy (#713)
* Update the middleware logging accuracy

* Update the log style
2025-07-05 19:52:37 +08:00
Wu Clan a2fa59285a Update the granian env to command params (#712) 2025-07-04 19:46:40 +08:00
Wu Clan 494942e87e Add CLI support for execute sql scripts (#711)
* Add CLI support for execute sql scripts

* Update the arg helps
2025-07-02 20:19:29 +08:00
Wu Clan aa2b76673f Update the reload excludes for CLI run (#709) 2025-07-02 18:16:43 +08:00
Wu Clan 099880dd1c Update the refresh token verify mechanism (#710)
* Remove the logout interface auth dependency

* Update the refresh token check
2025-07-02 18:08:28 +08:00
Wu Clan d906a103af Fix the code generation delete schema template (#708) 2025-07-01 22:28:37 +08:00
Wu Clan 4ed49d2d13 Replace gunicorn deployment to granian (#705)
* Replace gunicorn deployment to granian

* Update env comments

* Fix the granian env
2025-07-01 21:33:49 +08:00
Wu Clan 7b5ae4696f Fix the code generation schema template (#706) 2025-07-01 18:18:47 +08:00
Wu Clan 54ea301152 Update the CLI to be executed async (#704) 2025-07-01 18:18:29 +08:00
Wu Clan 69c27232ac Update the changelog for v1.6.0 (#703) 2025-06-30 17:02:55 +08:00
Wu Clan f36dcb3f5c Update the version number to 1.6.0 (#702) 2025-06-30 17:00:19 +08:00
Wu Clan 63d088d62c Update the Dockerfile to adapt the CLI (#701) 2025-06-30 16:53:16 +08:00
Wu Clan a461f78224 Optimize the installation of plugin dependencies (#700)
* Optimize the installation of plugin dependencies

* Remove enumerate

* Update main params
2025-06-30 09:29:35 +08:00
Wu Clan c306432708 Update the help for CLI run worker (#699) 2025-06-29 20:06:45 +08:00
Wu Clan 97f778cc90 Add CLI support for plugin install (#698)
* Add CLI support for plugin install

* Fix some usage errors

* Update prompt information
2025-06-29 17:41:12 +08:00
Wu Clan 88695ac6ad Add custom CLI for service startup (#697)
* Add custom CLI for service startup

* Remove redundant scripts

* Fix lint
2025-06-28 14:56:19 +08:00
Wu Clan 3b04329b04 Add the test user to SQL scripts (#696) 2025-06-27 18:43:42 +08:00
Wu Clan bd4acf8121 Update the extension plugin config (#695) 2025-06-27 18:14:25 +08:00
Wu Clan b9a9b1efe6 Update the SQL to adapt frontend plugin (#694) 2025-06-26 21:58:35 +08:00
Wu Clan c050f58ee9 Fix the OAuth2 redirect route names (#693) 2025-06-26 21:50:22 +08:00
Wu Clan d47375ae41 Optimize dict create and update logic (#691)
* Optimize dict create and update logic

* Update version
2025-06-25 10:41:09 +08:00
Wu Clan 0602c6144d Update the changelog for v1.5.2 (#690) 2025-06-24 17:40:09 +08:00
Wu Clan c84f0658fb Update the dict pagination query parameters (#689)
* Update the dict pagination query parameters

* Update the dict pagination query parameters

* Update version
2025-06-24 17:34:18 +08:00
Wu Clan a2902bd23a Add built-in plugin missing files (#688) 2025-06-24 17:34:06 +08:00
Wu Clan 69a9b90711 Optimize the zip plug-in file name parsing (#687) 2025-06-24 17:33:57 +08:00
Wu Clan b96402e11d Simplify custom response status codes (#686) 2025-06-24 10:42:15 +08:00
Wu Clan 0bc6f6a719 Update the init test data for SQL scripts (#685)
* Update the init test data for SQL scripts

* Update plugin sql scripts

* fix snowflake id
2025-06-24 10:15:23 +08:00
Wu Clan ebd65c8246 Update dict data label column config (#684) 2025-06-23 22:33:34 +08:00
Wu Clan 234bf708b3 Fix the code with outdated system config (#683) 2025-06-23 22:23:06 +08:00
Dylan 6d5e741d94 Optimize api with semantic HTTP status codes (#681) 2025-06-23 22:18:17 +08:00
Wu Clan f9bfe8f510 Add dictionary type and datas queries (#679) 2025-06-23 10:58:55 +08:00
Wu Clan 408c866dda Update cache cleanup for logout interface (#678)
* Update cache cleanup for logout interface

* Add dependencies of interface
2025-06-22 23:41:48 +08:00
Wu Clan bac41a46f8 Optimize token detection and caching logic (#677) 2025-06-21 20:18:13 +08:00
Wu Clan 8638c26db1 Add the snowflake ID sql script (#675) 2025-06-20 20:31:44 +08:00
Wu Clan 319ba13df1 Optimize routes to better align with RESTful (#673)
* Optimize routes to better align with RESTful

* Add codes endpoint description

* Update jinja templates

* fix typo

* fix sql
2025-06-19 11:06:34 +08:00
Wu Clan 0d1f05d307 Fix some error class import (#672) 2025-06-17 15:30:35 +08:00
Wu Clan e6608d18ce Update the changelog for v1.5.1 (#671) 2025-06-16 21:38:13 +08:00
Dylan 7afd8415cd Add support for snowflake ID primary key (#670)
* fix: 修复PostgreSQL SQL语法错误,将反引号替换为双引号

* feat: 新增雪花算法ID实现

* 优化雪花算法和主键类型

* 修复错误引用

* 添加雪花详情链接

* feat: add snowflake ID parser method

* 修复独立执行异常

* 更新系统时间错误类
2025-06-16 13:34:27 +08:00
Dylan 845f2f0ff8 Fix the postgresql sql script syntax error (#669) 2025-06-15 16:38:48 +08:00
Wu Clan 11d7792c0f Bump sqlalchemy crud plus version to 1.10.0 (#668)
* Bump sqlalchemy crud plus version to 1.10.0

* Update filter style

* Update filters style
2025-06-14 12:35:19 +08:00
Wu Clan 6883ec34c6 Fix the sidebar menu type filtering (#667) 2025-06-13 16:50:53 +08:00
Wu Clan 4c45e7ff27 Update the changelog for v1.5.0 (#664) 2025-06-09 21:36:54 +08:00
217 changed files with 6911 additions and 4391 deletions
+128
View File
@@ -1,3 +1,126 @@
<a id="v1.7.0"></a>
# [v1.7.0](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.7.0) - 2025-07-16
## What's Changed
* Update the changelog for v1.6.0 by [@wu-clan](https://github.com/wu-clan) in [#703](https://github.com/fastapi-practices/fastapi_best_architecture/pull/703)
* Update the CLI to be executed async by [@wu-clan](https://github.com/wu-clan) in [#704](https://github.com/fastapi-practices/fastapi_best_architecture/pull/704)
* Fix the code generation schema template by [@wu-clan](https://github.com/wu-clan) in [#706](https://github.com/fastapi-practices/fastapi_best_architecture/pull/706)
* Replace gunicorn deployment to granian by [@wu-clan](https://github.com/wu-clan) in [#705](https://github.com/fastapi-practices/fastapi_best_architecture/pull/705)
* Fix the code generation delete schema template by [@wu-clan](https://github.com/wu-clan) in [#708](https://github.com/fastapi-practices/fastapi_best_architecture/pull/708)
* Update the refresh token verify mechanism by [@wu-clan](https://github.com/wu-clan) in [#710](https://github.com/fastapi-practices/fastapi_best_architecture/pull/710)
* Update the reload excludes for CLI run by [@wu-clan](https://github.com/wu-clan) in [#709](https://github.com/fastapi-practices/fastapi_best_architecture/pull/709)
* Add CLI support for execute sql scripts by [@wu-clan](https://github.com/wu-clan) in [#711](https://github.com/fastapi-practices/fastapi_best_architecture/pull/711)
* Update the granian env to command params by [@wu-clan](https://github.com/wu-clan) in [#712](https://github.com/fastapi-practices/fastapi_best_architecture/pull/712)
* Update the middleware logging accuracy by [@wu-clan](https://github.com/wu-clan) in [#713](https://github.com/fastapi-practices/fastapi_best_architecture/pull/713)
* Update the log output default style by [@wu-clan](https://github.com/wu-clan) in [#714](https://github.com/fastapi-practices/fastapi_best_architecture/pull/714)
* Optimize the analysis of get plugins by [@wu-clan](https://github.com/wu-clan) in [#716](https://github.com/fastapi-practices/fastapi_best_architecture/pull/716)
* Simplify user permission database queries by [@wu-clan](https://github.com/wu-clan) in [#717](https://github.com/fastapi-practices/fastapi_best_architecture/pull/717)
* Update the CLI startup service mode by [@wu-clan](https://github.com/wu-clan) in [#718](https://github.com/fastapi-practices/fastapi_best_architecture/pull/718)
* Add support for celery dynamic tasks by [@wu-clan](https://github.com/wu-clan) in [#715](https://github.com/fastapi-practices/fastapi_best_architecture/pull/715)
* Fix the celery task scheduler query by [@wu-clan](https://github.com/wu-clan) in [#719](https://github.com/fastapi-practices/fastapi_best_architecture/pull/719)
* Update the celery task comment and name by [@wu-clan](https://github.com/wu-clan) in [#720](https://github.com/fastapi-practices/fastapi_best_architecture/pull/720)
* Optimize celery integrations and events by [@wu-clan](https://github.com/wu-clan) in [#721](https://github.com/fastapi-practices/fastapi_best_architecture/pull/721)
* Simplify celery task crontab config by [@wu-clan](https://github.com/wu-clan) in [#722](https://github.com/fastapi-practices/fastapi_best_architecture/pull/722)
* Delete the default value of schema enum data by [@wu-clan](https://github.com/wu-clan) in [#723](https://github.com/fastapi-practices/fastapi_best_architecture/pull/723)
* Fix the parsing of execution task params by [@wu-clan](https://github.com/wu-clan) in [#725](https://github.com/fastapi-practices/fastapi_best_architecture/pull/725)
* Bump granian from 2.4.0 to 2.4.2 by [@wu-clan](https://github.com/wu-clan) in [#727](https://github.com/fastapi-practices/fastapi_best_architecture/pull/727)
* Add CLI support for startup celery services by [@wu-clan](https://github.com/wu-clan) in [#724](https://github.com/fastapi-practices/fastapi_best_architecture/pull/724)
* Fix login and operation log clearing by [@wu-clan](https://github.com/wu-clan) in [#728](https://github.com/fastapi-practices/fastapi_best_architecture/pull/728)
**Full Changelog**: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.6.0...v1.7.0
[Changes][v1.7.0]
<a id="v1.6.0"></a>
# [v1.6.0](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.6.0) - 2025-06-30
## What's Changed
* Update changelog for v1.5.2 by [@wu-clan](https://github.com/wu-clan) in [#690](https://github.com/fastapi-practices/fastapi_best_architecture/pull/690)
* Optimize dict create and update logic by [@wu-clan](https://github.com/wu-clan) in [#691](https://github.com/fastapi-practices/fastapi_best_architecture/pull/691)
* Fix the OAuth2 redirect route names by [@wu-clan](https://github.com/wu-clan) in [#693](https://github.com/fastapi-practices/fastapi_best_architecture/pull/693)
* Update the SQL to adapt frontend plugin by [@wu-clan](https://github.com/wu-clan) in [#694](https://github.com/fastapi-practices/fastapi_best_architecture/pull/694)
* Update the extension plugin config by [@wu-clan](https://github.com/wu-clan) in [#695](https://github.com/fastapi-practices/fastapi_best_architecture/pull/695)
* Add the test user to SQL scripts by [@wu-clan](https://github.com/wu-clan) in [#696](https://github.com/fastapi-practices/fastapi_best_architecture/pull/696)
* Add custom CLI for service startup by [@wu-clan](https://github.com/wu-clan) in [#697](https://github.com/fastapi-practices/fastapi_best_architecture/pull/697)
* Add CLI support for plugin install by [@wu-clan](https://github.com/wu-clan) in [#698](https://github.com/fastapi-practices/fastapi_best_architecture/pull/698)
* Update the help for CLI run worker by [@wu-clan](https://github.com/wu-clan) in [#699](https://github.com/fastapi-practices/fastapi_best_architecture/pull/699)
* Optimize the installation of plugin dependencies by [@wu-clan](https://github.com/wu-clan) in [#700](https://github.com/fastapi-practices/fastapi_best_architecture/pull/700)
* Update the Dockerfile to adapt latest code by [@wu-clan](https://github.com/wu-clan) in [#701](https://github.com/fastapi-practices/fastapi_best_architecture/pull/701)
* Update the version number to 1.6.0 by [@wu-clan](https://github.com/wu-clan) in [#702](https://github.com/fastapi-practices/fastapi_best_architecture/pull/702)
**Full Changelog**: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.5.2...v1.6.0
[Changes][v1.6.0]
<a id="v1.5.2"></a>
# [v1.5.2](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.5.2) - 2025-06-24
## What's Changed
* Update changelog for v1.5.1 by [@wu-clan](https://github.com/wu-clan) in [#671](https://github.com/fastapi-practices/fastapi_best_architecture/pull/671)
* Fix some error class import by [@wu-clan](https://github.com/wu-clan) in [#672](https://github.com/fastapi-practices/fastapi_best_architecture/pull/672)
* Optimize routes to better align with RESTful by [@wu-clan](https://github.com/wu-clan) in [#673](https://github.com/fastapi-practices/fastapi_best_architecture/pull/673)
* Add the snowflake ID sql script by [@wu-clan](https://github.com/wu-clan) in [#675](https://github.com/fastapi-practices/fastapi_best_architecture/pull/675)
* Optimize token detection and caching logic by [@wu-clan](https://github.com/wu-clan) in [#677](https://github.com/fastapi-practices/fastapi_best_architecture/pull/677)
* Update cache cleanup for logout interface by [@wu-clan](https://github.com/wu-clan) in [#678](https://github.com/fastapi-practices/fastapi_best_architecture/pull/678)
* Add dictionary type and datas queries by [@wu-clan](https://github.com/wu-clan) in [#679](https://github.com/fastapi-practices/fastapi_best_architecture/pull/679)
* Optimize api with semantic HTTP status codes by [@downdawn](https://github.com/downdawn) in [#681](https://github.com/fastapi-practices/fastapi_best_architecture/pull/681)
* Fix the code with outdated system config by [@wu-clan](https://github.com/wu-clan) in [#683](https://github.com/fastapi-practices/fastapi_best_architecture/pull/683)
* Update dict data label column config by [@wu-clan](https://github.com/wu-clan) in [#684](https://github.com/fastapi-practices/fastapi_best_architecture/pull/684)
* Update the init test data for SQL scripts by [@wu-clan](https://github.com/wu-clan) in [#685](https://github.com/fastapi-practices/fastapi_best_architecture/pull/685)
* Simplify custom response status codes by [@wu-clan](https://github.com/wu-clan) in [#686](https://github.com/fastapi-practices/fastapi_best_architecture/pull/686)
* Optimize the zip plug-in file name parsing by [@wu-clan](https://github.com/wu-clan) in [#687](https://github.com/fastapi-practices/fastapi_best_architecture/pull/687)
* Add built-in plugin missing files by [@wu-clan](https://github.com/wu-clan) in [#688](https://github.com/fastapi-practices/fastapi_best_architecture/pull/688)
* Update the dict pagination query parameters by [@wu-clan](https://github.com/wu-clan) in [#689](https://github.com/fastapi-practices/fastapi_best_architecture/pull/689)
**Full Changelog**: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.5.1...v1.5.2
[Changes][v1.5.2]
<a id="v1.5.1"></a>
# [v1.5.1](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.5.1) - 2025-06-16
## What's Changed
* Update changelog for v1.5.0 by [@wu-clan](https://github.com/wu-clan) in [#664](https://github.com/fastapi-practices/fastapi_best_architecture/pull/664)
* Fix the sidebar menu type filtering by [@wu-clan](https://github.com/wu-clan) in [#667](https://github.com/fastapi-practices/fastapi_best_architecture/pull/667)
* Bump sqlalchemy crud plus version to 1.10.0 by [@wu-clan](https://github.com/wu-clan) in [#668](https://github.com/fastapi-practices/fastapi_best_architecture/pull/668)
* Fix the postgresql sql script syntax error by [@downdawn](https://github.com/downdawn) in [#669](https://github.com/fastapi-practices/fastapi_best_architecture/pull/669)
* Add Initial Snowflake ID Support by [@downdawn](https://github.com/downdawn) in [#670](https://github.com/fastapi-practices/fastapi_best_architecture/pull/670)
**Full Changelog**: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.5.0...v1.5.1
[Changes][v1.5.1]
<a id="v1.5.0"></a>
# [v1.5.0](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.5.0) - 2025-06-09
## What's Changed
* Update changelog for v1.4.3 by [@wu-clan](https://github.com/wu-clan) in [#651](https://github.com/fastapi-practices/fastapi_best_architecture/pull/651)
* Update OAuth2 callback interface return by [@wu-clan](https://github.com/wu-clan) in [#653](https://github.com/fastapi-practices/fastapi_best_architecture/pull/653)
* Update user email and phone operation logic by [@wu-clan](https://github.com/wu-clan) in [#654](https://github.com/fastapi-practices/fastapi_best_architecture/pull/654)
* Simplify OAuth2 model and optimize auth service by [@wu-clan](https://github.com/wu-clan) in [#655](https://github.com/fastapi-practices/fastapi_best_architecture/pull/655)
* Add OAuth2 user to auto bind a role by [@wu-clan](https://github.com/wu-clan) in [#656](https://github.com/fastapi-practices/fastapi_best_architecture/pull/656)
* Update data scope and rule to m2m by [@wu-clan](https://github.com/wu-clan) in [#657](https://github.com/fastapi-practices/fastapi_best_architecture/pull/657)
* Update code generate interface permission by [@wu-clan](https://github.com/wu-clan) in [#658](https://github.com/fastapi-practices/fastapi_best_architecture/pull/658)
* Update the plugin download interface permission by [@wu-clan](https://github.com/wu-clan) in [#659](https://github.com/fastapi-practices/fastapi_best_architecture/pull/659)
* Update auth failed default status code by [@wu-clan](https://github.com/wu-clan) in [#660](https://github.com/fastapi-practices/fastapi_best_architecture/pull/660)
* Update menu sort in init test sql by [@wu-clan](https://github.com/wu-clan) in [#661](https://github.com/fastapi-practices/fastapi_best_architecture/pull/661)
* Add data permission in init test sql by [@wu-clan](https://github.com/wu-clan) in [#662](https://github.com/fastapi-practices/fastapi_best_architecture/pull/662)
* Update the version to 1.5.0 by [@wu-clan](https://github.com/wu-clan) in [#663](https://github.com/fastapi-practices/fastapi_best_architecture/pull/663)
**Full Changelog**: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.3...v1.5.0
[Changes][v1.5.0]
<a id="v1.4.3"></a> <a id="v1.4.3"></a>
# [v1.4.3](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.4.3) - 2025-06-02 # [v1.4.3](https://github.com/fastapi-practices/fastapi_best_architecture/releases/tag/v1.4.3) - 2025-06-02
@@ -591,6 +714,11 @@
[Changes][v1.0.0] [Changes][v1.0.0]
[v1.7.0]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.6.0...v1.7.0
[v1.6.0]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.5.2...v1.6.0
[v1.5.2]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.5.1...v1.5.2
[v1.5.1]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.5.0...v1.5.1
[v1.5.0]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.3...v1.5.0
[v1.4.3]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.2...v1.4.3 [v1.4.3]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.2...v1.4.3
[v1.4.2]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.1...v1.4.2 [v1.4.2]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.1...v1.4.2
[v1.4.1]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.0...v1.4.1 [v1.4.1]: https://github.com/fastapi-practices/fastapi_best_architecture/compare/v1.4.0...v1.4.1
+9 -16
View File
@@ -10,6 +10,10 @@ RUN sed -i 's/deb.debian.org/mirrors.ustc.edu.cn/g' /etc/apt/sources.list.d/debi
&& apt-get install -y --no-install-recommends gcc python3-dev \ && apt-get install -y --no-install-recommends gcc python3-dev \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
COPY . /fba
WORKDIR /fba
# Configure uv environment # Configure uv environment
ENV UV_COMPILE_BYTECODE=1 \ ENV UV_COMPILE_BYTECODE=1 \
UV_NO_CACHE=1 \ UV_NO_CACHE=1 \
@@ -18,49 +22,38 @@ ENV UV_COMPILE_BYTECODE=1 \
# Install dependencies with cache # Install dependencies with cache
RUN --mount=type=cache,target=/root/.cache/uv \ RUN --mount=type=cache,target=/root/.cache/uv \
--mount=type=bind,source=uv.lock,target=uv.lock \
--mount=type=bind,source=pyproject.toml,target=pyproject.toml \
uv sync --frozen --no-default-groups --group server uv sync --frozen --no-default-groups --group server
# === Runtime base server image === # === Runtime base server image ===
FROM python:3.10-slim AS base_server FROM python:3.10-slim AS base_server
SHELL ["/bin/bash", "-c"]
RUN sed -i 's/deb.debian.org/mirrors.ustc.edu.cn/g' /etc/apt/sources.list.d/debian.sources \ RUN sed -i 's/deb.debian.org/mirrors.ustc.edu.cn/g' /etc/apt/sources.list.d/debian.sources \
&& apt-get update \ && apt-get update \
&& apt-get install -y --no-install-recommends supervisor \ && apt-get install -y --no-install-recommends supervisor \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
COPY . /fba COPY --from=builder /fba /fba
COPY --from=builder /usr/local /usr/local COPY --from=builder /usr/local /usr/local
# Install plugin dependencies COPY deploy/backend/supervisord.conf /etc/supervisor/supervisord.conf
WORKDIR /fba
ENV PYTHONPATH=/fba WORKDIR /fba/backend
RUN python3 backend/scripts/init_plugin.py
# === FastAPI server image === # === FastAPI server image ===
FROM base_server AS fastapi_server FROM base_server AS fastapi_server
WORKDIR /fba
COPY deploy/backend/supervisord.conf /etc/supervisor/supervisord.conf
COPY deploy/backend/fastapi_server.conf /etc/supervisor/conf.d/ COPY deploy/backend/fastapi_server.conf /etc/supervisor/conf.d/
RUN mkdir -p /var/log/fastapi_server RUN mkdir -p /var/log/fastapi_server
EXPOSE 8001 EXPOSE 8001
CMD ["uvicorn", "backend.main:app", "--host", "0.0.0.0", "--port","8000"] CMD ["/usr/local/bin/granian", "main:app", "--interface", "asgi", "--host", "0.0.0.0", "--port","8000"]
# === Celery server image === # === Celery server image ===
FROM base_server AS celery FROM base_server AS celery
WORKDIR /fba/backend/
COPY deploy/backend/supervisord.conf /etc/supervisor/supervisord.conf
COPY deploy/backend/celery.conf /etc/supervisor/conf.d/ COPY deploy/backend/celery.conf /etc/supervisor/conf.d/
RUN mkdir -p /var/log/celery RUN mkdir -p /var/log/celery
+9 -8
View File
@@ -15,18 +15,19 @@ REDIS_DATABASE=0
TOKEN_SECRET_KEY='1VkVF75nsNABBjK_7-qz7GtzNy3AMvktc9TCPwKczCk' TOKEN_SECRET_KEY='1VkVF75nsNABBjK_7-qz7GtzNy3AMvktc9TCPwKczCk'
# Opera Log # Opera Log
OPERA_LOG_ENCRYPT_SECRET_KEY='d77b25790a804c2b4a339dd0207941e4cefa5751935a33735bc73bb7071a005b' OPERA_LOG_ENCRYPT_SECRET_KEY='d77b25790a804c2b4a339dd0207941e4cefa5751935a33735bc73bb7071a005b'
# App Admin # [ App ] task
# OAuth2
OAUTH2_GITHUB_CLIENT_ID='test'
OAUTH2_GITHUB_CLIENT_SECRET='test'
OAUTH2_LINUX_DO_CLIENT_ID='test'
OAUTH2_LINUX_DO_CLIENT_SECRET='test'
# App Task
# Celery # Celery
CELERY_BROKER_REDIS_DATABASE=1 CELERY_BROKER_REDIS_DATABASE=1
CELERY_BACKEND_REDIS_DATABASE=2
# Rabbitmq # Rabbitmq
CELERY_RABBITMQ_HOST='127.0.0.1' CELERY_RABBITMQ_HOST='127.0.0.1'
CELERY_RABBITMQ_PORT=5672 CELERY_RABBITMQ_PORT=5672
CELERY_RABBITMQ_USERNAME='guest' CELERY_RABBITMQ_USERNAME='guest'
CELERY_RABBITMQ_PASSWORD='guest' CELERY_RABBITMQ_PASSWORD='guest'
# [ Plugin ] oauth2
OAUTH2_GITHUB_CLIENT_ID='test'
OAUTH2_GITHUB_CLIENT_SECRET='test'
OAUTH2_LINUX_DO_CLIENT_ID='test'
OAUTH2_LINUX_DO_CLIENT_SECRET='test'
# [ Plugin ] email
EMAIL_USERNAME=''
EMAIL_PASSWORD=''
+1 -1
View File
@@ -72,7 +72,7 @@
> >
> It is recommended to execute under the backend directory, and chmod authorization may be required > It is recommended to execute under the backend directory, and chmod authorization may be required
- `pre_start.sh`: Perform automatic database migration and create database tables - `pre_start.sh`: Perform automatic database migration
- `celery-start.sh`: For celery docker script, implementation is not recommended - `celery-start.sh`: For celery docker script, implementation is not recommended
+12
View File
@@ -1,2 +1,14 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from backend.common.i18n import i18n
from backend.utils.console import console
__version__ = '1.8.0'
def get_version() -> str | None:
console.print(f'[cyan]{__version__}[/]')
# 初始化 i18n
i18n.load_locales()
+4 -9
View File
@@ -13,21 +13,16 @@ from sqlalchemy.ext.asyncio import async_engine_from_config
sys.path.append('../') sys.path.append('../')
from backend.app import get_app_models
from backend.common.model import MappedBase from backend.common.model import MappedBase
from backend.core import path_conf from backend.core import path_conf
from backend.database.db import SQLALCHEMY_DATABASE_URL from backend.database.db import SQLALCHEMY_DATABASE_URL
from backend.plugin.tools import get_plugin_models from backend.plugin.tools import get_plugin_models
# import your new model here # import models
from backend.app.admin.model import * # noqa: F401 for cls in get_app_models() + get_plugin_models():
from backend.plugin.code_generator.model import * # noqa: F401
# import plugin model
for cls in get_plugin_models():
class_name = cls.__name__ class_name = cls.__name__
if class_name in globals(): if class_name not in globals():
print(f'\nWarning: Class "{class_name}" already exists in global namespace.')
else:
globals()[class_name] = cls globals()[class_name] = cls
if not os.path.exists(path_conf.ALEMBIC_VERSION_DIR): if not os.path.exists(path_conf.ALEMBIC_VERSION_DIR):
+43
View File
@@ -1,2 +1,45 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
import inspect
import os.path
from backend.common.log import log
from backend.core.path_conf import BASE_PATH
from backend.utils.import_parse import import_module_cached
def get_app_models():
"""获取 app 所有模型类"""
app_path = os.path.join(BASE_PATH, 'app')
list_dirs = os.listdir(app_path)
apps = []
for d in list_dirs:
if os.path.isdir(os.path.join(app_path, d)) and d != '__pycache__':
apps.append(d)
classes = []
for app in apps:
try:
module_path = f'backend.app.{app}.model'
module = import_module_cached(module_path)
except ModuleNotFoundError as e:
log.warning(f'应用 {app} 中不包含 model 相关配置: {e}')
continue
except Exception as e:
raise e
for name, obj in inspect.getmembers(module):
if inspect.isclass(obj):
classes.append(obj)
return classes
# import all app models for auto create db tables
for cls in get_app_models():
class_name = cls.__name__
if class_name not in globals():
globals()[class_name] = cls
+1 -1
View File
@@ -8,4 +8,4 @@ from backend.app.admin.api.v1.auth.captcha import router as captcha_router
router = APIRouter(prefix='/auth') router = APIRouter(prefix='/auth')
router.include_router(auth_router, tags=['授权']) router.include_router(auth_router, tags=['授权'])
router.include_router(captcha_router, prefix='/captcha', tags=['验证码']) router.include_router(captcha_router, tags=['验证码'])
+13 -6
View File
@@ -11,14 +11,15 @@ from backend.app.admin.schema.token import GetLoginToken, GetNewToken, GetSwagge
from backend.app.admin.schema.user import AuthLoginParam from backend.app.admin.schema.user import AuthLoginParam
from backend.app.admin.service.auth_service import auth_service from backend.app.admin.service.auth_service import auth_service
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
from backend.common.security.jwt import DependsJwtAuth
router = APIRouter() router = APIRouter()
@router.post('/login/swagger', summary='swagger 调试专用', description='用于快捷获取 token 进行 swagger 认证') @router.post('/login/swagger', summary='swagger 调试专用', description='用于快捷获取 token 进行 swagger 认证')
async def swagger_login(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetSwaggerToken: async def login_swagger(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetSwaggerToken:
token, user = await auth_service.swagger_login(obj=obj) token, user = await auth_service.swagger_login(obj=obj)
return GetSwaggerToken(access_token=token, user=user) # type: ignore return GetSwaggerToken(access_token=token, user=user)
@router.post( @router.post(
@@ -27,20 +28,26 @@ async def swagger_login(obj: Annotated[HTTPBasicCredentials, Depends()]) -> GetS
description='json 格式登录, 仅支持在第三方api工具调试, 例如: postman', description='json 格式登录, 仅支持在第三方api工具调试, 例如: postman',
dependencies=[Depends(RateLimiter(times=5, minutes=1))], dependencies=[Depends(RateLimiter(times=5, minutes=1))],
) )
async def user_login( async def login(
request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks
) -> ResponseSchemaModel[GetLoginToken]: ) -> ResponseSchemaModel[GetLoginToken]:
data = await auth_service.login(request=request, response=response, obj=obj, background_tasks=background_tasks) data = await auth_service.login(request=request, response=response, obj=obj, background_tasks=background_tasks)
return response_base.success(data=data) return response_base.success(data=data)
@router.post('/tokens/refresh', summary='刷新 token') @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)
return response_base.success(data=codes)
@router.post('/tokens', summary='刷新 token')
async def refresh_token(request: Request) -> ResponseSchemaModel[GetNewToken]: async def refresh_token(request: Request) -> ResponseSchemaModel[GetNewToken]:
data = await auth_service.new_token(request=request) data = await auth_service.refresh_token(request=request)
return response_base.success(data=data) return response_base.success(data=data)
@router.post('/logout', summary='用户登出') @router.post('/logout', summary='用户登出')
async def user_logout(request: Request, response: Response) -> ResponseModel: async def logout(request: Request, response: Response) -> ResponseModel:
await auth_service.logout(request=request, response=response) await auth_service.logout(request=request, response=response)
return response_base.success() return response_base.success()
+1 -1
View File
@@ -14,7 +14,7 @@ router = APIRouter()
@router.get( @router.get(
'', '/captcha',
summary='获取登录验证码', summary='获取登录验证码',
dependencies=[Depends(RateLimiter(times=5, seconds=10))], dependencies=[Depends(RateLimiter(times=5, seconds=10))],
) )
+7 -9
View File
@@ -4,7 +4,7 @@ from typing import Annotated
from fastapi import APIRouter, Depends, Query from fastapi import APIRouter, Depends, Query
from backend.app.admin.schema.login_log import GetLoginLogDetail from backend.app.admin.schema.login_log import DeleteLoginLogParam, GetLoginLogDetail
from backend.app.admin.service.login_log_service import login_log_service 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, paging_data
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
@@ -24,7 +24,7 @@ router = APIRouter()
DependsPagination, DependsPagination,
], ],
) )
async def get_pagination_login_logs( async def get_login_logs_paged(
db: CurrentSession, db: CurrentSession,
username: Annotated[str | None, Query(description='用户名')] = None, username: Annotated[str | None, Query(description='用户名')] = None,
status: Annotated[int | None, Query(description='状态')] = None, status: Annotated[int | None, Query(description='状态')] = None,
@@ -43,8 +43,8 @@ async def get_pagination_login_logs(
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_login_log(pk: Annotated[list[int], Query(description='登录日志 ID 列表')]) -> ResponseModel: async def delete_login_logs(obj: DeleteLoginLogParam) -> ResponseModel:
count = await login_log_service.delete(pk=pk) count = await login_log_service.delete(obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@@ -54,12 +54,10 @@ async def delete_login_log(pk: Annotated[list[int], Query(description='登录日
'/all', '/all',
summary='清空登录日志', summary='清空登录日志',
dependencies=[ dependencies=[
Depends(RequestPermission('log:login:empty')), Depends(RequestPermission('log:login:clear')),
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_all_login_logs() -> ResponseModel: async def delete_all_login_logs() -> ResponseModel:
count = await login_log_service.delete_all() await login_log_service.delete_all()
if count > 0: return response_base.success()
return response_base.success()
return response_base.fail()
+7 -9
View File
@@ -4,7 +4,7 @@ from typing import Annotated
from fastapi import APIRouter, Depends, Query from fastapi import APIRouter, Depends, Query
from backend.app.admin.schema.opera_log import GetOperaLogDetail from backend.app.admin.schema.opera_log import DeleteOperaLogParam, GetOperaLogDetail
from backend.app.admin.service.opera_log_service import opera_log_service 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, paging_data
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
@@ -24,7 +24,7 @@ router = APIRouter()
DependsPagination, DependsPagination,
], ],
) )
async def get_pagination_opera_logs( async def get_opera_logs_paged(
db: CurrentSession, db: CurrentSession,
username: Annotated[str | None, Query(description='用户名')] = None, username: Annotated[str | None, Query(description='用户名')] = None,
status: Annotated[int | None, Query(description='状态')] = None, status: Annotated[int | None, Query(description='状态')] = None,
@@ -43,8 +43,8 @@ async def get_pagination_opera_logs(
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_opera_log(pk: Annotated[list[int], Query(description='操作日志 ID 列表')]) -> ResponseModel: async def delete_opera_logs(obj: DeleteOperaLogParam) -> ResponseModel:
count = await opera_log_service.delete(pk=pk) count = await opera_log_service.delete(obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@@ -54,12 +54,10 @@ async def delete_opera_log(pk: Annotated[list[int], Query(description='操作日
'/all', '/all',
summary='清空操作日志', summary='清空操作日志',
dependencies=[ dependencies=[
Depends(RequestPermission('log:opera:empty')), Depends(RequestPermission('log:opera:clear')),
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_all_opera_logs() -> ResponseModel: async def delete_all_opera_logs() -> ResponseModel:
count = await opera_log_service.delete_all() await opera_log_service.delete_all()
if count > 0: return response_base.success()
return response_base.success()
return response_base.fail()
+1 -1
View File
@@ -10,4 +10,4 @@ router = APIRouter(prefix='/monitors')
router.include_router(redis_router, prefix='/redis', tags=['redis监控']) router.include_router(redis_router, prefix='/redis', tags=['redis监控'])
router.include_router(server_router, prefix='/server', tags=['服务器监控']) router.include_router(server_router, prefix='/server', tags=['服务器监控'])
router.include_router(token_router, prefix='/online', tags=['在线用户']) router.include_router(token_router, prefix='/sessions', tags=['会话监控'])
+8 -7
View File
@@ -19,7 +19,7 @@ router = APIRouter()
@router.get('', summary='获取在线用户', dependencies=[DependsJwtAuth]) @router.get('', summary='获取在线用户', dependencies=[DependsJwtAuth])
async def get_online( async def get_sessions(
username: Annotated[str | None, Query(description='用户名')] = None, username: Annotated[str | None, Query(description='用户名')] = None,
) -> ResponseSchemaModel[list[GetTokenDetail]]: ) -> ResponseSchemaModel[list[GetTokenDetail]]:
token_keys = await redis_client.keys(f'{settings.TOKEN_REDIS_PREFIX}:*') token_keys = await redis_client.keys(f'{settings.TOKEN_REDIS_PREFIX}:*')
@@ -44,9 +44,10 @@ async def get_online(
for key in token_keys: for key in token_keys:
token = await redis_client.get(key) token = await redis_client.get(key)
token_payload = jwt_decode(token) token_payload = jwt_decode(token)
user_id = token_payload.id
session_uuid = token_payload.session_uuid session_uuid = token_payload.session_uuid
token_detail = GetTokenDetail( token_detail = GetTokenDetail(
id=token_payload.id, id=user_id,
session_uuid=session_uuid, session_uuid=session_uuid,
username='未知', username='未知',
nickname='未知', nickname='未知',
@@ -58,7 +59,7 @@ async def get_online(
last_login_time='未知', last_login_time='未知',
expire_time=token_payload.expire_time, expire_time=token_payload.expire_time,
) )
extra_info = await redis_client.get(f'{settings.TOKEN_EXTRA_INFO_REDIS_PREFIX}:{session_uuid}') extra_info = await redis_client.get(f'{settings.TOKEN_EXTRA_INFO_REDIS_PREFIX}:{user_id}:{session_uuid}')
if extra_info: if extra_info:
extra_info = json.loads(extra_info) extra_info = json.loads(extra_info)
# 排除 swagger 登录生成的 token # 排除 swagger 登录生成的 token
@@ -75,17 +76,17 @@ async def get_online(
@router.delete( @router.delete(
'/{pk}', '/{pk}',
summary='下线', summary='强制下线',
dependencies=[ dependencies=[
Depends(RequestPermission('sys:token:kick')), Depends(RequestPermission('sys:session:delete')),
DependsRBAC, DependsRBAC,
], ],
) )
async def kick_out( async def delete_session(
request: Request, request: Request,
pk: Annotated[int, Path(description='用户 ID')], pk: Annotated[int, Path(description='用户 ID')],
session_uuid: Annotated[str, Query(description='会话 UUID')], session_uuid: Annotated[str, Query(description='会话 UUID')],
) -> ResponseModel: ) -> ResponseModel:
superuser_verify(request) superuser_verify(request)
await revoke_token(str(pk), session_uuid) await revoke_token(pk, session_uuid)
return response_base.success() return response_base.success()
+2 -10
View File
@@ -1,23 +1,15 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from fastapi import APIRouter, Depends from fastapi import APIRouter
from backend.common.response.response_schema import ResponseModel, response_base from backend.common.response.response_schema import ResponseModel, response_base
from backend.common.security.jwt import DependsJwtAuth from backend.common.security.jwt import DependsJwtAuth
from backend.common.security.permission import RequestPermission
from backend.utils.redis_info import redis_info from backend.utils.redis_info import redis_info
router = APIRouter() router = APIRouter()
@router.get( @router.get('', summary='redis 监控', dependencies=[DependsJwtAuth])
'',
summary='redis 监控',
dependencies=[
Depends(RequestPermission('sys:monitor:redis')),
DependsJwtAuth,
],
)
async def get_redis_info() -> ResponseModel: async def get_redis_info() -> ResponseModel:
data = { data = {
'info': await redis_info.get_info(), 'info': await redis_info.get_info(),
+2 -10
View File
@@ -1,24 +1,16 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from fastapi import APIRouter, Depends from fastapi import APIRouter
from starlette.concurrency import run_in_threadpool from starlette.concurrency import run_in_threadpool
from backend.common.response.response_schema import ResponseModel, response_base from backend.common.response.response_schema import ResponseModel, response_base
from backend.common.security.jwt import DependsJwtAuth from backend.common.security.jwt import DependsJwtAuth
from backend.common.security.permission import RequestPermission
from backend.utils.server_info import server_info from backend.utils.server_info import server_info
router = APIRouter() router = APIRouter()
@router.get( @router.get('', summary='server 监控', dependencies=[DependsJwtAuth])
'',
summary='server 监控',
dependencies=[
Depends(RequestPermission('sys:monitor:server')),
DependsJwtAuth,
],
)
async def get_server_info() -> ResponseModel: async def get_server_info() -> ResponseModel:
data = { data = {
# 扔到线程池,避免阻塞 # 扔到线程池,避免阻塞
+2 -2
View File
@@ -5,10 +5,10 @@ from fastapi import APIRouter
from backend.app.admin.api.v1.sys.data_rule import router as data_rule_router from backend.app.admin.api.v1.sys.data_rule import router as data_rule_router
from backend.app.admin.api.v1.sys.data_scope import router as data_scope_router from backend.app.admin.api.v1.sys.data_scope import router as data_scope_router
from backend.app.admin.api.v1.sys.dept import router as dept_router from backend.app.admin.api.v1.sys.dept import router as dept_router
from backend.app.admin.api.v1.sys.files import router as file_router
from backend.app.admin.api.v1.sys.menu import router as menu_router from backend.app.admin.api.v1.sys.menu import router as menu_router
from backend.app.admin.api.v1.sys.plugin import router as plugin_router from backend.app.admin.api.v1.sys.plugin import router as plugin_router
from backend.app.admin.api.v1.sys.role import router as role_router from backend.app.admin.api.v1.sys.role import router as role_router
from backend.app.admin.api.v1.sys.upload import router as upload_router
from backend.app.admin.api.v1.sys.user import router as user_router from backend.app.admin.api.v1.sys.user import router as user_router
router = APIRouter(prefix='/sys') router = APIRouter(prefix='/sys')
@@ -19,5 +19,5 @@ router.include_router(role_router, prefix='/roles', tags=['系统角色'])
router.include_router(user_router, prefix='/users', tags=['系统用户']) router.include_router(user_router, prefix='/users', tags=['系统用户'])
router.include_router(data_rule_router, prefix='/data-rules', tags=['系统数据规则']) router.include_router(data_rule_router, prefix='/data-rules', tags=['系统数据规则'])
router.include_router(data_scope_router, prefix='/data-scopes', tags=['系统数据范围']) router.include_router(data_scope_router, prefix='/data-scopes', tags=['系统数据范围'])
router.include_router(upload_router, prefix='/upload', tags=['系统上传']) router.include_router(file_router, prefix='/files', tags=['系统文件'])
router.include_router(plugin_router, prefix='/plugins', tags=['系统插件']) router.include_router(plugin_router, prefix='/plugins', tags=['系统插件'])
+4 -3
View File
@@ -6,6 +6,7 @@ from fastapi import APIRouter, Depends, Path, Query
from backend.app.admin.schema.data_rule import ( from backend.app.admin.schema.data_rule import (
CreateDataRuleParam, CreateDataRuleParam,
DeleteDataRuleParam,
GetDataRuleColumnDetail, GetDataRuleColumnDetail,
GetDataRuleDetail, GetDataRuleDetail,
UpdateDataRuleParam, UpdateDataRuleParam,
@@ -57,7 +58,7 @@ async def get_data_rule(
DependsPagination, DependsPagination,
], ],
) )
async def get_pagination_data_rules( async def get_data_rules_paged(
db: CurrentSession, name: Annotated[str | None, Query(description='规则名称')] = None db: CurrentSession, name: Annotated[str | None, Query(description='规则名称')] = None
) -> ResponseSchemaModel[PageData[GetDataRuleDetail]]: ) -> ResponseSchemaModel[PageData[GetDataRuleDetail]]:
data_rule_select = await data_rule_service.get_select(name=name) data_rule_select = await data_rule_service.get_select(name=name)
@@ -103,8 +104,8 @@ async def update_data_rule(
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_data_rule(pk: Annotated[list[int], Query(description='数据规则 ID 列表')]) -> ResponseModel: async def delete_data_rules(obj: DeleteDataRuleParam) -> ResponseModel:
count = await data_rule_service.delete(pk=pk) count = await data_rule_service.delete(obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
+4 -3
View File
@@ -6,6 +6,7 @@ from fastapi import APIRouter, Depends, Path, Query
from backend.app.admin.schema.data_scope import ( from backend.app.admin.schema.data_scope import (
CreateDataScopeParam, CreateDataScopeParam,
DeleteDataScopeParam,
GetDataScopeDetail, GetDataScopeDetail,
GetDataScopeWithRelationDetail, GetDataScopeWithRelationDetail,
UpdateDataScopeParam, UpdateDataScopeParam,
@@ -52,7 +53,7 @@ async def get_data_scope_rules(
DependsPagination, DependsPagination,
], ],
) )
async def get_pagination_data_scopes( async def get_data_scopes_paged(
db: CurrentSession, db: CurrentSession,
name: Annotated[str | None, Query(description='范围名称')] = None, name: Annotated[str | None, Query(description='范围名称')] = None,
status: Annotated[int | None, Query(description='状态')] = None, status: Annotated[int | None, Query(description='状态')] = None,
@@ -117,8 +118,8 @@ async def update_data_scope_rules(
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_data_scope(pk: Annotated[list[int], Query(description='数据范围 ID 列表')]) -> ResponseModel: async def delete_data_scopes(obj: DeleteDataScopeParam) -> ResponseModel:
count = await data_scope_service.delete(pk=pk) count = await data_scope_service.delete(obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
+3 -3
View File
@@ -20,15 +20,15 @@ async def get_dept(pk: Annotated[int, Path(description='部门 ID')]) -> Respons
return response_base.success(data=data) return response_base.success(data=data)
@router.get('', summary='获取所有部门展示', dependencies=[DependsJwtAuth]) @router.get('', summary='获取部门', dependencies=[DependsJwtAuth])
async def get_all_depts( async def get_dept_tree(
request: Request, request: Request,
name: Annotated[str | None, Query(description='部门名称')] = None, name: Annotated[str | None, Query(description='部门名称')] = None,
leader: Annotated[str | None, Query(description='部门负责人')] = None, leader: Annotated[str | None, Query(description='部门负责人')] = None,
phone: Annotated[str | None, Query(description='联系电话')] = None, phone: Annotated[str | None, Query(description='联系电话')] = None,
status: Annotated[int | None, Query(description='状态')] = None, status: Annotated[int | None, Query(description='状态')] = None,
) -> ResponseSchemaModel[list[dict[str, Any]]]: ) -> ResponseSchemaModel[list[dict[str, Any]]]:
dept = await dept_service.get_dept_tree(request=request, name=name, leader=leader, phone=phone, status=status) dept = await dept_service.get_tree(request=request, name=name, leader=leader, phone=phone, status=status)
return response_base.success(data=dept) return response_base.success(data=dept)
+27
View File
@@ -0,0 +1,27 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, File, UploadFile
from backend.common.dataclasses import UploadUrl
from backend.common.response.response_schema import ResponseSchemaModel, response_base
from backend.common.security.permission import RequestPermission
from backend.common.security.rbac import DependsRBAC
from backend.utils.file_ops import upload_file, upload_file_verify
router = APIRouter()
@router.post(
'/upload',
summary='文件上传',
dependencies=[
Depends(RequestPermission('sys:file:upload')),
DependsRBAC,
],
)
async def upload_files(file: Annotated[UploadFile, File()]) -> ResponseSchemaModel[UploadUrl]:
upload_file_verify(file)
filename = await upload_file(file)
return response_base.success(data={'url': f'/static/upload/{filename}'})
+4 -4
View File
@@ -14,7 +14,7 @@ from backend.common.security.rbac import DependsRBAC
router = APIRouter() router = APIRouter()
@router.get('/sidebar', summary='获取用户菜单侧边栏', description='适配 vben5', dependencies=[DependsJwtAuth]) @router.get('/sidebar', summary='获取用户菜单侧边栏', description='适配 vben admin v5', dependencies=[DependsJwtAuth])
async def get_user_sidebar(request: Request) -> ResponseSchemaModel[list[dict[str, Any] | None]]: async def get_user_sidebar(request: Request) -> ResponseSchemaModel[list[dict[str, Any] | None]]:
menu = await menu_service.get_sidebar(request=request) menu = await menu_service.get_sidebar(request=request)
return response_base.success(data=menu) return response_base.success(data=menu)
@@ -26,12 +26,12 @@ async def get_menu(pk: Annotated[int, Path(description='菜单 ID')]) -> Respons
return response_base.success(data=data) return response_base.success(data=data)
@router.get('', summary='获取所有菜单展示', dependencies=[DependsJwtAuth]) @router.get('', summary='获取菜单', dependencies=[DependsJwtAuth])
async def get_all_menus( async def get_menu_tree(
title: Annotated[str | None, Query(description='菜单标题')] = None, title: Annotated[str | None, Query(description='菜单标题')] = None,
status: Annotated[int | None, Query(description='状体')] = None, status: Annotated[int | None, Query(description='状体')] = None,
) -> ResponseSchemaModel[list[dict[str, Any]]]: ) -> ResponseSchemaModel[list[dict[str, Any]]]:
menu = await menu_service.get_menu_tree(title=title, status=status) menu = await menu_service.get_tree(title=title, status=status)
return response_base.success(data=menu) return response_base.success(data=menu)
+26 -29
View File
@@ -7,7 +7,8 @@ from fastapi.params import Query
from starlette.responses import StreamingResponse from starlette.responses import StreamingResponse
from backend.app.admin.service.plugin_service import plugin_service from backend.app.admin.service.plugin_service import plugin_service
from backend.common.response.response_code import CustomResponseCode from backend.common.enums import PluginType
from backend.common.response.response_code import CustomResponse
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
from backend.common.security.jwt import DependsJwtAuth from backend.common.security.jwt import DependsJwtAuth
from backend.common.security.permission import RequestPermission from backend.common.security.permission import RequestPermission
@@ -22,38 +23,32 @@ async def get_all_plugins() -> ResponseSchemaModel[list[dict[str, Any]]]:
return response_base.success(data=plugins) return response_base.success(data=plugins)
@router.get('/changes', summary='插件状态是否变更', dependencies=[DependsJwtAuth]) @router.get('/changed', summary='是否存在插件变更', dependencies=[DependsJwtAuth])
async def plugin_changed() -> ResponseSchemaModel[bool]: async def plugin_changed() -> ResponseSchemaModel[bool]:
plugins = await plugin_service.changed() plugins = await plugin_service.changed()
return response_base.success(data=bool(plugins)) return response_base.success(data=bool(plugins))
@router.post( @router.post(
'/zip', '',
summary='安装 zip 插件', summary='安装插件',
description='使用插件 zip 压缩包进行安装', description='使用插件 zip 压缩包或 git 仓库地址进行安装',
dependencies=[ dependencies=[
Depends(RequestPermission('sys:plugin:zip')), Depends(RequestPermission('sys:plugin:install')),
DependsRBAC, DependsRBAC,
], ],
) )
async def install_zip_plugin(file: Annotated[UploadFile, File()]) -> ResponseModel: async def install_plugin(
await plugin_service.install_zip(file=file) type: Annotated[PluginType, Query(description='插件类型')],
return response_base.success(res=CustomResponseCode.PLUGIN_INSTALL_SUCCESS) file: Annotated[UploadFile | None, File()] = None,
repo_url: Annotated[str | None, Query(description='插件 git 仓库地址')] = None,
) -> ResponseModel:
@router.post( plugin_name = await plugin_service.install(type=type, file=file, repo_url=repo_url)
'/git', return response_base.success(
summary='安装 git 插件', res=CustomResponse(
description='使用插件 git 仓库地址进行安装,不限制平台;如果需要凭证,需在 git 仓库地址中添加凭证信息', code=200, msg=f'插件 {plugin_name} 安装成功,请根据插件说明(README.md)进行相关配置并重启服务'
dependencies=[ )
Depends(RequestPermission('sys:plugin:git')), )
DependsRBAC,
],
)
async def install_git_plugin(repo_url: Annotated[str, Query(description='插件 git 仓库地址')]) -> ResponseModel:
await plugin_service.install_git(repo_url=repo_url)
return response_base.success(res=CustomResponseCode.PLUGIN_INSTALL_SUCCESS)
@router.delete( @router.delete(
@@ -61,20 +56,22 @@ async def install_git_plugin(repo_url: Annotated[str, Query(description='插件
summary='卸载插件', summary='卸载插件',
description='此操作会直接删除插件依赖,但不会直接删除插件,而是将插件移动到备份目录', description='此操作会直接删除插件依赖,但不会直接删除插件,而是将插件移动到备份目录',
dependencies=[ dependencies=[
Depends(RequestPermission('sys:plugin:del')), Depends(RequestPermission('sys:plugin:uninstall')),
DependsRBAC, DependsRBAC,
], ],
) )
async def uninstall_plugin(plugin: Annotated[str, Path(description='插件名称')]) -> ResponseModel: async def uninstall_plugin(plugin: Annotated[str, Path(description='插件名称')]) -> ResponseModel:
await plugin_service.uninstall(plugin=plugin) await plugin_service.uninstall(plugin=plugin)
return response_base.success(res=CustomResponseCode.PLUGIN_UNINSTALL_SUCCESS) return response_base.success(
res=CustomResponse(code=200, msg=f'插件 {plugin} 卸载成功,请根据插件说明(README.md)移除相关配置并重启服务')
)
@router.post( @router.put(
'/{plugin}/status', '/{plugin}/status',
summary='更新插件状态', summary='更新插件状态',
dependencies=[ dependencies=[
Depends(RequestPermission('sys:plugin:status')), Depends(RequestPermission('sys:plugin:edit')),
DependsRBAC, DependsRBAC,
], ],
) )
@@ -83,8 +80,8 @@ async def update_plugin_status(plugin: Annotated[str, Path(description='插件
return response_base.success() return response_base.success()
@router.get('/{plugin}', summary='打包并下载插件', dependencies=[DependsJwtAuth]) @router.get('/{plugin}', summary='下载插件', dependencies=[DependsJwtAuth])
async def build_plugin(plugin: Annotated[str, Path(description='插件名称')]) -> StreamingResponse: async def download_plugin(plugin: Annotated[str, Path(description='插件名称')]) -> StreamingResponse:
bio = await plugin_service.build(plugin=plugin) bio = await plugin_service.build(plugin=plugin)
return StreamingResponse( return StreamingResponse(
bio, bio,
+8 -9
View File
@@ -6,6 +6,7 @@ from fastapi import APIRouter, Depends, Path, Query
from backend.app.admin.schema.role import ( from backend.app.admin.schema.role import (
CreateRoleParam, CreateRoleParam,
DeleteRoleParam,
GetRoleDetail, GetRoleDetail,
GetRoleWithRelationDetail, GetRoleWithRelationDetail,
UpdateRoleMenuParam, UpdateRoleMenuParam,
@@ -29,8 +30,8 @@ async def get_all_roles() -> ResponseSchemaModel[list[GetRoleDetail]]:
return response_base.success(data=data) return response_base.success(data=data)
@router.get('/{pk}/menus', summary='获取角色所有菜单', dependencies=[DependsJwtAuth]) @router.get('/{pk}/menus', summary='获取角色菜单', dependencies=[DependsJwtAuth])
async def get_role_all_menus( async def get_role_menu_tree(
pk: Annotated[int, Path(description='角色 ID')], pk: Annotated[int, Path(description='角色 ID')],
) -> ResponseSchemaModel[list[dict[str, Any] | None]]: ) -> ResponseSchemaModel[list[dict[str, Any] | None]]:
menu = await role_service.get_menu_tree(pk=pk) menu = await role_service.get_menu_tree(pk=pk)
@@ -38,15 +39,13 @@ async def get_role_all_menus(
@router.get('/{pk}/scopes', summary='获取角色所有数据范围', dependencies=[DependsJwtAuth]) @router.get('/{pk}/scopes', summary='获取角色所有数据范围', dependencies=[DependsJwtAuth])
async def get_role_all_scopes(pk: Annotated[int, Path(description='角色 ID')]) -> ResponseSchemaModel[list[int]]: async def get_role_scopes(pk: Annotated[int, Path(description='角色 ID')]) -> ResponseSchemaModel[list[int]]:
rule = await role_service.get_scopes(pk=pk) rule = await role_service.get_scopes(pk=pk)
return response_base.success(data=rule) return response_base.success(data=rule)
@router.get('/{pk}', summary='获取角色详情', dependencies=[DependsJwtAuth]) @router.get('/{pk}', summary='获取角色详情', dependencies=[DependsJwtAuth])
async def get_role( async def get_role(pk: Annotated[int, Path(description='角色 ID')]) -> ResponseSchemaModel[GetRoleWithRelationDetail]:
pk: Annotated[int, Path(description='角色 ID')],
) -> ResponseSchemaModel[GetRoleWithRelationDetail]:
data = await role_service.get(pk=pk) data = await role_service.get(pk=pk)
return response_base.success(data=data) return response_base.success(data=data)
@@ -59,7 +58,7 @@ async def get_role(
DependsPagination, DependsPagination,
], ],
) )
async def get_pagination_roles( async def get_roles_paged(
db: CurrentSession, db: CurrentSession,
name: Annotated[str | None, Query(description='角色名称')] = None, name: Annotated[str | None, Query(description='角色名称')] = None,
status: Annotated[int | None, Query(description='状态')] = None, status: Annotated[int | None, Query(description='状态')] = None,
@@ -139,8 +138,8 @@ async def update_role_scopes(
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_role(pk: Annotated[list[int], Query(description='角色 ID 列表')]) -> ResponseModel: async def delete_roles(obj: DeleteRoleParam) -> ResponseModel:
count = await role_service.delete(pk=pk) count = await role_service.delete(obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
-27
View File
@@ -1,27 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, File, UploadFile
from backend.common.dataclasses import UploadUrl
from backend.common.enums import FileType
from backend.common.response.response_schema import ResponseSchemaModel, response_base
from backend.common.security.jwt import DependsJwtAuth
from backend.utils.file_ops import file_verify, upload_file
router = APIRouter()
@router.post('/image', summary='上传图片', dependencies=[DependsJwtAuth])
async def upload_image(file: Annotated[UploadFile, File()]) -> ResponseSchemaModel[UploadUrl]:
file_verify(file, FileType.image)
filename = await upload_file(file)
return response_base.success(data={'url': f'/static/upload/{filename}'})
@router.post('/video', summary='上传视频', dependencies=[DependsJwtAuth])
async def upload_video(file: Annotated[UploadFile, File()]) -> ResponseSchemaModel[UploadUrl]:
file_verify(file, FileType.video)
filename = await upload_file(file)
return response_base.success(data={'url': f'/static/upload/{filename}'})
+67 -46
View File
@@ -2,7 +2,7 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Annotated from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query, Request from fastapi import APIRouter, Body, Depends, Path, Query, Request
from backend.app.admin.schema.role import GetRoleDetail from backend.app.admin.schema.role import GetRoleDetail
from backend.app.admin.schema.user import ( from backend.app.admin.schema.user import (
@@ -13,6 +13,7 @@ from backend.app.admin.schema.user import (
UpdateUserParam, UpdateUserParam,
) )
from backend.app.admin.service.user_service import user_service 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, paging_data
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
from backend.common.security.jwt import DependsJwtAuth from backend.common.security.jwt import DependsJwtAuth
@@ -23,42 +24,23 @@ from backend.database.db import CurrentSession
router = APIRouter() router = APIRouter()
@router.post('/add', summary='添加用户', dependencies=[DependsRBAC])
async def add_user(request: Request, obj: AddUserParam) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]:
await user_service.add(request=request, obj=obj)
data = await user_service.get_userinfo(username=obj.username)
return response_base.success(data=data)
@router.post('/{username}/password', summary='密码重置', dependencies=[DependsJwtAuth])
async def password_reset(
username: Annotated[str, Path(description='用户名')], obj: ResetPasswordParam
) -> ResponseModel:
count = await user_service.pwd_reset(username=username, obj=obj)
if count > 0:
return response_base.success()
return response_base.fail()
@router.get('/me', summary='获取当前用户信息', dependencies=[DependsJwtAuth]) @router.get('/me', summary='获取当前用户信息', dependencies=[DependsJwtAuth])
async def get_current_user(request: Request) -> ResponseSchemaModel[GetCurrentUserInfoWithRelationDetail]: async def get_current_user(request: Request) -> ResponseSchemaModel[GetCurrentUserInfoWithRelationDetail]:
data = request.user.model_dump() data = request.user.model_dump()
return response_base.success(data=data) return response_base.success(data=data)
@router.get('/{username}', summary='查看用户信息', dependencies=[DependsJwtAuth]) @router.get('/{pk}', summary='获取用户信息', dependencies=[DependsJwtAuth])
async def get_user( async def get_userinfo(
username: Annotated[str, Path(description='用户')], pk: Annotated[int, Path(description='用户 ID')],
) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]: ) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]:
data = await user_service.get_userinfo(username=username) data = await user_service.get_userinfo(pk=pk)
return response_base.success(data=data) return response_base.success(data=data)
@router.get('/{username}/roles', summary='获取用户所有角色', dependencies=[DependsJwtAuth]) @router.get('/{pk}/roles', summary='获取用户所有角色', dependencies=[DependsJwtAuth])
async def get_user_all_roles( async def get_user_roles(pk: Annotated[int, Path(description='用户 ID')]) -> ResponseSchemaModel[list[GetRoleDetail]]:
username: Annotated[str, Path(description='用户名')], data = await user_service.get_roles(pk=pk)
) -> ResponseSchemaModel[list[GetRoleDetail]]:
data = await user_service.get_roles(username=username)
return response_base.success(data=data) return response_base.success(data=data)
@@ -70,7 +52,7 @@ async def get_user_all_roles(
DependsPagination, DependsPagination,
], ],
) )
async def get_pagination_users( async def get_users_paged(
db: CurrentSession, db: CurrentSession,
dept: Annotated[int | None, Query(description='部门 ID')] = None, dept: Annotated[int | None, Query(description='部门 ID')] = None,
username: Annotated[str | None, Query(description='用户名')] = None, username: Annotated[str | None, Query(description='用户名')] = None,
@@ -82,58 +64,97 @@ async def get_pagination_users(
return response_base.success(data=page_data) return response_base.success(data=page_data)
@router.put('/{username}', summary='更新用户信息', dependencies=[DependsJwtAuth]) @router.post('', summary='创建用户', dependencies=[DependsRBAC])
async def create_user(request: Request, obj: AddUserParam) -> ResponseSchemaModel[GetUserInfoWithRelationDetail]:
await user_service.create(request=request, obj=obj)
data = await user_service.get_userinfo(username=obj.username)
return response_base.success(data=data)
@router.put('/{pk}', summary='更新用户信息', dependencies=[DependsRBAC])
async def update_user( async def update_user(
request: Request, username: Annotated[str, Path(description='用户')], obj: UpdateUserParam request: Request, pk: Annotated[int, Path(description='用户 ID')], obj: UpdateUserParam
) -> ResponseModel: ) -> ResponseModel:
count = await user_service.update(request=request, username=username, obj=obj) count = await user_service.update(request=request, pk=pk, obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@router.put('/{pk}/super', summary='修改用户超级权限', dependencies=[DependsRBAC]) @router.put('/{pk}/permissions', summary='更新用户权限', dependencies=[DependsRBAC])
async def super_set(request: Request, pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel: async def update_user_permission(
count = await user_service.update_permission(request=request, pk=pk) 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)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@router.put('/{pk}/staff', summary='修改用户后台登录权限', dependencies=[DependsRBAC]) @router.put('/me/password', summary='更新当前用户密码', dependencies=[DependsJwtAuth])
async def staff_set(request: Request, pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel: async def update_user_password(request: Request, obj: ResetPasswordParam) -> ResponseModel:
count = await user_service.update_staff(request=request, pk=pk) count = await user_service.update_password(request=request, obj=obj)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@router.put('/{pk}/status', summary='修改用户状态', dependencies=[DependsRBAC]) @router.put('/{pk}/password', summary='重置用户密码', dependencies=[DependsRBAC])
async def status_set(request: Request, pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel: async def reset_user_password(
count = await user_service.update_status(request=request, pk=pk) request: Request,
pk: Annotated[int, Path(description='用户 ID')],
password: Annotated[str, Body(embed=True, description='新密码')],
) -> ResponseModel:
count = await user_service.reset_password(request=request, pk=pk, password=password)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@router.put('/{pk}/multi', summary='修改用户多端登录状态', dependencies=[DependsRBAC]) @router.put('/me/nickname', summary='更新当前用户昵称', dependencies=[DependsJwtAuth])
async def multi_set(request: Request, pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel: async def update_user_nickname(
count = await user_service.update_multi_login(request=request, pk=pk) request: Request, nickname: Annotated[str, Body(embed=True, description='用户昵称')]
) -> ResponseModel:
count = await user_service.update_nickname(request=request, nickname=nickname)
if count > 0:
return response_base.success()
return response_base.fail()
@router.put('/me/avatar', summary='更新当前用户头像', dependencies=[DependsJwtAuth])
async def update_user_avatar(
request: Request, avatar: Annotated[str, Body(embed=True, description='用户头像地址')]
) -> ResponseModel:
count = await user_service.update_avatar(request=request, avatar=avatar)
if count > 0:
return response_base.success()
return response_base.fail()
@router.put('/me/email', summary='更新当前用户邮箱', dependencies=[DependsJwtAuth])
async def update_user_email(
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)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
@router.delete( @router.delete(
path='/{username}', path='/{pk}',
summary='删除用户', summary='删除用户',
dependencies=[ dependencies=[
Depends(RequestPermission('sys:user:del')), Depends(RequestPermission('sys:user:del')),
DependsRBAC, DependsRBAC,
], ],
) )
async def delete_user(username: Annotated[str, Path(description='用户')]) -> ResponseModel: async def delete_user(pk: Annotated[int, Path(description='用户 ID')]) -> ResponseModel:
count = await user_service.delete(username=username) count = await user_service.delete(pk=pk)
if count > 0: if count > 0:
return response_base.success() return response_base.success()
return response_base.fail() return response_base.fail()
+8 -13
View File
@@ -2,9 +2,8 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Sequence from typing import Sequence
from sqlalchemy import Select, and_, desc, select from sqlalchemy import Select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import noload
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
from backend.app.admin.model import DataRule from backend.app.admin.model import DataRule
@@ -31,16 +30,12 @@ class CRUDDataRule(CRUDPlus[DataRule]):
:param name: 规则名称 :param name: 规则名称
:return: :return:
""" """
stmt = select(self.model).options(noload(self.model.scopes)).order_by(desc(self.model.created_time)) filters = {}
filters = []
if name is not None: if name is not None:
filters.append(self.model.name.like(f'%{name}%')) filters['name__like'] = f'%{name}%'
if filters: return await self.select_order('id', load_strategies={'scopes': 'noload'}, **filters)
stmt = stmt.where(and_(*filters))
return stmt
async def get_by_name(self, db: AsyncSession, name: str) -> DataRule | None: async def get_by_name(self, db: AsyncSession, name: str) -> DataRule | None:
""" """
@@ -82,15 +77,15 @@ class CRUDDataRule(CRUDPlus[DataRule]):
""" """
return await self.update_model(db, pk, obj) return await self.update_model(db, pk, obj)
async def delete(self, db: AsyncSession, pk: list[int]) -> int: async def delete(self, db: AsyncSession, pks: list[int]) -> int:
""" """
删除规则 批量删除规则
:param db: 数据库会话 :param db: 数据库会话
:param pk: 规则 ID 列表 :param pks: 规则 ID 列表
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pk) return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
data_rule_dao: CRUDDataRule = CRUDDataRule(DataRule) data_rule_dao: CRUDDataRule = CRUDDataRule(DataRule)
+10 -21
View File
@@ -2,9 +2,8 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Sequence from typing import Sequence
from sqlalchemy import Select, and_, desc, select from sqlalchemy import Select, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import noload, selectinload
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
from backend.app.admin.model import DataRule, DataScope from backend.app.admin.model import DataRule, DataScope
@@ -42,9 +41,7 @@ class CRUDDataScope(CRUDPlus[DataScope]):
:param pk: 范围 ID :param pk: 范围 ID
:return: :return:
""" """
stmt = select(self.model).options(selectinload(self.model.rules)).where(self.model.id == pk) return await self.select_model(db, pk, load_strategies=['rules'])
data_scope = await db.execute(stmt)
return data_scope.scalars().first()
async def get_all(self, db: AsyncSession) -> Sequence[DataScope]: async def get_all(self, db: AsyncSession) -> Sequence[DataScope]:
""" """
@@ -63,22 +60,14 @@ class CRUDDataScope(CRUDPlus[DataScope]):
:param status: 范围状态 :param status: 范围状态
:return: :return:
""" """
stmt = ( filters = {}
select(self.model)
.options(noload(self.model.rules), noload(self.model.roles))
.order_by(desc(self.model.created_time))
)
filters = []
if name is not None: if name is not None:
filters.append(self.model.name.like(f'%{name}%')) filters['name__like'] = f'%{name}%'
if status is not None: if status is not None:
filters.append(self.model.status == status) filters['status'] = status
if filters: return await self.select_order('id', load_strategies={'rules': 'noload', 'roles': 'noload'}, **filters)
stmt = stmt.where(and_(*filters))
return stmt
async def create(self, db: AsyncSession, obj: CreateDataScopeParam) -> None: async def create(self, db: AsyncSession, obj: CreateDataScopeParam) -> None:
""" """
@@ -116,15 +105,15 @@ class CRUDDataScope(CRUDPlus[DataScope]):
current_data_scope.rules = rules.scalars().all() current_data_scope.rules = rules.scalars().all()
return len(current_data_scope.rules) return len(current_data_scope.rules)
async def delete(self, db: AsyncSession, pk: list[int]) -> int: async def delete(self, db: AsyncSession, pks: list[int]) -> int:
""" """
删除数据范围 批量删除数据范围
:param db: 数据库会话 :param db: 数据库会话
:param pk: 范围 ID 列表 :param pks: 范围 ID 列表
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pk) return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
data_scope_dao: CRUDDataScope = CRUDDataScope(DataScope) data_scope_dao: CRUDDataScope = CRUDDataScope(DataScope)
+11 -14
View File
@@ -3,9 +3,7 @@
from typing import Sequence from typing import Sequence
from fastapi import Request from fastapi import Request
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
from backend.app.admin.model import Dept from backend.app.admin.model import Dept
@@ -56,16 +54,19 @@ class CRUDDept(CRUDPlus[Dept]):
:param status: 部门状态 :param status: 部门状态
:return: :return:
""" """
filters = {'del_flag__eq': 0} filters = {'del_flag': 0}
if name is not None: if name is not None:
filters.update(name__like=f'%{name}%') filters['name__like'] = f'%{name}%'
if leader is not None: if leader is not None:
filters.update(leader__like=f'%{leader}%') filters['leader__like'] = f'%{leader}%'
if phone is not None: if phone is not None:
filters.update(phone__startswith=phone) filters['phone__startswith'] = phone
if status is not None: if status is not None:
filters.update(status=status) filters['status'] = status
return await self.select_models_order(db, 'sort', None, await filter_data_permission(db, request), **filters)
data_filtered = await filter_data_permission(db, request)
return await self.select_models_order(db, 'sort', 'desc', data_filtered, **filters)
async def create(self, db: AsyncSession, obj: CreateDeptParam) -> None: async def create(self, db: AsyncSession, obj: CreateDeptParam) -> None:
""" """
@@ -106,9 +107,7 @@ class CRUDDept(CRUDPlus[Dept]):
:param dept_id: 部门 ID :param dept_id: 部门 ID
:return: :return:
""" """
stmt = select(self.model).options(selectinload(self.model.users)).where(self.model.id == dept_id) return await self.select_model(db, dept_id, load_strategies=['users'])
result = await db.execute(stmt)
return result.scalars().first()
async def get_children(self, db: AsyncSession, dept_id: int) -> Sequence[Dept | None]: async def get_children(self, db: AsyncSession, dept_id: int) -> Sequence[Dept | None]:
""" """
@@ -118,9 +117,7 @@ class CRUDDept(CRUDPlus[Dept]):
:param dept_id: 部门 ID :param dept_id: 部门 ID
:return: :return:
""" """
stmt = select(self.model).where(self.model.parent_id == dept_id, self.model.del_flag == 0) return await self.select_models(db, parent_id=dept_id, del_flag=0)
result = await db.execute(stmt)
return result.scalars().all()
dept_dao: CRUDDept = CRUDDept(Dept) dept_dao: CRUDDept = CRUDDept(Dept)
+13 -9
View File
@@ -1,6 +1,7 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from sqlalchemy import Select from sqlalchemy import Select
from sqlalchemy import delete as sa_delete
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
@@ -21,12 +22,14 @@ class CRUDLoginLog(CRUDPlus[LoginLog]):
:return: :return:
""" """
filters = {} filters = {}
if username is not None: if username is not None:
filters.update(username__like=f'%{username}%') filters['username__like'] = f'%{username}%'
if status is not None: if status is not None:
filters.update(status=status) filters['status'] = status
if ip is not None: if ip is not None:
filters.update(ip__like=f'%{ip}%') filters['ip__like'] = f'%{ip}%'
return await self.select_order('created_time', 'desc', **filters) return await self.select_order('created_time', 'desc', **filters)
async def create(self, db: AsyncSession, obj: CreateLoginLogParam) -> None: async def create(self, db: AsyncSession, obj: CreateLoginLogParam) -> None:
@@ -39,24 +42,25 @@ class CRUDLoginLog(CRUDPlus[LoginLog]):
""" """
await self.create_model(db, obj, commit=True) await self.create_model(db, obj, commit=True)
async def delete(self, db: AsyncSession, pk: list[int]) -> int: async def delete(self, db: AsyncSession, pks: list[int]) -> int:
""" """
删除登录日志 批量删除登录日志
:param db: 数据库会话 :param db: 数据库会话
:param pk: 登录日志 ID 列表 :param pks: 登录日志 ID 列表
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pk) return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
async def delete_all(self, db: AsyncSession) -> int: @staticmethod
async def delete_all(db: AsyncSession) -> None:
""" """
删除所有日志 删除所有日志
:param db: 数据库会话 :param db: 数据库会话
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True) await db.execute(sa_delete(LoginLog))
login_log_dao: CRUDLoginLog = CRUDLoginLog(LoginLog) login_log_dao: CRUDLoginLog = CRUDLoginLog(LoginLog)
+13 -17
View File
@@ -2,9 +2,7 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Sequence from typing import Sequence
from sqlalchemy import and_, asc, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
from backend.app.admin.model import Menu from backend.app.admin.model import Menu
@@ -44,28 +42,28 @@ class CRUDMenu(CRUDPlus[Menu]):
:return: :return:
""" """
filters = {} filters = {}
if title is not None: if title is not None:
filters.update(title__like=f'%{title}%') filters['title__like'] = f'%{title}%'
if status is not None: if status is not None:
filters.update(status=status) filters['status'] = status
return await self.select_models_order(db, 'sort', **filters) return await self.select_models_order(db, 'sort', **filters)
async def get_sidebar(self, db: AsyncSession, superuser: bool, menu_ids: list[int | None]) -> Sequence[Menu]: async def get_sidebar(self, db: AsyncSession, menu_ids: list[int] | None) -> Sequence[Menu]:
""" """
获取角色菜单列表 获取用户的菜单侧边栏
:param db: 数据库会话 :param db: 数据库会话
:param superuser: 是否超级管理员
:param menu_ids: 菜单 ID 列表 :param menu_ids: 菜单 ID 列表
:return: :return:
""" """
stmt = select(self.model).order_by(asc(self.model.sort)) filters = {'type__in': [0, 1, 3, 4]}
filters = [self.model.type.in_([0, 1])]
if not superuser: if menu_ids:
filters.append(self.model.id.in_(menu_ids)) filters['id__in'] = menu_ids
stmt = stmt.where(and_(*filters))
menu = await db.execute(stmt) return await self.select_models_order(db, 'sort', 'asc', **filters)
return menu.scalars().all()
async def create(self, db: AsyncSession, obj: CreateMenuParam) -> None: async def create(self, db: AsyncSession, obj: CreateMenuParam) -> None:
""" """
@@ -106,9 +104,7 @@ class CRUDMenu(CRUDPlus[Menu]):
:param menu_id: 菜单 ID :param menu_id: 菜单 ID
:return: :return:
""" """
stmt = select(self.model).options(selectinload(self.model.children)).where(self.model.id == menu_id) menu = await self.select_model(db, menu_id, load_strategies=['children'])
result = await db.execute(stmt)
menu = result.scalars().first()
return menu.children return menu.children
+24 -10
View File
@@ -1,6 +1,7 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from sqlalchemy import Select from sqlalchemy import Select
from sqlalchemy import delete as sa_delete
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
@@ -21,12 +22,14 @@ class CRUDOperaLogDao(CRUDPlus[OperaLog]):
:return: :return:
""" """
filters = {} filters = {}
if username is not None: if username is not None:
filters.update(username__like=f'%{username}%') filters['username__like'] = f'%{username}%'
if status is not None: if status is not None:
filters.update(status=status) filters['status__eq'] = status
if ip is not None: if ip is not None:
filters.update(ip__like=f'%{ip}%') filters['ip__like'] = f'%{ip}%'
return await self.select_order('created_time', 'desc', **filters) return await self.select_order('created_time', 'desc', **filters)
async def create(self, db: AsyncSession, obj: CreateOperaLogParam) -> None: async def create(self, db: AsyncSession, obj: CreateOperaLogParam) -> None:
@@ -34,29 +37,40 @@ class CRUDOperaLogDao(CRUDPlus[OperaLog]):
创建操作日志 创建操作日志
:param db: 数据库会话 :param db: 数据库会话
:param obj: 创建操作日志参数 :param obj: 操作日志创建参数
:return: :return:
""" """
await self.create_model(db, obj) await self.create_model(db, obj)
async def delete(self, db: AsyncSession, pk: list[int]) -> int: async def bulk_create(self, db: AsyncSession, objs: list[CreateOperaLogParam]) -> None:
""" """
删除操作日志 批量创建操作日志
:param db: 数据库会话 :param db: 数据库会话
:param pk: 操作日志 ID 列表 :param objs: 操作日志创建参数列表
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pk) await self.create_models(db, objs)
async def delete_all(self, db: AsyncSession) -> int: async def delete(self, db: AsyncSession, pks: list[int]) -> int:
"""
批量删除操作日志
:param db: 数据库会话
:param pks: 操作日志 ID 列表
:return:
"""
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
@staticmethod
async def delete_all(db: AsyncSession) -> None:
""" """
删除所有日志 删除所有日志
:param db: 数据库会话 :param db: 数据库会话
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True) await db.execute(sa_delete(OperaLog))
opera_log_dao: CRUDOperaLogDao = CRUDOperaLogDao(OperaLog) opera_log_dao: CRUDOperaLogDao = CRUDOperaLogDao(OperaLog)
+19 -25
View File
@@ -2,9 +2,8 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Sequence from typing import Sequence
from sqlalchemy import Select, and_, desc, select from sqlalchemy import Select, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import noload, selectinload
from sqlalchemy_crud_plus import CRUDPlus from sqlalchemy_crud_plus import CRUDPlus
from backend.app.admin.model import DataScope, Menu, Role from backend.app.admin.model import DataScope, Menu, Role
@@ -37,13 +36,7 @@ class CRUDRole(CRUDPlus[Role]):
:param role_id: 角色 ID :param role_id: 角色 ID
:return: :return:
""" """
stmt = ( return await self.select_model(db, role_id, load_strategies=['menus', 'scopes'])
select(self.model)
.options(selectinload(self.model.menus), selectinload(self.model.scopes))
.where(self.model.id == role_id)
)
role = await db.execute(stmt)
return role.scalars().first()
async def get_all(self, db: AsyncSession) -> Sequence[Role]: async def get_all(self, db: AsyncSession) -> Sequence[Role]:
""" """
@@ -62,22 +55,23 @@ class CRUDRole(CRUDPlus[Role]):
:param status: 角色状态 :param status: 角色状态
:return: :return:
""" """
stmt = (
select(self.model)
.options(noload(self.model.users), noload(self.model.menus), noload(self.model.scopes))
.order_by(desc(self.model.created_time))
)
filters = [] filters = {}
if name is not None: if name is not None:
filters.append(self.model.name.like(f'%{name}%')) filters['name__like'] = f'%{name}%'
if status is not None: if status is not None:
filters.append(self.model.status == status) filters['status'] = status
if filters: return await self.select_order(
stmt = stmt.where(and_(*filters)) 'id',
load_strategies={
return stmt 'users': 'noload',
'menus': 'noload',
'scopes': 'noload',
},
**filters,
)
async def get_by_name(self, db: AsyncSession, name: str) -> Role | None: async def get_by_name(self, db: AsyncSession, name: str) -> Role | None:
""" """
@@ -140,15 +134,15 @@ class CRUDRole(CRUDPlus[Role]):
current_role.scopes = scopes.scalars().all() current_role.scopes = scopes.scalars().all()
return len(current_role.scopes) return len(current_role.scopes)
async def delete(self, db: AsyncSession, role_id: list[int]) -> int: async def delete(self, db: AsyncSession, role_ids: list[int]) -> int:
""" """
删除角色 批量删除角色
:param db: 数据库会话 :param db: 数据库会话
:param role_id: 角色 ID 列表 :param role_ids: 角色 ID 列表
:return: :return:
""" """
return await self.delete_model_by_column(db, allow_multiple=True, id__in=role_id) return await self.delete_model_by_column(db, allow_multiple=True, id__in=role_ids)
role_dao: CRUDRole = CRUDRole(Role) role_dao: CRUDRole = CRUDRole(Role)
+55 -83
View File
@@ -2,7 +2,7 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
import bcrypt import bcrypt
from sqlalchemy import and_, desc, select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import noload, selectinload from sqlalchemy.orm import noload, selectinload
from sqlalchemy.sql import Select from sqlalchemy.sql import Select
@@ -89,10 +89,8 @@ class CRUDUser(CRUDPlus[User]):
:param obj: 注册用户参数 :param obj: 注册用户参数
:return: :return:
""" """
salt = bcrypt.gensalt()
obj.password = get_hash_password(obj.password, salt)
dict_obj = obj.model_dump() dict_obj = obj.model_dump()
dict_obj.update({'is_staff': True, 'salt': salt}) dict_obj.update({'is_staff': True, 'salt': None})
new_user = self.model(**dict_obj) new_user = self.model(**dict_obj)
stmt = select(Role) stmt = select(Role)
@@ -119,6 +117,17 @@ class CRUDUser(CRUDPlus[User]):
input_user.roles = roles.scalars().all() input_user.roles = roles.scalars().all()
return count return count
async def update_nickname(self, db: AsyncSession, user_id: int, nickname: str) -> int:
"""
更新用户昵称
:param db: 数据库会话
:param user_id: 用户 ID
:param nickname: 用户昵称
:return:
"""
return await self.update_model(db, user_id, {'nickname': nickname})
async def update_avatar(self, db: AsyncSession, user_id: int, avatar: str) -> int: async def update_avatar(self, db: AsyncSession, user_id: int, avatar: str) -> int:
""" """
更新用户头像 更新用户头像
@@ -130,6 +139,17 @@ class CRUDUser(CRUDPlus[User]):
""" """
return await self.update_model(db, user_id, {'avatar': avatar}) return await self.update_model(db, user_id, {'avatar': avatar})
async def update_email(self, db: AsyncSession, user_id: int, email: str) -> int:
"""
更新用户邮箱
:param db: 数据库会话
:param user_id: 用户 ID
:param email: 邮箱
:return:
"""
return await self.update_model(db, user_id, {'email': email})
async def delete(self, db: AsyncSession, user_id: int) -> int: async def delete(self, db: AsyncSession, user_id: int) -> int:
""" """
删除用户 删除用户
@@ -150,16 +170,18 @@ class CRUDUser(CRUDPlus[User]):
""" """
return await self.select_model_by_column(db, email=email) return await self.select_model_by_column(db, email=email)
async def reset_password(self, db: AsyncSession, pk: int, new_pwd: str) -> int: async def reset_password(self, db: AsyncSession, pk: int, password: str) -> int:
""" """
重置用户密码 重置用户密码
:param db: 数据库会话 :param db: 数据库会话
:param pk: 用户 ID :param pk: 用户 ID
:param new_pwd: 新密码已加密 :param password: 新密码
:return: :return:
""" """
return await self.update_model(db, pk, {'password': new_pwd}) salt = bcrypt.gensalt()
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_list(self, dept: int | None, username: str | None, phone: str | None, status: int | None) -> Select:
""" """
@@ -171,74 +193,27 @@ class CRUDUser(CRUDPlus[User]):
:param status: 用户状态 :param status: 用户状态
:return: :return:
""" """
stmt = ( filters = {}
select(self.model)
.options( if dept:
filters['dept_id'] = dept
if username:
filters['username__like'] = f'%{username}%'
if phone:
filters['phone_like'] = f'%{phone}%'
if status is not None:
filters['status'] = status
return await self.select_order(
'id',
'desc',
load_options=[
selectinload(self.model.dept).options(noload(Dept.parent), noload(Dept.children), noload(Dept.users)), selectinload(self.model.dept).options(noload(Dept.parent), noload(Dept.children), noload(Dept.users)),
selectinload(self.model.roles).options(noload(Role.users), noload(Role.menus), noload(Role.scopes)), selectinload(self.model.roles).options(noload(Role.users), noload(Role.menus), noload(Role.scopes)),
) ],
.order_by(desc(self.model.join_time)) **filters,
) )
filters = []
if dept:
filters.append(self.model.dept_id == dept)
if username:
filters.append(self.model.username.like(f'%{username}%'))
if phone:
filters.append(self.model.phone.like(f'%{phone}%'))
if status is not None:
filters.append(self.model.status == status)
if filters:
stmt = stmt.where(and_(*filters))
return stmt
async def get_super(self, db: AsyncSession, user_id: int) -> bool:
"""
获取用户是否为超级管理员
:param db: 数据库会话
:param user_id: 用户 ID
:return:
"""
user = await self.get(db, user_id)
return user.is_superuser
async def get_staff(self, db: AsyncSession, user_id: int) -> bool:
"""
获取用户是否可以登录后台
:param db: 数据库会话
:param user_id: 用户 ID
:return:
"""
user = await self.get(db, user_id)
return user.is_staff
async def get_status(self, db: AsyncSession, user_id: int) -> int:
"""
获取用户状态
:param db: 数据库会话
:param user_id: 用户 ID
:return:
"""
user = await self.get(db, user_id)
return user.status
async def get_multi_login(self, db: AsyncSession, user_id: int) -> bool:
"""
获取用户是否允许多端登录
:param db: 数据库会话
:param user_id: 用户 ID
:return:
"""
user = await self.get(db, user_id)
return user.is_multi_login
async def set_super(self, db: AsyncSession, user_id: int, is_super: bool) -> int: async def set_super(self, db: AsyncSession, user_id: int, is_super: bool) -> int:
""" """
设置用户超级管理员状态 设置用户超级管理员状态
@@ -294,22 +269,19 @@ class CRUDUser(CRUDPlus[User]):
:param username: 用户名 :param username: 用户名
:return: :return:
""" """
stmt = select(self.model).options( filters = {}
selectinload(self.model.dept),
selectinload(self.model.roles).options(selectinload(Role.menus), selectinload(Role.scopes)),
)
filters = []
if user_id: if user_id:
filters.append(self.model.id == user_id) filters['id'] = user_id
if username: if username:
filters.append(self.model.username == username) filters['username'] = username
if filters: return await self.select_model_by_column(
stmt = stmt.where(and_(*filters)) db,
load_options=[selectinload(self.model.roles).options(selectinload(Role.menus), selectinload(Role.scopes))],
user = await db.execute(stmt) load_strategies=['dept'],
return user.scalars().first() **filters,
)
user_dao: CRUDUser = CRUDUser(User) user_dao: CRUDUser = CRUDUser(User)
+2 -2
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
from sqlalchemy import Boolean, ForeignKey, String from sqlalchemy import BigInteger, Boolean, ForeignKey, String
from sqlalchemy.dialects.postgresql import INTEGER from sqlalchemy.dialects.postgresql import INTEGER
from sqlalchemy.orm import Mapped, mapped_column, relationship from sqlalchemy.orm import Mapped, mapped_column, relationship
@@ -32,7 +32,7 @@ class Dept(Base):
# 父级部门一对多 # 父级部门一对多
parent_id: Mapped[int | None] = mapped_column( parent_id: Mapped[int | None] = mapped_column(
ForeignKey('sys_dept.id', ondelete='SET NULL'), default=None, index=True, comment='父部门ID' BigInteger, ForeignKey('sys_dept.id', ondelete='SET NULL'), default=None, index=True, comment='父部门ID'
) )
parent: Mapped[Optional['Dept']] = relationship(init=False, back_populates='children', remote_side=[id]) parent: Mapped[Optional['Dept']] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[Optional[list['Dept']]] = relationship(init=False, back_populates='parent') children: Mapped[Optional[list['Dept']]] = relationship(init=False, back_populates='parent')
+13 -13
View File
@@ -1,33 +1,33 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from sqlalchemy import INT, Column, ForeignKey, Integer, Table from sqlalchemy import BigInteger, Column, ForeignKey, Table
from backend.common.model import MappedBase from backend.common.model import MappedBase
sys_user_role = Table( sys_user_role = Table(
'sys_user_role', 'sys_user_role',
MappedBase.metadata, MappedBase.metadata,
Column('id', INT, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'), Column('id', BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'),
Column('user_id', Integer, ForeignKey('sys_user.id', ondelete='CASCADE'), primary_key=True, comment='用户ID'), Column('user_id', BigInteger, ForeignKey('sys_user.id', ondelete='CASCADE'), primary_key=True, comment='用户ID'),
Column('role_id', Integer, ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色ID'), Column('role_id', BigInteger, ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色ID'),
) )
sys_role_menu = Table( sys_role_menu = Table(
'sys_role_menu', 'sys_role_menu',
MappedBase.metadata, MappedBase.metadata,
Column('id', INT, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'), Column('id', BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'),
Column('role_id', Integer, ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色ID'), Column('role_id', BigInteger, ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色ID'),
Column('menu_id', Integer, ForeignKey('sys_menu.id', ondelete='CASCADE'), primary_key=True, comment='菜单ID'), Column('menu_id', BigInteger, ForeignKey('sys_menu.id', ondelete='CASCADE'), primary_key=True, comment='菜单ID'),
) )
sys_role_data_scope = Table( sys_role_data_scope = Table(
'sys_role_data_scope', 'sys_role_data_scope',
MappedBase.metadata, MappedBase.metadata,
Column('id', INT, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键 ID'), Column('id', BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键 ID'),
Column('role_id', Integer, ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色 ID'), Column('role_id', BigInteger, ForeignKey('sys_role.id', ondelete='CASCADE'), primary_key=True, comment='角色 ID'),
Column( Column(
'data_scope_id', 'data_scope_id',
Integer, BigInteger,
ForeignKey('sys_data_scope.id', ondelete='CASCADE'), ForeignKey('sys_data_scope.id', ondelete='CASCADE'),
primary_key=True, primary_key=True,
comment='数据范围 ID', comment='数据范围 ID',
@@ -37,17 +37,17 @@ sys_role_data_scope = Table(
sys_data_scope_rule = Table( sys_data_scope_rule = Table(
'sys_data_scope_rule', 'sys_data_scope_rule',
MappedBase.metadata, MappedBase.metadata,
Column('id', INT, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'), Column('id', BigInteger, primary_key=True, unique=True, index=True, autoincrement=True, comment='主键ID'),
Column( Column(
'data_scope_id', 'data_scope_id',
Integer, BigInteger,
ForeignKey('sys_data_scope.id', ondelete='CASCADE'), ForeignKey('sys_data_scope.id', ondelete='CASCADE'),
primary_key=True, primary_key=True,
comment='数据范围 ID', comment='数据范围 ID',
), ),
Column( Column(
'data_rule_id', 'data_rule_id',
Integer, BigInteger,
ForeignKey('sys_data_rule.id', ondelete='CASCADE'), ForeignKey('sys_data_rule.id', ondelete='CASCADE'),
primary_key=True, primary_key=True,
comment='数据规则 ID', comment='数据规则 ID',
+2 -2
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
from sqlalchemy import ForeignKey, String from sqlalchemy import BigInteger, ForeignKey, String
from sqlalchemy.dialects.mysql import LONGTEXT from sqlalchemy.dialects.mysql import LONGTEXT
from sqlalchemy.dialects.postgresql import TEXT from sqlalchemy.dialects.postgresql import TEXT
from sqlalchemy.orm import Mapped, mapped_column, relationship from sqlalchemy.orm import Mapped, mapped_column, relationship
@@ -42,7 +42,7 @@ class Menu(Base):
# 父级菜单一对多 # 父级菜单一对多
parent_id: Mapped[int | None] = mapped_column( parent_id: Mapped[int | None] = mapped_column(
ForeignKey('sys_menu.id', ondelete='SET NULL'), default=None, index=True, comment='父菜单ID' BigInteger, ForeignKey('sys_menu.id', ondelete='SET NULL'), default=None, index=True, comment='父菜单ID'
) )
parent: Mapped[Optional['Menu']] = relationship(init=False, back_populates='children', remote_side=[id]) parent: Mapped[Optional['Menu']] = relationship(init=False, back_populates='children', remote_side=[id])
children: Mapped[Optional[list['Menu']]] = relationship(init=False, back_populates='parent') children: Mapped[Optional[list['Menu']]] = relationship(init=False, back_populates='parent')
+2 -2
View File
@@ -27,8 +27,8 @@ class User(Base):
uuid: Mapped[str] = mapped_column(String(50), init=False, default_factory=uuid4_str, unique=True) uuid: Mapped[str] = mapped_column(String(50), init=False, default_factory=uuid4_str, unique=True)
username: Mapped[str] = mapped_column(String(20), unique=True, index=True, comment='用户名') username: Mapped[str] = mapped_column(String(20), unique=True, index=True, comment='用户名')
nickname: Mapped[str] = mapped_column(String(20), comment='昵称') nickname: Mapped[str] = mapped_column(String(20), comment='昵称')
password: Mapped[str] = mapped_column(String(255), comment='密码') password: Mapped[str | None] = mapped_column(String(255), comment='密码')
salt: Mapped[bytes] = mapped_column(VARBINARY(255).with_variant(BYTEA(255), 'postgresql'), comment='加密盐') salt: Mapped[bytes | None] = mapped_column(VARBINARY(255).with_variant(BYTEA(255), 'postgresql'), comment='加密盐')
email: Mapped[str | None] = mapped_column(String(50), default=None, unique=True, index=True, comment='邮箱') email: Mapped[str | None] = mapped_column(String(50), default=None, unique=True, index=True, comment='邮箱')
phone: Mapped[str | None] = mapped_column(String(11), default=None, comment='手机号') phone: Mapped[str | None] = mapped_column(String(11), default=None, comment='手机号')
avatar: Mapped[str | None] = mapped_column(String(255), default=None, comment='头像') avatar: Mapped[str | None] = mapped_column(String(255), default=None, comment='头像')
+8 -2
View File
@@ -14,8 +14,8 @@ class DataRuleSchemaBase(SchemaBase):
name: str = Field(description='规则名称') name: str = Field(description='规则名称')
model: str = Field(description='模型名称') model: str = Field(description='模型名称')
column: str = Field(description='字段名称') column: str = Field(description='字段名称')
operator: RoleDataRuleOperatorType = Field(RoleDataRuleOperatorType.AND, description='操作符(AND/OR') operator: RoleDataRuleOperatorType = Field(description='操作符(AND/OR')
expression: RoleDataRuleExpressionType = Field(RoleDataRuleExpressionType.eq, description='表达式类型') expression: RoleDataRuleExpressionType = Field(description='表达式类型')
value: str = Field(description='规则值') value: str = Field(description='规则值')
@@ -27,6 +27,12 @@ class UpdateDataRuleParam(DataRuleSchemaBase):
"""更新数据规则参数""" """更新数据规则参数"""
class DeleteDataRuleParam(SchemaBase):
"""删除数据规则参数"""
pks: list[int] = Field(description='规则 ID 列表')
class GetDataRuleDetail(DataRuleSchemaBase): class GetDataRuleDetail(DataRuleSchemaBase):
"""数据规则详情""" """数据规则详情"""
+7 -1
View File
@@ -13,7 +13,7 @@ class DataScopeBase(SchemaBase):
"""数据范围基础模型""" """数据范围基础模型"""
name: str = Field(description='名称') name: str = Field(description='名称')
status: StatusType = Field(StatusType.enable, description='状态') status: StatusType = Field(description='状态')
class CreateDataScopeParam(DataScopeBase): class CreateDataScopeParam(DataScopeBase):
@@ -30,6 +30,12 @@ class UpdateDataScopeRuleParam(SchemaBase):
rules: list[int] = Field(description='数据规则 ID 列表') rules: list[int] = Field(description='数据规则 ID 列表')
class DeleteDataScopeParam(SchemaBase):
"""删除数据范围参数"""
pks: list[int] = Field(description='数据范围 ID 列表')
class GetDataScopeDetail(DataScopeBase): class GetDataScopeDetail(DataScopeBase):
"""数据范围详情""" """数据范围详情"""
+1 -1
View File
@@ -17,7 +17,7 @@ class DeptSchemaBase(SchemaBase):
leader: str | None = Field(None, description='负责人') leader: str | None = Field(None, description='负责人')
phone: CustomPhoneNumber | None = Field(None, description='联系电话') phone: CustomPhoneNumber | None = Field(None, description='联系电话')
email: CustomEmailStr | None = Field(None, description='邮箱') email: CustomEmailStr | None = Field(None, description='邮箱')
status: StatusType = Field(StatusType.enable, description='状态') status: StatusType = Field(description='状态')
class CreateDeptParam(DeptSchemaBase): class CreateDeptParam(DeptSchemaBase):
+6
View File
@@ -33,6 +33,12 @@ class UpdateLoginLogParam(LoginLogSchemaBase):
"""更新登录日志参数""" """更新登录日志参数"""
class DeleteLoginLogParam(SchemaBase):
"""删除登录日志参数"""
pks: list[int] = Field(description='登录日志 ID 列表')
class GetLoginLogDetail(LoginLogSchemaBase): class GetLoginLogDetail(LoginLogSchemaBase):
"""登录日志详情""" """登录日志详情"""
+4 -4
View File
@@ -17,12 +17,12 @@ class MenuSchemaBase(SchemaBase):
parent_id: int | None = Field(None, description='菜单父级 ID') parent_id: int | None = Field(None, description='菜单父级 ID')
sort: int = Field(0, ge=0, description='排序') sort: int = Field(0, ge=0, description='排序')
icon: str | None = Field(None, description='图标') icon: str | None = Field(None, description='图标')
type: MenuType = Field(MenuType.directory, description='菜单类型(0目录 1菜单 2按钮 3内嵌 4外链)') type: MenuType = Field(description='菜单类型(0目录 1菜单 2按钮 3内嵌 4外链)')
component: str | None = Field(None, description='组件路径') component: str | None = Field(None, description='组件路径')
perms: str | None = Field(None, description='权限标识') perms: str | None = Field(None, description='权限标识')
status: StatusType = Field(StatusType.enable, description='状态') status: StatusType = Field(description='状态')
display: StatusType = Field(StatusType.enable, description='是否显示') display: StatusType = Field(description='是否显示')
cache: StatusType = Field(StatusType.enable, description='是否缓存') cache: StatusType = Field(description='是否缓存')
link: str | None = Field(None, description='外链地址') link: str | None = Field(None, description='外链地址')
remark: str | None = Field(None, description='备注') remark: str | None = Field(None, description='备注')
+7 -1
View File
@@ -26,7 +26,7 @@ class OperaLogSchemaBase(SchemaBase):
browser: str | None = Field(None, description='浏览器') browser: str | None = Field(None, description='浏览器')
device: str | None = Field(None, description='设备') device: str | None = Field(None, description='设备')
args: dict[str, Any] | None = Field(None, description='请求参数') args: dict[str, Any] | None = Field(None, description='请求参数')
status: StatusType = Field(StatusType.enable, description='状态') status: StatusType = Field(description='状态')
code: str = Field(description='状态码') code: str = Field(description='状态码')
msg: str | None = Field(None, description='消息') msg: str | None = Field(None, description='消息')
cost_time: float = Field(description='耗时') cost_time: float = Field(description='耗时')
@@ -41,6 +41,12 @@ class UpdateOperaLogParam(OperaLogSchemaBase):
"""更新操作日志参数""" """更新操作日志参数"""
class DeleteOperaLogParam(SchemaBase):
"""删除操作日志参数"""
pks: list[int] = Field(description='操作日志 ID 列表')
class GetOperaLogDetail(OperaLogSchemaBase): class GetOperaLogDetail(OperaLogSchemaBase):
"""操作日志详情""" """操作日志详情"""
+7 -1
View File
@@ -14,7 +14,7 @@ class RoleSchemaBase(SchemaBase):
"""角色基础模型""" """角色基础模型"""
name: str = Field(description='角色名称') name: str = Field(description='角色名称')
status: StatusType = Field(StatusType.enable, description='状态') status: StatusType = Field(description='状态')
is_filter_scopes: bool = Field(True, description='过滤数据权限') is_filter_scopes: bool = Field(True, description='过滤数据权限')
remark: str | None = Field(None, description='备注') remark: str | None = Field(None, description='备注')
@@ -27,6 +27,12 @@ class UpdateRoleParam(RoleSchemaBase):
"""更新角色参数""" """更新角色参数"""
class DeleteRoleParam(SchemaBase):
"""删除角色参数"""
pks: list[int] = Field(description='角色 ID 列表')
class UpdateRoleMenuParam(SchemaBase): class UpdateRoleMenuParam(SchemaBase):
"""更新角色菜单参数""" """更新角色菜单参数"""
+10 -7
View File
@@ -3,7 +3,7 @@
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from pydantic import ConfigDict, EmailStr, Field, HttpUrl, model_validator from pydantic import ConfigDict, Field, HttpUrl, model_validator
from typing_extensions import Self from typing_extensions import Self
from backend.app.admin.schema.dept import GetDeptDetail from backend.app.admin.schema.dept import GetDeptDetail
@@ -16,7 +16,7 @@ class AuthSchemaBase(SchemaBase):
"""用户认证基础模型""" """用户认证基础模型"""
username: str = Field(description='用户名') username: str = Field(description='用户名')
password: str | None = Field(description='密码') password: str = Field(description='密码')
class AuthLoginParam(AuthSchemaBase): class AuthLoginParam(AuthSchemaBase):
@@ -28,16 +28,19 @@ class AuthLoginParam(AuthSchemaBase):
class AddUserParam(AuthSchemaBase): class AddUserParam(AuthSchemaBase):
"""添加用户参数""" """添加用户参数"""
nickname: str | None = Field(None, description='昵称')
email: CustomEmailStr | None = Field(None, description='邮箱')
phone: CustomPhoneNumber | None = Field(None, description='手机号码')
dept_id: int = Field(description='部门 ID') dept_id: int = Field(description='部门 ID')
roles: list[int] = Field(description='角色 ID 列表') roles: list[int] = Field(description='角色 ID 列表')
nickname: str | None = Field(None, description='昵称')
class AddOAuth2UserParam(AuthSchemaBase): class AddOAuth2UserParam(AuthSchemaBase):
"""添加 OAuth2 用户参数""" """添加 OAuth2 用户参数"""
password: str | None = Field(None, description='密码')
nickname: str | None = Field(None, description='昵称') nickname: str | None = Field(None, description='昵称')
email: EmailStr = Field(description='邮箱') email: CustomEmailStr | None = Field(None, description='邮箱')
avatar: HttpUrl | None = Field(None, description='头像地址') avatar: HttpUrl | None = Field(None, description='头像地址')
@@ -56,6 +59,8 @@ class UserInfoSchemaBase(SchemaBase):
username: str = Field(description='用户名') username: str = Field(description='用户名')
nickname: str = Field(description='昵称') nickname: str = Field(description='昵称')
avatar: HttpUrl | None = Field(None, description='头像地址') avatar: HttpUrl | None = Field(None, description='头像地址')
email: CustomEmailStr | None = Field(None, description='邮箱')
phone: CustomPhoneNumber | None = Field(None, description='手机号')
class UpdateUserParam(UserInfoSchemaBase): class UpdateUserParam(UserInfoSchemaBase):
@@ -72,9 +77,7 @@ class GetUserInfoDetail(UserInfoSchemaBase):
dept_id: int | None = Field(None, description='部门 ID') dept_id: int | None = Field(None, description='部门 ID')
id: int = Field(description='用户 ID') id: int = Field(description='用户 ID')
uuid: str = Field(description='用户 UUID') uuid: str = Field(description='用户 UUID')
email: CustomEmailStr | None = Field(None, description='邮箱') status: StatusType = Field(description='状态')
phone: CustomPhoneNumber | None = Field(None, description='手机号')
status: StatusType = Field(StatusType.enable, description='状态')
is_superuser: bool = Field(description='是否超级管理员') is_superuser: bool = Field(description='是否超级管理员')
is_staff: bool = Field(description='是否管理员') is_staff: bool = Field(description='是否管理员')
is_multi_login: bool = Field(description='是否允许多端登录') is_multi_login: bool = Field(description='是否允许多端登录')
+63 -42
View File
@@ -5,6 +5,7 @@ from fastapi.security import HTTPBasicCredentials
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from starlette.background import BackgroundTask, BackgroundTasks from starlette.background import BackgroundTask, BackgroundTasks
from backend.app.admin.crud.crud_menu import menu_dao
from backend.app.admin.crud.crud_user import user_dao from backend.app.admin.crud.crud_user import user_dao
from backend.app.admin.model import User from backend.app.admin.model import User
from backend.app.admin.schema.token import GetLoginToken, GetNewToken from backend.app.admin.schema.token import GetLoginToken, GetNewToken
@@ -12,6 +13,7 @@ from backend.app.admin.schema.user import AuthLoginParam
from backend.app.admin.service.login_log_service import login_log_service from backend.app.admin.service.login_log_service import login_log_service
from backend.common.enums import LoginLogStatusType from backend.common.enums import LoginLogStatusType
from backend.common.exception import errors from backend.common.exception import errors
from backend.common.i18n import t
from backend.common.log import log from backend.common.log import log
from backend.common.response.response_code import CustomErrorCode from backend.common.response.response_code import CustomErrorCode
from backend.common.security.jwt import ( from backend.common.security.jwt import (
@@ -32,7 +34,7 @@ class AuthService:
"""认证服务类""" """认证服务类"""
@staticmethod @staticmethod
async def user_verify(db: AsyncSession, username: str, password: str | None) -> User: async def user_verify(db: AsyncSession, username: str, password: str) -> User:
""" """
验证用户名和密码 验证用户名和密码
@@ -66,13 +68,13 @@ class AuthService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
user = await self.user_verify(db, obj.username, obj.password) user = await self.user_verify(db, obj.username, obj.password)
await user_dao.update_login_time(db, obj.username) await user_dao.update_login_time(db, obj.username)
a_token = await create_access_token( access_token = await create_access_token(
str(user.id), user.id,
user.is_multi_login, user.is_multi_login,
# extra info # extra info
swagger=True, swagger=True,
) )
return a_token.access_token, user return access_token.access_token, user
async def login( async def login(
self, *, request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks self, *, request: Request, response: Response, obj: AuthLoginParam, background_tasks: BackgroundTasks
@@ -92,36 +94,36 @@ class AuthService:
user = await self.user_verify(db, obj.username, obj.password) user = await self.user_verify(db, obj.username, obj.password)
captcha_code = await redis_client.get(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}') captcha_code = await redis_client.get(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}')
if not captcha_code: if not captcha_code:
raise errors.ForbiddenError(msg='验证码失效,请重新获取') raise errors.RequestError(msg=t('error.captcha.expired'))
if captcha_code.lower() != obj.captcha.lower(): if captcha_code.lower() != obj.captcha.lower():
raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR)
await redis_client.delete(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}') await redis_client.delete(f'{settings.CAPTCHA_LOGIN_REDIS_PREFIX}:{request.state.ip}')
await user_dao.update_login_time(db, obj.username) await user_dao.update_login_time(db, obj.username)
await db.refresh(user) await db.refresh(user)
a_token = await create_access_token( access_token = await create_access_token(
str(user.id), user.id,
user.is_multi_login, user.is_multi_login,
# extra info # extra info
username=user.username, username=user.username,
nickname=user.nickname, nickname=user.nickname,
last_login_time=timezone.t_str(user.last_login_time), last_login_time=timezone.to_str(user.last_login_time),
ip=request.state.ip, ip=request.state.ip,
os=request.state.os, os=request.state.os,
browser=request.state.browser, browser=request.state.browser,
device=request.state.device, device=request.state.device,
) )
r_token = await create_refresh_token(str(user.id), user.is_multi_login) refresh_token = await create_refresh_token(access_token.session_uuid, user.id, user.is_multi_login)
response.set_cookie( response.set_cookie(
key=settings.COOKIE_REFRESH_TOKEN_KEY, key=settings.COOKIE_REFRESH_TOKEN_KEY,
value=r_token.refresh_token, value=refresh_token.refresh_token,
max_age=settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS, max_age=settings.COOKIE_REFRESH_TOKEN_EXPIRE_SECONDS,
expires=timezone.f_utc(r_token.refresh_token_expire_time), expires=timezone.to_utc(refresh_token.refresh_token_expire_time),
httponly=True, httponly=True,
) )
except errors.NotFoundError as e: except errors.NotFoundError as e:
log.error('登陆错误: 用户名不存在') log.error('登陆错误: 用户名不存在')
raise errors.NotFoundError(msg=e.msg) raise errors.NotFoundError(msg=e.msg)
except (errors.ForbiddenError, errors.CustomError) as e: except (errors.RequestError, errors.CustomError) as e:
if not user: if not user:
log.error('登陆错误: 用户密码有误') log.error('登陆错误: 用户密码有误')
task = BackgroundTask( task = BackgroundTask(
@@ -136,7 +138,7 @@ class AuthService:
msg=e.msg, msg=e.msg,
), ),
) )
raise errors.RequestError(msg=e.msg, background=task) raise errors.RequestError(code=e.code, msg=e.msg, background=task)
except Exception as e: except Exception as e:
log.error(f'登陆错误: {e}') log.error(f'登陆错误: {e}')
raise e raise e
@@ -150,46 +152,72 @@ class AuthService:
username=obj.username, username=obj.username,
login_time=timezone.now(), login_time=timezone.now(),
status=LoginLogStatusType.success.value, status=LoginLogStatusType.success.value,
msg='登录成功', msg=t('success.login.success'),
), ),
) )
data = GetLoginToken( data = GetLoginToken(
access_token=a_token.access_token, access_token=access_token.access_token,
access_token_expire_time=a_token.access_token_expire_time, access_token_expire_time=access_token.access_token_expire_time,
session_uuid=a_token.session_uuid, session_uuid=access_token.session_uuid,
user=user, # type: ignore user=user, # type: ignore
) )
return data return data
@staticmethod @staticmethod
async def new_token(*, request: Request) -> GetNewToken: async def get_codes(*, request: Request) -> list[str]:
""" """
获取新的访问令牌 获取用户权限码
: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(','))
else:
roles = request.user.roles
if roles:
for role in roles:
for menu in role.menus:
if menu.perms:
codes.add(*menu.perms.split(','))
return list(codes)
@staticmethod
async def refresh_token(*, request: Request) -> GetNewToken:
"""
刷新令牌
:param request: FastAPI 请求对象 :param request: FastAPI 请求对象
:return: :return:
""" """
refresh_token = request.cookies.get(settings.COOKIE_REFRESH_TOKEN_KEY) refresh_token = request.cookies.get(settings.COOKIE_REFRESH_TOKEN_KEY)
if not refresh_token: if not refresh_token:
raise errors.TokenError(msg='Refresh Token 已过期,请重新登录') raise errors.RequestError(msg='Refresh Token 已过期,请重新登录')
try: token_payload = jwt_decode(refresh_token)
user_id = jwt_decode(refresh_token).id
except Exception:
raise errors.TokenError(msg='Refresh Token 无效')
async with async_db_session() as db: async with async_db_session() as db:
user = await user_dao.get(db, user_id) user = await user_dao.get(db, token_payload.id)
if not user: if not user:
raise errors.NotFoundError(msg='用户名或密码有误') raise errors.NotFoundError(msg='用户不存在')
elif not user.status: elif not user.status:
raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员') raise errors.AuthorizationError(msg='用户已被锁定, 请联系统管理员')
if not user.is_multi_login:
if await redis_client.keys(match=f'{settings.TOKEN_REDIS_PREFIX}:{user.id}:*'):
raise errors.ForbiddenError(msg='此用户已在异地登录,请重新登录并及时修改密码')
new_token = await create_new_token( new_token = await create_new_token(
user_id=str(user.id), refresh_token,
refresh_token=refresh_token, token_payload.session_uuid,
multi_login=user.is_multi_login, user.id,
user.is_multi_login,
# extra info # extra info
username=user.username, username=user.username,
nickname=user.nickname, nickname=user.nickname,
last_login_time=timezone.t_str(user.last_login_time), last_login_time=timezone.to_str(user.last_login_time),
ip=request.state.ip, ip=request.state.ip,
os=request.state.os, os=request.state.os,
browser=request.state.browser, browser=request.state.browser,
@@ -215,24 +243,17 @@ class AuthService:
token = get_token(request) token = get_token(request)
token_payload = jwt_decode(token) token_payload = jwt_decode(token)
user_id = token_payload.id user_id = token_payload.id
session_uuid = token_payload.session_uuid
refresh_token = request.cookies.get(settings.COOKIE_REFRESH_TOKEN_KEY) refresh_token = request.cookies.get(settings.COOKIE_REFRESH_TOKEN_KEY)
except errors.TokenError: except errors.TokenError:
return return
finally: finally:
response.delete_cookie(settings.COOKIE_REFRESH_TOKEN_KEY) response.delete_cookie(settings.COOKIE_REFRESH_TOKEN_KEY)
# 清理缓存 await redis_client.delete(f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:{session_uuid}')
if request.user.is_multi_login: await redis_client.delete(f'{settings.TOKEN_EXTRA_INFO_REDIS_PREFIX}:{user_id}:{session_uuid}')
await redis_client.delete(f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:{token_payload.session_uuid}') if refresh_token:
if refresh_token: await redis_client.delete(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}:{refresh_token}')
await redis_client.delete(f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}:{refresh_token}')
else:
key_prefix = [
f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:',
f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user_id}:',
]
for prefix in key_prefix:
await redis_client.delete_prefix(prefix)
auth_service: AuthService = AuthService() auth_service: AuthService = AuthService()
+12 -7
View File
@@ -6,7 +6,12 @@ from sqlalchemy import Select
from backend.app.admin.crud.crud_data_rule import data_rule_dao from backend.app.admin.crud.crud_data_rule import data_rule_dao
from backend.app.admin.model import DataRule from backend.app.admin.model import DataRule
from backend.app.admin.schema.data_rule import CreateDataRuleParam, GetDataRuleColumnDetail, UpdateDataRuleParam from backend.app.admin.schema.data_rule import (
CreateDataRuleParam,
DeleteDataRuleParam,
GetDataRuleColumnDetail,
UpdateDataRuleParam,
)
from backend.common.exception import errors from backend.common.exception import errors
from backend.core.conf import settings from backend.core.conf import settings
from backend.database.db import async_db_session from backend.database.db import async_db_session
@@ -82,7 +87,7 @@ class DataRuleService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
data_rule = await data_rule_dao.get_by_name(db, obj.name) data_rule = await data_rule_dao.get_by_name(db, obj.name)
if data_rule: if data_rule:
raise errors.ForbiddenError(msg='数据规则已存在') raise errors.ConflictError(msg='数据规则已存在')
await data_rule_dao.create(db, obj) await data_rule_dao.create(db, obj)
@staticmethod @staticmethod
@@ -100,20 +105,20 @@ class DataRuleService:
raise errors.NotFoundError(msg='数据规则不存在') raise errors.NotFoundError(msg='数据规则不存在')
if data_rule.name != obj.name: if data_rule.name != obj.name:
if await data_rule_dao.get_by_name(db, obj.name): if await data_rule_dao.get_by_name(db, obj.name):
raise errors.ForbiddenError(msg='数据规则已存在') raise errors.ConflictError(msg='数据规则已存在')
count = await data_rule_dao.update(db, pk, obj) count = await data_rule_dao.update(db, pk, obj)
return count return count
@staticmethod @staticmethod
async def delete(*, pk: list[int]) -> int: async def delete(*, obj: DeleteDataRuleParam) -> int:
""" """
删除数据规则 批量删除数据规则
:param pk: 规则 ID 列表 :param obj: 规则 ID 列表
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await data_rule_dao.delete(db, pk) count = await data_rule_dao.delete(db, obj.pks)
return count return count
@@ -6,7 +6,12 @@ from sqlalchemy import Select
from backend.app.admin.crud.crud_data_scope import data_scope_dao from backend.app.admin.crud.crud_data_scope import data_scope_dao
from backend.app.admin.model import DataScope from backend.app.admin.model import DataScope
from backend.app.admin.schema.data_scope import CreateDataScopeParam, UpdateDataScopeParam, UpdateDataScopeRuleParam from backend.app.admin.schema.data_scope import (
CreateDataScopeParam,
DeleteDataScopeParam,
UpdateDataScopeParam,
UpdateDataScopeRuleParam,
)
from backend.common.exception import errors from backend.common.exception import errors
from backend.core.conf import settings from backend.core.conf import settings
from backend.database.db import async_db_session from backend.database.db import async_db_session
@@ -73,7 +78,7 @@ class DataScopeService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
data_scope = await data_scope_dao.get_by_name(db, obj.name) data_scope = await data_scope_dao.get_by_name(db, obj.name)
if data_scope: if data_scope:
raise errors.ForbiddenError(msg='数据范围已存在') raise errors.ConflictError(msg='数据范围已存在')
await data_scope_dao.create(db, obj) await data_scope_dao.create(db, obj)
@staticmethod @staticmethod
@@ -91,7 +96,7 @@ class DataScopeService:
raise errors.NotFoundError(msg='数据范围不存在') raise errors.NotFoundError(msg='数据范围不存在')
if data_scope.name != obj.name: if data_scope.name != obj.name:
if await data_scope_dao.get_by_name(db, obj.name): if await data_scope_dao.get_by_name(db, obj.name):
raise errors.ForbiddenError(msg='数据范围已存在') raise errors.ConflictError(msg='数据范围已存在')
count = await data_scope_dao.update(db, pk, obj) count = await data_scope_dao.update(db, pk, obj)
for role in await data_scope.awaitable_attrs.roles: for role in await data_scope.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users: for user in await role.awaitable_attrs.users:
@@ -112,17 +117,17 @@ class DataScopeService:
return count return count
@staticmethod @staticmethod
async def delete(*, pk: list[int]) -> int: async def delete(*, obj: DeleteDataScopeParam) -> int:
""" """
删除数据范围 批量删除数据范围
:param pk: 范围 ID 列表 :param obj: 范围 ID 列表
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await data_scope_dao.delete(db, pk) count = await data_scope_dao.delete(db, obj.pks)
for _pk in pk: for pk in obj.pks:
data_rule = await data_scope_dao.get(db, _pk) data_rule = await data_scope_dao.get(db, pk)
if data_rule: if data_rule:
for role in await data_rule.awaitable_attrs.roles: for role in await data_rule.awaitable_attrs.roles:
for user in await role.awaitable_attrs.users: for user in await role.awaitable_attrs.users:
+5 -5
View File
@@ -32,7 +32,7 @@ class DeptService:
return dept return dept
@staticmethod @staticmethod
async def get_dept_tree( async def get_tree(
*, request: Request, name: str | None, leader: str | None, phone: str | None, status: int | None *, request: Request, name: str | None, leader: str | None, phone: str | None, status: int | None
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
""" """
@@ -61,7 +61,7 @@ class DeptService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
dept = await dept_dao.get_by_name(db, obj.name) dept = await dept_dao.get_by_name(db, obj.name)
if dept: if dept:
raise errors.ForbiddenError(msg='部门名称已存在') raise errors.ConflictError(msg='部门名称已存在')
if obj.parent_id: if obj.parent_id:
parent_dept = await dept_dao.get(db, obj.parent_id) parent_dept = await dept_dao.get(db, obj.parent_id)
if not parent_dept: if not parent_dept:
@@ -83,7 +83,7 @@ class DeptService:
raise errors.NotFoundError(msg='部门不存在') raise errors.NotFoundError(msg='部门不存在')
if dept.name != obj.name: if dept.name != obj.name:
if await dept_dao.get_by_name(db, obj.name): if await dept_dao.get_by_name(db, obj.name):
raise errors.ForbiddenError(msg='部门名称已存在') raise errors.ConflictError(msg='部门名称已存在')
if obj.parent_id: if obj.parent_id:
parent_dept = await dept_dao.get(db, obj.parent_id) parent_dept = await dept_dao.get(db, obj.parent_id)
if not parent_dept: if not parent_dept:
@@ -104,10 +104,10 @@ class DeptService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
dept = await dept_dao.get_with_relation(db, pk) dept = await dept_dao.get_with_relation(db, pk)
if dept.users: if dept.users:
raise errors.ForbiddenError(msg='部门下存在用户,无法删除') raise errors.ConflictError(msg='部门下存在用户,无法删除')
children = await dept_dao.get_children(db, pk) children = await dept_dao.get_children(db, pk)
if children: if children:
raise errors.ForbiddenError(msg='部门下存在子部门,无法删除') raise errors.ConflictError(msg='部门下存在子部门,无法删除')
count = await dept_dao.delete(db, pk) count = await dept_dao.delete(db, pk)
for user in dept.users: for user in dept.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
@@ -7,7 +7,7 @@ from sqlalchemy import Select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from backend.app.admin.crud.crud_login_log import login_log_dao from backend.app.admin.crud.crud_login_log import login_log_dao
from backend.app.admin.schema.login_log import CreateLoginLogParam from backend.app.admin.schema.login_log import CreateLoginLogParam, DeleteLoginLogParam
from backend.common.log import log from backend.common.log import log
from backend.database.db import async_db_session from backend.database.db import async_db_session
@@ -71,23 +71,22 @@ class LoginLogService:
log.error(f'登录日志创建失败: {e}') log.error(f'登录日志创建失败: {e}')
@staticmethod @staticmethod
async def delete(*, pk: list[int]) -> int: async def delete(*, obj: DeleteLoginLogParam) -> int:
""" """
删除登录日志 批量删除登录日志
:param pk: 日志 ID 列表 :param obj: 日志 ID 列表
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await login_log_dao.delete(db, pk) count = await login_log_dao.delete(db, obj.pks)
return count return count
@staticmethod @staticmethod
async def delete_all() -> int: async def delete_all() -> None:
"""清空所有登录日志""" """清空所有登录日志"""
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await login_log_dao.delete_all(db) await login_log_dao.delete_all(db)
return count
login_log_service: LoginLogService = LoginLogService() login_log_service: LoginLogService = LoginLogService()
+15 -19
View File
@@ -32,7 +32,7 @@ class MenuService:
return menu return menu
@staticmethod @staticmethod
async def get_menu_tree(*, title: str | None, status: int | None) -> list[dict[str, Any]]: async def get_tree(*, title: str | None, status: int | None) -> list[dict[str, Any]]:
""" """
获取菜单树形结构 获取菜单树形结构
@@ -54,21 +54,17 @@ class MenuService:
:return: :return:
""" """
async with async_db_session() as db: async with async_db_session() as db:
roles = request.user.roles if request.user.is_superuser:
menu_tree = [] menu_data = await menu_dao.get_sidebar(db, None)
if roles: else:
unique_menus = {} roles = request.user.roles
for role in roles: menu_ids = set()
for menu in role.menus: if roles:
unique_menus[menu.id] = menu for role in roles:
all_ids = set(unique_menus.keys()) for menu in role.menus:
valid_menu_ids = [ menu_ids.add(menu.id)
menu_id menu_data = await menu_dao.get_sidebar(db, list(menu_ids))
for menu_id, menu in unique_menus.items() menu_tree = get_vben5_tree_data(menu_data)
if menu.parent_id is None or menu.parent_id in all_ids
]
menu_data = await menu_dao.get_sidebar(db, request.user.is_superuser, valid_menu_ids)
menu_tree = get_vben5_tree_data(menu_data)
return menu_tree return menu_tree
@staticmethod @staticmethod
@@ -82,7 +78,7 @@ class MenuService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
title = await menu_dao.get_by_title(db, obj.title) title = await menu_dao.get_by_title(db, obj.title)
if title: if title:
raise errors.ForbiddenError(msg='菜单标题已存在') raise errors.ConflictError(msg='菜单标题已存在')
if obj.parent_id: if obj.parent_id:
parent_menu = await menu_dao.get(db, obj.parent_id) parent_menu = await menu_dao.get(db, obj.parent_id)
if not parent_menu: if not parent_menu:
@@ -104,7 +100,7 @@ class MenuService:
raise errors.NotFoundError(msg='菜单不存在') raise errors.NotFoundError(msg='菜单不存在')
if menu.title != obj.title: if menu.title != obj.title:
if await menu_dao.get_by_title(db, obj.title): if await menu_dao.get_by_title(db, obj.title):
raise errors.ForbiddenError(msg='菜单标题已存在') raise errors.ConflictError(msg='菜单标题已存在')
if obj.parent_id: if obj.parent_id:
parent_menu = await menu_dao.get(db, obj.parent_id) parent_menu = await menu_dao.get(db, obj.parent_id)
if not parent_menu: if not parent_menu:
@@ -128,7 +124,7 @@ class MenuService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
children = await menu_dao.get_children(db, pk) children = await menu_dao.get_children(db, pk)
if children: if children:
raise errors.ForbiddenError(msg='菜单下存在子菜单,无法删除') raise errors.ConflictError(msg='菜单下存在子菜单,无法删除')
menu = await menu_dao.get(db, pk) menu = await menu_dao.get(db, pk)
count = await menu_dao.delete(db, pk) count = await menu_dao.delete(db, pk)
if menu: if menu:
+18 -8
View File
@@ -3,7 +3,7 @@
from sqlalchemy import Select from sqlalchemy import Select
from backend.app.admin.crud.crud_opera_log import opera_log_dao from backend.app.admin.crud.crud_opera_log import opera_log_dao
from backend.app.admin.schema.opera_log import CreateOperaLogParam from backend.app.admin.schema.opera_log import CreateOperaLogParam, DeleteOperaLogParam
from backend.database.db import async_db_session from backend.database.db import async_db_session
@@ -34,23 +34,33 @@ class OperaLogService:
await opera_log_dao.create(db, obj) await opera_log_dao.create(db, obj)
@staticmethod @staticmethod
async def delete(*, pk: list[int]) -> int: async def bulk_create(*, objs: list[CreateOperaLogParam]) -> None:
""" """
删除操作日志 批量创建操作日志
:param pk: 日志 ID 列表 :param objs: 操作日志创建参数列表
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await opera_log_dao.delete(db, pk) await opera_log_dao.bulk_create(db, objs)
@staticmethod
async def delete(*, obj: DeleteOperaLogParam) -> int:
"""
批量删除操作日志
:param obj: 日志 ID 列表
:return:
"""
async with async_db_session.begin() as db:
count = await opera_log_dao.delete(db, obj.pks)
return count return count
@staticmethod @staticmethod
async def delete_all() -> int: async def delete_all() -> None:
"""清空所有操作日志""" """清空所有操作日志"""
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await opera_log_dao.delete_all(db) await opera_log_dao.delete_all(db)
return count
opera_log_service: OperaLogService = OperaLogService() opera_log_service: OperaLogService = OperaLogService()
+22 -79
View File
@@ -8,17 +8,15 @@ import zipfile
from typing import Any from typing import Any
from dulwich import porcelain
from fastapi import UploadFile from fastapi import UploadFile
from backend.common.enums import StatusType from backend.common.enums import PluginType, StatusType
from backend.common.exception import errors from backend.common.exception import errors
from backend.common.log import log
from backend.core.conf import settings from backend.core.conf import settings
from backend.core.path_conf import PLUGIN_DIR from backend.core.path_conf import PLUGIN_DIR
from backend.database.redis import redis_client from backend.database.redis import redis_client
from backend.plugin.tools import install_requirements_async, uninstall_requirements_async from backend.plugin.tools import uninstall_requirements_async
from backend.utils.re_verify import is_git_url from backend.utils.file_ops import install_git_plugin, install_zip_plugin
from backend.utils.timezone import timezone from backend.utils.timezone import timezone
@@ -31,7 +29,7 @@ class PluginService:
keys = [] keys = []
result = [] result = []
async for key in redis_client.scan_iter(f'{settings.PLUGIN_REDIS_PREFIX}:info:*'): async for key in redis_client.scan_iter(f'{settings.PLUGIN_REDIS_PREFIX}:*'):
keys.append(key) keys.append(key)
for info in await redis_client.mget(*keys): for info in await redis_client.mget(*keys):
@@ -41,77 +39,26 @@ class PluginService:
@staticmethod @staticmethod
async def changed() -> str | None: async def changed() -> str | None:
"""插件状态是否变更""" """检查插件是否发生变更"""
return await redis_client.get(f'{settings.PLUGIN_REDIS_PREFIX}:changed') return await redis_client.get(f'{settings.PLUGIN_REDIS_PREFIX}:changed')
@staticmethod @staticmethod
async def install_zip(*, file: UploadFile) -> None: async def install(*, type: PluginType, file: UploadFile | None = None, repo_url: str | None = None) -> str:
""" """
通过 zip 压缩包安装插件 安装插件
:param type: 插件类型
:param file: 插件 zip 压缩包 :param file: 插件 zip 压缩包
:param repo_url: git 仓库地址
:return: :return:
""" """
contents = await file.read() if type == PluginType.zip:
file_bytes = io.BytesIO(contents) if not file:
if not zipfile.is_zipfile(file_bytes): raise errors.RequestError(msg='ZIP 压缩包不能为空')
raise errors.ForbiddenError(msg='插件压缩包格式非法') return await install_zip_plugin(file)
with zipfile.ZipFile(file_bytes) as zf: if not repo_url:
# 校验压缩包 raise errors.RequestError(msg='Git 仓库地址不能为空')
plugin_namelist = zf.namelist() return await install_git_plugin(repo_url)
plugin_name = plugin_namelist[0].split('/')[0]
if not plugin_namelist or plugin_name not in file.filename:
raise errors.ForbiddenError(msg='插件压缩包内容非法')
if (
len(plugin_namelist) <= 3
or f'{plugin_name}/plugin.toml' not in plugin_namelist
or f'{plugin_name}/README.md' not in plugin_namelist
):
raise errors.ForbiddenError(msg='插件压缩包内缺少必要文件')
# 插件是否可安装
full_plugin_path = os.path.join(PLUGIN_DIR, plugin_name)
if os.path.exists(full_plugin_path):
raise errors.ForbiddenError(msg='此插件已安装')
else:
os.makedirs(full_plugin_path, exist_ok=True)
# 解压(安装)
members = []
for member in zf.infolist():
if member.filename.startswith(plugin_name):
new_filename = member.filename.replace(plugin_name, '')
if new_filename:
member.filename = new_filename
members.append(member)
zf.extractall(os.path.join(PLUGIN_DIR, plugin_name), members)
await install_requirements_async(plugin_name)
await redis_client.set(f'{settings.PLUGIN_REDIS_PREFIX}:changed', 'ture')
@staticmethod
async def install_git(*, repo_url: str):
"""
通过 git 安装插件
:param repo_url: git 存储库的 URL
:return:
"""
match = is_git_url(repo_url)
if not match:
raise errors.ForbiddenError(msg='Git 仓库地址格式非法')
repo_name = match.group('repo')
plugins = await redis_client.lrange(settings.PLUGIN_REDIS_PREFIX, 0, -1)
if repo_name in plugins:
raise errors.ForbiddenError(msg=f'{repo_name} 插件已安装')
try:
porcelain.clone(repo_url, os.path.join(PLUGIN_DIR, repo_name), checkout=True)
except Exception as e:
log.error(f'插件安装失败: {e}')
raise errors.ServerError(msg='插件安装失败,请稍后重试') from e
else:
await install_requirements_async(repo_name)
await redis_client.set(f'{settings.PLUGIN_REDIS_PREFIX}:changed', 'ture')
@staticmethod @staticmethod
async def uninstall(*, plugin: str): async def uninstall(*, plugin: str):
@@ -123,12 +70,11 @@ class PluginService:
""" """
plugin_dir = os.path.join(PLUGIN_DIR, plugin) plugin_dir = os.path.join(PLUGIN_DIR, plugin)
if not os.path.exists(plugin_dir): if not os.path.exists(plugin_dir):
raise errors.ForbiddenError(msg='插件不存在') raise errors.NotFoundError(msg='插件不存在')
await uninstall_requirements_async(plugin) await uninstall_requirements_async(plugin)
bacup_dir = os.path.join(PLUGIN_DIR, f'{plugin}.{timezone.now().strftime("%Y%m%d%H%M%S")}.backup') bacup_dir = os.path.join(PLUGIN_DIR, f'{plugin}.{timezone.now().strftime("%Y%m%d%H%M%S")}.backup')
shutil.move(plugin_dir, bacup_dir) shutil.move(plugin_dir, bacup_dir)
await redis_client.delete(f'{settings.PLUGIN_REDIS_PREFIX}:info:{plugin}') await redis_client.delete(f'{settings.PLUGIN_REDIS_PREFIX}:{plugin}')
await redis_client.hdel(f'{settings.PLUGIN_REDIS_PREFIX}:status', plugin)
await redis_client.set(f'{settings.PLUGIN_REDIS_PREFIX}:changed', 'ture') await redis_client.set(f'{settings.PLUGIN_REDIS_PREFIX}:changed', 'ture')
@staticmethod @staticmethod
@@ -139,9 +85,9 @@ class PluginService:
:param plugin: 插件名称 :param plugin: 插件名称
:return: :return:
""" """
plugin_info = await redis_client.get(f'{settings.PLUGIN_REDIS_PREFIX}:info:{plugin}') plugin_info = await redis_client.get(f'{settings.PLUGIN_REDIS_PREFIX}:{plugin}')
if not plugin_info: if not plugin_info:
raise errors.ForbiddenError(msg='插件不存在') raise errors.NotFoundError(msg='插件不存在')
plugin_info = json.loads(plugin_info) plugin_info = json.loads(plugin_info)
# 更新持久缓存状态 # 更新持久缓存状态
@@ -151,10 +97,7 @@ class PluginService:
else str(StatusType.disable.value) else str(StatusType.disable.value)
) )
plugin_info['plugin']['enable'] = new_status plugin_info['plugin']['enable'] = new_status
await redis_client.set( await redis_client.set(f'{settings.PLUGIN_REDIS_PREFIX}:{plugin}', json.dumps(plugin_info, ensure_ascii=False))
f'{settings.PLUGIN_REDIS_PREFIX}:info:{plugin}', json.dumps(plugin_info, ensure_ascii=False)
)
await redis_client.hset(f'{settings.PLUGIN_REDIS_PREFIX}:status', plugin, new_status)
@staticmethod @staticmethod
async def build(*, plugin: str) -> io.BytesIO: async def build(*, plugin: str) -> io.BytesIO:
@@ -166,7 +109,7 @@ class PluginService:
""" """
plugin_dir = os.path.join(PLUGIN_DIR, plugin) plugin_dir = os.path.join(PLUGIN_DIR, plugin)
if not os.path.exists(plugin_dir): if not os.path.exists(plugin_dir):
raise errors.ForbiddenError(msg='插件不存在') raise errors.NotFoundError(msg='插件不存在')
bio = io.BytesIO() bio = io.BytesIO()
with zipfile.ZipFile(bio, 'w') as zf: with zipfile.ZipFile(bio, 'w') as zf:
+9 -8
View File
@@ -10,6 +10,7 @@ from backend.app.admin.crud.crud_role import role_dao
from backend.app.admin.model import Role from backend.app.admin.model import Role
from backend.app.admin.schema.role import ( from backend.app.admin.schema.role import (
CreateRoleParam, CreateRoleParam,
DeleteRoleParam,
UpdateRoleMenuParam, UpdateRoleMenuParam,
UpdateRoleParam, UpdateRoleParam,
UpdateRoleScopeParam, UpdateRoleScopeParam,
@@ -97,7 +98,7 @@ class RoleService:
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
role = await role_dao.get_by_name(db, obj.name) role = await role_dao.get_by_name(db, obj.name)
if role: if role:
raise errors.ForbiddenError(msg='角色已存在') raise errors.ConflictError(msg='角色已存在')
await role_dao.create(db, obj) await role_dao.create(db, obj)
@staticmethod @staticmethod
@@ -115,7 +116,7 @@ class RoleService:
raise errors.NotFoundError(msg='角色不存在') raise errors.NotFoundError(msg='角色不存在')
if role.name != obj.name: if role.name != obj.name:
if await role_dao.get_by_name(db, obj.name): if await role_dao.get_by_name(db, obj.name):
raise errors.ForbiddenError(msg='角色已存在') raise errors.ConflictError(msg='角色已存在')
count = await role_dao.update(db, pk, obj) count = await role_dao.update(db, pk, obj)
for user in await role.awaitable_attrs.users: for user in await role.awaitable_attrs.users:
await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') await redis_client.delete_prefix(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
@@ -166,17 +167,17 @@ class RoleService:
return count return count
@staticmethod @staticmethod
async def delete(*, pk: list[int]) -> int: async def delete(*, obj: DeleteRoleParam) -> int:
""" """
删除角色 批量删除角色
:param pk: 角色 ID 列表 :param obj: 角色 ID 列表
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
count = await role_dao.delete(db, pk) count = await role_dao.delete(db, obj.pks)
for _pk in pk: for pk in obj.pks:
role = await role_dao.get(db, _pk) role = await role_dao.get(db, pk)
if role: if role:
for user in await role.awaitable_attrs.users: for user in await role.awaitable_attrs.users:
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
+186 -140
View File
@@ -16,8 +16,10 @@ from backend.app.admin.schema.user import (
ResetPasswordParam, ResetPasswordParam,
UpdateUserParam, UpdateUserParam,
) )
from backend.common.enums import UserPermissionType
from backend.common.exception import errors from backend.common.exception import errors
from backend.common.security.jwt import get_hash_password, get_token, jwt_decode, password_verify, superuser_verify from backend.common.response.response_code import CustomErrorCode
from backend.common.security.jwt import get_token, jwt_decode, password_verify, superuser_verify
from backend.core.conf import settings from backend.core.conf import settings
from backend.database.db import async_db_session from backend.database.db import async_db_session
from backend.database.redis import redis_client from backend.database.redis import redis_client
@@ -27,86 +29,30 @@ class UserService:
"""用户服务类""" """用户服务类"""
@staticmethod @staticmethod
async def add(*, request: Request, obj: AddUserParam) -> None: async def get_userinfo(*, pk: int | None = None, username: str | None = None) -> User:
"""
添加新用户
:param request: FastAPI 请求对象
:param obj: 用户添加参数
:return:
"""
async with async_db_session.begin() as db:
superuser_verify(request)
username = await user_dao.get_by_username(db, obj.username)
if username:
raise errors.ForbiddenError(msg='用户已注册')
obj.nickname = obj.nickname if obj.nickname else f'#{random.randrange(88888, 99999)}'
nickname = await user_dao.get_by_nickname(db, obj.nickname)
if nickname:
raise errors.ForbiddenError(msg='昵称已注册')
if not obj.password:
raise errors.ForbiddenError(msg='密码为空')
dept = await dept_dao.get(db, obj.dept_id)
if not dept:
raise errors.NotFoundError(msg='部门不存在')
for role_id in obj.roles:
role = await role_dao.get(db, role_id)
if not role:
raise errors.NotFoundError(msg='角色不存在')
await user_dao.add(db, obj)
@staticmethod
async def pwd_reset(*, username: str, obj: ResetPasswordParam) -> int:
"""
重置用户密码
:param username: 用户名
:param obj: 密码重置参数
:return:
"""
async with async_db_session.begin() as db:
user = await user_dao.get_by_username(db, username)
if not user:
raise errors.NotFoundError(msg='用户不存在')
if not password_verify(obj.old_password, user.password):
raise errors.ForbiddenError(msg='原密码错误')
if obj.new_password != obj.confirm_password:
raise errors.ForbiddenError(msg='密码输入不一致')
new_pwd = get_hash_password(obj.new_password, user.salt)
count = await user_dao.reset_password(db, user.id, new_pwd)
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 get_userinfo(*, username: str) -> User:
""" """
获取用户信息 获取用户信息
:param pk: 用户 ID
:param username: 用户名 :param username: 用户名
:return: :return:
""" """
async with async_db_session() as db: async with async_db_session() as db:
user = await user_dao.get_with_relation(db, username=username) user = await user_dao.get_with_relation(db, user_id=pk, username=username)
if not user: if not user:
raise errors.NotFoundError(msg='用户不存在') raise errors.NotFoundError(msg='用户不存在')
return user return user
@staticmethod @staticmethod
async def get_roles(*, username: str) -> Sequence[Role]: async def get_roles(*, pk: int) -> Sequence[Role]:
""" """
获取用户所有角色 获取用户所有角色
:param username: 用户 :param pk: 用户 ID
:return: :return:
""" """
async with async_db_session() as db: async with async_db_session() as db:
user = await user_dao.get_with_relation(db, username=username) user = await user_dao.get_with_relation(db, user_id=pk)
if not user: if not user:
raise errors.NotFoundError(msg='用户不存在') raise errors.NotFoundError(msg='用户不存在')
return user.roles return user.roles
@@ -125,65 +71,122 @@ class UserService:
return await user_dao.get_list(dept=dept, username=username, phone=phone, status=status) return await user_dao.get_list(dept=dept, username=username, phone=phone, status=status)
@staticmethod @staticmethod
async def update(*, request: Request, username: str, obj: UpdateUserParam) -> int: async def create(*, request: Request, obj: AddUserParam) -> None:
"""
创建用户
:param request: FastAPI 请求对象
:param obj: 用户添加参数
:return:
"""
async with async_db_session.begin() as db:
superuser_verify(request)
if await user_dao.get_by_username(db, obj.username):
raise errors.ConflictError(msg='用户名已注册')
obj.nickname = obj.nickname if obj.nickname else f'#{random.randrange(88888, 99999)}'
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(*, request: Request, pk: int, obj: UpdateUserParam) -> int:
""" """
更新用户信息 更新用户信息
:param request: FastAPI 请求对象 :param request: FastAPI 请求对象
:param username: 用户 :param pk: 用户 ID
:param obj: 用户更新参数 :param obj: 用户更新参数
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
if request.user.username != username: superuser_verify(request)
raise errors.ForbiddenError(msg='你只能修改自己的信息') user = await user_dao.get_with_relation(db, user_id=pk)
user = await user_dao.get_with_relation(db, username=username)
if not user: if not user:
raise errors.NotFoundError(msg='用户不存在') raise errors.NotFoundError(msg='用户不存在')
if user.username != obj.username: if obj.username != user.username:
_username = await user_dao.get_by_username(db, obj.username) if await user_dao.get_by_username(db, obj.username):
if _username: raise errors.ConflictError(msg='用户名已注册')
raise errors.ForbiddenError(msg='用户名已注册')
if user.nickname != obj.nickname:
nickname = await user_dao.get_by_nickname(db, obj.nickname)
if nickname:
raise errors.ForbiddenError(msg='昵称已注册')
for role_id in obj.roles: for role_id in obj.roles:
role = await role_dao.get(db, role_id) if not await role_dao.get(db, role_id):
if not role:
raise errors.NotFoundError(msg='角色不存在') raise errors.NotFoundError(msg='角色不存在')
count = await user_dao.update(db, user, obj) count = await user_dao.update(db, user, obj)
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
return count return count
@staticmethod @staticmethod
async def update_permission(*, request: Request, pk: int) -> int: async def update_permission(*, request: Request, pk: int, type: UserPermissionType) -> int:
""" """
更新用户权限 更新用户权限
:param request: FastAPI 请求对象 :param request: FastAPI 请求对象
:param pk: 用户 ID :param pk: 用户 ID
:param type: 权限类型
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
superuser_verify(request) superuser_verify(request)
user = await user_dao.get(db, pk) match type:
if not user: case UserPermissionType.superuser:
raise errors.NotFoundError(msg='用户不存在') user = await user_dao.get(db, pk)
if pk == request.user.id: if not user:
raise errors.ForbiddenError(msg='非法操作') raise errors.NotFoundError(msg='用户不存在')
super_status = await user_dao.get_super(db, pk) if pk == request.user.id:
count = await user_dao.set_super(db, pk, not super_status) raise errors.ForbiddenError(msg='禁止修改自身权限')
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') count = await user_dao.set_super(db, pk, not user.status)
return count case UserPermissionType.staff:
user = await user_dao.get(db, pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
if pk == request.user.id:
raise errors.ForbiddenError(msg='禁止修改自身权限')
count = await user_dao.set_staff(db, pk, not user.is_staff)
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, 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 @staticmethod
async def update_staff(*, request: Request, pk: int) -> int: async def reset_password(*, request: Request, pk: int, password: str) -> int:
""" """
更新用户职员状态 重置用户密码
:param request: FastAPI 请求对象 :param request: FastAPI 请求对象
:param pk: 用户 ID :param pk: 用户 ID
:param password: 新密码
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
@@ -191,76 +194,119 @@ class UserService:
user = await user_dao.get(db, pk) user = await user_dao.get(db, pk)
if not user: if not user:
raise errors.NotFoundError(msg='用户不存在') raise errors.NotFoundError(msg='用户不存在')
if pk == request.user.id: count = await user_dao.reset_password(db, user.id, password)
raise errors.ForbiddenError(msg='非法操作') key_prefix = [
staff_status = await user_dao.get_staff(db, pk) f'{settings.TOKEN_REDIS_PREFIX}:{user.id}',
count = await user_dao.set_staff(db, pk, not staff_status) f'{settings.TOKEN_REFRESH_REDIS_PREFIX}:{user.id}',
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}',
]
for prefix in key_prefix:
await redis_client.delete(prefix)
return count return count
@staticmethod @staticmethod
async def update_status(*, request: Request, pk: int) -> int: async def update_nickname(*, request: Request, nickname: str) -> int:
""" """
更新用户状态 更新当前用户昵称
:param request: FastAPI 请求对象 :param request: FastAPI 请求对象
:param pk: 用户 ID :param nickname: 用户昵称
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
superuser_verify(request)
user = await user_dao.get(db, pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
if pk == request.user.id:
raise errors.ForbiddenError(msg='非法操作')
status = await user_dao.get_status(db, pk)
count = await user_dao.set_status(db, pk, 0 if status == 1 else 1)
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
return count
@staticmethod
async def update_multi_login(*, request: Request, pk: int) -> int:
"""
更新用户多端登录状态
:param request: FastAPI 请求对象
:param pk: 用户 ID
:return:
"""
async with async_db_session.begin() as db:
superuser_verify(request)
user = await user_dao.get(db, pk)
if not user:
raise errors.NotFoundError(msg='用户不存在')
multi_login = await user_dao.get_multi_login(db, pk) if pk != user.id else request.user.is_multi_login
new_multi_login = not multi_login
count = await user_dao.set_multi_login(db, pk, new_multi_login)
await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
token = get_token(request) token = get_token(request)
token_payload = jwt_decode(token) token_payload = jwt_decode(token)
if pk == user.id: user = await user_dao.get(db, token_payload.id)
# 系统管理员修改自身时,除当前 token 外,其他 token 失效 if not user:
if not new_multi_login: raise errors.NotFoundError(msg='用户不存在')
key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{user.id}' count = await user_dao.update_nickname(db, token_payload.id, nickname)
await redis_client.delete_prefix(key_prefix, exclude=f'{key_prefix}:{token_payload.session_uuid}') await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}')
else:
# 系统管理员修改他人时,他人 token 全部失效
if not new_multi_login:
key_prefix = f'{settings.TOKEN_REDIS_PREFIX}:{user.id}'
await redis_client.delete_prefix(key_prefix)
return count return count
@staticmethod @staticmethod
async def delete(*, username: str) -> int: async def update_avatar(*, request: Request, avatar: str) -> int:
"""
更新当前用户头像
: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
@staticmethod
async def update_email(*, request: Request, captcha: str, email: str) -> int:
"""
更新当前用户邮箱
: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}:{request.state.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}:{request.state.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:
"""
更新当前用户密码
: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
@staticmethod
async def delete(*, pk: int) -> int:
""" """
删除用户 删除用户
:param username: 用户 :param pk: 用户 ID
:return: :return:
""" """
async with async_db_session.begin() as db: async with async_db_session.begin() as db:
user = await user_dao.get_by_username(db, username) user = await user_dao.get(db, pk)
if not user: if not user:
raise errors.NotFoundError(msg='用户不存在') raise errors.NotFoundError(msg='用户不存在')
count = await user_dao.delete(db, user.id) count = await user_dao.delete(db, user.id)
+6 -6
View File
@@ -3,21 +3,21 @@
当前任务使用 Celery 当前任务使用 Celery
实现,实施方案请查看 [#225](https://github.com/fastapi-practices/fastapi_best_architecture/discussions/225) 实现,实施方案请查看 [#225](https://github.com/fastapi-practices/fastapi_best_architecture/discussions/225)
## 添加任务 ## 定时任务
> [!IMPORTANT] `backend/app/task/tasks/beat.py` 文件内编写相关定时任务
> 由于 Celery 任务扫描规则,使其对任务的目录结构要求及其严格,务必在 celery_task 目录下添加任务
### 简单任务 ### 简单任务
可以直接`tasks.py` 文件内编写相关任务代码 `backend/app/task/tasks/tasks.py` 文件内编写相关任务代码
### 层级任务 ### 层级任务
如果你想对任务进行目录层级划分,使任务结构更加清晰,你可以新建任意目录,但必须注意的是 如果你想对任务进行目录层级划分,使任务结构更加清晰,你可以新建任意目录,但必须注意的是
1. 新建目录后,务必更新任务配置 `CELERY_TASKS_PACKAGES`,将新建目录添加到此列表 1. `backend/app/task/tasks` 目录下新建 python 包目录
2. 新建目录,务必添加 `tasks.py` 文件,并在此文件中编写相关任务代码 2. 新建目录,务必更新 `conf.py` 配置中的 `CELERY_TASKS_PACKAGES`,将新建目录模块路径添加到此列表
3. 在新建目录下,务必添加 `tasks.py` 文件,并在此文件中编写相关任务代码
## 消息代理 ## 消息代理
+4 -2
View File
@@ -2,7 +2,9 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
import sys import sys
from pathlib import Path from backend.core.path_conf import BASE_PATH
from .actions import * # noqa: F403
# 导入项目根目录 # 导入项目根目录
sys.path.append(str(Path(__file__).resolve().parent.parent.parent.parent)) sys.path.append(str(BASE_PATH.parent))
+13
View File
@@ -0,0 +1,13 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from starlette.concurrency import run_in_threadpool
from backend.app.task.celery import celery_app
from backend.common.socketio.server import sio
@sio.event
async def task_worker_status(sid, data):
"""任务 Worker 状态事件"""
worker = await run_in_threadpool(celery_app.control.ping)
await sio.emit('task_worker_status', worker, sid)
+7 -3
View File
@@ -2,9 +2,13 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from fastapi import APIRouter from fastapi import APIRouter
from backend.app.task.api.v1.task import router as task_router from backend.app.task.api.v1.control import router as task_control_router
from backend.app.task.api.v1.result import router as task_result_router
from backend.app.task.api.v1.scheduler import router as task_scheduler_router
from backend.core.conf import settings from backend.core.conf import settings
v1 = APIRouter(prefix=settings.FASTAPI_API_V1_PATH) v1 = APIRouter(prefix=f'{settings.FASTAPI_API_V1_PATH}/tasks', tags=['任务'])
v1.include_router(task_router, prefix='/tasks', tags=['任务']) v1.include_router(task_control_router)
v1.include_router(task_result_router, prefix='/results')
v1.include_router(task_scheduler_router, prefix='/schedulers')
+51
View File
@@ -0,0 +1,51 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path
from starlette.concurrency import run_in_threadpool
from backend.app.task import celery_app
from backend.app.task.schema.control import TaskRegisteredDetail
from backend.common.exception import errors
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
router = APIRouter()
@router.get('/registered', summary='获取已注册的任务', dependencies=[DependsJwtAuth])
async def get_task_registered() -> ResponseSchemaModel[list[TaskRegisteredDetail]]:
inspector = celery_app.control.inspect(timeout=0.5)
registered = await run_in_threadpool(inspector.registered)
if not registered:
raise errors.ServerError(msg='Celery Worker 暂不可用,请稍后重试')
task_registered = []
celery_app_tasks = celery_app.tasks
for _, tasks in registered.items():
for task in tasks:
task_ins = celery_app_tasks.get(task)
if task_ins:
task_doc = task_ins.__doc__
task_registered.append({'name': task_doc or task_ins, 'task': task_ins})
else:
task_registered.append({'name': task, 'task': task})
return response_base.success(data=task_registered)
@router.delete(
'/{task_id}/cancel',
summary='撤销任务',
dependencies=[
Depends(RequestPermission('sys:task:revoke')),
DependsRBAC,
],
)
async def revoke_task(task_id: Annotated[str, Path(description='任务 UUID')]) -> ResponseModel:
workers = await run_in_threadpool(celery_app.control.ping, timeout=0.5)
if not workers:
raise errors.ServerError(msg='Celery Worker 暂不可用,请稍后重试')
celery_app.control.revoke(task_id)
return response_base.success()
+57
View File
@@ -0,0 +1,57 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
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.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
router = APIRouter()
@router.get('/{pk}', summary='获取任务结果详情', dependencies=[DependsJwtAuth])
async def get_task_result(
pk: Annotated[int, Path(description='任务结果 ID')],
) -> ResponseSchemaModel[GetTaskResultDetail]:
result = await task_result_service.get(pk=pk)
return response_base.success(data=result)
@router.get(
'',
summary='分页获取所有任务结果',
dependencies=[
DependsJwtAuth,
DependsPagination,
],
)
async def get_task_results_paged(
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)
return response_base.success(data=page_data)
@router.delete(
'',
summary='批量删除任务结果',
dependencies=[
Depends(RequestPermission('sys:task:del')),
DependsRBAC,
],
)
async def delete_task_result(obj: DeleteTaskResultParam) -> ResponseModel:
count = await task_result_service.delete(obj=obj)
if count > 0:
return response_base.success()
return response_base.fail()
+121
View File
@@ -0,0 +1,121 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path, Query
from backend.app.task.schema.scheduler import CreateTaskSchedulerParam, GetTaskSchedulerDetail, UpdateTaskSchedulerParam
from backend.app.task.service.scheduler_service import task_scheduler_service
from backend.common.pagination import DependsPagination, PageData, paging_data
from backend.common.response.response_schema import ResponseModel, ResponseSchemaModel, response_base
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
router = APIRouter()
@router.get('/all', summary='获取所有任务调度', dependencies=[DependsJwtAuth])
async def get_all_task_schedulers() -> ResponseSchemaModel[list[GetTaskSchedulerDetail]]:
schedulers = await task_scheduler_service.get_all()
return response_base.success(data=schedulers)
@router.get('/{pk}', summary='获取任务调度详情', dependencies=[DependsJwtAuth])
async def get_task_scheduler(
pk: Annotated[int, Path(description='任务调度 ID')],
) -> ResponseSchemaModel[GetTaskSchedulerDetail]:
task_scheduler = await task_scheduler_service.get(pk=pk)
return response_base.success(data=task_scheduler)
@router.get(
'',
summary='分页获取所有任务调度',
dependencies=[
DependsJwtAuth,
DependsPagination,
],
)
async def get_task_scheduler_paged(
db: CurrentSession,
name: Annotated[int, 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)
return response_base.success(data=page_data)
@router.post(
'',
summary='创建任务调度',
dependencies=[
Depends(RequestPermission('sys:task:add')),
DependsRBAC,
],
)
async def create_task_scheduler(obj: CreateTaskSchedulerParam) -> ResponseModel:
await task_scheduler_service.create(obj=obj)
return response_base.success()
@router.put(
'/{pk}',
summary='更新任务调度',
dependencies=[
Depends(RequestPermission('sys:task:edit')),
DependsRBAC,
],
)
async def update_task_scheduler(
pk: Annotated[int, Path(description='任务调度 ID')], obj: UpdateTaskSchedulerParam
) -> ResponseModel:
count = await task_scheduler_service.update(pk=pk, obj=obj)
if count > 0:
return response_base.success()
return response_base.fail()
@router.put(
'/{pk}/status',
summary='更新任务调度状态',
dependencies=[
Depends(RequestPermission('sys:task:edit')),
DependsRBAC,
],
)
async def update_task_scheduler_status(pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel:
count = await task_scheduler_service.update_status(pk=pk)
if count > 0:
return response_base.success()
return response_base.fail()
@router.delete(
'/{pk}',
summary='删除任务调度',
dependencies=[
Depends(RequestPermission('sys:task:del')),
DependsRBAC,
],
)
async def delete_task_scheduler(pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel:
count = await task_scheduler_service.delete(pk=pk)
if count > 0:
return response_base.success()
return response_base.fail()
@router.post(
'/{pk}/executions',
summary='手动执行任务',
dependencies=[
Depends(RequestPermission('sys:task:exec')),
DependsRBAC,
],
)
async def execute_task(pk: Annotated[int, Path(description='任务调度 ID')]) -> ResponseModel:
await task_scheduler_service.execute(pk=pk)
return response_base.success()
-58
View File
@@ -1,58 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Annotated
from fastapi import APIRouter, Depends, Path
from backend.app.task.schema.task import RunParam, TaskResult
from backend.app.task.service.task_service import task_service
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
router = APIRouter()
@router.get('', summary='获取可执行任务', dependencies=[DependsJwtAuth])
async def get_all_tasks() -> ResponseSchemaModel[list[str]]:
tasks = await task_service.get_list()
return response_base.success(data=tasks)
@router.get(
'/{tid}',
summary='获取任务详情',
deprecated=True,
description='此接口被视为作废,建议使用 flower 查看任务详情',
dependencies=[DependsJwtAuth],
)
async def get_task_detail(tid: Annotated[str, Path(description='任务 UUID')]) -> ResponseSchemaModel[TaskResult]:
status = task_service.get_detail(tid=tid)
return response_base.success(data=status)
@router.post(
'/{tid}',
summary='撤销任务',
dependencies=[
Depends(RequestPermission('sys:task:revoke')),
DependsRBAC,
],
)
async def revoke_task(tid: Annotated[str, Path(description='任务 UUID')]) -> ResponseModel:
task_service.revoke(tid=tid)
return response_base.success()
@router.post(
'',
summary='执行任务',
dependencies=[
Depends(RequestPermission('sys:task:run')),
DependsRBAC,
],
)
async def run_task(obj: RunParam) -> ResponseSchemaModel[str]:
task = task_service.run(obj=obj)
return response_base.success(data=task)
+32 -40
View File
@@ -1,44 +1,24 @@
#!/usr/bin/env python3 #!/usr/bin/env python3
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Any import os
import celery import celery
import celery_aio_pool import celery_aio_pool
from backend.app.task.model.result import OVERWRITE_CELERY_RESULT_GROUP_TABLE_NAME, OVERWRITE_CELERY_RESULT_TABLE_NAME
from backend.app.task.tasks.beat import LOCAL_BEAT_SCHEDULE
from backend.core.conf import settings from backend.core.conf import settings
from backend.core.path_conf import BASE_PATH
__all__ = ['celery_app']
def get_broker_url() -> str: def find_task_packages():
"""获取消息代理 URL""" packages = []
if settings.CELERY_BROKER == 'redis': task_dir = os.path.join(BASE_PATH, 'app', 'task', 'tasks')
return ( for root, dirs, files in os.walk(task_dir):
f'redis://:{settings.REDIS_PASSWORD}@{settings.REDIS_HOST}:' if 'tasks.py' in files:
f'{settings.REDIS_PORT}/{settings.CELERY_BROKER_REDIS_DATABASE}' package = root.replace(str(BASE_PATH.parent) + os.path.sep, '').replace(os.path.sep, '.')
) packages.append(package)
return ( return packages
f'amqp://{settings.CELERY_RABBITMQ_USERNAME}:{settings.CELERY_RABBITMQ_PASSWORD}@'
f'{settings.CELERY_RABBITMQ_HOST}:{settings.CELERY_RABBITMQ_PORT}'
)
def get_result_backend() -> str:
"""获取结果后端 URL"""
return (
f'redis://:{settings.REDIS_PASSWORD}@{settings.REDIS_HOST}:'
f'{settings.REDIS_PORT}/{settings.CELERY_BACKEND_REDIS_DATABASE}'
)
def get_result_backend_transport_options() -> dict[str, Any]:
"""获取结果后端传输选项"""
return {
'global_keyprefix': settings.CELERY_BACKEND_REDIS_PREFIX,
'retry_policy': {
'timeout': settings.CELERY_BACKEND_REDIS_TIMEOUT,
},
}
def init_celery() -> celery.Celery: def init_celery() -> celery.Celery:
@@ -52,19 +32,31 @@ def init_celery() -> celery.Celery:
app = celery.Celery( app = celery.Celery(
'fba_celery', 'fba_celery',
broker=f'redis://:{settings.REDIS_PASSWORD}@{settings.REDIS_HOST}:{settings.REDIS_PORT}/{settings.CELERY_BROKER_REDIS_DATABASE}'
if settings.CELERY_BROKER == 'redis'
else f'amqp://{settings.CELERY_RABBITMQ_USERNAME}:{settings.CELERY_RABBITMQ_PASSWORD}@{settings.CELERY_RABBITMQ_HOST}:{settings.CELERY_RABBITMQ_PORT}',
broker_connection_retry_on_startup=True,
backend=f'db+{settings.DATABASE_TYPE}+{"pymysql" if settings.DATABASE_TYPE == "mysql" else "psycopg"}'
f'://{settings.DATABASE_USER}:{settings.DATABASE_PASSWORD}@{settings.DATABASE_HOST}:{settings.DATABASE_PORT}/{settings.DATABASE_SCHEMA}',
database_engine_options={'echo': settings.DATABASE_ECHO},
database_table_names={
'task': OVERWRITE_CELERY_RESULT_TABLE_NAME,
'group': OVERWRITE_CELERY_RESULT_GROUP_TABLE_NAME,
},
result_extended=True,
# result_expires=0, # 清理任务结果,默认每天凌晨 4 点,0 或 None 表示不清理
# beat_sync_every=1, # 保存任务状态周期,默认 3 * 60 秒
beat_schedule=LOCAL_BEAT_SCHEDULE,
beat_scheduler='backend.app.task.utils.schedulers:DatabaseScheduler',
task_cls='backend.app.task.tasks.base:TaskBase',
task_track_started=True,
enable_utc=False, enable_utc=False,
timezone=settings.DATETIME_TIMEZONE, timezone=settings.DATETIME_TIMEZONE,
beat_schedule=settings.CELERY_SCHEDULE,
broker_url=get_broker_url(),
broker_connection_retry_on_startup=True,
result_backend=get_result_backend(),
result_backend_transport_options=get_result_backend_transport_options(),
task_cls='app.task.celery_task.base:TaskBase',
task_track_started=True,
) )
# 自动发现任务 # 自动发现任务
app.autodiscover_tasks(settings.CELERY_TASK_PACKAGES) packages = find_task_packages()
app.autodiscover_tasks(packages)
return app return app
@@ -1,19 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.app.admin.service.login_log_service import login_log_service
from backend.app.admin.service.opera_log_service import opera_log_service
from backend.app.task.celery import celery_app
@celery_app.task(name='delete_db_opera_log')
async def delete_db_opera_log() -> int:
"""自动删除数据库操作日志"""
result = await opera_log_service.delete_all()
return result
@celery_app.task(name='delete_db_login_log')
async def delete_db_login_log() -> int:
"""自动删除数据库登录日志"""
result = await login_log_service.delete_all()
return result
-12
View File
@@ -1,12 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from anyio import sleep
from backend.app.task.celery import celery_app
@celery_app.task(name='task_demo_async')
async def task_demo_async() -> str:
"""异步示例任务,模拟耗时操作"""
await sleep(20)
return 'test async'
+51
View File
@@ -0,0 +1,51 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from sqlalchemy import Select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from backend.app.task.model.result import TaskResult
class CRUDTaskResult(CRUDPlus[TaskResult]):
"""任务结果数据库操作类"""
async def get(self, db: AsyncSession, pk: int) -> TaskResult | None:
"""
获取任务结果详情
:param db: 数据库会话
:param pk: 任务 ID
:return:
"""
return await self.select_model(db, pk)
async def get_list(self, name: str | None, task_id: str | None) -> Select:
"""
获取任务结果列表
:param name: 任务名称
:param task_id: 任务 ID
:return:
"""
filters = {}
if name is not None:
filters['name__like'] = f'%{name}%'
if task_id is not None:
filters['task_id'] = task_id
return await self.select_order('id', 'desc', **filters)
async def delete(self, db: AsyncSession, pks: list[int]) -> int:
"""
批量删除任务结果
:param db: 数据库会话
:param pks: 任务结果 ID 列表
:return:
"""
return await self.delete_model_by_column(db, allow_multiple=True, id__in=pks)
task_result_dao: CRUDTaskResult = CRUDTaskResult(TaskResult)
+117
View File
@@ -0,0 +1,117 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Sequence
from sqlalchemy import Select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy_crud_plus import CRUDPlus
from backend.app.task.model import TaskScheduler
from backend.app.task.schema.scheduler import CreateTaskSchedulerParam, UpdateTaskSchedulerParam
class CRUDTaskScheduler(CRUDPlus[TaskScheduler]):
"""任务调度数据库操作类"""
@staticmethod
async def get(db: AsyncSession, pk: int) -> TaskScheduler | None:
"""
获取任务调度
:param db: 数据库会话
:param pk: 任务调度 ID
:return:
"""
return await task_scheduler_dao.select_model(db, pk)
async def get_all(self, db: AsyncSession) -> Sequence[TaskScheduler]:
"""
获取所有任务调度
:param db: 数据库会话
:return:
"""
return await self.select_models(db)
async def get_list(self, name: str | None, type: int | None) -> Select:
"""
获取任务调度列表
:param name: 任务调度名称
:param type: 任务调度类型
:return:
"""
filters = {}
if name is not None:
filters['name__like'] = f'%{name}%'
if type is not None:
filters['type'] = type
return await self.select_order('id', **filters)
async def get_by_name(self, db: AsyncSession, name: str) -> TaskScheduler | None:
"""
通过名称获取任务调度
:param db: 数据库会话
:param name: 任务调度名称
:return:
"""
return await self.select_model_by_column(db, name=name)
async def create(self, db: AsyncSession, obj: CreateTaskSchedulerParam) -> None:
"""
创建任务调度
:param db: 数据库会话
:param obj: 创建任务调度参数
:return:
"""
await self.create_model(db, obj, flush=True)
TaskScheduler.no_changes = False
async def update(self, db: AsyncSession, pk: int, obj: UpdateTaskSchedulerParam) -> int:
"""
更新任务调度
:param db: 数据库会话
:param pk: 任务调度 ID
:param obj: 更新任务调度参数
:return:
"""
task_scheduler = await self.get(db, pk)
for key, value in obj.model_dump(exclude_unset=True).items():
setattr(task_scheduler, key, value)
TaskScheduler.no_changes = False
return 1
async def set_status(self, db: AsyncSession, pk: int, status: bool) -> int:
"""
设置任务调度状态
:param db: 数据库会话
:param pk: 任务调度 ID
:param status: 状态
:return:
"""
task_scheduler = await self.get(db, pk)
setattr(task_scheduler, 'enabled', status)
TaskScheduler.no_changes = False
return 1
async def delete(self, db: AsyncSession, pk: int) -> int:
"""
删除任务调度
:param db: 数据库会话
:param pk: 任务调度 ID
:return:
"""
task_scheduler = await self.get(db, pk)
await db.delete(task_scheduler)
TaskScheduler.no_changes = False
return 1
task_scheduler_dao: CRUDTaskScheduler = CRUDTaskScheduler(TaskScheduler)
+20
View File
@@ -0,0 +1,20 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.common.enums import IntEnum, StrEnum
class TaskSchedulerType(IntEnum):
"""任务调度类型"""
INTERVAL = 0
CRONTAB = 1
class PeriodType(StrEnum):
"""周期类型"""
DAYS = 'days'
HOURS = 'hours'
MINUTES = 'minutes'
SECONDS = 'seconds'
MICROSECONDS = 'microseconds'
+3
View File
@@ -0,0 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.app.task.model.scheduler import TaskScheduler
+9
View File
@@ -0,0 +1,9 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from celery.backends.database.models import TaskExtended as TaskResult
OVERWRITE_CELERY_RESULT_TABLE_NAME = 'task_result'
OVERWRITE_CELERY_RESULT_GROUP_TABLE_NAME = 'task_group_result'
# 重写表名配置
TaskResult.configure(name=OVERWRITE_CELERY_RESULT_TABLE_NAME)
+86
View File
@@ -0,0 +1,86 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import asyncio
from datetime import datetime
from sqlalchemy import (
JSON,
Boolean,
DateTime,
String,
event,
)
from sqlalchemy.dialects.mysql import LONGTEXT
from sqlalchemy.dialects.postgresql import INTEGER, TEXT
from sqlalchemy.orm import Mapped, mapped_column
from backend.common.exception import errors
from backend.common.model import Base, id_key
from backend.core.conf import settings
from backend.database.redis import redis_client
from backend.utils.timezone import timezone
class TaskScheduler(Base):
"""任务调度表"""
__tablename__ = 'task_scheduler'
id: Mapped[id_key] = mapped_column(init=False)
name: Mapped[str] = mapped_column(String(50), unique=True, comment='任务名称')
task: Mapped[str] = mapped_column(String(255), comment='要运行的 Celery 任务')
args: Mapped[str | None] = mapped_column(JSON(), comment='任务可接收的位置参数')
kwargs: Mapped[str | None] = mapped_column(JSON(), comment='任务可接收的关键字参数')
queue: Mapped[str | None] = mapped_column(String(255), comment='CELERY_TASK_QUEUES 中定义的队列')
exchange: Mapped[str | None] = mapped_column(String(255), comment='低级别 AMQP 路由的交换机')
routing_key: Mapped[str | None] = mapped_column(String(255), comment='低级别 AMQP 路由的路由密钥')
start_time: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), comment='任务开始触发的时间')
expire_time: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), comment='任务不再触发的截止时间')
expire_seconds: Mapped[int | None] = mapped_column(comment='任务不再触发的秒数时间差')
type: Mapped[int] = mapped_column(comment='调度类型(0间隔 1定时)')
interval_every: Mapped[int | None] = mapped_column(comment='任务再次运行前的间隔周期数')
interval_period: Mapped[str | None] = mapped_column(String(255), comment='任务运行之间的周期类型')
crontab: Mapped[str | None] = mapped_column(String(50), default='* * * * *', comment='任务运行的 Crontab 计划')
one_off: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=False, comment='是否仅运行一次'
)
enabled: Mapped[bool] = mapped_column(
Boolean().with_variant(INTEGER, 'postgresql'), default=True, comment='是否启用任务'
)
total_run_count: Mapped[int] = mapped_column(default=0, comment='任务触发的总次数')
last_run_time: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), default=None, comment='任务最后触发的时间'
)
remark: Mapped[str | None] = mapped_column(
LONGTEXT().with_variant(TEXT, 'postgresql'), default=None, comment='备注'
)
no_changes: bool = False
@staticmethod
def before_insert_or_update(mapper, connection, target):
if target.expire_seconds is not None and target.expire_time:
raise errors.ConflictError(msg='expires 和 expire_seconds 只能设置一个')
@classmethod
def changed(cls, mapper, connection, target):
if not target.no_changes:
cls.update_changed(mapper, connection, target)
@classmethod
async def update_changed_async(cls):
now = timezone.now()
await redis_client.set(f'{settings.CELERY_REDIS_PREFIX}:last_update', timezone.to_str(now))
@classmethod
def update_changed(cls, mapper, connection, target):
asyncio.create_task(cls.update_changed_async())
# 事件监听器
event.listen(TaskScheduler, 'before_insert', TaskScheduler.before_insert_or_update)
event.listen(TaskScheduler, 'before_update', TaskScheduler.before_insert_or_update)
event.listen(TaskScheduler, 'after_insert', TaskScheduler.update_changed)
event.listen(TaskScheduler, 'after_delete', TaskScheduler.update_changed)
event.listen(TaskScheduler, 'after_update', TaskScheduler.changed)
+8
View File
@@ -0,0 +1,8 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from backend.common.schema import SchemaBase
class TaskRegisteredDetail(SchemaBase):
name: str
task: str
+43
View File
@@ -0,0 +1,43 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from typing import Any
from pydantic import ConfigDict, Field, field_serializer
from backend.app.task import celery_app
from backend.common.schema import SchemaBase
class TaskResultSchemaBase(SchemaBase):
"""任务结果基础模型"""
task_id: str = Field(description='任务 ID')
status: str = Field(description='执行状态')
result: Any | None = Field(description='执行结果')
date_done: datetime | None = Field(description='结束时间')
traceback: str | None = Field(description='错误回溯')
name: str | None = Field(description='任务名称')
args: bytes | None = Field(description='任务位置参数')
kwargs: bytes | None = Field(description='任务关键字参数')
worker: str | None = Field(description='运行 Worker')
retries: int | None = Field(description='重试次数')
queue: str | None = Field(description='运行队列')
class DeleteTaskResultParam(SchemaBase):
"""删除任务结果参数"""
pks: list[int] = Field(description='任务结果 ID 列表')
class GetTaskResultDetail(TaskResultSchemaBase):
"""任务结果详情"""
model_config = ConfigDict(from_attributes=True)
id: int = Field(description='任务结果 ID')
@field_serializer('args', 'kwargs', when_used='unless-none')
def serialize_params(self, value: bytes | None, _info) -> Any:
return celery_app.backend.decode(value)
+51
View File
@@ -0,0 +1,51 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from pydantic import ConfigDict, Field
from pydantic.types import JsonValue
from backend.app.task.enums import PeriodType, TaskSchedulerType
from backend.common.schema import SchemaBase
class TaskSchedulerSchemeBase(SchemaBase):
"""任务调度参数"""
name: str = Field(description='任务名称')
task: str = Field(description='要运行的 Celery 任务')
args: JsonValue | None = Field(default=None, description='任务可接收的位置参数')
kwargs: JsonValue | None = Field(default=None, description='任务可接收的关键字参数')
queue: str | None = Field(default=None, description='CELERY_TASK_QUEUES 中定义的队列')
exchange: str | None = Field(default=None, description='低级别 AMQP 路由的交换机')
routing_key: str | None = Field(default=None, description='低级别 AMQP 路由的路由密钥')
start_time: datetime | None = Field(default=None, description='任务开始触发的时间')
expire_time: datetime | None = Field(default=None, description='任务不再触发的截止时间')
expire_seconds: int | None = Field(default=None, description='任务不再触发的秒数时间差')
type: TaskSchedulerType = Field(description='任务调度类型(0间隔 1定时)')
interval_every: int | None = Field(default=None, description='任务再次运行前的间隔周期数')
interval_period: PeriodType | None = Field(default=None, description='任务运行之间的周期类型')
crontab: str = Field(default='* * * * *', description='运行的 Crontab 表达式')
one_off: bool = Field(default=False, description='是否仅运行一次')
remark: str | None = Field(default=None, description='备注')
class CreateTaskSchedulerParam(TaskSchedulerSchemeBase):
"""创建任务调度参数"""
class UpdateTaskSchedulerParam(TaskSchedulerSchemeBase):
"""更新任务调度参数"""
class GetTaskSchedulerDetail(TaskSchedulerSchemeBase):
"""任务调度详情"""
model_config = ConfigDict(from_attributes=True)
id: int = Field(description='任务调度 ID')
enabled: bool = Field(description='是否启用任务')
total_run_count: int = Field(description='已运行总次数')
last_run_time: datetime | None = Field(None, description='最后运行时间')
created_time: datetime = Field(description='创建时间')
updated_time: datetime | None = Field(None, description='更新时间')
-29
View File
@@ -1,29 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Any
from pydantic import Field
from backend.common.schema import SchemaBase
class RunParam(SchemaBase):
"""任务运行参数"""
name: str = Field(description='任务名称')
args: list[Any] | None = Field(None, description='任务函数位置参数')
kwargs: dict[str, Any] | None = Field(None, description='任务函数关键字参数')
class TaskResult(SchemaBase):
"""任务执行结果"""
result: str = Field(description='任务执行结果')
traceback: str | None = Field(None, description='错误堆栈信息')
status: str = Field(description='任务状态')
name: str | None = Field(None, description='任务名称')
args: list[Any] | None = Field(None, description='任务函数位置参数')
kwargs: dict[str, Any] | None = Field(None, description='任务函数关键字参数')
worker: str | None = Field(None, description='执行任务的 worker')
retries: int | None = Field(None, description='重试次数')
queue: str | None = Field(None, description='任务队列')
@@ -0,0 +1,51 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from sqlalchemy import Select
from backend.app.task.crud.crud_result import task_result_dao
from backend.app.task.model.result import TaskResult
from backend.app.task.schema.result import DeleteTaskResultParam
from backend.common.exception import errors
from backend.database.db import async_db_session
class TaskResultService:
@staticmethod
async def get(*, pk: int) -> TaskResult:
"""
获取任务结果详情
: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
@staticmethod
async def get_select(*, name: str | None, task_id: str | None) -> Select:
"""
获取任务结果列表查询条件
:param name: 任务名称
:param task_id: 任务 ID
:return:
"""
return await task_result_dao.get_list(name, task_id)
@staticmethod
async def delete(*, obj: DeleteTaskResultParam) -> int:
"""
批量删除任务结果
:param obj: 任务结果 ID 列表
:return:
"""
async with async_db_session.begin() as db:
count = await task_result_dao.delete(db, obj.pks)
return count
task_result_service: TaskResultService = TaskResultService()
@@ -0,0 +1,146 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import json
from typing import Sequence
from sqlalchemy import Select
from starlette.concurrency import run_in_threadpool
from backend.app.task.celery import celery_app
from backend.app.task.crud.crud_scheduler import task_scheduler_dao
from backend.app.task.enums import TaskSchedulerType
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
class TaskSchedulerService:
"""任务调度服务类"""
@staticmethod
async def get(*, pk) -> TaskScheduler | None:
"""
获取任务调度详情
: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
@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:
"""
获取任务调度列表查询条件
:param name: 任务调度名称
:param type: 任务调度类型
:return:
"""
return await task_scheduler_dao.get_list(name=name, type=type)
@staticmethod
async def create(*, obj: CreateTaskSchedulerParam) -> None:
"""
创建任务调度
: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)
@staticmethod
async def update(*, pk: int, obj: UpdateTaskSchedulerParam) -> int:
"""
更新任务调度
: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:
if 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:
"""
更新任务调度状态
: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, not task_scheduler.enabled)
return count
@staticmethod
async def delete(*, pk) -> int:
"""
删除任务调度
: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
@staticmethod
async def execute(*, pk: int) -> None:
"""
执行任务
: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)
task_scheduler_service: TaskSchedulerService = TaskSchedulerService()
-73
View File
@@ -1,73 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from celery.exceptions import NotRegistered
from celery.result import AsyncResult
from starlette.concurrency import run_in_threadpool
from backend.app.task.celery import celery_app
from backend.app.task.schema.task import RunParam, TaskResult
from backend.common.exception import errors
from backend.common.exception.errors import NotFoundError
class TaskService:
@staticmethod
async def get_list() -> list[str]:
"""获取所有已注册的 Celery 任务列表"""
registered_tasks = await run_in_threadpool(celery_app.control.inspect().registered)
if not registered_tasks:
raise errors.ForbiddenError(msg='Celery 服务未启动')
tasks = list(registered_tasks.values())[0]
return tasks
@staticmethod
def get_detail(*, tid: str) -> TaskResult:
"""
获取指定任务的详细信息
:param tid: 任务 UUID
:return:
"""
try:
result = AsyncResult(id=tid, app=celery_app)
except NotRegistered:
raise NotFoundError(msg='任务不存在')
return TaskResult(
result=result.result,
traceback=result.traceback,
status=result.state,
name=result.name,
args=result.args,
kwargs=result.kwargs,
worker=result.worker,
retries=result.retries,
queue=result.queue,
)
@staticmethod
def revoke(*, tid: str) -> None:
"""
撤销指定的任务
:param tid: 任务 UUID
:return:
"""
try:
result = AsyncResult(id=tid, app=celery_app)
except NotRegistered:
raise NotFoundError(msg='任务不存在')
result.revoke(terminate=True)
@staticmethod
def run(*, obj: RunParam) -> str:
"""
运行指定的任务
:param obj: 任务运行参数
:return:
"""
task: AsyncResult = celery_app.send_task(name=obj.name, args=obj.args, kwargs=obj.kwargs)
return task.task_id
task_service: TaskService = TaskService()
@@ -45,5 +45,4 @@ class TaskBase(Task):
:param einfo: 异常信息 :param einfo: 异常信息
:return: :return:
""" """
loop = asyncio.get_event_loop() asyncio.create_task(task_notification(msg=f'任务 {task_id} 执行失败'))
loop.create_task(task_notification(msg=f'任务 {task_id} 执行失败'))
+31
View File
@@ -0,0 +1,31 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from celery.schedules import schedule
from backend.app.task.utils.tzcrontab import TzAwareCrontab
# 参考:https://docs.celeryq.dev/en/stable/userguide/periodic-tasks.html
LOCAL_BEAT_SCHEDULE = {
'测试同步任务': {
'task': 'task_demo',
'schedule': schedule(30),
},
'测试异步任务': {
'task': 'task_demo_async',
'schedule': TzAwareCrontab('1'),
},
'测试传参任务': {
'task': 'task_demo_params',
'schedule': TzAwareCrontab('1'),
'args': ['你好,'],
'kwargs': {'world': '世界'},
},
'清理操作日志': {
'task': 'backend.app.task.tasks.db_log.tasks.delete_db_opera_log',
'schedule': TzAwareCrontab('0', '0', day_of_week='6'),
},
'清理登录日志': {
'task': 'backend.app.task.tasks.db_log.tasks.delete_db_login_log',
'schedule': TzAwareCrontab('0', '0', day_of_month='15'),
},
}
@@ -0,0 +1,2 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
+20
View File
@@ -0,0 +1,20 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from celery import shared_task
from backend.app.admin.service.login_log_service import login_log_service
from backend.app.admin.service.opera_log_service import opera_log_service
@shared_task
async def delete_db_opera_log() -> str:
"""自动删除数据库操作日志"""
await opera_log_service.delete_all()
return 'Success'
@shared_task
async def delete_db_login_log() -> str:
"""自动删除数据库登录日志"""
await login_log_service.delete_all()
return 'Success'
+27
View File
@@ -0,0 +1,27 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from time import sleep
from anyio import sleep as asleep
from backend.app.task.celery import celery_app
@celery_app.task(name='task_demo')
def task_demo() -> str:
"""示例任务,模拟耗时操作"""
sleep(30)
return 'test async'
@celery_app.task(name='task_demo_async')
async def task_demo_async() -> str:
"""异步示例任务,模拟耗时操作"""
await asleep(30)
return 'test async'
@celery_app.task(name='task_demo_params')
async def task_demo_params(hello: str, world: str | None = None) -> str:
"""参数示例任务,模拟传参操作"""
return hello + world
+2
View File
@@ -0,0 +1,2 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
+497
View File
@@ -0,0 +1,497 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import asyncio
import json
import math
from datetime import datetime, timedelta
from multiprocessing.util import Finalize
from celery import current_app, schedules
from celery.beat import ScheduleEntry, Scheduler
from celery.signals import beat_init
from celery.utils.log import get_logger
from redis.asyncio.lock import Lock
from sqlalchemy import select
from sqlalchemy.exc import DatabaseError, InterfaceError
from backend.app.task.enums import PeriodType, TaskSchedulerType
from backend.app.task.model.scheduler import TaskScheduler
from backend.app.task.schema.scheduler import CreateTaskSchedulerParam
from backend.app.task.utils.tzcrontab import TzAwareCrontab, crontab_verify
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._await import run_await
from backend.utils.serializers import select_as_dict
from backend.utils.timezone import timezone
# 此计划程序必须比常规的 5 分钟更频繁地唤醒,因为它需要考虑对计划的外部更改
DEFAULT_MAX_INTERVAL = 5 # seconds
# 计划锁时长,避免重复创建
DEFAULT_MAX_LOCK_TIMEOUT = 300 # seconds
# 锁检测周期,应小于计划锁时长
DEFAULT_LOCK_INTERVAL = 60 # seconds
# Copied from:
# https://github.com/andymccurdy/redis-py/blob/master/redis/lock.py#L33
# Changes:
# The second line from the bottom: The original Lua script intends
# to extend time to (lock remaining time + additional time); while
# the script here extend time to an expected expiration time.
# KEYS[1] - lock name
# ARGS[1] - token
# ARGS[2] - additional milliseconds
# return 1 if the locks time was extended, otherwise 0
LUA_EXTEND_TO_SCRIPT = """
local token = redis.call('get', KEYS[1])
if not token or token ~= ARGV[1] then
return 0
end
local expiration = redis.call('pttl', KEYS[1])
if not expiration then
expiration = 0
end
if expiration < 0 then
return 0
end
redis.call('pexpire', KEYS[1], ARGV[2])
return 1
"""
logger = get_logger('fba.schedulers')
class ModelEntry(ScheduleEntry):
"""任务调度实体"""
def __init__(self, model: TaskScheduler, app=None):
super().__init__(
app=app or current_app._get_current_object(),
name=model.name,
task=model.task,
)
try:
if (
model.type == TaskSchedulerType.INTERVAL
and model.interval_every is not None
and model.interval_period is not None
):
self.schedule = schedules.schedule(timedelta(**{model.interval_period: model.interval_every}))
elif model.type == TaskSchedulerType.CRONTAB and model.crontab is not None:
crontab_split = model.crontab.split(' ')
self.schedule = TzAwareCrontab(
minute=crontab_split[0],
hour=crontab_split[1],
day_of_week=crontab_split[2],
day_of_month=crontab_split[3],
month_of_year=crontab_split[4],
)
else:
raise errors.NotFoundError(msg=f'{self.name} 计划为空!')
# logger.debug('Schedule: {}'.format(self.schedule))
except Exception as e:
logger.error(f'禁用计划为空的任务 {self.name},详情:{e}')
asyncio.create_task(self._disable(model))
try:
self.args = json.loads(model.args) if model.args else None
self.kwargs = json.loads(model.kwargs) if model.kwargs else None
except ValueError as exc:
logger.error(f'禁用参数错误的任务:{self.name}error: {str(exc)}')
asyncio.create_task(self._disable(model))
self.options = {}
for option in ['queue', 'exchange', 'routing_key']:
value = getattr(model, option)
if value is None:
continue
self.options[option] = value
expires = getattr(model, 'expires_', None)
if expires:
if isinstance(expires, int):
self.options['expires'] = expires
elif isinstance(expires, datetime):
self.options['expires'] = timezone.from_datetime(expires)
if not model.last_run_time:
model.last_run_time = timezone.now()
if model.start_time:
model.last_run_time = timezone.from_datetime(model.start_time) - timedelta(days=365)
self.last_run_at = timezone.from_datetime(model.last_run_time)
self.options['periodic_task_name'] = model.name
self.model = model
async def _disable(self, model: TaskScheduler) -> None:
"""禁用任务"""
model.no_changes = True
self.model.enabled = self.enabled = model.enabled = False
async with async_db_session.begin():
setattr(model, 'enabled', False)
def is_due(self) -> tuple[bool, int | float]:
"""任务到期状态"""
if not self.model.enabled:
# 重新启用时延迟 5 秒
return schedules.schedstate(is_due=False, next=5)
# 仅在 'start_time' 之后运行
if self.model.start_time is not None:
now = timezone.now()
start_time = timezone.from_datetime(self.model.start_time)
if now < start_time:
delay = math.ceil((start_time - now).total_seconds())
return schedules.schedstate(is_due=False, next=delay)
# 一次性任务
if self.model.one_off and self.model.enabled and self.model.total_run_count > 0:
self.model.enabled = False
self.model.total_run_count = 0
self.model.no_changes = False
save_fields = ('enabled',)
run_await(self.save)(save_fields)
return schedules.schedstate(is_due=False, next=1000000000) # 高延迟,避免重新检查
return self.schedule.is_due(self.last_run_at)
def __next__(self):
self.model.last_run_time = timezone.now()
self.model.total_run_count += 1
self.model.no_changes = True
return self.__class__(self.model)
next = __next__
async def save(self, fields: tuple = ()):
"""
保存任务状态字段
:param fields: 要保存的其他字段
:return:
"""
async with async_db_session.begin() as db:
stmt = select(TaskScheduler).where(TaskScheduler.id == self.model.id).with_for_update()
query = await db.execute(stmt)
task = query.scalars().first()
if task:
for field in ['last_run_time', 'total_run_count', 'no_changes']:
setattr(task, field, getattr(self.model, field))
for field in fields:
setattr(task, field, getattr(self.model, field))
else:
logger.warning(f'任务 {self.model.name} 不存在,跳过更新')
@classmethod
async def from_entry(cls, name, app=None, **entry):
"""保存或更新本地任务调度"""
async with async_db_session.begin() as db:
stmt = select(TaskScheduler).where(TaskScheduler.name == name)
query = await db.execute(stmt)
task = query.scalars().first()
temp = await cls._unpack_fields(name, **entry)
if not task:
task = TaskScheduler(**temp)
db.add(task)
else:
for key, value in temp.items():
setattr(task, key, value)
res = cls(task, app=app)
return res
@staticmethod
async def to_model_schedule(name: str, task: str, schedule: schedules.schedule | TzAwareCrontab):
schedule = schedules.maybe_schedule(schedule)
async with async_db_session() as db:
if isinstance(schedule, schedules.schedule):
every = max(schedule.run_every.total_seconds(), 0)
spec = {
'name': name,
'type': TaskSchedulerType.INTERVAL.value,
'interval_every': every,
'interval_period': PeriodType.SECONDS.value,
}
stmt = select(TaskScheduler).filter_by(**spec)
query = await db.execute(stmt)
obj = query.scalars().first()
if not obj:
obj = TaskScheduler(**CreateTaskSchedulerParam(task=task, **spec).model_dump())
elif isinstance(schedule, schedules.crontab):
crontab = f'{schedule._orig_minute} {schedule._orig_hour} {schedule._orig_day_of_week} {schedule._orig_day_of_month} {schedule._orig_month_of_year}' # noqa: E501
crontab_verify(crontab)
spec = {
'name': name,
'type': TaskSchedulerType.CRONTAB.value,
'crontab': crontab,
}
stmt = select(TaskScheduler).filter_by(**spec)
query = await db.execute(stmt)
obj = query.scalars().first()
if not obj:
obj = TaskScheduler(**CreateTaskSchedulerParam(task=task, **spec).model_dump())
else:
raise errors.NotFoundError(msg=f'暂不支持的计划类型:{schedule}')
return obj
@classmethod
async def _unpack_fields(
cls,
name: str,
task: str,
schedule: schedules.schedule | TzAwareCrontab,
args: tuple | None = None,
kwargs: dict | None = None,
options: dict = None,
**entry,
) -> dict:
model_schedule = await cls.to_model_schedule(name, task, schedule)
model_dict = select_as_dict(model_schedule)
for k in ['id', 'created_time', 'updated_time']:
try:
del model_dict[k]
except KeyError:
continue
model_dict.update(
args=json.dumps(args, ensure_ascii=False) if args else None,
kwargs=json.dumps(kwargs, ensure_ascii=False) if kwargs else None,
**cls._unpack_options(**options or {}),
**entry,
)
return model_dict
@classmethod
def _unpack_options(
cls,
queue: str = None,
exchange: str = None,
routing_key: str = None,
start_time: datetime = None,
expires: datetime = None,
expire_seconds: int = None,
one_off: bool = False,
) -> dict:
data = {
'queue': queue,
'exchange': exchange,
'routing_key': routing_key,
'start_time': start_time,
'expire_time': expires,
'expire_seconds': expire_seconds,
'one_off': one_off,
}
if expires:
if isinstance(expires, int):
data['expire_seconds'] = expires
elif isinstance(expires, timedelta):
data['expire_time'] = timezone.now() + expires
return data
class DatabaseScheduler(Scheduler):
"""数据库调度程序"""
Entry = ModelEntry
_schedule = None
_last_update = None
_initial_read = True
_heap_invalidated = False
lock: Lock | None = None
lock_key = f'{settings.CELERY_REDIS_PREFIX}:beat_lock'
def __init__(self, *args, **kwargs):
self.app = kwargs['app']
self._dirty = set()
super().__init__(*args, **kwargs)
self._finalize = Finalize(self, self.sync, exitpriority=5)
self.max_interval = kwargs.get('max_interval') or self.app.conf.beat_max_loop_interval or DEFAULT_MAX_INTERVAL
def setup_schedule(self):
"""重写父函数"""
logger.info('setup_schedule')
tasks = self.schedule
self.install_default_entries(tasks)
self.update_from_dict(self.app.conf.beat_schedule)
async def get_all_task_schedulers(self):
"""获取所有任务调度"""
async with async_db_session() as db:
logger.debug('DatabaseScheduler: Fetching database schedule')
stmt = select(TaskScheduler).where(TaskScheduler.enabled == 1)
query = await db.execute(stmt)
tasks = query.scalars().all()
s = {}
for task in tasks:
s[task.name] = self.Entry(task, app=self.app)
return s
def schedule_changed(self) -> bool:
"""任务调度变更状态"""
now = timezone.now()
last_update = run_await(redis_client.get)(f'{settings.CELERY_REDIS_PREFIX}:last_update')
if not last_update:
run_await(redis_client.set)(f'{settings.CELERY_REDIS_PREFIX}:last_update', timezone.to_str(now))
return False
last, ts = self._last_update, timezone.from_str(last_update)
try:
if ts and ts > (last if last else ts):
return True
finally:
self._last_update = now
def reserve(self, entry):
"""重写父函数"""
new_entry = next(entry)
# 需要按名称存储条目,因为条目可能会发生变化
self._dirty.add(new_entry.name)
return new_entry
def close(self):
"""重写父函数"""
if self.lock:
logger.info('beat: Releasing lock')
if run_await(self.lock.owned)():
run_await(self.lock.release)()
self.lock = None
super().close()
def sync(self):
"""重写父函数"""
_tried = set()
_failed = set()
try:
while self._dirty:
name = self._dirty.pop()
try:
tasks = self.schedule
run_await(tasks[name].save)()
logger.debug(f'保存任务 {name} 最新状态到数据库')
_tried.add(name)
except KeyError as e:
logger.error(f'保存任务 {name} 最新状态失败:{e} ')
_failed.add(name)
except DatabaseError as e:
logger.exception('同步时出现数据库错误: %r', e)
except InterfaceError as e:
logger.warning(f'DatabaseScheduler InterfaceError{str(e)},等待下次调用时重试...')
finally:
# 请稍后重试(仅针对失败的)
self._dirty |= _failed
def update_from_dict(self, beat_dict: dict):
"""重写父函数"""
s = {}
for name, entry_fields in beat_dict.items():
try:
entry = run_await(self.Entry.from_entry)(name, app=self.app, **entry_fields)
if entry.model.enabled:
s[name] = entry
except Exception as e:
logger.error(f'添加任务 {name} 到数据库失败')
raise e
tasks = self.schedule
tasks.update(s)
def install_default_entries(self, data):
"""重写父函数"""
entries = {}
if self.app.conf.result_expires:
entries.setdefault(
'celery.backend_cleanup',
{
'task': 'celery.backend_cleanup',
'schedule': schedules.crontab('0', '4', '*'),
'options': {'expire_seconds': 12 * 3600},
},
)
self.update_from_dict(entries)
def schedules_equal(self, *args, **kwargs):
"""重写父函数"""
if self._heap_invalidated:
self._heap_invalidated = False
return False
return super().schedules_equal(*args, **kwargs)
@property
def schedule(self) -> dict[str, ModelEntry]:
"""获取任务调度"""
initial = update = False
if self._initial_read:
logger.debug('DatabaseScheduler: initial read')
initial = update = True
self._initial_read = False
elif self.schedule_changed():
logger.info('DatabaseScheduler: Schedule changed.')
update = True
if update:
logger.debug('beat: Synchronizing schedule...')
self.sync()
self._schedule = run_await(self.get_all_task_schedulers)()
# 计划已更改,使 Scheduler.tick 中的堆无效
if not initial:
self._heap = []
self._heap_invalidated = True
logger.debug(
'Current schedule:\n%s',
'\n'.join(repr(entry) for entry in self._schedule.values()),
)
# logger.debug(self._schedule)
return self._schedule
async def extend_scheduler_lock(lock):
"""
延长调度程序锁
:param lock: 计划程序锁
:return:
"""
while True:
await asyncio.sleep(DEFAULT_LOCK_INTERVAL)
if lock:
try:
await lock.extend(DEFAULT_MAX_LOCK_TIMEOUT)
except Exception as e:
logger.error(f'Failed to extend lock: {e}')
@beat_init.connect
def acquire_distributed_beat_lock(sender=None, *args, **kwargs):
"""
尝试在启动时获取锁
:param sender: 接收方应响应的发送方
:return:
"""
scheduler = sender.scheduler
if not scheduler.lock_key:
return
logger.debug('beat: Acquiring lock...')
lock = redis_client.lock(
scheduler.lock_key,
timeout=DEFAULT_MAX_LOCK_TIMEOUT,
sleep=scheduler.max_interval,
)
# overwrite redis-py's extend script
# which will add additional timeout instead of extend to a new timeout
lock.lua_extend = redis_client.register_script(LUA_EXTEND_TO_SCRIPT)
run_await(lock.acquire)()
logger.info('beat: Acquired lock')
scheduler.lock = lock
loop = asyncio.get_event_loop()
loop.create_task(extend_scheduler_lock(scheduler.lock))
+73
View File
@@ -0,0 +1,73 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from datetime import datetime
from celery import schedules
from celery.schedules import ParseException, crontab_parser
from backend.common.exception import errors
from backend.utils.timezone import timezone
class TzAwareCrontab(schedules.crontab):
"""时区感知 Crontab"""
def __init__(self, minute='*', hour='*', day_of_week='*', day_of_month='*', month_of_year='*', app=None):
super().__init__(
minute=minute,
hour=hour,
day_of_week=day_of_week,
day_of_month=day_of_month,
month_of_year=month_of_year,
nowfun=timezone.now,
app=app,
)
def is_due(self, last_run_at: datetime) -> tuple[bool, int | float]:
"""
任务到期状态
:param last_run_at: 最后运行时间
:return:
"""
rem_delta = self.remaining_estimate(last_run_at)
rem = max(rem_delta.total_seconds(), 0)
due = rem == 0
if due:
rem_delta = self.remaining_estimate(self.now())
rem = max(rem_delta.total_seconds(), 0)
return schedules.schedstate(is_due=due, next=rem)
def __reduce__(self) -> tuple[type, tuple[str, str, str, str, str], None]:
return (
self.__class__,
(
self._orig_minute,
self._orig_hour,
self._orig_day_of_week,
self._orig_day_of_month,
self._orig_month_of_year,
),
None,
)
def crontab_verify(crontab: str) -> None:
"""
验证 Celery crontab 表达式
:param crontab: 计划表达式
:return:
"""
crontab_split = crontab.split(' ')
if len(crontab_split) != 5:
raise errors.RequestError(msg='Crontab 表达式非法')
try:
crontab_parser(60, 0).parse(crontab_split[0]) # minute
crontab_parser(24, 0).parse(crontab_split[1]) # hour
crontab_parser(7, 0).parse(crontab_split[2]) # day_of_week
crontab_parser(31, 1).parse(crontab_split[3]) # day_of_month
crontab_parser(12, 1).parse(crontab_split[4]) # month_of_year
except ParseException:
raise errors.RequestError(msg='Crontab 表达式非法')
+3 -3
View File
@@ -1,10 +1,10 @@
#!/usr/bin/env bash #!/usr/bin/env bash
# work && beat # work && beat
celery -A app.task.celery worker -l info -P gevent -c 100 & celery -A backend.app.task.celery worker -l info -P gevent -c 100 &
# beat # beat
celery -A app.task.celery beat -l info & celery -A backend.app.task.celery beat -l info &
# flower # flower
celery -A app.task.celery flower --port=8555 --basic-auth=admin:123456 celery -A backend.app.task.celery flower --port=8555 --basic-auth=admin:123456
+247
View File
@@ -0,0 +1,247 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import asyncio
import subprocess
from dataclasses import dataclass
from typing import Annotated, Literal
import cappa
import granian
from rich.panel import Panel
from rich.text import Text
from sqlalchemy import text
from watchfiles import PythonFilter
from backend import console, get_version
from backend.common.enums import DataBaseType, PrimaryKeyType
from backend.common.exception.errors import BaseExceptionMixin
from backend.core.conf import settings
from backend.database.db import async_db_session
from backend.plugin.tools import get_plugin_sql
from backend.utils.file_ops import install_git_plugin, install_zip_plugin, parse_sql_script
class CustomReloadFilter(PythonFilter):
"""自定义重载过滤器"""
def __init__(self):
super().__init__(extra_extensions=['.json', '.yaml', '.yml'])
def run(host: str, port: int, reload: bool, workers: int | None) -> None:
url = f'http://{host}:{port}'
docs_url = url + settings.FASTAPI_DOCS_URL
redoc_url = url + settings.FASTAPI_REDOC_URL
openapi_url = url + settings.FASTAPI_OPENAPI_URL
panel_content = Text()
panel_content.append(f'📝 Swagger 文档: {docs_url}\n', style='blue')
panel_content.append(f'📚 Redoc 文档: {redoc_url}\n', style='yellow')
panel_content.append(f'📡 OpenAPI JSON: {openapi_url}\n', style='green')
panel_content.append(
'🌍 fba 官方文档: https://fastapi-practices.github.io/fastapi_best_architecture_docs/',
style='cyan',
)
console.print(Panel(panel_content, title='fba 服务信息', border_style='purple', padding=(1, 2)))
granian.Granian(
target='backend.main:app',
interface='asgi',
address=host,
port=port,
reload=not reload,
reload_filter=CustomReloadFilter,
workers=workers or 1,
).serve()
def run_celery_worker(log_level: Literal['info', 'debug']) -> None:
try:
subprocess.run(['celery', '-A', 'backend.app.task.celery', 'worker', '-l', f'{log_level}', '-P', 'gevent'])
except KeyboardInterrupt:
pass
def run_celery_beat(log_level: Literal['info', 'debug']) -> None:
try:
subprocess.run(['celery', '-A', 'backend.app.task.celery', 'beat', '-l', f'{log_level}'])
except KeyboardInterrupt:
pass
def run_celery_flower(port: int, basic_auth: str) -> None:
try:
subprocess.run([
'celery',
'-A',
'backend.app.task.celery',
'flower',
f'--port={port}',
f'--basic-auth={basic_auth}',
])
except KeyboardInterrupt:
pass
async def install_plugin(
path: str, repo_url: str, no_sql: bool, db_type: DataBaseType, pk_type: PrimaryKeyType
) -> None:
if not path and not repo_url:
raise cappa.Exit('path 或 repo_url 必须指定其中一项', code=1)
if path and repo_url:
raise cappa.Exit('path 和 repo_url 不能同时指定', code=1)
plugin_name = None
console.print(Text('开始安装插件...', style='bold cyan'))
try:
if path:
plugin_name = await install_zip_plugin(file=path)
if repo_url:
plugin_name = await install_git_plugin(repo_url=repo_url)
console.print(Text(f'插件 {plugin_name} 安装成功', style='bold green'))
sql_file = await get_plugin_sql(plugin_name, db_type, pk_type)
if sql_file and not no_sql:
console.print(Text('开始自动执行插件 SQL 脚本...', style='bold cyan'))
await execute_sql_scripts(sql_file)
except Exception as e:
raise cappa.Exit(e.msg if isinstance(e, BaseExceptionMixin) else str(e), code=1)
async def execute_sql_scripts(sql_scripts: str) -> None:
async with async_db_session.begin() as db:
try:
stmts = await parse_sql_script(sql_scripts)
for stmt in stmts:
await db.execute(text(stmt))
except Exception as e:
raise cappa.Exit(f'SQL 脚本执行失败:{e}', code=1)
console.print(Text('SQL 脚本已执行完成', style='bold green'))
@cappa.command(help='运行 API 服务')
@dataclass
class Run:
host: Annotated[
str,
cappa.Arg(
long=True,
default='127.0.0.1',
help='提供服务的主机 IP 地址,对于本地开发,请使用 `127.0.0.1`。'
'要启用公共访问,例如在局域网中,请使用 `0.0.0.0`',
),
]
port: Annotated[
int,
cappa.Arg(long=True, default=8000, help='提供服务的主机端口号'),
]
no_reload: Annotated[
bool,
cappa.Arg(long=True, default=False, help='禁用在(代码)文件更改时自动重新加载服务器'),
]
workers: Annotated[
int | None,
cappa.Arg(long=True, default=None, help='使用多个工作进程,必须与 `--no-reload` 同时使用'),
]
def __call__(self):
run(host=self.host, port=self.port, reload=self.no_reload, workers=self.workers)
@cappa.command(help='从当前主机启动 Celery worker 服务')
@dataclass
class Worker:
log_level: Annotated[
Literal['info', 'debug'],
cappa.Arg(long=True, short='-l', default='info', help='日志输出级别'),
]
def __call__(self):
run_celery_worker(log_level=self.log_level)
@cappa.command(help='从当前主机启动 Celery beat 服务')
@dataclass
class Beat:
log_level: Annotated[
Literal['info', 'debug'],
cappa.Arg(long=True, short='-l', default='info', help='日志输出级别'),
]
def __call__(self):
run_celery_beat(log_level=self.log_level)
@cappa.command(help='从当前主机启动 Celery flower 服务')
@dataclass
class Flower:
port: Annotated[int, cappa.Arg(long=True, default=8555, help='提供服务的主机端口号')]
basic_auth: Annotated[str, cappa.Arg(long=True, default='admin:123456', help='页面登录的用户名和密码')]
def __call__(self):
run_celery_flower(port=self.port, basic_auth=self.basic_auth)
@cappa.command(help='运行 Celery 服务')
@dataclass
class Celery:
subcmd: cappa.Subcommands[Worker | Beat | Flower]
@cappa.command(help='新增插件')
@dataclass
class Add:
path: Annotated[
str | None,
cappa.Arg(long=True, help='ZIP 插件的本地完整路径'),
]
repo_url: Annotated[
str | None,
cappa.Arg(long=True, help='Git 插件的仓库地址'),
]
no_sql: Annotated[
bool,
cappa.Arg(long=True, default=False, help='禁用插件 SQL 脚本自动执行'),
]
db_type: Annotated[
DataBaseType,
cappa.Arg(long=True, default='mysql', help='执行插件 SQL 脚本的数据库类型'),
]
pk_type: Annotated[
PrimaryKeyType,
cappa.Arg(long=True, default='autoincrement', help='执行插件 SQL 脚本数据库主键类型'),
]
async def __call__(self):
await install_plugin(self.path, self.repo_url, self.no_sql, self.db_type, self.pk_type)
@cappa.command(help='一个高效的 fba 命令行界面')
@dataclass
class FbaCli:
version: Annotated[
bool,
cappa.Arg(short='-V', long=True, default=False, show_default=False, help='打印当前版本号'),
]
sql: Annotated[
str,
cappa.Arg(value_name='PATH', long=True, default='', show_default=False, help='在事务中执行 SQL 脚本'),
]
subcmd: cappa.Subcommands[Run | Celery | Add | None] = None
async def __call__(self):
if self.version:
get_version()
if self.sql:
await execute_sql_scripts(self.sql)
def main() -> None:
output = cappa.Output(error_format='[red]Error[/]: {message}\n\n更多信息,尝试 "[cyan]--help[/]"')
asyncio.run(cappa.invoke_async(FbaCli, output=output))
+18 -7
View File
@@ -34,13 +34,6 @@ class RequestCallNext:
response: Response response: Response
@dataclasses.dataclass
class NewToken:
new_access_token: str
new_access_token_expire_time: datetime
session_uuid: str
@dataclasses.dataclass @dataclasses.dataclass
class AccessToken: class AccessToken:
access_token: str access_token: str
@@ -54,6 +47,15 @@ class RefreshToken:
refresh_token_expire_time: datetime refresh_token_expire_time: datetime
@dataclasses.dataclass
class NewToken:
new_access_token: str
new_access_token_expire_time: datetime
new_refresh_token: str
new_refresh_token_expire_time: datetime
session_uuid: str
@dataclasses.dataclass @dataclasses.dataclass
class TokenPayload: class TokenPayload:
id: int id: int
@@ -64,3 +66,12 @@ class TokenPayload:
@dataclasses.dataclass @dataclasses.dataclass
class UploadUrl: class UploadUrl:
url: str url: str
@dataclasses.dataclass
class SnowflakeInfo:
timestamp: int
datetime: str
cluster_id: int
node_id: int
sequence: int
+30
View File
@@ -121,3 +121,33 @@ class FileType(StrEnum):
image = 'image' image = 'image'
video = 'video' video = 'video'
class PluginType(StrEnum):
"""插件类型"""
zip = 'zip'
git = 'git'
class UserPermissionType(StrEnum):
"""用户权限类型"""
superuser = 'superuser'
staff = 'staff'
status = 'status'
multi_login = 'multi_login'
class DataBaseType(StrEnum):
"""数据库类型"""
mysql = 'mysql'
postgresql = 'postgresql'
class PrimaryKeyType(StrEnum):
"""主键类型"""
autoincrement = 'autoincrement'
snowflake = 'snowflake'
+18 -3
View File
@@ -38,9 +38,15 @@ class CustomError(BaseExceptionMixin):
class RequestError(BaseExceptionMixin): class RequestError(BaseExceptionMixin):
"""请求异常""" """请求异常"""
code = StandardResponseCode.HTTP_400 def __init__(
self,
def __init__(self, *, msg: str = 'Bad Request', data: Any = None, background: BackgroundTask | None = None): *,
code: int = StandardResponseCode.HTTP_400,
msg: str = 'Bad Request',
data: Any = None,
background: BackgroundTask | None = None,
):
self.code = code
super().__init__(msg=msg, data=data, background=background) super().__init__(msg=msg, data=data, background=background)
@@ -98,3 +104,12 @@ class TokenError(HTTPError):
def __init__(self, *, msg: str = 'Not Authenticated', headers: dict[str, Any] | None = None): def __init__(self, *, msg: str = 'Not Authenticated', headers: dict[str, Any] | None = None):
super().__init__(code=self.code, msg=msg, headers=headers or {'WWW-Authenticate': 'Bearer'}) super().__init__(code=self.code, msg=msg, headers=headers or {'WWW-Authenticate': 'Bearer'})
class ConflictError(BaseExceptionMixin):
"""资源冲突异常"""
code = StandardResponseCode.HTTP_409
def __init__(self, *, msg: str = 'Conflict', data: Any = None, background: BackgroundTask | None = None):
super().__init__(msg=msg, data=data, background=background)
+16 -16
View File
@@ -8,11 +8,9 @@ from starlette.middleware.cors import CORSMiddleware
from uvicorn.protocols.http.h11_impl import STATUS_PHRASES from uvicorn.protocols.http.h11_impl import STATUS_PHRASES
from backend.common.exception.errors import BaseExceptionMixin from backend.common.exception.errors import BaseExceptionMixin
from backend.common.i18n import i18n, t
from backend.common.response.response_code import CustomResponseCode, StandardResponseCode from backend.common.response.response_code import CustomResponseCode, StandardResponseCode
from backend.common.response.response_schema import response_base from backend.common.response.response_schema import response_base
from backend.common.schema import (
CUSTOM_VALIDATION_ERROR_MESSAGES,
)
from backend.core.conf import settings from backend.core.conf import settings
from backend.utils.serializers import MsgSpecJSONResponse from backend.utils.serializers import MsgSpecJSONResponse
from backend.utils.trace_id import get_request_trace_id from backend.utils.trace_id import get_request_trace_id
@@ -46,18 +44,20 @@ async def _validation_exception_handler(request: Request, exc: RequestValidation
""" """
errors = [] errors = []
for error in exc.errors(): for error in exc.errors():
custom_message = CUSTOM_VALIDATION_ERROR_MESSAGES.get(error['type']) # 非 en-US 语言下,使用自定义错误信息
if custom_message: if i18n.current_language != 'en-US':
ctx = error.get('ctx') custom_message = t(f'pydantic.{error["type"]}')
if not ctx: if custom_message:
error['msg'] = custom_message ctx = error.get('ctx')
else: if not ctx:
error['msg'] = custom_message.format(**ctx) error['msg'] = custom_message
ctx_error = ctx.get('error') else:
if ctx_error: ctx_error = ctx.get('error')
error['ctx']['error'] = ( if ctx_error:
ctx_error.__str__().replace("'", '"') if isinstance(ctx_error, Exception) else None error['msg'] = custom_message.format(**ctx)
) error['ctx']['error'] = (
ctx_error.__str__().replace("'", '"') if isinstance(ctx_error, Exception) else None
)
errors.append(error) errors.append(error)
error = errors[0] error = errors[0]
if error.get('type') == 'json_invalid': if error.get('type') == 'json_invalid':
@@ -76,7 +76,7 @@ async def _validation_exception_handler(request: Request, exc: RequestValidation
} }
request.state.__request_validation_exception__ = content # 用于在中间件中获取异常信息 request.state.__request_validation_exception__ = content # 用于在中间件中获取异常信息
content.update(trace_id=get_request_trace_id(request)) content.update(trace_id=get_request_trace_id(request))
return MsgSpecJSONResponse(status_code=422, content=content) return MsgSpecJSONResponse(status_code=StandardResponseCode.HTTP_422, content=content)
def register_exception(app: FastAPI): def register_exception(app: FastAPI):
+83
View File
@@ -0,0 +1,83 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import glob
import json
import os
from pathlib import Path
from typing import Any
import yaml
from backend.core.conf import settings
from backend.core.path_conf import LOCALE_DIR
class I18n:
"""国际化管理器"""
def __init__(self):
self.locales: dict[str, dict[str, Any]] = {}
self.current_language: str = settings.I18N_DEFAULT_LANGUAGE
def load_locales(self):
"""加载语言文本"""
patterns = [
os.path.join(LOCALE_DIR, '*.json'),
os.path.join(LOCALE_DIR, '*.yaml'),
os.path.join(LOCALE_DIR, '*.yml'),
]
lang_files = []
for pattern in patterns:
lang_files.extend(glob.glob(pattern))
for lang_file in lang_files:
with open(lang_file, 'r', encoding='utf-8') as f:
lang = Path(lang_file).stem
file_type = Path(lang_file).suffix[1:]
match file_type:
case 'json':
self.locales[lang] = json.loads(f.read())
case 'yaml' | 'yml':
self.locales[lang] = yaml.full_load(f.read())
def t(self, key: str, default: Any | None = None, **kwargs) -> str:
"""
翻译函数
:param key: 目标文本键支持点分隔例如 'response.success'
:param default: 目标语言文本不存在时的默认文本
:param kwargs: 目标文本中的变量参数
:return:
"""
keys = key.split('.')
try:
translation = self.locales[self.current_language]
except KeyError:
keys = 'error.language_not_found'
translation = self.locales[settings.I18N_DEFAULT_LANGUAGE]
for k in keys:
if isinstance(translation, dict) and k in list(translation.keys()):
translation = translation[k]
else:
# Pydantic 兼容
if keys[0] == 'pydantic':
translation = None
else:
translation = key
if translation and kwargs:
translation = translation.format(**kwargs)
return translation or default
# 创建 i18n 单例
i18n = I18n()
# 创建翻译函数实例
t = i18n.t
+42 -17
View File
@@ -3,13 +3,15 @@
import inspect import inspect
import logging import logging
import os import os
import re
import sys import sys
from asgi_correlation_id import correlation_id from asgi_correlation_id import correlation_id
from loguru import logger from loguru import logger
from backend.core import path_conf
from backend.core.conf import settings from backend.core.conf import settings
from backend.core.path_conf import LOG_DIR
from backend.utils.timezone import timezone
class InterceptHandler(logging.Handler): class InterceptHandler(logging.Handler):
@@ -35,6 +37,18 @@ class InterceptHandler(logging.Handler):
logger.opt(depth=depth, exception=record.exc_info).log(level, record.getMessage()) logger.opt(depth=depth, exception=record.exc_info).log(level, record.getMessage())
def default_formatter(record):
"""默认日志格式化程序"""
# 重写 sqlalchemy echo 输出
# https://github.com/sqlalchemy/sqlalchemy/discussions/12791
record_name = record['name'] or ''
if record_name.startswith('sqlalchemy'):
record['message'] = re.sub(r'\s+', ' ', record['message']).strip()
return settings.LOG_FORMAT if settings.LOG_FORMAT.endswith('\n') else f'{settings.LOG_FORMAT}\n'
def setup_logging() -> None: def setup_logging() -> None:
""" """
设置日志处理器 设置日志处理器
@@ -47,9 +61,11 @@ def setup_logging() -> None:
logging.root.handlers = [InterceptHandler()] logging.root.handlers = [InterceptHandler()]
logging.root.setLevel(settings.LOG_STD_LEVEL) logging.root.setLevel(settings.LOG_STD_LEVEL)
# 配置日志传播规则
for name in logging.root.manager.loggerDict.keys(): for name in logging.root.manager.loggerDict.keys():
# 清空所有默认日志处理器
logging.getLogger(name).handlers = [] logging.getLogger(name).handlers = []
# 配置日志传播规则
if 'uvicorn.access' in name or 'watchfiles.main' in name: if 'uvicorn.access' in name or 'watchfiles.main' in name:
logging.getLogger(name).propagate = False logging.getLogger(name).propagate = False
else: else:
@@ -58,22 +74,24 @@ def setup_logging() -> None:
# Debug log handlers # Debug log handlers
# logging.debug(f'{logging.getLogger(name)}, {logging.getLogger(name).propagate}') # logging.debug(f'{logging.getLogger(name)}, {logging.getLogger(name).propagate}')
# 定义 correlation_id 默认过滤函数 # 移除 loguru 默认处理器
logger.remove()
# correlation_id 过滤器
# https://github.com/snok/asgi-correlation-id/issues/7 # https://github.com/snok/asgi-correlation-id/issues/7
def correlation_id_filter(record): def correlation_id_filter(record):
cid = correlation_id.get(settings.LOG_CID_DEFAULT_VALUE) cid = correlation_id.get(settings.TRACE_ID_LOG_DEFAULT_VALUE)
record['correlation_id'] = cid[: settings.LOG_CID_UUID_LENGTH] record['correlation_id'] = cid[: settings.TRACE_ID_LOG_LENGTH]
return record return record
# 配置 loguru 处理器 # 配置 loguru 处理器
logger.remove() # 移除默认处理器
logger.configure( logger.configure(
handlers=[ handlers=[
{ {
'sink': sys.stdout, 'sink': sys.stdout,
'level': settings.LOG_STD_LEVEL, 'level': settings.LOG_STD_LEVEL,
'format': default_formatter,
'filter': lambda record: correlation_id_filter(record), 'filter': lambda record: correlation_id_filter(record),
'format': settings.LOG_STD_FORMAT,
} }
] ]
) )
@@ -81,28 +99,35 @@ def setup_logging() -> None:
def set_custom_logfile() -> None: def set_custom_logfile() -> None:
"""设置自定义日志文件""" """设置自定义日志文件"""
log_path = path_conf.LOG_DIR if not os.path.exists(LOG_DIR):
if not os.path.exists(log_path): os.mkdir(LOG_DIR)
os.mkdir(log_path)
# 日志文件 # 日志文件
log_access_file = os.path.join(log_path, settings.LOG_ACCESS_FILENAME) log_access_file = os.path.join(LOG_DIR, settings.LOG_ACCESS_FILENAME)
log_error_file = os.path.join(log_path, settings.LOG_ERROR_FILENAME) log_error_file = os.path.join(LOG_DIR, settings.LOG_ERROR_FILENAME)
# 日志压缩回调
def compression(filepath):
filename = filepath.split(os.sep)[-1]
original_filename = filename.split('.')[0]
if '-' in original_filename:
return os.path.join(LOG_DIR, f'{original_filename}.log')
return os.path.join(LOG_DIR, f'{original_filename}_{timezone.now().strftime("%Y-%m-%d")}.log')
# 日志文件通用配置 # 日志文件通用配置
# https://loguru.readthedocs.io/en/stable/api/logger.html#loguru._logger.Logger.add # https://loguru.readthedocs.io/en/stable/api/logger.html#loguru._logger.Logger.add
log_config = { log_config = {
'format': settings.LOG_FILE_FORMAT, 'format': default_formatter,
'enqueue': True, 'enqueue': True,
'rotation': '5 MB', 'rotation': '00:00',
'retention': '7 days', 'retention': '7 days',
'compression': 'tar.gz', 'compression': lambda filepath: os.rename(filepath, compression(filepath)),
} }
# 标准输出文件 # 标准输出文件
logger.add( logger.add(
str(log_access_file), str(log_access_file),
level=settings.LOG_ACCESS_FILE_LEVEL, level=settings.LOG_FILE_ACCESS_LEVEL,
filter=lambda record: record['level'].no <= 25, filter=lambda record: record['level'].no <= 25,
backtrace=False, backtrace=False,
diagnose=False, diagnose=False,
@@ -112,7 +137,7 @@ def set_custom_logfile() -> None:
# 标准错误文件 # 标准错误文件
logger.add( logger.add(
str(log_error_file), str(log_error_file),
level=settings.LOG_ERROR_FILE_LEVEL, level=settings.LOG_FILE_ERROR_LEVEL,
filter=lambda record: record['level'].no >= 30, filter=lambda record: record['level'].no >= 30,
backtrace=True, backtrace=True,
diagnose=True, diagnose=True,
+28 -2
View File
@@ -3,17 +3,43 @@
from datetime import datetime from datetime import datetime
from typing import Annotated from typing import Annotated
from sqlalchemy import DateTime from sqlalchemy import BigInteger, DateTime
from sqlalchemy.ext.asyncio import AsyncAttrs from sqlalchemy.ext.asyncio import AsyncAttrs
from sqlalchemy.orm import DeclarativeBase, Mapped, MappedAsDataclass, declared_attr, mapped_column from sqlalchemy.orm import DeclarativeBase, Mapped, MappedAsDataclass, declared_attr, mapped_column
from backend.utils.snowflake import snowflake
from backend.utils.timezone import timezone from backend.utils.timezone import timezone
# 通用 Mapped 类型主键, 需手动添加,参考以下使用方式 # 通用 Mapped 类型主键, 需手动添加,参考以下使用方式
# MappedBase -> id: Mapped[id_key] # MappedBase -> id: Mapped[id_key]
# DataClassBase && Base -> id: Mapped[id_key] = mapped_column(init=False) # DataClassBase && Base -> id: Mapped[id_key] = mapped_column(init=False)
id_key = Annotated[ id_key = Annotated[
int, mapped_column(primary_key=True, index=True, autoincrement=True, sort_order=-999, comment='主键 ID') int,
mapped_column(
BigInteger,
primary_key=True,
unique=True,
index=True,
autoincrement=True,
sort_order=-999,
comment='主键 ID',
),
]
# 雪花算法 Mapped 类型主键,使用方法与 id_key 相同
# 详情:https://fastapi-practices.github.io/fastapi_best_architecture_docs/backend/reference/pk.html
snowflake_id_key = Annotated[
int,
mapped_column(
BigInteger,
primary_key=True,
unique=True,
index=True,
default=snowflake.generate,
sort_order=-999,
comment='雪花算法主键 ID',
),
] ]
+29
View File
@@ -0,0 +1,29 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import asyncio
from asyncio import Queue
async def batch_dequeue(queue: Queue, max_items: int, timeout: float) -> list:
"""
从异步队列中获取多个项目
:param queue: 用于获取项目的 `asyncio.Queue` 队列
:param max_items: 从队列中获取的最大项目数量
:param timeout: 总的等待超时时间
:return:
"""
items = []
async def collector():
while len(items) < max_items:
item = await queue.get()
items.append(item)
try:
await asyncio.wait_for(collector(), timeout=timeout)
except asyncio.TimeoutError:
pass
return items

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