Merge pull request #175 from 1014TaoTao/v2.0.0

V2.0.0
This commit is contained in:
fastapiadmin
2025-09-20 01:54:44 +08:00
committed by GitHub
91 changed files with 2144 additions and 1493 deletions
+15 -40
View File
@@ -23,21 +23,19 @@
<img src="https://img.shields.io/badge/-CSS3-1572B6?style=flat-square&logo=css3"/> <img src="https://img.shields.io/badge/-CSS3-1572B6?style=flat-square&logo=css3"/>
<img src="https://img.shields.io/badge/-JavaScript-563D7C?style=flat-square&logo=bootstrap"/> <img src="https://img.shields.io/badge/-JavaScript-563D7C?style=flat-square&logo=bootstrap"/>
</p> </p>
</div>
---
English | [Chinese](./README.md) English | [Chinese](./README.md)
--- </div>
## 📘 Project Introduction (Author: @1014TaoTao) ## 📘 Project Introduction
**Fastapi-Vue3-Admin** is a **completely open-source, highly modular, and technologically advanced modern rapid development platform**. Its aim is to help developers efficiently build high-quality enterprise-level mid - and back - end systems. This project adopts a **front - end and back - end separation architecture**, integrating the Python back - end framework `FastAPI` and the mainstream front - end framework `Vue3` to achieve unified development across multiple terminals, providing a one - stop out - of - the - box development experience. **Fastapi-Vue3-Admin** is a **completely open-source, highly modular, and technologically advanced modern rapid development platform**. Its aim is to help developers efficiently build high-quality enterprise-level mid - and back - end systems. This project adopts a **front - end and back - end separation architecture**, integrating the Python back - end framework `FastAPI` and the mainstream front - end framework `Vue3` to achieve unified development across multiple terminals, providing a one - stop out - of - the - box development experience.
> **Original Design Concept**: With modularity and loose coupling at its core, it pursues rich functional modules, simple and easy - to - use interfaces, detailed development documentation, and convenient maintenance methods. By unifying frameworks and components, it reduces the cost of technology selection, follows development specifications and design patterns, builds a powerful code hierarchical model, and comes with comprehensive local Chinese support. It is specifically tailored for team and enterprise development scenarios. > **Original Design Concept**: With modularity and loose coupling at its core, it pursues rich functional modules, simple and easy - to - use interfaces, detailed development documentation, and convenient maintenance methods. By unifying frameworks and components, it reduces the cost of technology selection, follows development specifications and design patterns, builds a powerful code hierarchical model, and comes with comprehensive local Chinese support. It is specifically tailored for team and enterprise development scenarios.
## 📦Engineering Structures ## 📦Engineering Structures
```sh ```sh
fastapi_vue3_admin fastapi_vue3_admin
├─ backend # Backend project ├─ backend # Backend project
@@ -52,8 +50,6 @@ fastapi_vue3_admin
└─ README.md # Chinese documentation └─ README.md # Chinese documentation
``` ```
---
## ✨ Core Highlights ## ✨ Core Highlights
| Feature | Description | | Feature | Description |
@@ -68,8 +64,6 @@ fastapi_vue3_admin
| 📄 Developer-friendly | Provides comprehensive Chinese documentation, a Chinese interface, and a visual toolchain to reduce the learning curve. | | 📄 Developer-friendly | Provides comprehensive Chinese documentation, a Chinese interface, and a visual toolchain to reduce the learning curve. |
| 🧩 Quick Access | Based on mainstream front-end technology stacks such as Vue3, Vite5, Pinia, and ElementPlus, it's ready to use out of the box. | | 🧩 Quick Access | Based on mainstream front-end technology stacks such as Vue3, Vite5, Pinia, and ElementPlus, it's ready to use out of the box. |
---
## 🛠️ Technology Stack Overview ## 🛠️ Technology Stack Overview
| Type | Technology Selection | Description | | Type | Technology Selection | Description |
@@ -85,8 +79,6 @@ fastapi_vue3_admin
| Documentation | Swagger / Redoc | Automatically generate API documentation. | | Documentation | Swagger / Redoc | Automatically generate API documentation. |
| Deployment | Docker / Nginx / Docker Compose | Rapidly deploy projects. | | Deployment | Docker / Nginx / Docker Compose | Rapidly deploy projects. |
---
## 📌 Built-in Modules ## 📌 Built-in Modules
| Module Name | Sub-module Name | Description | | Module Name | Sub-module Name | Description |
@@ -96,17 +88,12 @@ fastapi_vue3_admin
| Monitoring Management | Online users, server monitoring, cache monitoring | System monitoring and management functions | | Monitoring Management | Online users, server monitoring, cache monitoring | System monitoring and management functions |
| Public Management | Interface management, document management | Project interface documentation | | Public Management | Interface management, document management | Project interface documentation |
---
## 🍪 Demo Environment ## 🍪 Demo Environment
- Official website: <https://service.fastapiadmin.com> - Official website: <https://service.fastapiadmin.com>
- Web address: <https://service.fastapiadmin.com/web> - Web address: <https://service.fastapiadmin.com/web>
- App address: <https://service.fastapiadmin.com/app> - App address: <https://service.fastapiadmin.com/app>
- Admin account: `admin` Password: `123456` - Account: `admin` Password: `123456`
- Demo account: `demo` Password: `123456`
---
## 👷 Installation and Usage ## 👷 Installation and Usage
@@ -125,8 +112,6 @@ fastapi_vue3_admin
| Database | MySQL | 8.0 (It is recommended to use the latest version.) | | Database | MySQL | 8.0 (It is recommended to use the latest version.) |
| Middleware | Redis | 7.0 (It is recommended to use the latest version.) | | Middleware | Redis | 7.0 (It is recommended to use the latest version.) |
---
### Get Code ### Get Code
```sh ```sh
@@ -136,8 +121,6 @@ git clone https://gitee.com/tao__tao/fastapi_vue3_admin.git
git clone https://github.com/1014TaoTao/fastapi_vue3_admin.git git clone https://github.com/1014TaoTao/fastapi_vue3_admin.git
``` ```
---
### Local Backend Start ### Local Backend Start
```sh ```sh
@@ -155,8 +138,6 @@ python3 main.py revision "Initial migration" --env=dev (default is dev if not sp
python3 main.py upgrade --env=dev (default is dev if not specified) python3 main.py upgrade --env=dev (default is dev if not specified)
``` ```
---
### Local Frontend Start ### Local Frontend Start
```sh ```sh
@@ -170,8 +151,6 @@ pnpm run dev
pnpm run build pnpm run build
``` ```
---
### Local App H5 Start ### Local App H5 Start
```sh ```sh
@@ -185,8 +164,6 @@ pnpm run dev:h5
pnpm run build:h5 pnpm run build:h5
``` ```
---
### Local Project Website Start ### Local Project Website Start
```sh ```sh
@@ -200,8 +177,6 @@ pnpm run docs:dev
pnpm run docs:build pnpm run docs:build
``` ```
---
### Local Access Address ### Local Access Address
- Project official website address: <http://localhost:5180> - Project official website address: <http://localhost:5180>
@@ -210,8 +185,6 @@ pnpm run docs:build
- Admin account: `admin` Password: `123456` - Admin account: `admin` Password: `123456`
- Demo account: `demo` Password: `123456` - Demo account: `demo` Password: `123456`
---
### Docker Build ### Docker Build
```sh ```sh
@@ -244,8 +217,6 @@ fastapi_vue3_admin/devops/devops/nginx/nginx.conf
``` ```
---
## 🔧 Models ## 🔧 Models
| Module | Screenshot | | Module | Screenshot |
@@ -270,13 +241,12 @@ fastapi_vue3_admin/devops/devops/nginx/nginx.conf
| Document | ![Document](./fastdocs/src/public/help.png) | | Document | ![Document](./fastdocs/src/public/help.png) |
| Lock | ![Lock](./fastdocs/src/public/lock.png) | | Lock | ![Lock](./fastdocs/src/public/lock.png) |
### 移动端 ### Mobile
| Module <div style="width:60px"/> | Details | Module <div style="width:60px"/> | Details | Module <div style="width:60px"/> | Details | | Module <div style="width:60px"/> | Details | Module <div style="width:60px"/> | Details | Module <div style="width:60px"/> | Details |
|----------|------|----------|------|----------|------| |----------|------|----------|------|----------|------|
| Login | ![Mobile Login](./fastdocs/src/public/app_login.png) | Home | ![Mobile Home](./fastdocs/src/public/app_home.png) | Profile | ![Mobile Profile](./fastdocs/src/public/app_mine.png) | | Login | ![Mobile Login](./fastdocs/src/public/app_login.png) | Home | ![Mobile Home](./fastdocs/src/public/app_home.png) | Profile | ![Mobile Profile](./fastdocs/src/public/app_mine.png) |
| Personal | ![Mobile Personal Info](./fastdocs/src/public/app_profile.png) | Settings | ![Mobile Settings](./fastdocs/src/public/app_setting.png) | Workbench | ![Mobile Workbench](./fastdocs/src/public/app_work.png) | | Personal | ![Mobile Personal Info](./fastdocs/src/public/app_profile.png) | Settings | ![Mobile Settings](./fastdocs/src/public/app_setting.png) | Workbench | ![Mobile Workbench](./fastdocs/src/public/app_work.png) |
---
## 🛠️ Secondary Development Tutorial ## 🛠️ Secondary Development Tutorial
@@ -303,8 +273,15 @@ fastapi_vue3_admin/devops/devops/nginx/nginx.conf
1. **Configure the mobile access address for backend interfaces**: Write the code in `fastapp/src/api`. 1. **Configure the mobile access address for backend interfaces**: Write the code in `fastapp/src/api`.
2. **Write mobile pages**: Write the code in `fastapp/src/pages`. 2. **Write mobile pages**: Write the code in `fastapp/src/pages`.
## ℹ️ Help
--- For more details, please check the [Official Documentation](https://service.fastapiadmin.com)
## 👥 Contributors
<a href="https://github.com/1014TaoTao/fastapi_vue3_admin/graphs/contributors">
<img src="https://contrib.rocks/image?repo=1014TaoTao/fastapi_vue3_admin"/>
</a>
## 🙏 Thanks ## 🙏 Thanks
@@ -321,16 +298,14 @@ Thanks to the contributions and support of the following projects, which have en
- [Vue3-element-admin Project](https://gitee.com/youlaiorg/vue3-element-admin) - [Vue3-element-admin Project](https://gitee.com/youlaiorg/vue3-element-admin)
- [Vue3-element-plus-admin Project](https://gitee.com/kailong110120130/vue-element-plus-admin) - [Vue3-element-plus-admin Project](https://gitee.com/kailong110120130/vue-element-plus-admin)
---
## 🎨 Community ## 🎨 Community
| WeChat QR Code | Group QR Code | WeChat Pay QR Code | | WeChat QR Code | Group QR Code | WeChat Pay QR Code |
| --- | --- | --- | | --- | --- | --- |
| ![WeChat QR Code](./fastdocs/src/public/wechat.jpg) | ![Group QR Code](./fastdocs/src/public/group.jpg) | ![WeChat Pay QR Code](./fastdocs/src/public/wechatPay.jpg) | | ![WeChat QR Code](./fastdocs/src/public/wechat.jpg) | ![Group QR Code](./fastdocs/src/public/group.jpg) | ![WeChat Pay QR Code](./fastdocs/src/public/wechatPay.jpg) |
---
## ❤️ Star ## ❤️ Star
If you like this project, please give it a ⭐️ Star to show your support! Thank you very much! If you like this project, please give it a ⭐️ Star to show your support! Thank you very much!
[![Stargazers over time](https://starchart.cc/1014TaoTao/fastapi_vue3_admin.svg?variant=adaptive)](https://starchart.cc/1014TaoTao/fastapi_vue3_admin)
+15 -38
View File
@@ -23,21 +23,19 @@
<img src="https://img.shields.io/badge/-CSS3-1572B6?style=flat-square&logo=css3"/> <img src="https://img.shields.io/badge/-CSS3-1572B6?style=flat-square&logo=css3"/>
<img src="https://img.shields.io/badge/-JavaScript-563D7C?style=flat-square&logo=bootstrap"/> <img src="https://img.shields.io/badge/-JavaScript-563D7C?style=flat-square&logo=bootstrap"/>
</p> </p>
</div>
---
简体中文 | [English](./README.en.md) 简体中文 | [English](./README.en.md)
--- </div>
## 📘 项目介绍(作者:@1014TaoTao) ## 📘 项目介绍
**Fastapi-Vue3-Admin** 是一套 **完全开源、高度模块化、技术先进的现代化快速开发平台**,旨在帮助开发者高效搭建高质量的企业级中后台系统。该项目采用 **前后端分离架构**,融合 Python 后端框架 `FastAPI` 和前端主流框架 `Vue3` 实现多端统一开发,提供了一站式开箱即用的开发体验。 **Fastapi-Vue3-Admin** 是一套 **完全开源、高度模块化、技术先进的现代化快速开发平台**,旨在帮助开发者高效搭建高质量的企业级中后台系统。该项目采用 **前后端分离架构**,融合 Python 后端框架 `FastAPI` 和前端主流框架 `Vue3` 实现多端统一开发,提供了一站式开箱即用的开发体验。
> **设计初心**: 以模块化、松耦合为核心,追求丰富的功能模块、简洁易用的接口、详尽的开发文档和便捷的维护方式。通过统一框架和组件,降低技术选型成本,遵循开发规范和设计模式,构建强大的代码分层模型,搭配完善的本地中文化支持,专为团队和企业开发场景量身定制。 > **设计初心**: 以模块化、松耦合为核心,追求丰富的功能模块、简洁易用的接口、详尽的开发文档和便捷的维护方式。通过统一框架和组件,降低技术选型成本,遵循开发规范和设计模式,构建强大的代码分层模型,搭配完善的本地中文化支持,专为团队和企业开发场景量身定制。
## 📦工程结构概览 ## 📦工程结构概览
```sh ```sh
fastapi_vue3_admin fastapi_vue3_admin
├─ backend # 后端工程 ├─ backend # 后端工程
@@ -52,8 +50,6 @@ fastapi_vue3_admin
└─ README.md # 中文文档 └─ README.md # 中文文档
``` ```
---
## ✨ 核心亮点 ## ✨ 核心亮点
| 特性 | 描述 | | 特性 | 描述 |
@@ -68,8 +64,6 @@ fastapi_vue3_admin
| 📄 开发友好 | 提供完善的中文文档 + 中文化界面 + 可视化工具链,降低学习成本 | | 📄 开发友好 | 提供完善的中文文档 + 中文化界面 + 可视化工具链,降低学习成本 |
| 🧩 快速接入 |基于 Vue3、Vite5、Pinia、ElementPlus 等主流前端技术栈,开箱即用。| | 🧩 快速接入 |基于 Vue3、Vite5、Pinia、ElementPlus 等主流前端技术栈,开箱即用。|
---
## 🛠️ 技术栈概览 ## 🛠️ 技术栈概览
| 类型 | 技术选型 | 描述 | | 类型 | 技术选型 | 描述 |
@@ -85,8 +79,6 @@ fastapi_vue3_admin
| 文档 | Swagger / Redoc | 自动生成 API 文档。 | | 文档 | Swagger / Redoc | 自动生成 API 文档。 |
| 部署 | Docker / Nginx / Docker Compose | 快速部署项目。 | | 部署 | Docker / Nginx / Docker Compose | 快速部署项目。 |
---
## 📌 内置模块 ## 📌 内置模块
| 模块名 | 子模块名 | 描述 | | 模块名 | 子模块名 | 描述 |
@@ -96,17 +88,12 @@ fastapi_vue3_admin
| 监控管理 | 在线用户、服务器监控、缓存监控 |系统监控管理功能 | | 监控管理 | 在线用户、服务器监控、缓存监控 |系统监控管理功能 |
| 公共管理 | 接口管理、文档管理|项目接口文档 | | 公共管理 | 接口管理、文档管理|项目接口文档 |
---
## 🍪 演示环境 ## 🍪 演示环境
- 官网地址:<https://service.fastapiadmin.com> - 官网地址:<https://service.fastapiadmin.com>
- 演示地址:<https://service.fastapiadmin.com/web> - 演示地址:<https://service.fastapiadmin.com/web>
- 小程序地址:<https://service.fastapiadmin.com/app> - 小程序地址:<https://service.fastapiadmin.com/app>
- 管理员账号:`admin` 密码:`123456` - 登录账号:`admin` 密码:`123456`
- 演示账号:`demo` 密码:`123456`
---
## 👷 安装和使用 ## 👷 安装和使用
@@ -125,8 +112,6 @@ fastapi_vue3_admin
| 数据库 | MySQL | 8.0 (推荐使用最新版)| | 数据库 | MySQL | 8.0 (推荐使用最新版)|
| 中间件 | Redis | 7.0 (推荐使用最新版)| | 中间件 | Redis | 7.0 (推荐使用最新版)|
---
### 获取代码 ### 获取代码
```sh ```sh
@@ -136,8 +121,6 @@ git clone https://gitee.com/tao__tao/fastapi_vue3_admin.git
git clone https://github.com/1014TaoTao/fastapi_vue3_admin.git git clone https://github.com/1014TaoTao/fastapi_vue3_admin.git
``` ```
---
### 本地后端启动 ### 本地后端启动
```sh ```sh
@@ -155,8 +138,6 @@ python3 main.py revision "初始化迁移" --env=dev(不加默认为dev)
python3 main.py upgrade --env=dev(不加默认为dev) python3 main.py upgrade --env=dev(不加默认为dev)
``` ```
---
### 本地前端启动 ### 本地前端启动
```sh ```sh
@@ -170,8 +151,6 @@ pnpm run dev
pnpm run build pnpm run build
``` ```
---
### 本地小程序h5启动 ### 本地小程序h5启动
```sh ```sh
@@ -185,8 +164,6 @@ pnpm run dev:h5
pnpm run build:h5 pnpm run build:h5
``` ```
---
### 本地项目官网启动 ### 本地项目官网启动
```sh ```sh
@@ -200,8 +177,6 @@ pnpm run docs:dev
pnpm run docs:build pnpm run docs:build
``` ```
---
### 本地访问地址 ### 本地访问地址
- 项目官网地址: <http://localhost:5180> - 项目官网地址: <http://localhost:5180>
@@ -210,8 +185,6 @@ pnpm run docs:build
- 管理员账号:`admin` 密码:`123456` - 管理员账号:`admin` 密码:`123456`
- 演示账号:`demo` 密码:`123456` - 演示账号:`demo` 密码:`123456`
---
### docker 部署 ### docker 部署
```sh ```sh
@@ -244,8 +217,6 @@ fastapi_vue3_amdin/devops/devops/nginx/nginx.conf
``` ```
---
## 🔧 模块展示 ## 🔧 模块展示
### web 端 ### web 端
@@ -279,8 +250,6 @@ fastapi_vue3_amdin/devops/devops/nginx/nginx.conf
| 登录 | ![移动端登录](./fastdocs/src/public/app_login.png) | 首页 | ![移动端首页](./fastdocs/src/public/app_home.png) | 我的 | ![移动端个人中心](./fastdocs/src/public/app_mine.png) | | 登录 | ![移动端登录](./fastdocs/src/public/app_login.png) | 首页 | ![移动端首页](./fastdocs/src/public/app_home.png) | 我的 | ![移动端个人中心](./fastdocs/src/public/app_mine.png) |
| 个人 | ![移动端个人信息](./fastdocs/src/public/app_profile.png) | 设置 | ![移动端设置](./fastdocs/src/public/app_setting.png) | 工作台 | ![移动端工作台](./fastdocs/src/public/app_work.png) | | 个人 | ![移动端个人信息](./fastdocs/src/public/app_profile.png) | 设置 | ![移动端设置](./fastdocs/src/public/app_setting.png) | 工作台 | ![移动端工作台](./fastdocs/src/public/app_work.png) |
---
## 🛠️ 二开教程 ## 🛠️ 二开教程
### 后端部分 ### 后端部分
@@ -306,7 +275,15 @@ fastapi_vue3_amdin/devops/devops/nginx/nginx.conf
1. **移动端接入后端接口地址**:在 `fastapp/src/api` 中编写 1. **移动端接入后端接口地址**:在 `fastapp/src/api` 中编写
2. **编写移动端页面**:在 `fastapp/src/pages` 中编写 2. **编写移动端页面**:在 `fastapp/src/pages` 中编写
--- ## ℹ️ 帮助
更多详情请查看 [官方文档](https://service.fastapiadmin.com)
## 👥 贡献者
<a href="https://github.com/1014TaoTao/fastapi_vue3_admin/graphs/contributors">
<img src="https://contrib.rocks/image?repo=1014TaoTao/fastapi_vue3_admin"/>
</a>
## 🙏 特别鸣谢 ## 🙏 特别鸣谢
@@ -329,8 +306,8 @@ fastapi_vue3_amdin/devops/devops/nginx/nginx.conf
| --- | --- | --- | | --- | --- | --- |
| ![微信二维码](./fastdocs/src/public/wechat.jpg) | ![群组二维码](./fastdocs/src/public/group.jpg) | ![微信支付二维码](./fastdocs/src/public/wechatPay.jpg) | | ![微信二维码](./fastdocs/src/public/wechat.jpg) | ![群组二维码](./fastdocs/src/public/group.jpg) | ![微信支付二维码](./fastdocs/src/public/wechatPay.jpg) |
---
## ❤️ Star 支持我 ## ❤️ Star 支持我
如果你喜欢这个项目,请给我一个 ⭐️ Star 支持一下吧!非常感谢! 如果你喜欢这个项目,请给我一个 ⭐️ Star 支持一下吧!非常感谢!
[![Stargazers over time](https://starchart.cc/1014TaoTao/fastapi_vue3_admin.svg?variant=adaptive)](https://starchart.cc/1014TaoTao/fastapi_vue3_admin)
+13 -9
View File
@@ -1,9 +1,7 @@
import asyncio import asyncio
from logging.config import fileConfig from logging.config import fileConfig
from sqlalchemy.ext.asyncio import ( from sqlalchemy.ext.asyncio import create_async_engine
create_async_engine
)
from sqlalchemy import pool from sqlalchemy import pool
from alembic import context from alembic import context
@@ -12,10 +10,6 @@ from alembic import context
# access to the values within the .ini file in use. # access to the values within the .ini file in use.
config = context.config config = context.config
from app.config.setting import settings
config.set_main_option("sqlalchemy.url", settings.ASYNC_DB_URI)
# Interpret the config file for Python logging. # Interpret the config file for Python logging.
# This line sets up loggers basically. # This line sets up loggers basically.
if config.config_file_name is not None: if config.config_file_name is not None:
@@ -25,7 +19,6 @@ if config.config_file_name is not None:
# for 'autogenerate' support # for 'autogenerate' support
# from myapp import mymodel # from myapp import mymodel
# target_metadata = mymodel.Base.metadata # target_metadata = mymodel.Base.metadata
from app.core.base_model import MappedBase from app.core.base_model import MappedBase
target_metadata = MappedBase.metadata target_metadata = MappedBase.metadata
@@ -33,6 +26,8 @@ target_metadata = MappedBase.metadata
# can be acquired: # can be acquired:
# my_important_option = config.get_main_option("my_important_option") # my_important_option = config.get_main_option("my_important_option")
# ... etc. # ... etc.
from app.config.setting import settings
config.set_main_option("sqlalchemy.url", settings.ASYNC_DB_URI)
def run_migrations_offline() -> None: def run_migrations_offline() -> None:
@@ -48,6 +43,10 @@ def run_migrations_offline() -> None:
""" """
url = config.get_main_option("sqlalchemy.url") url = config.get_main_option("sqlalchemy.url")
# 确保URL不为None
if url is None:
raise ValueError("数据库URL未正确配置,请检查环境配置文件")
context.configure( context.configure(
url=url, url=url,
target_metadata=target_metadata, target_metadata=target_metadata,
@@ -66,7 +65,12 @@ def run_migrations_online() -> None:
and associate a connection with the context. and associate a connection with the context.
""" """
connectable = create_async_engine(config.get_main_option("sqlalchemy.url"), poolclass=pool.NullPool) url = config.get_main_option("sqlalchemy.url")
# 确保URL不为None
if url is None:
raise ValueError("数据库URL未正确配置,请检查环境配置文件")
connectable = create_async_engine(url, poolclass=pool.NullPool)
async def run_async_migrations(): async def run_async_migrations():
async with connectable.connect() as connection: async with connectable.connect() as connection:
+205
View File
@@ -0,0 +1,205 @@
# -*- coding: utf-8 -*-
import asyncio
import httpx, aiofiles
import numpy as np
from openai import AsyncOpenAI
from app.core.logger import logger
class AIClient:
def __init__(self, kb_filepath=None, model="qwen3:4b", embedding_model="nomic-embed-text"):
# AI模型配置
self.model = model
self.embedding_model = embedding_model
# 创建HTTP客户端
self.http_client = httpx.AsyncClient(
timeout=30.0,
follow_redirects=True
)
# 初始化OpenAI客户端(用于与Ollama交互)
self.client = AsyncOpenAI(
api_key="ollama",
base_url="http://127.0.0.1:11434/v1",
http_client=self.http_client
)
# 知识库相关属性
self.docs = []
self.embeds = None
# 如果提供了知识库文件路径,则加载知识库
self.kb_loaded = False
self.kb_filepath = kb_filepath
# RAG提示词模板
self.prompt_template = """
基于以下知识回答用户的问题:
1: %s
2: %s
3: %s
4: %s
5: %s
用户的问题: %s
请根据提供的知识,用中文简洁准确地回答问题。如果提供的知识不足以回答,请说明这一点。
"""
# 知识库相关异步方法
async def load_kb(self):
"""异步加载知识库文件"""
if not self.kb_filepath or self.kb_loaded:
return
try:
# 异步读取文件
async with aiofiles.open(self.kb_filepath, 'r', encoding='utf-8') as f:
content = await f.read()
self.docs = self.split_content(content)
self.embeds = await self.encode(self.docs)
self.kb_loaded = True
logger.info(f"成功加载知识库,包含 {len(self.docs)} 个文档片段")
except Exception as e:
logger.error(f"加载知识库失败: {str(e)}")
raise
@staticmethod
def split_content(content):
"""将内容分割成文档块"""
chunks = []
# 按换行符分割成行
lines = content.splitlines()
for line in lines:
stripped_line = line.strip()
if stripped_line:
chunks.append(stripped_line)
return chunks
async def encode(self, texts):
"""异步使用Ollama生成嵌入向量"""
embeds = []
for text in texts:
try:
# 使用AsyncOpenAI客户端异步生成嵌入
response = await self.client.embeddings.create(
model=self.embedding_model,
input=text
)
embeds.append(response.data[0].embedding)
except Exception as e:
logger.error(f"生成嵌入向量失败 for text: {text[:30]}...: {str(e)}")
# 对于失败的嵌入,添加一个零向量
embeds.append([0.0] * 768) # 假设nomic-embed-text生成768维向量
return np.array(embeds)
@staticmethod
def similarity(e1, e2):
"""计算余弦相似度"""
dot_product = np.dot(e1, e2)
norm_e1 = np.linalg.norm(e1)
norm_e2 = np.linalg.norm(e2)
if norm_e1 == 0 or norm_e2 == 0:
return 0.0 # 避免除以零
return dot_product / (norm_e1 * norm_e2)
async def search(self, text, top_k=5):
"""异步在知识库中搜索相似文本"""
# 确保知识库已加载
if not self.kb_loaded:
await self.load_kb()
if not self.embeds.any():
logger.warning("知识库为空,无法进行搜索")
return []
# 生成查询文本的嵌入向量
query_embed = (await self.encode([text]))[0]
# 计算与所有文档的相似度
sims = [(idx, self.similarity(query_embed, doc_embed))
for idx, doc_embed in enumerate(self.embeds)]
# 按相似度排序
sims.sort(key=lambda x: x[1], reverse=True)
# 返回前top_k个匹配结果
top_matches = [self.docs[idx] for idx, _ in sims[:top_k]]
return top_matches
# RAG相关异步方法
async def build_rag_prompt(self, query):
"""异步构建RAG提示词"""
# 搜索知识库获取相关上下文
context = await self.search(query)
# 确保上下文有5个元素,不足的用空字符串填充
context += [""] * (5 - len(context))
# 构建提示词
return self.prompt_template % (
context[0], context[1], context[2], context[3], context[4], query
)
# AI处理相关方法
async def process(self, query: str, use_rag=True):
"""处理查询并返回流式响应,支持RAG模式"""
system_prompt = """你是一个有用的AI助手,可以帮助用户回答问题和提供帮助。请用中文回答用户的问题。"""
# 如果启用RAG,构建增强提示词
if use_rag and self.kb_filepath:
user_query = await self.build_rag_prompt(query)
else:
user_query = query
try:
# 使用 await 调用异步客户端
response = await self.client.chat.completions.create(
model=self.model,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_query}
],
stream=True
)
# 流式返回响应
async for chunk in response:
if chunk.choices and chunk.choices[0].delta.content:
yield chunk.choices[0].delta.content
except Exception as e:
logger.error(f"AI处理查询失败: {str(e)}")
yield f"抱歉,处理您的请求时出现了错误: {str(e)}"
async def close(self):
"""关闭客户端连接"""
if hasattr(self, 'client'):
await self.client.close()
if hasattr(self, 'http_client'):
await self.http_client.aclose()
async def chat_query(message: str, kb_filepath=None):
"""处理聊天查询的异步函数"""
# 创建AI客户端实例,传入知识库文件路径
# message = message + "/no_think"
ai_client = AIClient(kb_filepath=kb_filepath)
try:
# 处理消息,启用RAG
async for response in ai_client.process(message, use_rag=True):
print(response, end='', flush=True)
finally:
# 确保关闭客户端连接
await ai_client.close()
if __name__ == "__main__":
# 在异步事件循环中运行聊天查询
asyncio.run(chat_query("帕金森氏症介绍,怎么治疗", kb_filepath='帕金森氏症en.txt'))
@@ -5,13 +5,13 @@ from fastapi.responses import JSONResponse
from app.common.response import SuccessResponse from app.common.response import SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .param import ApplicationQueryParams from .param import ApplicationQueryParam
from .service import ApplicationService from .service import ApplicationService
from .schema import ( from .schema import (
ApplicationCreateSchema, ApplicationCreateSchema,
@@ -32,12 +32,12 @@ async def get_obj_detail_controller(
@MyAppRouter.get("/list", summary="查询应用列表", description="查询应用列表") @MyAppRouter.get("/list", summary="查询应用列表", description="查询应用列表")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: ApplicationQueryParams = Depends(), search: ApplicationQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["application:myapp:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["application:myapp:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await ApplicationService.get_application_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await ApplicationService.get_application_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size) result_dict = await PaginationService.paginate(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size)
logger.info(f"查询应用列表成功") logger.info(f"查询应用列表成功")
return SuccessResponse(data=result_dict, msg="查询应用列表成功") return SuccessResponse(data=result_dict, msg="查询应用列表成功")
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class ApplicationQueryParams: class ApplicationQueryParam:
"""应用系统查询参数""" """应用系统查询参数"""
def __init__( def __init__(
@@ -7,7 +7,7 @@ from app.core.exceptions import CustomException
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .schema import ApplicationCreateSchema, ApplicationUpdateSchema, ApplicationOutSchema from .schema import ApplicationCreateSchema, ApplicationUpdateSchema, ApplicationOutSchema
from .param import ApplicationQueryParams from .param import ApplicationQueryParam
from .crud import ApplicationCRUD from .crud import ApplicationCRUD
@@ -25,7 +25,7 @@ class ApplicationService:
return ApplicationOutSchema.model_validate(obj).model_dump() return ApplicationOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_application_list_service(cls, auth: AuthSchema, search: ApplicationQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_application_list_service(cls, auth: AuthSchema, search: ApplicationQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
"""应用列表查询""" """应用列表查询"""
if order_by: if order_by:
order_by = eval(order_by) if isinstance(order_by, str) else order_by order_by = eval(order_by) if isinstance(order_by, str) else order_by
@@ -7,13 +7,13 @@ import urllib.parse
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .param import DemoQueryParams from .param import DemoQueryParam
from .service import DemoService from .service import DemoService
from .schema import ( from .schema import (
DemoCreateSchema, DemoCreateSchema,
@@ -34,12 +34,12 @@ async def get_obj_detail_controller(
@DemoRouter.get("/list", summary="查询示例列表", description="查询示例列表") @DemoRouter.get("/list", summary="查询示例列表", description="查询示例列表")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: DemoQueryParams = Depends(), search: DemoQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["demo:example:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["demo:example:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await DemoService.get_demo_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await DemoService.get_demo_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询示例列表成功") logger.info(f"查询示例列表成功")
return SuccessResponse(data=result_dict, msg="查询公告列表成功") return SuccessResponse(data=result_dict, msg="查询公告列表成功")
@@ -82,7 +82,7 @@ async def batch_set_available_obj_controller(
@DemoRouter.post('/export', summary="导出示例", description="导出示例") @DemoRouter.post('/export', summary="导出示例", description="导出示例")
async def export_obj_list_controller( async def export_obj_list_controller(
search: DemoQueryParams = Depends(), search: DemoQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["demo:example:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["demo:example:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class DemoQueryParams: class DemoQueryParam:
"""示例查询参数""" """示例查询参数"""
def __init__( def __init__(
@@ -12,7 +12,7 @@ from app.common.response import ErrorResponse
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .schema import DemoCreateSchema, DemoUpdateSchema, DemoOutSchema from .schema import DemoCreateSchema, DemoUpdateSchema, DemoOutSchema
from .param import DemoQueryParams from .param import DemoQueryParam
from .crud import DemoCRUD from .crud import DemoCRUD
@@ -28,7 +28,7 @@ class DemoService:
return DemoOutSchema.model_validate(obj).model_dump() return DemoOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_demo_list_service(cls, auth: AuthSchema, search: DemoQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_demo_list_service(cls, auth: AuthSchema, search: DemoQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
"""列表查询""" """列表查询"""
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
@@ -1,158 +1,159 @@
# -*- coding:utf-8 -*-
from datetime import datetime from datetime import datetime
from fastapi import APIRouter, Depends, Query, Request from fastapi import APIRouter, Depends, Query, Request, Body
from fastapi.responses import StreamingResponse
from pydantic_validation_decorator import ValidateFields from pydantic_validation_decorator import ValidateFields
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from config.enums import BusinessType from app.common.enums import BusinessType
from config.env import GenConfig from app.common.response import SuccessResponse, ErrorResponse, StreamResponse
from config.get_db import get_db from app.core.dependencies import AuthPermission
from module_admin.annotation.log_annotation import Log from app.core.router_class import OperationLogRoute
from module_admin.aspect.interface_auth import CheckRoleInterfaceAuth, CheckUserInterfaceAuth from app.core.base_params import PaginationQueryParam
from module_admin.service.login_service import LoginService from app.api.v1.module_system.auth.schema import AuthSchema
from module_admin.entity.vo.user_vo import CurrentUserModel from app.api.v1.module_system.user.schema import UserOutSchema
from module_generator.entity.vo.gen_vo import DeleteGenTableModel, EditGenTableModel, GenTablePageQueryModel from .param import GenTableQueryParam
from module_generator.service.gen_service import GenTableColumnService, GenTableService from .schema import DeleteGenTableSchema, EditGenTableSchema, GenTableSchema
from utils.common_util import bytes2file_response from .service import GenTableColumnService, GenTableService
from utils.log_util import logger from app.utils.common_util import bytes2file_response
from utils.page_util import PageResponseModel from app.core.logger import logger
from utils.response_util import ResponseUtil
genController = APIRouter(prefix='/tool/gen', dependencies=[Depends(LoginService.get_current_user)]) genController = APIRouter(route_class=OperationLogRoute, prefix='/tool/gen', tags=["代码生成模块"])
@genController.get( @genController.get('/list', summary="查询代码生成业务表列表", description="查询代码生成业务表列表")
'/list', response_model=PageResponseModel, dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:list'))]
)
async def get_gen_table_list( async def get_gen_table_list(
request: Request, page: PaginationQueryParam = Depends(),
gen_page_query: GenTablePageQueryModel = Depends(GenTablePageQueryModel.as_query), search: GenTableQueryParam = Depends(),
query_db: AsyncSession = Depends(get_db), auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:list"]))
): ):
# 获取分页数据 # 获取分页数据
gen_page_query_result = await GenTableService.get_gen_table_list_services(query_db, gen_page_query, is_page=True) gen_page_query_result = await GenTableService.get_gen_table_list_services(auth, search, is_page=True)
logger.info('获取成功') logger.info('获取代码生成业务表列表成功')
return SuccessResponse(data=gen_page_query_result, msg="获取代码生成业务表列表成功")
return ResponseUtil.success(model_content=gen_page_query_result)
@genController.get( @genController.get('/db/list', summary="查询数据库表列表", description="查询数据库表列表")
'/db/list', response_model=PageResponseModel, dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:list'))]
)
async def get_gen_db_table_list( async def get_gen_db_table_list(
request: Request, page: PaginationQueryParam = Depends(),
gen_page_query: GenTablePageQueryModel = Depends(GenTablePageQueryModel.as_query), search: GenTableQueryParam = Depends(),
query_db: AsyncSession = Depends(get_db), auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:list"]))
): ):
# 获取分页数据 # 获取分页数据
gen_page_query_result = await GenTableService.get_gen_db_table_list_services(query_db, gen_page_query, is_page=True) gen_page_query_result = await GenTableService.get_gen_db_table_list_services(auth, search, is_page=True)
logger.info('获取成功') logger.info('获取数据库表列表成功')
return SuccessResponse(data=gen_page_query_result, msg="获取数据库表列表成功")
return ResponseUtil.success(model_content=gen_page_query_result)
@genController.post('/importTable', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:import'))]) @genController.post('/importTable', summary="导入表结构", description="导入表结构")
@Log(title='代码生成', business_type=BusinessType.IMPORT)
async def import_gen_table(
request: Request,
tables: str = Query(),
query_db: AsyncSession = Depends(get_db),
current_user: CurrentUserModel = Depends(LoginService.get_current_user),
):
table_names = tables.split(',') if tables else []
add_gen_table_list = await GenTableService.get_gen_db_table_list_by_name_services(query_db, table_names)
add_gen_table_result = await GenTableService.import_gen_table_services(query_db, add_gen_table_list, current_user)
logger.info(add_gen_table_result.message)
return ResponseUtil.success(msg=add_gen_table_result.message)
@genController.put('', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:edit'))])
@ValidateFields(validate_model='edit_gen_table') @ValidateFields(validate_model='edit_gen_table')
@Log(title='代码生成', business_type=BusinessType.UPDATE) async def import_gen_table(
async def edit_gen_table( tables: str = Query(..., description="表名列表"),
request: Request, auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:import"])),
edit_gen_table: EditGenTableModel, current_user: UserOutSchema = Depends(lambda auth: auth.user)
query_db: AsyncSession = Depends(get_db),
current_user: CurrentUserModel = Depends(LoginService.get_current_user),
): ):
edit_gen_table.update_by = current_user.user.user_name
edit_gen_table.update_time = datetime.now()
await GenTableService.validate_edit(edit_gen_table)
edit_gen_result = await GenTableService.edit_gen_table_services(query_db, edit_gen_table)
logger.info(edit_gen_result.message)
return ResponseUtil.success(msg=edit_gen_result.message)
@genController.delete('/{table_ids}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:remove'))])
@Log(title='代码生成', business_type=BusinessType.DELETE)
async def delete_gen_table(request: Request, table_ids: str, query_db: AsyncSession = Depends(get_db)):
delete_gen_table = DeleteGenTableModel(tableIds=table_ids)
delete_gen_table_result = await GenTableService.delete_gen_table_services(query_db, delete_gen_table)
logger.info(delete_gen_table_result.message)
return ResponseUtil.success(msg=delete_gen_table_result.message)
@genController.post('/createTable', dependencies=[Depends(CheckRoleInterfaceAuth('admin'))])
@Log(title='创建表', business_type=BusinessType.OTHER)
async def create_table(
request: Request,
sql: str = Query(),
query_db: AsyncSession = Depends(get_db),
current_user: CurrentUserModel = Depends(LoginService.get_current_user),
):
create_table_result = await GenTableService.create_table_services(query_db, sql, current_user)
logger.info(create_table_result.message)
return ResponseUtil.success(msg=create_table_result.message)
@genController.get('/batchGenCode', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:code'))])
@Log(title='代码生成', business_type=BusinessType.GENCODE)
async def batch_gen_code(request: Request, tables: str = Query(), query_db: AsyncSession = Depends(get_db)):
table_names = tables.split(',') if tables else [] table_names = tables.split(',') if tables else []
batch_gen_code_result = await GenTableService.batch_gen_code_services(query_db, table_names) add_gen_table_list = await GenTableService.get_gen_db_table_list_by_name_services(auth, table_names)
logger.info('生成代码成功') result = await GenTableService.import_gen_table_services(auth, add_gen_table_list, current_user)
logger.info('导入表结构成功')
return ResponseUtil.streaming(data=bytes2file_response(batch_gen_code_result)) return result
@genController.get('/genCode/{table_name}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:code'))]) @genController.put('', summary="编辑业务表信息", description="编辑业务表信息")
@Log(title='代码生成', business_type=BusinessType.GENCODE) @ValidateFields(validate_model='edit_gen_table')
async def gen_code_local(request: Request, table_name: str, query_db: AsyncSession = Depends(get_db)): async def edit_gen_table(
if not GenConfig.allow_overwrite: data: EditGenTableSchema,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:edit"])),
current_user: UserOutSchema = Depends(lambda auth: auth.user)
):
data.update_by = current_user.username
data.update_time = datetime.now()
await GenTableService.validate_edit(data)
edit_gen_result = await GenTableService.edit_gen_table_services(auth, data)
logger.info('编辑业务表信息成功')
return SuccessResponse(data=edit_gen_result, msg="编辑业务表信息成功")
@genController.delete('/{table_ids}', summary="删除业务表信息", description="删除业务表信息")
async def delete_gen_table(
table_ids: str,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:remove"]))
):
delete_gen_table = DeleteGenTableSchema(table_ids=table_ids)
result = await GenTableService.delete_gen_table_services(auth, delete_gen_table)
logger.info('删除业务表信息成功')
return result
@genController.post('/createTable', summary="创建表结构", description="创建表结构")
async def create_table(
sql: str = Query(..., description="SQL语句"),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:create"])),
current_user: UserOutSchema = Depends(lambda auth: auth.user)
):
result = await GenTableService.create_table_services(auth, sql, current_user)
logger.info('创建表结构成功')
return result
@genController.get('/batchGenCode', summary="批量生成代码", description="批量生成代码")
async def batch_gen_code(
tables: str = Query(..., description="表名列表"),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:code"]))
):
table_names = tables.split(',') if tables else []
batch_gen_code_result = await GenTableService.batch_gen_code_services(auth, table_names)
logger.info('批量生成代码成功')
return StreamResponse(
data=bytes2file_response(batch_gen_code_result),
media_type='application/zip',
headers={'Content-Disposition': 'attachment; filename=code.zip'}
)
@genController.get('/genCode/{table_name}', summary="生成代码到指定路径", description="生成代码到指定路径")
async def gen_code_local(
table_name: str,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:code"]))
):
from app.config.setting import settings
if not settings.allow_overwrite:
logger.error('【系统预设】不允许生成文件覆盖到本地') logger.error('【系统预设】不允许生成文件覆盖到本地')
return ResponseUtil.error('【系统预设】不允许生成文件覆盖到本地') return ErrorResponse(msg='【系统预设】不允许生成文件覆盖到本地')
gen_code_local_result = await GenTableService.generate_code_services(query_db, table_name) result = await GenTableService.generate_code_services(auth, table_name)
logger.info(gen_code_local_result.message) logger.info('生成代码到指定路径成功')
return result
return ResponseUtil.success(msg=gen_code_local_result.message)
@genController.get('/{table_id}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:query'))]) @genController.get('/{table_id}', summary="获取业务表详细信息", description="获取业务表详细信息")
async def query_detail_gen_table(request: Request, table_id: int, query_db: AsyncSession = Depends(get_db)): async def query_detail_gen_table(
gen_table = await GenTableService.get_gen_table_by_id_services(query_db, table_id) table_id: int,
gen_tables = await GenTableService.get_gen_table_all_services(query_db) auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:query"]))
gen_columns = await GenTableColumnService.get_gen_table_column_list_by_table_id_services(query_db, table_id) ):
gen_table = await GenTableService.get_gen_table_by_id_services(auth, table_id)
gen_tables = await GenTableService.get_gen_table_all_services(auth)
gen_columns = await GenTableColumnService.get_gen_table_column_list_by_table_id_services(auth, table_id)
gen_table_detail_result = dict(info=gen_table, rows=gen_columns, tables=gen_tables) gen_table_detail_result = dict(info=gen_table, rows=gen_columns, tables=gen_tables)
logger.info(f'获取table_id为{table_id}的信息成功') logger.info(f'获取table_id为{table_id}的信息成功')
return SuccessResponse(data=gen_table_detail_result, msg="获取业务表详细信息成功")
return ResponseUtil.success(data=gen_table_detail_result)
@genController.get('/preview/{table_id}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:preview'))]) @genController.get('/preview/{table_id}', summary="预览代码", description="预览代码")
async def preview_code(request: Request, table_id: int, query_db: AsyncSession = Depends(get_db)): async def preview_code(
preview_code_result = await GenTableService.preview_code_services(query_db, table_id) table_id: int,
logger.info('获取预览代码成功') auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:preview"]))
):
return ResponseUtil.success(data=preview_code_result) preview_code_result = await GenTableService.preview_code_services(auth, table_id)
logger.info('预览代码成功')
return SuccessResponse(data=preview_code_result, msg="预览代码成功")
@genController.get('/synchDb/{table_name}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:edit'))]) @genController.get('/synchDb/{table_name}', summary="同步数据库", description="同步数据库")
@Log(title='代码生成', business_type=BusinessType.UPDATE) async def sync_db(
async def sync_db(request: Request, table_name: str, query_db: AsyncSession = Depends(get_db)): table_name: str,
sync_db_result = await GenTableService.sync_db_services(query_db, table_name) auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:edit"]))
logger.info(sync_db_result.message) ):
result = await GenTableService.sync_db_services(auth, table_name)
return ResponseUtil.success(data=sync_db_result.message) logger.info('同步数据库成功')
return result
@@ -1,28 +1,35 @@
# -*- coding:utf-8 -*-
from datetime import datetime, time from datetime import datetime, time
from sqlalchemy import delete, func, select, text, update from sqlalchemy import delete, func, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
from sqlglot.expressions import Expression from typing import List, Optional, Sequence, Any, Dict
from typing import List
from config.env import DataBaseConfig from .model import GenTableModel, GenTableColumnModel
from module_generator.entity.do.gen_do import GenTable, GenTableColumn from app.config.setting import settings
from module_generator.entity.vo.gen_vo import ( from app.common.request import PaginationService
GenTableBaseModel, from .schema import (
GenTableColumnBaseModel, GenTableBaseSchema,
GenTableColumnModel, GenTableColumnBaseSchema,
GenTableModel, GenTableColumnSchema,
GenTablePageQueryModel, GenTableSchema,
) )
from utils.page_util import PageUtil from .param import GenTableQueryParam
from app.core.base_crud import CRUDBase
from app.api.v1.module_system.auth.schema import AuthSchema
class GenTableDao: class GenTableDao(CRUDBase[GenTableModel, GenTableBaseSchema, GenTableBaseSchema]):
""" """
代码生成业务表模块数据库操作层 代码生成业务表模块数据库操作层
""" """
@classmethod def __init__(self, auth: AuthSchema) -> None:
async def get_gen_table_by_id(cls, db: AsyncSession, table_id: int): """初始化CRUD"""
super().__init__(model=GenTableModel, auth=auth)
async def get_gen_table_by_id(self, db: AsyncSession, table_id: int) -> Optional[GenTableModel]:
""" """
根据业务表id获取需要生成的业务表信息 根据业务表id获取需要生成的业务表信息
@@ -33,7 +40,7 @@ class GenTableDao:
gen_table_info = ( gen_table_info = (
( (
await db.execute( await db.execute(
select(GenTable).options(selectinload(GenTable.columns)).where(GenTable.table_id == table_id) select(GenTableModel).options(selectinload(GenTableModel.columns)).where(GenTableModel.table_id == table_id)
) )
) )
.scalars() .scalars()
@@ -42,8 +49,7 @@ class GenTableDao:
return gen_table_info return gen_table_info
@classmethod async def get_gen_table_by_name(self, db: AsyncSession, table_name: str) -> Optional[GenTableModel]:
async def get_gen_table_by_name(cls, db: AsyncSession, table_name: str):
""" """
根据业务表名称获取需要生成的业务表信息 根据业务表名称获取需要生成的业务表信息
@@ -54,7 +60,7 @@ class GenTableDao:
gen_table_info = ( gen_table_info = (
( (
await db.execute( await db.execute(
select(GenTable).options(selectinload(GenTable.columns)).where(GenTable.table_name == table_name) select(GenTableModel).options(selectinload(GenTableModel.columns)).where(GenTableModel.table_name == table_name)
) )
) )
.scalars() .scalars()
@@ -63,20 +69,18 @@ class GenTableDao:
return gen_table_info return gen_table_info
@classmethod async def get_gen_table_all(self, db: AsyncSession) -> Sequence[GenTableModel]:
async def get_gen_table_all(cls, db: AsyncSession):
""" """
获取所有业务表信息 获取所有业务表信息
:param db: orm对象 :param db: orm对象
:return: 所有业务表信息 :return: 所有业务表信息
""" """
gen_table_all = (await db.execute(select(GenTable).options(selectinload(GenTable.columns)))).scalars().all() gen_table_all = (await db.execute(select(GenTableModel).options(selectinload(GenTableModel.columns)))).scalars().all()
return gen_table_all return gen_table_all
@classmethod async def create_table_by_sql_dao(self, db: AsyncSession, sql_statements: List) -> None:
async def create_table_by_sql_dao(cls, db: AsyncSession, sql_statements: List[Expression]):
""" """
根据sql语句创建表结构 根据sql语句创建表结构
@@ -85,11 +89,10 @@ class GenTableDao:
:return: :return:
""" """
for sql_statement in sql_statements: for sql_statement in sql_statements:
sql = sql_statement.sql(dialect=DataBaseConfig.sqlglot_parse_dialect) sql = sql_statement.sql(dialect=settings.DATABASE_TYPE)
await db.execute(text(sql)) await db.execute(text(sql))
@classmethod async def get_gen_table_list(self, db: AsyncSession, query_object: GenTableQueryParam, is_page: bool = False):
async def get_gen_table_list(cls, db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False):
""" """
根据查询参数获取代码生成业务表列表信息 根据查询参数获取代码生成业务表列表信息
@@ -98,31 +101,51 @@ class GenTableDao:
:param is_page: 是否开启分页 :param is_page: 是否开启分页
:return: 代码生成业务表列表信息对象 :return: 代码生成业务表列表信息对象
""" """
# 构建查询条件
conditions = []
# 访问name属性而不是table_name
if query_object.name:
conditions.append(func.lower(GenTableModel.table_name).like(f'%{str(query_object.name).lower()}%'))
# 访问table_comment属性
if query_object.table_comment:
conditions.append(func.lower(GenTableModel.table_comment).like(f'%{str(query_object.table_comment).lower()}%'))
# 访问created_at属性而不是start_time和end_time
if hasattr(query_object, 'created_at') and query_object.created_at:
if isinstance(query_object.created_at, tuple) and query_object.created_at[0] == "between":
conditions.append(GenTableModel.create_time.between(*query_object.created_at[1]))
query = ( query = (
select(GenTable) select(GenTableModel)
.options(selectinload(GenTable.columns)) .options(selectinload(GenTableModel.columns))
.where( .where(*conditions)
func.lower(GenTable.table_name).like(f'%{query_object.table_name.lower()}%')
if query_object.table_name
else True,
func.lower(GenTable.table_comment).like(f'%{query_object.table_comment.lower()}%')
if query_object.table_comment
else True,
GenTable.create_time.between(
datetime.combine(datetime.strptime(query_object.begin_time, '%Y-%m-%d'), time(00, 00, 00)),
datetime.combine(datetime.strptime(query_object.end_time, '%Y-%m-%d'), time(23, 59, 59)),
)
if query_object.begin_time and query_object.end_time
else True,
)
.distinct() .distinct()
) )
gen_table_list = await PageUtil.paginate(db, query, query_object.page_num, query_object.page_size, is_page)
return gen_table_list # 获取所有数据
result = await db.execute(query)
all_data = list(result.scalars().all())
@classmethod # 使用PaginationService.paginate进行分页
async def get_gen_db_table_list(cls, db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False): if is_page and query_object.page_no is not None and query_object.page_size is not None:
paginated_result = await PaginationService.paginate(
data_list=all_data,
page_no=query_object.page_no,
page_size=query_object.page_size
)
return paginated_result
else:
return {
"items": all_data,
"total": len(all_data),
"page_no": None,
"page_size": None,
"has_next": False
}
async def get_gen_db_table_list(self, db: AsyncSession, query_object: GenTableQueryParam, is_page: bool = False):
""" """
根据查询参数获取数据库列表信息 根据查询参数获取数据库列表信息
@@ -131,7 +154,7 @@ class GenTableDao:
:param is_page: 是否开启分页 :param is_page: 是否开启分页
:return: 数据库列表信息对象 :return: 数据库列表信息对象
""" """
if DataBaseConfig.db_type == 'postgresql': if settings.DATABASE_TYPE == 'postgresql':
query_sql = """ query_sql = """
table_name as table_name, table_name as table_name,
table_comment as table_comment, table_comment as table_comment,
@@ -158,35 +181,46 @@ class GenTableDao:
and table_name not like 'gen\_%' and table_name not like 'gen\_%'
and table_name not in (select table_name from gen_table) and table_name not in (select table_name from gen_table)
""" """
if query_object.table_name: if query_object.name:
query_sql += """and lower(table_name) like lower(concat('%', :table_name, '%'))""" query_sql += """and lower(table_name) like lower(concat('%', :table_name, '%'))"""
if query_object.table_comment: if query_object.table_comment:
query_sql += """and lower(table_comment) like lower(concat('%', :table_comment, '%'))""" query_sql += """and lower(table_comment) like lower(concat('%', :table_comment, '%'))"""
if query_object.begin_time: if hasattr(query_object, 'created_at') and query_object.created_at:
if DataBaseConfig.db_type == 'postgresql': if isinstance(query_object.created_at, tuple) and query_object.created_at[0] == "between":
query_sql += """and create_time::date >= to_date(:begin_time, 'yyyy-MM-dd')""" # 这里需要特殊处理时间范围查询
else: pass
query_sql += """and date_format(create_time, '%Y%m%d') >= date_format(:begin_time, '%Y%m%d')"""
if query_object.end_time:
if DataBaseConfig.db_type == 'postgresql':
query_sql += """and create_time::date <= to_date(:end_time, 'yyyy-MM-dd')"""
else:
query_sql += """and date_format(create_time, '%Y%m%d') >= date_format(:end_time, '%Y%m%d')"""
query_sql += """order by create_time desc""" query_sql += """order by create_time desc"""
query = select( query = select(
text(query_sql).bindparams( text(query_sql).bindparams(
**{ **{
k: v k: v
for k, v in query_object.model_dump(exclude_none=True, exclude={'page_num', 'page_size'}).items() for k, v in query_object.model_dump(exclude_none=True, exclude={'page_no', 'page_size'}).items()
} }
) )
) )
gen_db_table_list = await PageUtil.paginate(db, query, query_object.page_num, query_object.page_size, is_page)
return gen_db_table_list # 执行查询
result = await db.execute(query)
all_data = list(result.fetchall())
@classmethod # 使用PaginationService.paginate进行分页
async def get_gen_db_table_list_by_names(cls, db: AsyncSession, table_names: List[str]): if is_page and query_object.page_no is not None and query_object.page_size is not None:
paginated_result = await PaginationService.paginate(
data_list=all_data,
page_no=query_object.page_no,
page_size=query_object.page_size
)
return paginated_result
else:
return {
"items": all_data,
"total": len(all_data),
"page_no": None,
"page_size": None,
"has_next": False
}
async def get_gen_db_table_list_by_names(self, db: AsyncSession, table_names: List[str]):
""" """
根据业务表名称组获取数据库列表信息 根据业务表名称组获取数据库列表信息
@@ -194,7 +228,7 @@ class GenTableDao:
:param table_names: 业务表名称组 :param table_names: 业务表名称组
:return: 数据库列表信息对象 :return: 数据库列表信息对象
""" """
if DataBaseConfig.db_type == 'postgresql': if settings.DATABASE_TYPE == 'postgresql':
query_sql = """ query_sql = """
select select
table_name as table_name, table_name as table_name,
@@ -228,51 +262,17 @@ class GenTableDao:
return gen_db_table_list return gen_db_table_list
@classmethod
async def add_gen_table_dao(cls, db: AsyncSession, gen_table: GenTableModel):
"""
新增业务表数据库操作
:param db: orm对象 class GenTableColumnDao(CRUDBase[GenTableColumnModel, GenTableColumnBaseSchema, GenTableColumnBaseSchema]):
:param gen_table: 业务表对象
:return:
"""
db_gen_table = GenTable(**GenTableBaseModel(**gen_table.model_dump(by_alias=True)).model_dump())
db.add(db_gen_table)
await db.flush()
return db_gen_table
@classmethod
async def edit_gen_table_dao(cls, db: AsyncSession, gen_table: dict):
"""
编辑业务表数据库操作
:param db: orm对象
:param gen_table: 需要更新的业务表字典
:return:
"""
await db.execute(update(GenTable), [GenTableBaseModel(**gen_table).model_dump()])
@classmethod
async def delete_gen_table_dao(cls, db: AsyncSession, gen_table: GenTableModel):
"""
删除业务表数据库操作
:param db: orm对象
:param gen_table: 业务表对象
:return:
"""
await db.execute(delete(GenTable).where(GenTable.table_id.in_([gen_table.table_id])))
class GenTableColumnDao:
""" """
代码生成业务表字段模块数据库操作层 代码生成业务表字段模块数据库操作层
""" """
@classmethod def __init__(self, auth: AuthSchema) -> None:
async def get_gen_table_column_list_by_table_id(cls, db: AsyncSession, table_id: int): """初始化CRUD"""
super().__init__(model=GenTableColumnModel, auth=auth)
async def get_gen_table_column_list_by_table_id(self, db: AsyncSession, table_id: int) -> Sequence[GenTableColumnModel]:
""" """
根据业务表id获取需要生成的业务表字段列表信息 根据业务表id获取需要生成的业务表字段列表信息
@@ -283,7 +283,7 @@ class GenTableColumnDao:
gen_table_column_list = ( gen_table_column_list = (
( (
await db.execute( await db.execute(
select(GenTableColumn).where(GenTableColumn.table_id == table_id).order_by(GenTableColumn.sort) select(GenTableColumnModel).where(GenTableColumnModel.table_id == table_id).order_by(GenTableColumnModel.sort)
) )
) )
.scalars() .scalars()
@@ -292,8 +292,7 @@ class GenTableColumnDao:
return gen_table_column_list return gen_table_column_list
@classmethod async def get_gen_db_table_columns_by_name(self, db: AsyncSession, table_name: str):
async def get_gen_db_table_columns_by_name(cls, db: AsyncSession, table_name: str):
""" """
根据业务表名称获取业务表字段列表信息 根据业务表名称获取业务表字段列表信息
@@ -301,7 +300,7 @@ class GenTableColumnDao:
:param table_name: 业务表名称 :param table_name: 业务表名称
:return: 业务表字段列表信息对象 :return: 业务表字段列表信息对象
""" """
if DataBaseConfig.db_type == 'postgresql': if settings.DATABASE_TYPE == 'postgresql':
query_sql = """ query_sql = """
select select
column_name, is_required, is_pk, sort, column_comment, is_increment, column_type column_name, is_required, is_pk, sort, column_comment, is_increment, column_type
@@ -341,53 +340,3 @@ class GenTableColumnDao:
gen_db_table_columns = (await db.execute(query)).fetchall() gen_db_table_columns = (await db.execute(query)).fetchall()
return gen_db_table_columns return gen_db_table_columns
@classmethod
async def add_gen_table_column_dao(cls, db: AsyncSession, gen_table_column: GenTableColumnModel):
"""
新增业务表字段数据库操作
:param db: orm对象
:param gen_table_column: 岗位对象
:return:
"""
db_gen_table_column = GenTableColumn(
**GenTableColumnBaseModel(**gen_table_column.model_dump(by_alias=True)).model_dump()
)
db.add(db_gen_table_column)
await db.flush()
return db_gen_table_column
@classmethod
async def edit_gen_table_column_dao(cls, db: AsyncSession, gen_table_column: dict):
"""
编辑业务表字段数据库操作
:param db: orm对象
:param gen_table_column: 需要更新的业务表字段字典
:return:
"""
await db.execute(update(GenTableColumn), [GenTableColumnBaseModel(**gen_table_column).model_dump()])
@classmethod
async def delete_gen_table_column_by_table_id_dao(cls, db: AsyncSession, gen_table_column: GenTableColumnModel):
"""
通过业务表id删除业务表字段数据库操作
:param db: orm对象
:param gen_table_column: 业务表字段对象
:return:
"""
await db.execute(delete(GenTableColumn).where(GenTableColumn.table_id.in_([gen_table_column.table_id])))
@classmethod
async def delete_gen_table_column_by_column_id_dao(cls, db: AsyncSession, gen_table_column: GenTableColumnModel):
"""
通过业务字段id删除业务表字段数据库操作
:param db: orm对象
:param post: 业务表字段对象
:return:
"""
await db.execute(delete(GenTableColumn).where(GenTableColumn.column_id.in_([gen_table_column.column_id])))
@@ -1,75 +1,82 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from datetime import datetime from datetime import datetime
from sqlalchemy.orm import relationship from typing import Optional, List
from sqlalchemy import Boolean, Column, ForeignKey, String, Integer, Text, DateTime from sqlalchemy import String, Integer, Text, DateTime, Boolean, ForeignKey, text
from sqlalchemy.orm import Mapped, mapped_column, relationship, declared_attr
from app.core.base_model import CreatorMixin from app.core.base_model import CreatorMixin
class GenTable(CreatorMixin): class GenTableModel(CreatorMixin):
""" """
代码生成表 代码生成表
""" """
__tablename__ = 'gen_table'
__table_args__ = ({'comment': '代码生成表'})
table_id = Column(Integer, primary_key=True, autoincrement=True, comment='编号') @declared_attr.directive
table_name = Column(String(200), nullable=True, default='', comment='表名称') def __tablename__(cls) -> str:
table_comment = Column(String(500), nullable=True, default='', comment='表描述') return 'gen_table'
sub_table_name = Column(String(64), nullable=True, comment='关联子表的表名')
sub_table_fk_name = Column(String(64), nullable=True, comment='子表关联的外键名')
class_name = Column(String(100), nullable=True, default='', comment='实体类名称')
tpl_category = Column(String(200), nullable=True, default='crud', comment='使用的模板(crud单表操作 tree树表操作)')
tpl_web_type = Column(String(30), nullable=True, default='', comment='前端模板类型(element-ui模版 element-plus模版)')
package_name = Column(String(100), nullable=True, comment='生成包路径')
module_name = Column(String(30), nullable=True, comment='生成模块名')
business_name = Column(String(30), nullable=True, comment='生成业务名')
function_name = Column(String(100), nullable=True, comment='生成功能名')
function_author = Column(String(100), nullable=True, comment='生成功能作者')
gen_type = Column(String(1), nullable=True, default='0', comment='生成代码方式(0zip压缩包 1自定义路径)')
gen_path = Column(String(200), nullable=True, default='/', comment='生成路径(不填默认项目路径)')
options = Column(String(1000), nullable=True, comment='其它生成选项')
create_by = Column(String(64), default='', comment='创建者')
create_time = Column(DateTime, nullable=True, default=datetime.now(), comment='创建时间')
update_by = Column(String(64), default='', comment='更新者')
update_time = Column(DateTime, nullable=True, default=datetime.now(), comment='更新时间')
remark = Column(String(500), nullable=True, default=None, comment='备注')
columns = relationship('GenTableColumn', order_by='GenTableColumn.sort', back_populates='tables') @declared_attr.directive
def __table_args__(cls) -> dict:
return {'comment': '代码生成表'}
table_name: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='', comment='表名称')
table_comment: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, default='', comment='表描述')
sub_table_name: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, comment='关联子表的表名')
sub_table_fk_name: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, comment='子表关联的外键名')
class_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, default='', comment='实体类名称')
tpl_category: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='crud', comment='使用的模板(crud单表操作 tree树表操作)')
tpl_web_type: Mapped[Optional[str]] = mapped_column(String(30), nullable=True, default='', comment='前端模板类型(element-ui模版 element-plus模版)')
package_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='生成包路径')
module_name: Mapped[Optional[str]] = mapped_column(String(30), nullable=True, comment='生成模块名')
business_name: Mapped[Optional[str]] = mapped_column(String(30), nullable=True, comment='生成业务名')
function_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='生成功能名')
function_author: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='生成功能作者')
gen_type: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, default='0', comment='生成代码方式(0zip压缩包 1自定义路径)')
gen_path: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='/', comment='生成路径(不填默认项目路径)')
options: Mapped[Optional[str]] = mapped_column(String(1000), nullable=True, comment='其它生成选项')
del_flag: Mapped[str] = mapped_column(String(1), nullable=False, default='0', server_default=text("'0'"), comment='删除标志(0代表存在 2代表删除)')
create_by: Mapped[Optional[str]] = mapped_column(String(64), default='', comment='创建者')
create_time: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True, default=None, comment='创建时间')
update_by: Mapped[Optional[str]] = mapped_column(String(64), default='', comment='更新者')
update_time: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True, default=None, comment='更新时间')
remark: Mapped[Optional[str]] = mapped_column(Text, nullable=True, default=None, comment='备注')
columns: Mapped[List['GenTableColumnModel']] = relationship('GenTableColumnModel', order_by='GenTableColumnModel.sort', back_populates='table')
class GenTableColumn(CreatorMixin): class GenTableColumnModel(CreatorMixin):
""" """
代码生成业务表字段 代码生成表字段
""" """
__tablename__ = 'gen_table_column' @declared_attr.directive
def __tablename__(cls) -> str:
return 'gen_table_column'
column_id = Column(Integer, primary_key=True, autoincrement=True, comment='编号') @declared_attr.directive
table_id = Column(Integer, ForeignKey('gen_table.table_id'), nullable=True, comment='归属表编号') def __table_args__(cls) -> dict:
column_name = Column(String(200), nullable=True, comment='列名称') return {'comment': '代码生成表字段'}
column_comment = Column(String(500), nullable=True, comment='列描述')
column_type = Column(String(100), nullable=True, comment='列类型')
python_type = Column(String(500), nullable=True, comment='PYTHON类型')
python_field = Column(String(200), nullable=True, comment='PYTHON字段名')
is_pk = Column(String(1), nullable=True, comment='是否主键(1是)')
is_increment = Column(String(1), nullable=True, comment='是否自增(1是)')
is_required = Column(String(1), nullable=True, comment='是否必填(1是)')
is_unique = Column(String(1), nullable=True, comment='是否唯一(1是)')
is_insert = Column(String(1), nullable=True, comment='是否为插入字段(1是)')
is_edit = Column(String(1), nullable=True, comment='是否编辑字段(1是)')
is_list = Column(String(1), nullable=True, comment='是否列表字段(1是)')
is_query = Column(String(1), nullable=True, comment='是否查询字段(1是)')
query_type = Column(String(200), nullable=True, default='EQ', comment='查询方式(等于、不等于、大于、小于、范围)')
html_type = Column(
String(200), nullable=True, comment='显示类型(文本框、文本域、下拉框、复选框、单选框、日期控件)'
)
dict_type = Column(String(200), nullable=True, default='', comment='字典类型')
sort = Column(Integer, nullable=True, comment='排序')
create_by = Column(String(64), default='', comment='创建者')
create_time = Column(DateTime, nullable=True, default=datetime.now(), comment='创建时间')
update_by = Column(String(64), default='', comment='更新者')
update_time = Column(DateTime, nullable=True, default=datetime.now(), comment='更新时间')
tables = relationship('GenTable', back_populates='columns') table_id: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey('gen_table.id'), nullable=True, comment='归属表编号')
column_name: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment='列名称')
column_comment: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, comment='列描述')
column_type: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='列类型')
python_type: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, comment='PYTHON类型')
python_field: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment='PYTHON字段名')
is_pk: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否主键(1是)')
is_increment: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否自增(1是)')
is_required: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否必填(1是)')
is_unique: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否唯一(1是)')
is_insert: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否为插入字段(1是)')
is_edit: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否编辑字段(1是)')
is_list: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否列表字段(1是)')
is_query: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否查询字段(1是)')
query_type: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='EQ', comment='查询方式(等于、不等于、大于、小于、范围)')
html_type: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment='显示类型(文本框、文本域、下拉框、复选框、单选框、日期控件)')
dict_type: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='', comment='字典类型')
sort: Mapped[Optional[int]] = mapped_column(Integer, nullable=True, comment='排序')
table: Mapped['GenTableModel'] = relationship('GenTableModel', back_populates='columns')
@@ -0,0 +1,64 @@
# -*- coding: utf-8 -*-
from datetime import datetime
from typing import Optional
from fastapi import Query
from app.core.validator import DateTimeStr
from app.common.request import PageResultSchema
from .schema import GenTableBaseSchema, GenTableColumnBaseSchema
class GenTableQueryParam(PageResultSchema, GenTableBaseSchema):
"""数据库表查询参数"""
def __init__(
self,
name: Optional[str] = Query(None, description="名称"),
status: Optional[bool] = Query(None, description="是否启用"),
creator: Optional[int] = Query(None, description="创建人"),
start_time: Optional[DateTimeStr] = Query(None, description="开始时间", example="2023-01-01 00:00:00"),
end_time: Optional[DateTimeStr] = Query(None, description="结束时间", example="2023-12-31 23:59:59"),
) -> None:
super().__init__()
# 模糊查询字段
self.name = ("like", name)
# 精确查询字段
self.creator_id = creator
self.status = status
# 时间范围查询
if start_time and end_time:
start_datetime = datetime.strptime(str(start_time), '%Y-%m-%d %H:%M:%S')
end_datetime = datetime.strptime(str(end_time), '%Y-%m-%d %H:%M:%S')
self.created_at = ("between", (start_datetime, end_datetime))
class GenTableColumnQueryParam(PageResultSchema, GenTableColumnBaseSchema):
"""数据库表字段查询参数"""
def __init__(
self,
name: Optional[str] = Query(None, description="名称"),
status: Optional[bool] = Query(None, description="是否启用"),
creator: Optional[int] = Query(None, description="创建人"),
start_time: Optional[DateTimeStr] = Query(None, description="开始时间", example="2023-01-01 00:00:00"),
end_time: Optional[DateTimeStr] = Query(None, description="结束时间", example="2023-12-31 23:59:59"),
) -> None:
super().__init__()
# 模糊查询字段
self.name = ("like", name)
# 精确查询字段
self.creator_id = creator
self.status = status
# 时间范围查询
if start_time and end_time:
start_datetime = datetime.strptime(str(start_time), '%Y-%m-%d %H:%M:%S')
end_datetime = datetime.strptime(str(end_time), '%Y-%m-%d %H:%M:%S')
self.created_at = ("between", (start_datetime, end_datetime))
@@ -1,24 +1,16 @@
# -*- coding:utf-8 -*-
from datetime import datetime from datetime import datetime
from typing import List, Literal, Optional, Union
from pydantic import BaseModel, ConfigDict, Field, model_validator from pydantic import BaseModel, ConfigDict, Field, model_validator
from pydantic.alias_generators import to_camel from pydantic.alias_generators import to_camel
from pydantic_validation_decorator import NotBlank from pydantic_validation_decorator import NotBlank
from typing import List, Literal, Optional
from config.constant import GenConstant
from module_admin.annotation.pydantic_annotation import as_query
from utils.string_util import StringUtil from utils.string_util import StringUtil
from pydantic import BaseModel, ConfigDict, Field from app.common.constant import GenConstant
from pydantic.alias_generators import to_camel
from typing import List, Literal, Optional, Union
from module_admin.annotation.pydantic_annotation import as_query
# -*- coding:utf-8 -*-
from datetime import datetime
from pydantic import BaseModel, ConfigDict, Field
from pydantic.alias_generators import to_camel
from typing import List, Literal, Optional, Union
from module_admin.annotation.pydantic_annotation import as_query
class GenTableBaseModel(BaseModel): class GenTableBaseSchema(BaseModel):
""" """
代码生成业务表对应pydantic模型 代码生成业务表对应pydantic模型
""" """
@@ -41,6 +33,7 @@ class GenTableBaseModel(BaseModel):
gen_type: Optional[Literal['0', '1']] = Field(default=None, description='生成代码方式(0zip压缩包 1自定义路径)') gen_type: Optional[Literal['0', '1']] = Field(default=None, description='生成代码方式(0zip压缩包 1自定义路径)')
gen_path: Optional[str] = Field(default=None, description='生成路径(不填默认项目路径)') gen_path: Optional[str] = Field(default=None, description='生成路径(不填默认项目路径)')
options: Optional[str] = Field(default=None, description='其它生成选项') options: Optional[str] = Field(default=None, description='其它生成选项')
create_by: Optional[str] = Field(default=None, description='创建者') create_by: Optional[str] = Field(default=None, description='创建者')
create_time: Optional[datetime] = Field(default=None, description='创建时间') create_time: Optional[datetime] = Field(default=None, description='创建时间')
update_by: Optional[str] = Field(default=None, description='更新者') update_by: Optional[str] = Field(default=None, description='更新者')
@@ -90,14 +83,14 @@ class GenTableBaseModel(BaseModel):
self.get_function_author() self.get_function_author()
class GenTableModel(GenTableBaseModel): class GenTableSchema(GenTableBaseSchema):
""" """
代码生成业务表模型 代码生成业务表模型
""" """
pk_column: Optional['GenTableColumnModel'] = Field(default=None, description='主键信息') pk_column: Optional['GenTableColumnSchema'] = Field(default=None, description='主键信息')
sub_table: Optional['GenTableModel'] = Field(default=None, description='子表信息') sub_table: Optional['GenTableSchema'] = Field(default=None, description='子表信息')
columns: Optional[List['GenTableColumnModel']] = Field(default=None, description='表列信息') columns: Optional[List['GenTableColumnSchema']] = Field(default=None, description='表列信息')
tree_code: Optional[str] = Field(default=None, description='树编码字段') tree_code: Optional[str] = Field(default=None, description='树编码字段')
tree_parent_code: Optional[str] = Field(default=None, description='树父编码字段') tree_parent_code: Optional[str] = Field(default=None, description='树父编码字段')
tree_name: Optional[str] = Field(default=None, description='树名称字段') tree_name: Optional[str] = Field(default=None, description='树名称字段')
@@ -108,22 +101,22 @@ class GenTableModel(GenTableBaseModel):
crud: Optional[bool] = Field(default=None, description='是否为单表') crud: Optional[bool] = Field(default=None, description='是否为单表')
@model_validator(mode='after') @model_validator(mode='after')
def check_some_is(self) -> 'GenTableModel': def check_some_is(self) -> 'GenTableSchema':
self.sub = True if self.tpl_category and self.tpl_category == GenConstant.TPL_SUB else False self.sub = True if self.tpl_category and self.tpl_category == GenConstant.TPL_SUB else False
self.tree = True if self.tpl_category and self.tpl_category == GenConstant.TPL_TREE else False self.tree = True if self.tpl_category and self.tpl_category == GenConstant.TPL_TREE else False
self.crud = True if self.tpl_category and self.tpl_category == GenConstant.TPL_CRUD else False self.crud = True if self.tpl_category and self.tpl_category == GenConstant.TPL_CRUD else False
return self return self
class EditGenTableModel(GenTableModel): class EditGenTableSchema(GenTableSchema):
""" """
修改代码生成业务表模型 修改代码生成业务表模型
""" """
params: Optional['GenTableParamsModel'] = Field(default=None, description='业务表参数') params: Optional['GenTableParamsSchema'] = Field(default=None, description='业务表参数')
class GenTableParamsModel(BaseModel): class GenTableParamsSchema(BaseModel):
""" """
代码生成业务表参数模型 代码生成业务表参数模型
""" """
@@ -136,26 +129,7 @@ class GenTableParamsModel(BaseModel):
parent_menu_id: Optional[int] = Field(default=None, description='上级菜单ID字段') parent_menu_id: Optional[int] = Field(default=None, description='上级菜单ID字段')
class GenTableQueryModel(GenTableBaseModel): class DeleteGenTableSchema(BaseModel):
"""
代码生成业务表不分页查询模型
"""
begin_time: Optional[str] = Field(default=None, description='开始时间')
end_time: Optional[str] = Field(default=None, description='结束时间')
@as_query
class GenTablePageQueryModel(GenTableQueryModel):
"""
代码生成业务表分页查询模型
"""
page_num: int = Field(default=1, description='当前页码')
page_size: int = Field(default=10, description='每页记录数')
class DeleteGenTableModel(BaseModel):
""" """
删除代码生成业务表模型 删除代码生成业务表模型
""" """
@@ -165,7 +139,7 @@ class DeleteGenTableModel(BaseModel):
table_ids: str = Field(description='需要删除的代码生成业务表ID') table_ids: str = Field(description='需要删除的代码生成业务表ID')
class GenTableColumnBaseModel(BaseModel): class GenTableColumnBaseSchema(BaseModel):
""" """
代码生成业务表字段对应pydantic模型 代码生成业务表字段对应pydantic模型
""" """
@@ -188,9 +162,7 @@ class GenTableColumnBaseModel(BaseModel):
is_list: Optional[str] = Field(default=None, description='是否列表字段(1是)') is_list: Optional[str] = Field(default=None, description='是否列表字段(1是)')
is_query: Optional[str] = Field(default=None, description='是否查询字段(1是)') is_query: Optional[str] = Field(default=None, description='是否查询字段(1是)')
query_type: Optional[str] = Field(default=None, description='查询方式(等于、不等于、大于、小于、范围)') query_type: Optional[str] = Field(default=None, description='查询方式(等于、不等于、大于、小于、范围)')
html_type: Optional[str] = Field( html_type: Optional[str] = Field(default=None, description='显示类型(文本框、文本域、下拉框、复选框、单选框、日期控件)')
default=None, description='显示类型(文本框、文本域、下拉框、复选框、单选框、日期控件)'
)
dict_type: Optional[str] = Field(default=None, description='字典类型') dict_type: Optional[str] = Field(default=None, description='字典类型')
sort: Optional[int] = Field(default=None, description='排序') sort: Optional[int] = Field(default=None, description='排序')
create_by: Optional[str] = Field(default=None, description='创建者') create_by: Optional[str] = Field(default=None, description='创建者')
@@ -206,7 +178,7 @@ class GenTableColumnBaseModel(BaseModel):
self.get_python_field() self.get_python_field()
class GenTableColumnModel(GenTableColumnBaseModel): class GenTableColumnSchema(GenTableColumnBaseSchema):
""" """
代码生成业务表字段模型 代码生成业务表字段模型
""" """
@@ -224,7 +196,7 @@ class GenTableColumnModel(GenTableColumnBaseModel):
usable_column: Optional[bool] = Field(default=None, description='是否为基类字段白名单') usable_column: Optional[bool] = Field(default=None, description='是否为基类字段白名单')
@model_validator(mode='after') @model_validator(mode='after')
def check_some_is(self) -> 'GenTableModel': def check_some_is(self) -> 'GenTableSchema':
self.cap_python_field = self.python_field[0].upper() + self.python_field[1:] if self.python_field else None self.cap_python_field = self.python_field[0].upper() + self.python_field[1:] if self.python_field else None
self.pk = True if self.is_pk and self.is_pk == '1' else False self.pk = True if self.is_pk and self.is_pk == '1' else False
self.increment = True if self.is_increment and self.is_increment == '1' else False self.increment = True if self.is_increment and self.is_increment == '1' else False
@@ -245,26 +217,7 @@ class GenTableColumnModel(GenTableColumnBaseModel):
return self return self
class GenTableColumnQueryModel(GenTableColumnBaseModel): class DeleteGenTableColumnSchema(BaseModel):
"""
代码生成业务表字段不分页查询模型
"""
begin_time: Optional[str] = Field(default=None, description='开始时间')
end_time: Optional[str] = Field(default=None, description='结束时间')
@as_query
class GenTableColumnPageQueryModel(GenTableColumnQueryModel):
"""
代码生成业务表字段分页查询模型
"""
page_num: int = Field(default=1, description='当前页码')
page_size: int = Field(default=10, description='每页记录数')
class DeleteGenTableColumnModel(BaseModel):
""" """
删除代码生成业务表字段模型 删除代码生成业务表字段模型
""" """
@@ -1,28 +1,35 @@
# -*- coding:utf-8 -*-
import io import io
import json import json
import os import os
import zipfile import zipfile
from datetime import datetime from datetime import datetime
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlglot import parse as sqlglot_parse from typing import Any, List, Dict, Optional, Sequence
from sqlglot.expressions import Add, Alter, Create, Delete, Drop, Expression, Insert, Table, TruncateTable, Update
from typing import List from app.config.setting import settings
from config.constant import GenConstant from app.core.exceptions import CustomException
from config.env import DataBaseConfig, GenConfig from app.utils.common_util import CamelCaseUtil
from exceptions.exception import ServiceException from app.utils.gen_util import GenUtils
from module_admin.entity.vo.common_vo import CrudResponseModel from app.utils.template_util import TemplateInitializer, TemplateUtils
from module_admin.entity.vo.user_vo import CurrentUserModel from app.common.constant import GenConstant
from module_generator.entity.vo.gen_vo import ( from app.common.response import SuccessResponse
DeleteGenTableModel, from app.api.v1.module_system.user.schema import UserOutSchema
EditGenTableModel, from .schema import (
GenTableColumnModel, DeleteGenTableSchema,
GenTableModel, EditGenTableSchema,
GenTablePageQueryModel, GenTableColumnSchema,
GenTableSchema,
) )
from module_generator.dao.gen_dao import GenTableColumnDao, GenTableDao from .param import GenTableQueryParam
from utils.common_util import CamelCaseUtil from .crud import GenTableColumnDao, GenTableDao
from utils.gen_util import GenUtils from .model import GenTableModel, GenTableColumnModel
from utils.template_util import TemplateInitializer, TemplateUtils from app.api.v1.module_system.auth.schema import AuthSchema
# 定义默认的GenConfig值
GEN_PATH = "generated_code" # 默认生成路径
class GenTableService: class GenTableService:
@@ -32,237 +39,275 @@ class GenTableService:
@classmethod @classmethod
async def get_gen_table_list_services( async def get_gen_table_list_services(
cls, query_db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False cls, auth: AuthSchema, query_object: GenTableQueryParam, is_page: bool = False
): ):
""" """
获取代码生成业务表列表信息service 获取代码生成业务表列表信息service
:param query_db: orm对象 :param auth: 认证信息
:param query_object: 查询参数对象 :param query_object: 查询参数对象
:param is_page: 是否开启分页 :param is_page: 是否开启分页
:return: 代码生成业务列表信息对象 :return: 代码生成业务列表信息对象
""" """
gen_table_list_result = await GenTableDao.get_gen_table_list(query_db, query_object, is_page) # 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_list_result = await gen_table_dao.get_gen_table_list(auth.db, query_object, is_page)
return gen_table_list_result return gen_table_list_result
@classmethod @classmethod
async def get_gen_db_table_list_services( async def get_gen_db_table_list_services(
cls, query_db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False cls, auth: AuthSchema, query_object: GenTableQueryParam, is_page: bool = False
): ):
""" """
获取数据库列表信息service 获取数据库列表信息service
:param query_db: orm对象 :param auth: 认证信息
:param query_object: 查询参数对象 :param query_object: 查询参数对象
:param is_page: 是否开启分页 :param is_page: 是否开启分页
:return: 数据库列表信息对象 :return: 数据库列表信息对象
""" """
gen_db_table_list_result = await GenTableDao.get_gen_db_table_list(query_db, query_object, is_page) # 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_db_table_list_result = await gen_table_dao.get_gen_db_table_list(auth.db, query_object, is_page)
return gen_db_table_list_result return gen_db_table_list_result
@classmethod @classmethod
async def get_gen_db_table_list_by_name_services(cls, query_db: AsyncSession, table_names: List[str]): async def get_gen_db_table_list_by_name_services(cls, auth: AuthSchema, table_names: List[str]) -> list[GenTableSchema]:
""" """
根据表名称组获取数据库列表信息service 根据表名称组获取数据库列表信息service
:param query_db: orm对象 :param auth: 认证信息
:param table_names: 表名称组 :param table_names: 表名称组
:return: 数据库列表信息对象 :return: 数据库列表信息对象
""" """
gen_db_table_list_result = await GenTableDao.get_gen_db_table_list_by_names(query_db, table_names) # 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
return [GenTableModel(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_db_table_list_result)] gen_table_dao = GenTableDao(auth=auth)
gen_db_table_list_result = await gen_table_dao.get_gen_db_table_list_by_names(auth.db, table_names)
return [GenTableSchema(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_db_table_list_result)]
@classmethod @classmethod
async def import_gen_table_services( async def import_gen_table_services(
cls, query_db: AsyncSession, gen_table_list: List[GenTableModel], current_user: CurrentUserModel cls, auth: AuthSchema, gen_table_list: List[GenTableSchema], current_user: UserOutSchema
): ):
""" """
导入表结构service 导入表结构service
:param query_db: orm对象 :param auth: 认证信息
:param gen_table_list: 导入表列表 :param gen_table_list: 导入表列表
:param current_user: 当前用户信息对象 :param current_user: 当前用户信息对象
:return: 导入结果 :return: 导入结果
""" """
try: try:
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
for table in gen_table_list: for table in gen_table_list:
table_name = table.table_name table_name = table.table_name
GenUtils.init_table(table, current_user.user.user_name) GenUtils.init_table(table, current_user.username) # 使用username而不是user.user_name
add_gen_table = await GenTableDao.add_gen_table_dao(query_db, table) add_gen_table = await gen_table_dao.create(data=table.model_dump())
if add_gen_table: if add_gen_table:
table.table_id = add_gen_table.table_id table.table_id = add_gen_table.table_id
gen_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name(query_db, table_name) gen_table_columns = await gen_table_column_dao.get_gen_db_table_columns_by_name(auth.db, table_name or "")
for column in [ for column in [
GenTableColumnModel(**gen_table_column) GenTableColumnSchema(**gen_table_column)
for gen_table_column in CamelCaseUtil.transform_result(gen_table_columns) for gen_table_column in CamelCaseUtil.transform_result(gen_table_columns)
]: ]:
GenUtils.init_column_field(column, table) GenUtils.init_column_field(column, table)
await GenTableColumnDao.add_gen_table_column_dao(query_db, column) await gen_table_column_dao.create(data=column.model_dump())
await query_db.commit() if isinstance(auth.db, AsyncSession):
return CrudResponseModel(is_success=True, message='导入成功') await auth.db.commit()
return SuccessResponse(msg='导入成功')
except Exception as e: except Exception as e:
await query_db.rollback() if isinstance(auth.db, AsyncSession):
raise ServiceException(message=f'导入失败, {str(e)}') try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'导入失败, {str(e)}')
@classmethod @classmethod
async def edit_gen_table_services(cls, query_db: AsyncSession, page_object: EditGenTableModel): async def edit_gen_table_services(cls, auth: AuthSchema, page_object: EditGenTableSchema) -> Dict[str, Any]:
""" """
编辑业务表信息service 编辑业务表信息service
:param query_db: orm对象 :param auth: 认证信息
:param page_object: 编辑业务表对象 :param page_object: 编辑业务表对象
:return: 编辑业务表校验结果 :return: 编辑业务表校验结果
""" """
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
# 检查必要字段是否存在
if page_object.table_id is None:
raise CustomException(msg='业务表ID不能为空')
edit_gen_table = page_object.model_dump(exclude_unset=True, by_alias=True) edit_gen_table = page_object.model_dump(exclude_unset=True, by_alias=True)
gen_table_info = await cls.get_gen_table_by_id_services(query_db, page_object.table_id) gen_table_info = await cls.get_gen_table_by_id_services(auth, page_object.table_id)
if gen_table_info.table_id: if gen_table_info.table_id:
try: try:
edit_gen_table['options'] = json.dumps(edit_gen_table.get('params')) # 处理params字段,确保不为None
await GenTableDao.edit_gen_table_dao(query_db, edit_gen_table) params = edit_gen_table.get('params')
for gen_table_column in page_object.columns: if params is not None:
gen_table_column.update_by = page_object.update_by edit_gen_table['options'] = json.dumps(params)
gen_table_column.update_time = datetime.now() else:
await GenTableColumnDao.edit_gen_table_column_dao( edit_gen_table['options'] = '{}' # 默认空对象
query_db, gen_table_column.model_dump(by_alias=True)
) # 移除params字段,因为options字段已经包含了序列化的params
await query_db.commit() edit_gen_table.pop('params', None)
return CrudResponseModel(is_success=True, message='更新成功')
await gen_table_dao.update(id=page_object.table_id, data=edit_gen_table)
if page_object.columns:
for gen_table_column in page_object.columns:
gen_table_column.update_by = page_object.update_by
gen_table_column.update_time = datetime.now()
if gen_table_column.column_id is not None:
await gen_table_column_dao.update(
id=gen_table_column.column_id,
data=gen_table_column.model_dump(by_alias=True)
)
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return {"is_success": True, "message": "更新成功"}
except Exception as e: except Exception as e:
await query_db.rollback() if isinstance(auth.db, AsyncSession):
raise e try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'更新失败: {str(e)}')
else: else:
raise ServiceException(message='业务表不存在') raise CustomException(msg='业务表不存在')
@classmethod @classmethod
async def delete_gen_table_services(cls, query_db: AsyncSession, page_object: DeleteGenTableModel): async def delete_gen_table_services(cls, auth: AuthSchema, page_object: DeleteGenTableSchema) -> SuccessResponse:
""" """
删除业务表信息service 删除业务表信息service
:param query_db: orm对象 :param auth: 认证信息
:param page_object: 删除业务表对象 :param page_object: 删除业务表对象
:return: 删除业务表校验结果 :return: 删除业务表校验结果
""" """
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
if page_object.table_ids: if page_object.table_ids:
table_id_list = page_object.table_ids.split(',') table_id_list = page_object.table_ids.split(',')
try: try:
for table_id in table_id_list: for table_id in table_id_list:
await GenTableDao.delete_gen_table_dao(query_db, GenTableModel(tableId=table_id)) await gen_table_dao.delete(ids=[int(table_id)])
await GenTableColumnDao.delete_gen_table_column_by_table_id_dao( # 删除相关的字段信息
query_db, GenTableColumnModel(tableId=table_id) # 这里需要先查询出所有相关的column_id,然后删除
) columns = await gen_table_column_dao.get_gen_table_column_list_by_table_id(auth.db, int(table_id))
await query_db.commit() if columns:
return CrudResponseModel(is_success=True, message='删除成功') column_ids = [column.column_id for column in columns]
await gen_table_column_dao.delete(ids=column_ids)
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return SuccessResponse(msg='删除成功')
except Exception as e: except Exception as e:
await query_db.rollback() if isinstance(auth.db, AsyncSession):
raise e try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'删除失败: {str(e)}')
else: else:
raise ServiceException(message='传入业务表id为空') raise CustomException(msg='传入业务表id为空')
@classmethod @classmethod
async def get_gen_table_by_id_services(cls, query_db: AsyncSession, table_id: int): async def get_gen_table_by_id_services(cls, auth: AuthSchema, table_id: int) -> GenTableSchema:
""" """
获取需要生成的业务表详细信息service 获取需要生成的业务表详细信息service
:param query_db: orm对象 :param auth: 认证信息
:param table_id: 需要生成的业务表id :param table_id: 需要生成的业务表id
:return: 需要生成的业务表id对应的信息 :return: 需要生成的业务表id对应的信息
""" """
gen_table = await GenTableDao.get_gen_table_by_id(query_db, table_id) # 确保db是AsyncSession类型
result = await cls.set_table_from_options(GenTableModel(**CamelCaseUtil.transform_result(gen_table))) if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
return result gen_table_dao = GenTableDao(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_id(auth.db, table_id)
if gen_table:
result = await cls.set_table_from_options(GenTableSchema(**CamelCaseUtil.transform_result(gen_table)))
return result
else:
raise CustomException(msg='业务表不存在')
@classmethod @classmethod
async def get_gen_table_all_services(cls, query_db: AsyncSession): async def get_gen_table_all_services(cls, auth: AuthSchema) -> list[GenTableSchema]:
""" """
获取所有业务表信息service 获取所有业务表信息service
:param query_db: orm对象 :param auth: 认证信息
:return: 所有业务表信息 :return: 所有业务表信息
""" """
gen_table_all = await GenTableDao.get_gen_table_all(query_db) # 确保db是AsyncSession类型
result = [GenTableModel(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_table_all)] if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_all = await gen_table_dao.get_gen_table_all(auth.db)
result = [GenTableSchema(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_table_all)]
return result return result
@classmethod @classmethod
async def create_table_services(cls, query_db: AsyncSession, sql: str, current_user: CurrentUserModel): async def create_table_services(cls, auth: AuthSchema, sql: str, current_user: UserOutSchema) -> SuccessResponse:
""" """
创建表结构service 创建表结构service
:param query_db: orm对象 :param auth: 认证信息
:param sql: 建表语句 :param sql: 建表语句
:param current_user: 当前用户信息对象 :param current_user: 当前用户信息对象
:return: 创建表结构结果 :return: 创建表结构结果
""" """
sql_statements = sqlglot_parse(sql, dialect=DataBaseConfig.sqlglot_parse_dialect) # 移除sqlglot相关代码,因为导入失败
if cls.__is_valid_create_table(sql_statements): raise CustomException(msg='建表功能暂不可用')
try:
table_names = cls.__get_table_names(sql_statements)
await GenTableDao.create_table_by_sql_dao(query_db, sql_statements)
gen_table_list = await cls.get_gen_db_table_list_by_name_services(query_db, table_names)
await cls.import_gen_table_services(query_db, gen_table_list, current_user)
return CrudResponseModel(is_success=True, message='创建表结构成功')
except Exception as e:
raise ServiceException(message=f'创建表结构异常,详细错误信息:{str(e)}')
else:
raise ServiceException(message='建表语句不合法')
@classmethod @classmethod
def __is_valid_create_table(cls, sql_statements: List[Expression]): async def preview_code_services(cls, auth: AuthSchema, table_id: int) -> dict[Any, Any]:
"""
校验sql语句是否为合法的建表语句
:param sql_statements: sql语句的ast列表
:return: 校验结果
"""
validate_create = [isinstance(sql_statement, Create) for sql_statement in sql_statements]
validate_forbidden_keywords = [
isinstance(
sql_statement,
(Add, Alter, Delete, Drop, Insert, TruncateTable, Update),
)
for sql_statement in sql_statements
]
if not any(validate_create) or any(validate_forbidden_keywords):
return False
return True
@classmethod
def __get_table_names(cls, sql_statements: List[Expression]):
"""
获取sql语句中所有的建表表名
:param sql_statements: sql语句的ast列表
:return: 建表表名列表
"""
table_names = []
for sql_statement in sql_statements:
if isinstance(sql_statement, Create):
table_names.append(sql_statement.find(Table).name)
return table_names
@classmethod
async def preview_code_services(cls, query_db: AsyncSession, table_id: int):
""" """
预览代码service 预览代码service
:param query_db: orm对象 :param auth: 认证信息
:param table_id: 业务表id :param table_id: 业务表id
:return: 预览数据列表 :return: 预览数据列表
""" """
gen_table = GenTableModel( gen_table = await cls.get_gen_table_by_id_services(auth, table_id)
**CamelCaseUtil.transform_result(await GenTableDao.get_gen_table_by_id(query_db, table_id)) await cls.set_sub_table(auth, gen_table)
) await cls._set_pk_column(gen_table)
await cls.set_sub_table(query_db, gen_table)
await cls.set_pk_column(gen_table)
env = TemplateInitializer.init_jinja2() env = TemplateInitializer.init_jinja2()
context = TemplateUtils.prepare_context(gen_table) context = TemplateUtils.prepare_context(gen_table)
template_list = TemplateUtils.get_template_list(gen_table.tpl_category, gen_table.tpl_web_type) template_list = TemplateUtils.get_template_list(
gen_table.tpl_category or "",
gen_table.tpl_web_type or ""
)
preview_code_result = {} preview_code_result = {}
for template in template_list: for template in template_list:
render_content = env.get_template(template).render(**context) render_content = env.get_template(template).render(**context)
@@ -270,36 +315,35 @@ class GenTableService:
return preview_code_result return preview_code_result
@classmethod @classmethod
async def generate_code_services(cls, query_db: AsyncSession, table_name: str): async def generate_code_services(cls, auth: AuthSchema, table_name: str) -> SuccessResponse:
""" """
生成代码至指定路径service 生成代码至指定路径service
:param query_db: orm对象 :param auth: 认证信息
:param table_name: 业务表名称 :param table_name: 业务表名称
:return: 生成代码结果 :return: 生成代码结果
""" """
env = TemplateInitializer.init_jinja2() env = TemplateInitializer.init_jinja2()
render_info = await cls.__get_gen_render_info(query_db, table_name) render_info = await cls.__get_gen_render_info(auth, table_name)
for template in render_info[0]: for template in render_info[0]:
try: try:
render_content = env.get_template(template).render(**render_info[2]) render_content = env.get_template(template).render(**render_info[2])
gen_path = cls.__get_gen_path(render_info[3], template) gen_path = cls.__get_gen_path(render_info[3], template)
os.makedirs(os.path.dirname(gen_path), exist_ok=True) if gen_path:
with open(gen_path, 'w', encoding='utf-8') as f: os.makedirs(os.path.dirname(gen_path), exist_ok=True)
f.write(render_content) with open(gen_path, 'w', encoding='utf-8') as f:
f.write(render_content)
except Exception as e: except Exception as e:
raise ServiceException( raise CustomException(msg=f'渲染模板失败,表名:{render_info[3].table_name},详细错误信息:{str(e)}')
message=f'渲染模板失败,表名:{render_info[3].table_name},详细错误信息:{str(e)}'
)
return CrudResponseModel(is_success=True, message='生成代码成功') return SuccessResponse(msg='生成代码成功')
@classmethod @classmethod
async def batch_gen_code_services(cls, query_db: AsyncSession, table_names: List[str]): async def batch_gen_code_services(cls, auth: AuthSchema, table_names: List[str]) -> bytes:
""" """
批量生成代码service 批量生成代码service
:param query_db: orm对象 :param auth: 认证信息
:param table_names: 业务表名称组 :param table_names: 业务表名称组
:return: 下载代码结果 :return: 下载代码结果
""" """
@@ -307,7 +351,7 @@ class GenTableService:
with zipfile.ZipFile(zip_buffer, 'w', zipfile.ZIP_DEFLATED) as zip_file: with zipfile.ZipFile(zip_buffer, 'w', zipfile.ZIP_DEFLATED) as zip_file:
for table_name in table_names: for table_name in table_names:
env = TemplateInitializer.init_jinja2() env = TemplateInitializer.init_jinja2()
render_info = await cls.__get_gen_render_info(query_db, table_name) render_info = await cls.__get_gen_render_info(auth, table_name)
for template_file, output_file in zip(render_info[0], render_info[1]): for template_file, output_file in zip(render_info[0], render_info[1]):
render_content = env.get_template(template_file).render(**render_info[2]) render_content = env.get_template(template_file).render(**render_info[2])
zip_file.writestr(output_file, render_content) zip_file.writestr(output_file, render_content)
@@ -317,27 +361,37 @@ class GenTableService:
return zip_data return zip_data
@classmethod @classmethod
async def __get_gen_render_info(cls, query_db: AsyncSession, table_name: str): async def __get_gen_render_info(cls, auth: AuthSchema, table_name: str) -> list[Any]:
""" """
获取生成代码渲染模板相关信息 获取生成代码渲染模板相关信息
:param query_db: orm对象 :param auth: 认证信息
:param table_name: 业务表名称 :param table_name: 业务表名称
:return: 生成代码渲染模板相关信息 :return: 生成代码渲染模板相关信息
""" """
gen_table = GenTableModel( # 确保db是AsyncSession类型
**CamelCaseUtil.transform_result(await GenTableDao.get_gen_table_by_name(query_db, table_name)) if not isinstance(auth.db, AsyncSession):
) raise CustomException(msg='数据库会话类型不正确')
await cls.set_sub_table(query_db, gen_table)
await cls.set_pk_column(gen_table)
context = TemplateUtils.prepare_context(gen_table)
template_list = TemplateUtils.get_template_list(gen_table.tpl_category, gen_table.tpl_web_type)
output_files = [TemplateUtils.get_file_name(template, gen_table) for template in template_list]
return [template_list, output_files, context, gen_table] gen_table_dao = GenTableDao(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_name(auth.db, table_name)
if gen_table:
gen_table_schema = GenTableSchema(**CamelCaseUtil.transform_result(gen_table))
await cls.set_sub_table(auth, gen_table_schema)
await cls._set_pk_column(gen_table_schema)
context = TemplateUtils.prepare_context(gen_table_schema)
template_list = TemplateUtils.get_template_list(
gen_table_schema.tpl_category or "",
gen_table_schema.tpl_web_type or ""
)
output_files = [TemplateUtils.get_file_name([template], gen_table_schema)[0] for template in template_list]
return [template_list, output_files, context, gen_table_schema]
else:
raise CustomException(msg=f'业务表 {table_name} 不存在')
@classmethod @classmethod
def __get_gen_path(cls, gen_table: GenTableModel, template: str): def __get_gen_path(cls, gen_table: GenTableSchema, template: str) -> Optional[str]:
""" """
根据GenTableModel对象和模板名称生成路径 根据GenTableModel对象和模板名称生成路径
@@ -345,99 +399,129 @@ class GenTableService:
:param template: 模板名称 :param template: 模板名称
:return: 生成的路径 :return: 生成的路径
""" """
gen_path = gen_table.gen_path try:
if gen_path == '/': gen_path = gen_table.gen_path or ""
return os.path.join(os.getcwd(), GenConfig.GEN_PATH, TemplateUtils.get_file_name(template, gen_table)) if gen_path == '/':
else: file_name = TemplateUtils.get_file_name([template], gen_table)[0]
return os.path.join(gen_path, TemplateUtils.get_file_name(template, gen_table)) return os.path.join(os.getcwd(), GEN_PATH, file_name)
else:
file_name = TemplateUtils.get_file_name([template], gen_table)[0]
return os.path.join(gen_path, file_name)
except Exception:
return None
@classmethod @classmethod
async def sync_db_services(cls, query_db: AsyncSession, table_name: str): async def sync_db_services(cls, auth: AuthSchema, table_name: str) -> SuccessResponse:
""" """
同步数据库service 同步数据库service
:param query_db: orm对象 :param auth: 认证信息
:param table_name: 业务表名称 :param table_name: 业务表名称
:return: 同步数据库结果 :return: 同步数据库结果
""" """
gen_table = await GenTableDao.get_gen_table_by_name(query_db, table_name) # 确保db是AsyncSession类型
table = GenTableModel(**CamelCaseUtil.transform_result(gen_table)) if not isinstance(auth.db, AsyncSession):
table_columns = table.columns raise CustomException(msg='数据库会话类型不正确')
table_column_map = {column.column_name: column for column in table_columns}
query_db_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name(query_db, table_name) gen_table_dao = GenTableDao(auth=auth)
db_table_columns = [ gen_table_column_dao = GenTableColumnDao(auth=auth)
GenTableColumnModel(**column) for column in CamelCaseUtil.transform_result(query_db_table_columns)
] gen_table = await gen_table_dao.get_gen_table_by_name(auth.db, table_name)
if not db_table_columns: if gen_table:
raise ServiceException('同步数据失败,原表结构不存在') table = GenTableSchema(**CamelCaseUtil.transform_result(gen_table))
db_table_column_names = [column.column_name for column in db_table_columns] table_columns = table.columns or [] # 确保不为None
try: table_column_map = {column.column_name: column for column in table_columns}
for column in db_table_columns: query_db_table_columns = await gen_table_column_dao.get_gen_db_table_columns_by_name(auth.db, table_name)
GenUtils.init_column_field(column, table) db_table_columns = [
if column.column_name in table_column_map: GenTableColumnSchema(**column) for column in CamelCaseUtil.transform_result(query_db_table_columns)
prev_column = table_column_map[column.column_name] ]
column.column_id = prev_column.column_id if not db_table_columns:
if column.list: raise CustomException('同步数据失败,原表结构不存在')
column.dict_type = prev_column.dict_type db_table_column_names = [column.column_name for column in db_table_columns]
column.query_type = prev_column.query_type try:
if ( for column in db_table_columns:
prev_column.is_required != '' GenUtils.init_column_field(column, table)
and not column.pk if column.column_name in table_column_map:
and (column.insert or column.edit) prev_column = table_column_map[column.column_name]
and (column.usable_column or column.super_column) column.column_id = prev_column.column_id
): if getattr(column, 'list', False): # 使用getattr安全访问属性
column.is_required = prev_column.is_required column.dict_type = prev_column.dict_type
column.html_type = prev_column.html_type column.query_type = prev_column.query_type
await GenTableColumnDao.edit_gen_table_column_dao(query_db, column.model_dump(by_alias=True)) if (
else: prev_column.is_required != ''
await GenTableColumnDao.add_gen_table_column_dao(query_db, column) and not column.pk
del_columns = [column for column in table_columns if column.column_name not in db_table_column_names] and (column.insert or column.edit)
if del_columns: and (column.usable_column or column.super_column)
for column in del_columns: ):
await GenTableColumnDao.delete_gen_table_column_by_column_id_dao(query_db, column) column.is_required = prev_column.is_required
await query_db.commit() column.html_type = prev_column.html_type
return CrudResponseModel(is_success=True, message='同步成功') if column.column_id is not None:
except Exception as e: await gen_table_column_dao.update(id=column.column_id, data=column.model_dump(by_alias=True))
await query_db.rollback() else:
raise e await gen_table_column_dao.create(data=column.model_dump(by_alias=True))
del_columns = [column for column in table_columns if column.column_name not in db_table_column_names]
if del_columns:
for column in del_columns:
if column.column_id is not None:
await gen_table_column_dao.delete(ids=[column.column_id])
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return SuccessResponse(msg='同步成功')
except Exception as e:
if isinstance(auth.db, AsyncSession):
try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'同步失败: {str(e)}')
else:
raise CustomException('业务表不存在')
@classmethod @classmethod
async def set_sub_table(cls, query_db: AsyncSession, gen_table: GenTableModel): async def set_sub_table(cls, auth: AuthSchema, gen_table: GenTableSchema) -> None:
""" """
设置主子表信息 设置主子表信息
:param query_db: orm对象 :param auth: 认证信息
:param gen_table: 业务表信息 :param gen_table: 业务表信息
:return: :return:
""" """
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
if gen_table.sub_table_name: if gen_table.sub_table_name:
sub_table = await GenTableDao.get_gen_table_by_name(query_db, gen_table.sub_table_name) gen_table_dao = GenTableDao(auth=auth)
gen_table.sub_table = GenTableModel(**CamelCaseUtil.transform_result(sub_table)) sub_table = await gen_table_dao.get_gen_table_by_name(auth.db, gen_table.sub_table_name)
if sub_table:
gen_table.sub_table = GenTableSchema(**CamelCaseUtil.transform_result(sub_table))
@classmethod @classmethod
async def set_pk_column(cls, gen_table: GenTableModel): async def _set_pk_column(cls, gen_table: GenTableSchema) -> None:
""" """
设置主键列信息 设置主键列信息
:param gen_table: 业务表信息 :param gen_table: 业务表信息
:return: :return:
""" """
for column in gen_table.columns: if gen_table.columns:
if column.pk: for column in gen_table.columns:
gen_table.pk_column = column
break
if gen_table.pk_column is None:
gen_table.pk_column = gen_table.columns[0]
if gen_table.tpl_category == GenConstant.TPL_SUB:
for column in gen_table.sub_table.columns:
if column.pk: if column.pk:
gen_table.sub_table.pk_column = column gen_table.pk_column = column
break break
if gen_table.sub_table.columns is None: if gen_table.pk_column is None and gen_table.columns:
gen_table.pk_column = gen_table.columns[0]
if gen_table.tpl_category == GenConstant.TPL_SUB and gen_table.sub_table:
if gen_table.sub_table.columns:
for column in gen_table.sub_table.columns:
if column.pk:
gen_table.sub_table.pk_column = column
break
if gen_table.sub_table.pk_column is None and gen_table.sub_table.columns:
gen_table.sub_table.pk_column = gen_table.sub_table.columns[0] gen_table.sub_table.pk_column = gen_table.sub_table.columns[0]
@classmethod @classmethod
async def set_table_from_options(cls, gen_table: GenTableModel): async def set_table_from_options(cls, gen_table: GenTableSchema) -> GenTableSchema:
""" """
设置代码生成其他选项值 设置代码生成其他选项值
@@ -455,26 +539,30 @@ class GenTableService:
return gen_table return gen_table
@classmethod @classmethod
async def validate_edit(cls, edit_gen_table: EditGenTableModel): async def validate_edit(cls, edit_gen_table: EditGenTableSchema):
""" """
编辑保存参数校验 编辑保存参数校验
:param edit_gen_table: 编辑业务表对象 :param edit_gen_table: 编辑业务表对象
""" """
if edit_gen_table.tpl_category == GenConstant.TPL_TREE: if edit_gen_table.tpl_category == GenConstant.TPL_TREE:
# 检查params是否为None
if edit_gen_table.params is None:
raise CustomException(msg='树表参数不能为空')
params_obj = edit_gen_table.params.model_dump(by_alias=True) params_obj = edit_gen_table.params.model_dump(by_alias=True)
if GenConstant.TREE_CODE not in params_obj: if GenConstant.TREE_CODE not in params_obj:
raise ServiceException(message='树编码字段不能为空') raise CustomException(msg='树编码字段不能为空')
elif GenConstant.TREE_PARENT_CODE not in params_obj: elif GenConstant.TREE_PARENT_CODE not in params_obj:
raise ServiceException(message='树父编码字段不能为空') raise CustomException(msg='树父编码字段不能为空')
elif GenConstant.TREE_NAME not in params_obj: elif GenConstant.TREE_NAME not in params_obj:
raise ServiceException(message='树名称字段不能为空') raise CustomException(msg='树名称字段不能为空')
elif edit_gen_table.tpl_category == GenConstant.TPL_SUB: elif edit_gen_table.tpl_category == GenConstant.TPL_SUB:
if not edit_gen_table.sub_table_name: if not edit_gen_table.sub_table_name:
raise ServiceException(message='关联子表的表名不能为空') raise CustomException(msg='关联子表的表名不能为空')
elif not edit_gen_table.sub_table_fk_name: elif not edit_gen_table.sub_table_fk_name:
raise ServiceException(message='子表关联的外键名不能为空') raise CustomException(msg='子表关联的外键名不能为空')
class GenTableColumnService: class GenTableColumnService:
@@ -483,17 +571,22 @@ class GenTableColumnService:
""" """
@classmethod @classmethod
async def get_gen_table_column_list_by_table_id_services(cls, query_db: AsyncSession, table_id: int): async def get_gen_table_column_list_by_table_id_services(cls, auth: AuthSchema, table_id: int):
""" """
获取业务表字段列表信息service 获取业务表字段列表信息service
:param query_db: orm对象 :param auth: 认证信息
:param table_id: 业务表格id :param table_id: 业务表格id
:return: 业务表字段列表信息对象 :return: 业务表字段列表信息对象
""" """
gen_table_column_list_result = await GenTableColumnDao.get_gen_table_column_list_by_table_id(query_db, table_id) # 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_column_dao = GenTableColumnDao(auth=auth)
gen_table_column_list_result = await gen_table_column_dao.get_gen_table_column_list_by_table_id(auth.db, table_id)
return [ return [
GenTableColumnModel(**gen_table_column) GenTableColumnSchema(**gen_table_column)
for gen_table_column in CamelCaseUtil.transform_result(gen_table_column_list_result) for gen_table_column in CamelCaseUtil.transform_result(gen_table_column_list_result)
] ]
@@ -15,7 +15,7 @@ from utils.page_util import PageUtil, PageResponseModel
from utils.common_util import CamelCaseUtil from utils.common_util import CamelCaseUtil
class {{ tableName|snake_to_pascal_case }}Dao: class {{ tableName|snake_to_pascal_case }}CRUD:
@classmethod @classmethod
async def get_by_id(cls, db: AsyncSession, {{ tableName }}_id: int) -> {{ tableName|snake_to_pascal_case }}: async def get_by_id(cls, db: AsyncSession, {{ tableName }}_id: int) -> {{ tableName|snake_to_pascal_case }}:
@@ -6,7 +6,7 @@ from config.database import BaseMixin, Base
from sqlalchemy.orm import relationship from sqlalchemy.orm import relationship
{% endif %} {% endif %}
class {{ tableName|snake_to_pascal_case }}(Base, BaseMixin): class {{ tableName|snake_to_pascal_case }}Model(Base, BaseMixin):
""" """
{{ functionName }}表 {{ functionName }}表
""" """
@@ -0,0 +1,35 @@
# -*- coding: utf-8 -*-
from datetime import datetime
from typing import Optional
from fastapi import Query
from app.core.validator import DateTimeStr
class {{ tableName|snake_to_pascal_case }}QueryParam:
"""示例查询参数"""
def __init__(
self,
name: Optional[str] = Query(None, description="名称"),
status: Optional[bool] = Query(None, description="是否启用"),
creator: Optional[int] = Query(None, description="创建人"),
start_time: Optional[DateTimeStr] = Query(None, description="开始时间", example="2023-01-01 00:00:00"),
end_time: Optional[DateTimeStr] = Query(None, description="结束时间", example="2023-12-31 23:59:59"),
) -> None:
super().__init__()
# 模糊查询字段
self.name = ("like", name)
# 精确查询字段
self.creator_id = creator
self.status = status
# 时间范围查询
if start_time and end_time:
start_datetime = datetime.strptime(str(start_time), '%Y-%m-%d %H:%M:%S')
end_datetime = datetime.strptime(str(end_time), '%Y-%m-%d %H:%M:%S')
self.created_at = ("between", (start_datetime, end_datetime))
@@ -37,7 +37,7 @@ class {{ tableName|snake_to_pascal_case }}PageModel({{ tableName|snake_to_pascal
{% if subTable %} {% if subTable %}
class {{ subTable.table_name | snake_to_pascal_case }}Model(BaseModel): class {{ subTable.table_name | snake_to_pascal_case }}Schema(BaseModel):
""" """
{{ subTable.function_name }}表对应pydantic模型 {{ subTable.function_name }}表对应pydantic模型
""" """
@@ -13,7 +13,7 @@ from {{ packageName }}.entity.vo.{{ tableName }}_vo import {{ tableName|snake_to
class {{ tableName|snake_to_pascal_case }}Service: class {{ tableName|snake_to_pascal_case }}Service:
""" """
用户管理模块服务层 {{ tableName|snake_to_pascal_case }}管理模块服务层
""" """
@classmethod @classmethod
@@ -6,12 +6,12 @@ from fastapi.responses import JSONResponse, StreamingResponse
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .param import JobQueryParams, JobLogQueryParams from .param import JobQueryParam, JobLogQueryParam
from .service import JobService, JobLogService from .service import JobService, JobLogService
from .schema import ( from .schema import (
JobCreateSchema, JobCreateSchema,
@@ -33,12 +33,12 @@ async def get_obj_detail_controller(
@JobRouter.get("/list", summary="查询定时任务", description="查询定时任务") @JobRouter.get("/list", summary="查询定时任务", description="查询定时任务")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: JobQueryParams = Depends(), search: JobQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await JobService.get_job_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await JobService.get_job_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询定时任务列表成功") logger.info(f"查询定时任务列表成功")
return SuccessResponse(data=result_dict, msg="查询定时任务列表成功") return SuccessResponse(data=result_dict, msg="查询定时任务列表成功")
@@ -72,7 +72,7 @@ async def delete_obj_controller(
@JobRouter.post('/export', summary="导出定时任务", description="导出定时任务") @JobRouter.post('/export', summary="导出定时任务", description="导出定时任务")
async def export_obj_list_controller( async def export_obj_list_controller(
search: JobQueryParams = Depends(), search: JobQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -143,12 +143,12 @@ async def get_job_log_detail_controller(
@JobRouter.get("/log/list", summary="查询定时任务日志", description="查询定时任务日志") @JobRouter.get("/log/list", summary="查询定时任务日志", description="查询定时任务日志")
async def get_job_log_list_controller( async def get_job_log_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: JobLogQueryParams = Depends(), search: JobLogQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await JobLogService.get_job_log_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await JobLogService.get_job_log_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size) result_dict = await PaginationService.paginate(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size)
logger.info(f"查询定时任务日志列表成功") logger.info(f"查询定时任务日志列表成功")
return SuccessResponse(data=result_dict, msg="查询定时任务日志列表成功") return SuccessResponse(data=result_dict, msg="查询定时任务日志列表成功")
@@ -174,7 +174,7 @@ async def clear_job_log_controller(
@JobRouter.post('/log/export', summary="导出定时任务日志", description="导出定时任务日志") @JobRouter.post('/log/export', summary="导出定时任务日志", description="导出定时任务日志")
async def export_job_log_list_controller( async def export_job_log_list_controller(
search: JobLogQueryParams = Depends(), search: JobLogQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["monitor:job:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -7,7 +7,7 @@ from datetime import datetime
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class JobQueryParams: class JobQueryParam:
"""定时任务查询参数""" """定时任务查询参数"""
def __init__( def __init__(
@@ -34,7 +34,7 @@ class JobQueryParams:
self.created_at = ("between", (start_datetime, end_datetime)) self.created_at = ("between", (start_datetime, end_datetime))
class JobLogQueryParams: class JobLogQueryParam:
"""定时任务查询参数""" """定时任务查询参数"""
def __init__( def __init__(
@@ -8,7 +8,7 @@ from app.utils.cron_util import CronUtil
from app.utils.excel_util import ExcelUtil from app.utils.excel_util import ExcelUtil
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .schema import JobCreateSchema, JobUpdateSchema, JobOutSchema, JobLogOutSchema from .schema import JobCreateSchema, JobUpdateSchema, JobOutSchema, JobLogOutSchema
from .param import JobQueryParams, JobLogQueryParams from .param import JobQueryParam, JobLogQueryParam
from .crud import JobCRUD, JobLogCRUD from .crud import JobCRUD, JobLogCRUD
@@ -23,7 +23,7 @@ class JobService:
return JobOutSchema.model_validate(obj).model_dump() return JobOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_job_list_service(cls, auth: AuthSchema, search: JobQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_job_list_service(cls, auth: AuthSchema, search: JobQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
obj_list = await JobCRUD(auth).get_obj_list_crud(search=search.__dict__, order_by=order_by) obj_list = await JobCRUD(auth).get_obj_list_crud(search=search.__dict__, order_by=order_by)
@@ -130,7 +130,7 @@ class JobLogService:
return JobLogOutSchema.model_validate(obj).model_dump() return JobLogOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_job_log_list_service(cls, auth: AuthSchema, search: JobLogQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_job_log_list_service(cls, auth: AuthSchema, search: JobLogQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
"""获取定时任务日志列表""" """获取定时任务日志列表"""
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
@@ -7,10 +7,10 @@ from redis.asyncio.client import Redis
from app.common.request import PaginationService from app.common.request import PaginationService
from app.common.response import SuccessResponse,ErrorResponse from app.common.response import SuccessResponse,ErrorResponse
from app.core.dependencies import AuthPermission, redis_getter from app.core.dependencies import AuthPermission, redis_getter
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.logger import logger from app.core.logger import logger
from .param import OnlineQueryParams from .param import OnlineQueryParam
from .service import OnlineService from .service import OnlineService
@@ -25,12 +25,12 @@ OnlineRouter = APIRouter(route_class=OperationLogRoute, prefix="/online", tags=[
) )
async def get_online_list_controller( async def get_online_list_controller(
redis: Redis = Depends(redis_getter), redis: Redis = Depends(redis_getter),
paging_query: PaginationQueryParams = Depends(), paging_query: PaginationQueryParam = Depends(),
search: OnlineQueryParams = Depends() search: OnlineQueryParam = Depends()
)->JSONResponse: )->JSONResponse:
# 获取全量数据 # 获取全量数据
result_dict_list = await OnlineService.get_online_list_service(redis=redis, search=search) result_dict_list = await OnlineService.get_online_list_service(redis=redis, search=search)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= paging_query.page_no, page_size = paging_query.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= paging_query.page_no, page_size = paging_query.page_size)
logger.info('获取成功') logger.info('获取成功')
return SuccessResponse(data=result_dict,msg='获取成功') return SuccessResponse(data=result_dict,msg='获取成功')
@@ -4,7 +4,7 @@ from typing import Optional
from fastapi import Query from fastapi import Query
class OnlineQueryParams: class OnlineQueryParam:
"""在线用户查询参数""" """在线用户查询参数"""
def __init__( def __init__(
@@ -8,14 +8,14 @@ from app.common.enums import RedisInitKeyConfig
from app.core.redis_crud import RedisCURD from app.core.redis_crud import RedisCURD
from app.core.security import decode_access_token from app.core.security import decode_access_token
from app.core.logger import logger from app.core.logger import logger
from .param import OnlineQueryParams from .param import OnlineQueryParam
from .schema import OnlineOutSchema from .schema import OnlineOutSchema
class OnlineService: class OnlineService:
"""在线用户管理模块服务层""" """在线用户管理模块服务层"""
@classmethod @classmethod
async def get_online_list_service(cls, redis: Redis, search: Optional[OnlineQueryParams] = None) -> List[Dict]: async def get_online_list_service(cls, redis: Redis, search: Optional[OnlineQueryParam] = None) -> List[Dict]:
""" """
获取在线用户列表信息(支持分页和搜索) 获取在线用户列表信息(支持分页和搜索)
""" """
@@ -64,7 +64,7 @@ class OnlineService:
@staticmethod @staticmethod
def _match_search_conditions(online_info: Dict, search: Optional[OnlineQueryParams]) -> bool: def _match_search_conditions(online_info: Dict, search: Optional[OnlineQueryParam]) -> bool:
"""检查是否匹配搜索条件""" """检查是否匹配搜索条件"""
if not search: if not search:
return True return True
@@ -7,12 +7,12 @@ from typing import List, Optional
from app.common.request import PaginationService from app.common.request import PaginationService
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.logger import logger from app.core.logger import logger
from ...module_system.auth.schema import AuthSchema from ...module_system.auth.schema import AuthSchema
from .param import ResourceQueryParams from .param import ResourceQueryParam
from .schema import ( from .schema import (
ResourceSearchSchema, ResourceSearchSchema,
ResourceMoveSchema, ResourceMoveSchema,
@@ -15,7 +15,7 @@ class ResourceType(Enum):
OTHER = "other" # 其他 OTHER = "other" # 其他
class ResourceQueryParams(BaseModel): class ResourceQueryParam(BaseModel):
"""资源查询参数模型""" """资源查询参数模型"""
path: Optional[str] = Field(None, description="文件路径") path: Optional[str] = Field(None, description="文件路径")
keyword: Optional[str] = Field(None, description="关键词搜索") keyword: Optional[str] = Field(None, description="关键词搜索")
@@ -28,7 +28,7 @@ if not MAGIC_AVAILABLE:
from app.utils.excel_util import ExcelUtil from app.utils.excel_util import ExcelUtil
from app.config.setting import settings from app.config.setting import settings
from ...module_system.auth.schema import AuthSchema from ...module_system.auth.schema import AuthSchema
from .param import ResourceQueryParams from .param import ResourceQueryParam
from .schema import ( from .schema import (
ResourceItemSchema, ResourceItemSchema,
ResourceDirectorySchema, ResourceDirectorySchema,
@@ -9,9 +9,9 @@ from app.core.router_class import OperationLogRoute
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.logger import logger from app.core.logger import logger
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import DeptQueryParams from .param import DeptQueryParam
from .service import DeptService from .service import DeptService
from .schema import ( from .schema import (
DeptCreateSchema, DeptCreateSchema,
@@ -24,7 +24,7 @@ DeptRouter = APIRouter(route_class=OperationLogRoute, prefix="/dept", tags=["部
@DeptRouter.get("/tree", summary="查询部门树", description="查询部门树") @DeptRouter.get("/tree", summary="查询部门树", description="查询部门树")
async def get_dept_tree_controller( async def get_dept_tree_controller(
search: DeptQueryParams = Depends(), search: DeptQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:dept:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:dept:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await DeptService.get_dept_tree_service(search=search, auth=auth) result_dict_list = await DeptService.get_dept_tree_service(search=search, auth=auth)
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class DeptQueryParams: class DeptQueryParam:
"""部门管理查询参数""" """部门管理查询参数"""
def __init__( def __init__(
@@ -13,7 +13,7 @@ from app.utils.common_util import (
) )
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .crud import DeptCRUD from .crud import DeptCRUD
from .param import DeptQueryParams from .param import DeptQueryParam
from .schema import ( from .schema import (
DeptCreateSchema, DeptCreateSchema,
DeptUpdateSchema, DeptUpdateSchema,
@@ -39,7 +39,7 @@ class DeptService:
return DeptOutSchema.model_validate(dept).model_dump() return DeptOutSchema.model_validate(dept).model_dump()
@classmethod @classmethod
async def get_dept_tree_service(cls, auth: AuthSchema, search: DeptQueryParams, order_by: List[Dict] = None) -> List[Dict]: async def get_dept_tree_service(cls, auth: AuthSchema, search: DeptQueryParam, order_by: List[Dict] = None) -> List[Dict]:
""" """
获取部门树形列表service 获取部门树形列表service
@@ -6,7 +6,7 @@ from fastapi.responses import JSONResponse, StreamingResponse
from redis.asyncio.client import Redis from redis.asyncio.client import Redis
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.dependencies import AuthPermission, redis_getter from app.core.dependencies import AuthPermission, redis_getter
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
@@ -14,7 +14,7 @@ from app.core.logger import logger
from app.common.request import PaginationService from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import DictTypeQueryParams, DictDataQueryParams from .param import DictTypeQueryParam, DictDataQueryParam
from .service import DictTypeService, DictDataService from .service import DictTypeService, DictDataService
from .schema import ( from .schema import (
DictTypeCreateSchema, DictTypeCreateSchema,
@@ -37,17 +37,17 @@ async def get_type_detail_controller(
@DictRouter.get("/type/list", summary="查询字典类型", description="查询字典类型") @DictRouter.get("/type/list", summary="查询字典类型", description="查询字典类型")
async def get_type_list_controller( async def get_type_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: DictTypeQueryParams = Depends(), search: DictTypeQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_type:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_type:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await DictTypeService.get_obj_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await DictTypeService.get_obj_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询字典类型列表成功") logger.info(f"查询字典类型列表成功")
return SuccessResponse(data=result_dict, msg="查询字典类型列表成功") return SuccessResponse(data=result_dict, msg="查询字典类型列表成功")
@DictRouter.get("/type/optionselect", summary="获取全部字典类型", description="获取全部字典类型") @DictRouter.get("/type/optionselect", summary="获取全部字典类型", description="获取全部字典类型")
async def get_type_list_controller( async def get_type_loptionselect_controller(
auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_type:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_type:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await DictTypeService.get_obj_list_service(auth=auth) result_dict_list = await DictTypeService.get_obj_list_service(auth=auth)
@@ -96,7 +96,7 @@ async def batch_set_available_obj_controller(
@DictRouter.post('/type/export', summary="导出字典类型", description="导出字典类型") @DictRouter.post('/type/export', summary="导出字典类型", description="导出字典类型")
async def export_type_list_controller( async def export_type_list_controller(
search: DictTypeQueryParams = Depends(), search: DictTypeQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_type:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_type:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -123,12 +123,12 @@ async def get_data_detail_controller(
@DictRouter.get("/data/list", summary="查询字典数据", description="查询字典数据") @DictRouter.get("/data/list", summary="查询字典数据", description="查询字典数据")
async def get_data_list_controller( async def get_data_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: DictDataQueryParams = Depends(), search: DictDataQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_data:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_data:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await DictDataService.get_obj_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await DictDataService.get_obj_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询字典数据列表成功") logger.info(f"查询字典数据列表成功")
return SuccessResponse(data=result_dict, msg="查询字典数据列表成功") return SuccessResponse(data=result_dict, msg="查询字典数据列表成功")
@@ -174,8 +174,8 @@ async def batch_set_available_obj_controller(
@DictRouter.post('/data/export', summary="导出字典数据", description="导出字典数据") @DictRouter.post('/data/export', summary="导出字典数据", description="导出字典数据")
async def export_data_list_controller( async def export_data_list_controller(
search: DictDataQueryParams = Depends(), search: DictDataQueryParam = Depends(),
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_data:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:dict_data:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -7,7 +7,7 @@ from datetime import datetime
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class DictTypeQueryParams: class DictTypeQueryParam:
"""字典类型查询参数""" """字典类型查询参数"""
def __init__( def __init__(
@@ -36,7 +36,7 @@ class DictTypeQueryParams:
self.created_at = ("between", (start_datetime, end_datetime)) self.created_at = ("between", (start_datetime, end_datetime))
class DictDataQueryParams: class DictDataQueryParam:
"""字典数据查询参数""" """字典数据查询参数"""
def __init__( def __init__(
@@ -14,7 +14,7 @@ from app.core.exceptions import CustomException
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from app.api.v1.module_system.dict.schema import DictDataCreateSchema,DictDataOutSchema,DictDataUpdateSchema,DictTypeCreateSchema,DictTypeOutSchema,DictTypeUpdateSchema from app.api.v1.module_system.dict.schema import DictDataCreateSchema,DictDataOutSchema,DictDataUpdateSchema,DictTypeCreateSchema,DictTypeOutSchema,DictTypeUpdateSchema
from app.api.v1.module_system.dict.param import DictDataQueryParams, DictTypeQueryParams from app.api.v1.module_system.dict.param import DictDataQueryParam, DictTypeQueryParam
from app.api.v1.module_system.dict.crud import DictDataCRUD, DictTypeCRUD from app.api.v1.module_system.dict.crud import DictDataCRUD, DictTypeCRUD
@@ -29,7 +29,7 @@ class DictTypeService:
return DictTypeOutSchema.model_validate(obj).model_dump() return DictTypeOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_obj_list_service(cls, auth: AuthSchema, search: DictTypeQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_obj_list_service(cls, auth: AuthSchema, search: DictTypeQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
obj_list = None obj_list = None
@@ -179,7 +179,7 @@ class DictDataService:
return DictDataOutSchema.model_validate(obj).model_dump() return DictDataOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_obj_list_service(cls, auth: AuthSchema, search: DictDataQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_obj_list_service(cls, auth: AuthSchema, search: DictDataQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
obj_list = await DictDataCRUD(auth).get_obj_list_crud(search=search.__dict__, order_by=order_by) obj_list = await DictDataCRUD(auth).get_obj_list_crud(search=search.__dict__, order_by=order_by)
@@ -8,10 +8,10 @@ from app.common.response import SuccessResponse, StreamResponse
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import OperationLogQueryParams from .param import OperationLogQueryParam
from .service import OperationLogService from .service import OperationLogService
@@ -20,13 +20,13 @@ LogRouter = APIRouter(route_class=OperationLogRoute, prefix="/log", tags=["日
@LogRouter.get("/list", summary="查询日志", description="查询日志") @LogRouter.get("/list", summary="查询日志", description="查询日志")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: OperationLogQueryParams = Depends(), search: OperationLogQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:log:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:log:query"]))
) -> JSONResponse: ) -> JSONResponse:
""" 查询日志 """ """ 查询日志 """
result_dict_list = await OperationLogService.get_log_list_service(search=search, auth=auth, order_by=page.order_by) result_dict_list = await OperationLogService.get_log_list_service(search=search, auth=auth, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询日志成功") logger.info(f"查询日志成功")
return SuccessResponse(data=result_dict, msg="查询日志成功") return SuccessResponse(data=result_dict, msg="查询日志成功")
@@ -55,7 +55,7 @@ async def delete_obj_log_controller(
@LogRouter.post("/export", summary="导出日志", description="导出日志") @LogRouter.post("/export", summary="导出日志", description="导出日志")
async def export_obj_list_controller( async def export_obj_list_controller(
search: OperationLogQueryParams = Depends(), search: OperationLogQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:log:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:log:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
""" 导出日志 """ """ 导出日志 """
@@ -6,7 +6,7 @@ from datetime import datetime
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class OperationLogQueryParams: class OperationLogQueryParam:
"""操作日志查询参数""" """操作日志查询参数"""
def __init__( def __init__(
@@ -5,7 +5,7 @@ from typing import Any, Dict, List
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
from app.utils.excel_util import ExcelUtil from app.utils.excel_util import ExcelUtil
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import OperationLogQueryParams from .param import OperationLogQueryParam
from .crud import OperationLogCRUD from .crud import OperationLogCRUD
from .schema import ( from .schema import (
OperationLogCreateSchema, OperationLogCreateSchema,
@@ -26,7 +26,7 @@ class OperationLogService:
return log_dict return log_dict
@classmethod @classmethod
async def get_log_list_service(cls, auth: AuthSchema, search: OperationLogQueryParams, order_by: List[Dict] = None) -> List[Dict]: async def get_log_list_service(cls, auth: AuthSchema, search: OperationLogQueryParam, order_by: List[Dict] = None) -> List[Dict]:
"""获取日志列表""" """获取日志列表"""
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
@@ -9,7 +9,7 @@ from app.core.base_schema import BatchSetAvailable
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import MenuQueryParams from .param import MenuQueryParam
from .service import MenuService from .service import MenuService
from .schema import ( from .schema import (
MenuCreateSchema, MenuCreateSchema,
@@ -21,7 +21,7 @@ MenuRouter = APIRouter(route_class=OperationLogRoute, prefix="/menu", tags=["菜
@MenuRouter.get("/tree", summary="查询菜单树", description="查询菜单树") @MenuRouter.get("/tree", summary="查询菜单树", description="查询菜单树")
async def get_menu_tree_controller( async def get_menu_tree_controller(
search: MenuQueryParams = Depends(), search: MenuQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:menu:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:menu:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await MenuService.get_menu_tree_service(search=search, auth=auth) result_dict_list = await MenuService.get_menu_tree_service(search=search, auth=auth)
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class MenuQueryParams: class MenuQueryParam:
"""菜单管理查询参数""" """菜单管理查询参数"""
def __init__( def __init__(
@@ -12,7 +12,7 @@ from app.utils.common_util import (
traversal_to_tree traversal_to_tree
) )
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import MenuQueryParams from .param import MenuQueryParam
from .crud import MenuCRUD from .crud import MenuCRUD
from .schema import ( from .schema import (
MenuCreateSchema, MenuCreateSchema,
@@ -33,7 +33,7 @@ class MenuService:
return menu_dict return menu_dict
@classmethod @classmethod
async def get_menu_tree_service(cls, auth: AuthSchema, search: MenuQueryParams, order_by: List[Dict] = None) -> List[Dict]: async def get_menu_tree_service(cls, auth: AuthSchema, search: MenuQueryParam, order_by: List[Dict] = None) -> List[Dict]:
""" """
获取菜单树形列表service 获取菜单树形列表service
@@ -4,7 +4,7 @@ from fastapi import APIRouter, Body, Depends, Path, Query
from fastapi.responses import JSONResponse, StreamingResponse from fastapi.responses import JSONResponse, StreamingResponse
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission, get_current_user from app.core.dependencies import AuthPermission, get_current_user
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
@@ -12,7 +12,7 @@ from app.core.logger import logger
from app.common.request import PaginationService from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import NoticeQueryParams from .param import NoticeQueryParam
from .service import NoticeService from .service import NoticeService
from .schema import ( from .schema import (
NoticeCreateSchema, NoticeCreateSchema,
@@ -33,12 +33,12 @@ async def get_obj_detail_controller(
@NoticeRouter.get("/list", summary="查询公告", description="查询公告") @NoticeRouter.get("/list", summary="查询公告", description="查询公告")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: NoticeQueryParams = Depends(), search: NoticeQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:notice:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:notice:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await NoticeService.get_notice_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await NoticeService.get_notice_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询公告列表成功") logger.info(f"查询公告列表成功")
return SuccessResponse(data=result_dict, msg="查询公告列表成功") return SuccessResponse(data=result_dict, msg="查询公告列表成功")
@@ -81,7 +81,7 @@ async def batch_set_available_obj_controller(
@NoticeRouter.post('/export', summary="导出公告", description="导出公告") @NoticeRouter.post('/export', summary="导出公告", description="导出公告")
async def export_obj_list_controller( async def export_obj_list_controller(
search: NoticeQueryParams = Depends(), search: NoticeQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:notice:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:notice:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -103,6 +103,6 @@ async def get_obj_list_available_controller(
auth: AuthSchema = Depends(get_current_user) auth: AuthSchema = Depends(get_current_user)
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await NoticeService.get_notice_list_available_service(auth=auth) result_dict_list = await NoticeService.get_notice_list_available_service(auth=auth)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list) result_dict = await PaginationService.paginate(data_list= result_dict_list)
logger.info(f"查询已启用公告列表成功") logger.info(f"查询已启用公告列表成功")
return SuccessResponse(data=result_dict, msg="查询已启用公告列表成功") return SuccessResponse(data=result_dict, msg="查询已启用公告列表成功")
@@ -7,7 +7,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class NoticeQueryParams: class NoticeQueryParam:
"""公告通知查询参数""" """公告通知查询参数"""
def __init__( def __init__(
@@ -8,7 +8,7 @@ from app.core.exceptions import CustomException
from app.utils.excel_util import ExcelUtil from app.utils.excel_util import ExcelUtil
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .schema import NoticeCreateSchema, NoticeUpdateSchema, NoticeOutSchema from .schema import NoticeCreateSchema, NoticeUpdateSchema, NoticeOutSchema
from .param import NoticeQueryParams from .param import NoticeQueryParam
from .crud import NoticeCRUD from .crud import NoticeCRUD
@@ -28,7 +28,7 @@ class NoticeService:
return [NoticeOutSchema.model_validate(notice_obj).model_dump() for notice_obj in notice_obj_list] return [NoticeOutSchema.model_validate(notice_obj).model_dump() for notice_obj in notice_obj_list]
@classmethod @classmethod
async def get_notice_list_service(cls, auth: AuthSchema, search: NoticeQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_notice_list_service(cls, auth: AuthSchema, search: NoticeQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
notice_obj_list = await NoticeCRUD(auth).get_list_crud(search=search.__dict__, order_by=order_by) notice_obj_list = await NoticeCRUD(auth).get_list_crud(search=search.__dict__, order_by=order_by)
@@ -8,12 +8,12 @@ from redis.asyncio.client import Redis
from app.common.request import PaginationService from app.common.request import PaginationService
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission, redis_getter from app.core.dependencies import AuthPermission, redis_getter
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import ParamsQueryParams from .param import ParamsQueryParam
from .schema import ParamsCreateSchema, ParamsUpdateSchema from .schema import ParamsCreateSchema, ParamsUpdateSchema
from .service import ParamsService from .service import ParamsService
@@ -53,11 +53,11 @@ async def get_config_value_by_key_controller(
@ParamsRouter.get("/list", summary="获取参数列表", description="获取参数列表") @ParamsRouter.get("/list", summary="获取参数列表", description="获取参数列表")
async def get_obj_list_controller( async def get_obj_list_controller(
auth: AuthSchema = Depends(AuthPermission(permissions=["system:param:query"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:param:query"])),
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: ParamsQueryParams = Depends(), search: ParamsQueryParam = Depends(),
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await ParamsService.get_obj_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await ParamsService.get_obj_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"获取参数列表成功") logger.info(f"获取参数列表成功")
return SuccessResponse(data=result_dict, msg="查询参数列表成功") return SuccessResponse(data=result_dict, msg="查询参数列表成功")
@@ -98,7 +98,7 @@ async def delete_obj_controller(
@ParamsRouter.post('/export', summary="导出参数", description="导出参数") @ParamsRouter.post('/export', summary="导出参数", description="导出参数")
async def export_obj_list_controller( async def export_obj_list_controller(
search: ParamsQueryParams = Depends(), search: ParamsQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:param:export"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:param:export"]))
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class ParamsQueryParams: class ParamsQueryParam:
"""配置管理查询参数""" """配置管理查询参数"""
def __init__( def __init__(
@@ -16,7 +16,7 @@ from app.core.base_schema import UploadResponseSchema
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import ParamsQueryParams from .param import ParamsQueryParam
from .schema import ParamsOutSchema, ParamsUpdateSchema, ParamsCreateSchema, UpdateSystemParamsSchema from .schema import ParamsOutSchema, ParamsUpdateSchema, ParamsCreateSchema, UpdateSystemParamsSchema
from .crud import ParamsCRUD from .crud import ParamsCRUD
@@ -47,7 +47,7 @@ class ParamsService:
return obj.config_value return obj.config_value
@classmethod @classmethod
async def get_obj_list_service(cls, auth: AuthSchema, search: ParamsQueryParams = None, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_obj_list_service(cls, auth: AuthSchema, search: ParamsQueryParam = None, order_by: List[Dict[str, str]] = None) -> List[Dict]:
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
obj_list = None obj_list = None
@@ -6,14 +6,14 @@ from fastapi.responses import JSONResponse, StreamingResponse
from app.common.response import StreamResponse, SuccessResponse from app.common.response import StreamResponse, SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .service import PositionService from .service import PositionService
from .param import PositionQueryParams from .param import PositionQueryParam
from .schema import ( from .schema import (
PositionCreateSchema, PositionCreateSchema,
PositionUpdateSchema PositionUpdateSchema
@@ -25,12 +25,12 @@ PositionRouter = APIRouter(route_class=OperationLogRoute, prefix="/position", ta
@PositionRouter.get("/list", summary="查询岗位", description="查询岗位") @PositionRouter.get("/list", summary="查询岗位", description="查询岗位")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: PositionQueryParams = Depends(), search: PositionQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:position:query"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:position:query"])),
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await PositionService.get_position_list_service(search=search, auth=auth, order_by=page.order_by) result_dict_list = await PositionService.get_position_list_service(search=search, auth=auth, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询岗位列表成功") logger.info(f"查询岗位列表成功")
return SuccessResponse(data=result_dict, msg="查询岗位列表成功") return SuccessResponse(data=result_dict, msg="查询岗位列表成功")
@@ -88,7 +88,7 @@ async def batch_set_available_obj_controller(
@PositionRouter.post('/export', summary="导出岗位", description="导出岗位") @PositionRouter.post('/export', summary="导出岗位", description="导出岗位")
async def export_obj_list_controller( async def export_obj_list_controller(
search: PositionQueryParams = Depends(), search: PositionQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:position:export"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:position:export"])),
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class PositionQueryParams: class PositionQueryParam:
"""岗位管理查询参数""" """岗位管理查询参数"""
def __init__( def __init__(
@@ -6,7 +6,7 @@ from app.core.base_schema import BatchSetAvailable
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
from app.utils.excel_util import ExcelUtil from app.utils.excel_util import ExcelUtil
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .param import PositionQueryParams from .param import PositionQueryParam
from .crud import PositionCRUD from .crud import PositionCRUD
from .schema import ( from .schema import (
PositionCreateSchema, PositionCreateSchema,
@@ -25,7 +25,7 @@ class PositionService:
return PositionOutSchema.model_validate(position).model_dump() return PositionOutSchema.model_validate(position).model_dump()
@classmethod @classmethod
async def get_position_list_service(cls, auth: AuthSchema, search: PositionQueryParams, order_by: List[Dict] = None) -> List[Dict]: async def get_position_list_service(cls, auth: AuthSchema, search: PositionQueryParam, order_by: List[Dict] = None) -> List[Dict]:
"""获取岗位列表""" """获取岗位列表"""
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
@@ -7,13 +7,13 @@ from app.common.response import StreamResponse, SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .service import RoleService from .service import RoleService
from .param import RoleQueryParams from .param import RoleQueryParam
from .schema import ( from .schema import (
RoleCreateSchema, RoleCreateSchema,
RoleUpdateSchema, RoleUpdateSchema,
@@ -26,12 +26,12 @@ RoleRouter = APIRouter(route_class=OperationLogRoute, prefix="/role", tags=["角
@RoleRouter.get("/list", summary="查询角色", description="查询角色") @RoleRouter.get("/list", summary="查询角色", description="查询角色")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: RoleQueryParams = Depends(), search: RoleQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:role:query"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:role:query"])),
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await RoleService.get_role_list_service(search=search, auth=auth, order_by=page.order_by) result_dict_list = await RoleService.get_role_list_service(search=search, auth=auth, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询角色成功") logger.info(f"查询角色成功")
return SuccessResponse(data=result_dict, msg="查询角色成功") return SuccessResponse(data=result_dict, msg="查询角色成功")
@@ -99,7 +99,7 @@ async def set_role_permission_controller(
@RoleRouter.post('/export', summary="导出角色", description="导出角色") @RoleRouter.post('/export', summary="导出角色", description="导出角色")
async def export_obj_list_controller( async def export_obj_list_controller(
search: RoleQueryParams = Depends(), search: RoleQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:role:export"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:role:export"])),
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class RoleQueryParams: class RoleQueryParam:
"""角色管理查询参数""" """角色管理查询参数"""
def __init__( def __init__(
@@ -7,7 +7,7 @@ from app.core.exceptions import CustomException
from app.utils.excel_util import ExcelUtil from app.utils.excel_util import ExcelUtil
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .crud import RoleCRUD from .crud import RoleCRUD
from .param import RoleQueryParams from .param import RoleQueryParam
from .schema import ( from .schema import (
RoleCreateSchema, RoleCreateSchema,
RoleUpdateSchema, RoleUpdateSchema,
@@ -26,7 +26,7 @@ class RoleService:
return RoleOutSchema.model_validate(role).model_dump() return RoleOutSchema.model_validate(role).model_dump()
@classmethod @classmethod
async def get_role_list_service(cls, auth: AuthSchema, search: RoleQueryParams, order_by: List[Dict[str, str]] = None) -> List[Dict]: async def get_role_list_service(cls, auth: AuthSchema, search: RoleQueryParam, order_by: List[Dict[str, str]] = None) -> List[Dict]:
"""获取角色列表""" """获取角色列表"""
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
@@ -6,11 +6,11 @@ from fastapi.responses import JSONResponse
from app.common.response import SuccessResponse from app.common.response import SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .param import TicketQueryParams from .param import TicketQueryParam
from .service import TicketService from .service import TicketService
from .schema import TicketCreateSchema, TicketUpdateSchema from .schema import TicketCreateSchema, TicketUpdateSchema
@@ -30,12 +30,12 @@ async def get_ticket_detail_controller(
@TicketRouter.get("/list", summary="查询工单列表", description="查询工单列表") @TicketRouter.get("/list", summary="查询工单列表", description="查询工单列表")
async def get_ticket_list_controller( async def get_ticket_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: TicketQueryParams = Depends(), search: TicketQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:ticket:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:ticket:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await TicketService.get_ticket_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await TicketService.get_ticket_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size) result_dict = await PaginationService.paginate(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size)
logger.info("查询工单列表成功") logger.info("查询工单列表成功")
return SuccessResponse(data=result_dict, msg="查询工单列表成功") return SuccessResponse(data=result_dict, msg="查询工单列表成功")
@@ -7,7 +7,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class TicketQueryParams: class TicketQueryParam:
"""工单管理查询参数""" """工单管理查询参数"""
def __init__( def __init__(
@@ -5,7 +5,7 @@ from app.api.v1.module_system.auth.schema import AuthSchema
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
from .crud import TicketCRUD from .crud import TicketCRUD
from .schema import TicketCreateSchema, TicketUpdateSchema, TicketOutSchema from .schema import TicketCreateSchema, TicketUpdateSchema, TicketOutSchema
from .param import TicketQueryParams from .param import TicketQueryParam
class TicketService: class TicketService:
@@ -18,7 +18,7 @@ class TicketService:
return TicketOutSchema.model_validate(obj).model_dump() return TicketOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_ticket_list_service(cls, auth: AuthSchema, search: Optional[TicketQueryParams] = None, order_by: Optional[Union[str, List[Dict[str, str]]]] = None) -> List[Dict]: async def get_ticket_list_service(cls, auth: AuthSchema, search: Optional[TicketQueryParam] = None, order_by: Optional[Union[str, List[Dict[str, str]]]] = None) -> List[Dict]:
"""获取工单列表""" """获取工单列表"""
# 处理排序参数 # 处理排序参数
processed_order_by = None processed_order_by = None
@@ -10,12 +10,12 @@ from app.common.request import PaginationService
from app.utils.common_util import bytes2file_response from app.utils.common_util import bytes2file_response
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.dependencies import db_getter, get_current_user, AuthPermission from app.core.dependencies import db_getter, get_current_user, AuthPermission
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.base_schema import BatchSetAvailable from app.core.base_schema import BatchSetAvailable
from app.core.logger import logger from app.core.logger import logger
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from .service import UserService from .service import UserService
from .param import UserQueryParams from .param import UserQueryParam
from .schema import ( from .schema import (
CurrentUserUpdateSchema, CurrentUserUpdateSchema,
ResetPasswordSchema, ResetPasswordSchema,
@@ -101,12 +101,12 @@ async def forget_password_controller(
@UserRouter.get("/list", summary="查询用户", description="查询用户") @UserRouter.get("/list", summary="查询用户", description="查询用户")
async def get_obj_list_controller( async def get_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: UserQueryParams = Depends(), search: UserQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:user:query"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:user:query"])),
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await UserService.get_user_list_service(search=search, auth=auth, order_by=page.order_by) result_dict_list = await UserService.get_user_list_service(search=search, auth=auth, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size) result_dict = await PaginationService.paginate(data_list= result_dict_list, page_no= page.page_no, page_size = page.page_size)
logger.info(f"查询用户成功") logger.info(f"查询用户成功")
return SuccessResponse(data=result_dict, msg="查询用户成功") return SuccessResponse(data=result_dict, msg="查询用户成功")
@@ -179,8 +179,8 @@ async def export_obj_template_controller()-> StreamingResponse:
@UserRouter.post('/export', summary="导出用户", description="导出用户") @UserRouter.post('/export', summary="导出用户", description="导出用户")
async def export_obj_list_controller( async def export_obj_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: UserQueryParams = Depends(), search: UserQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:user:export"])), auth: AuthSchema = Depends(AuthPermission(permissions=["system:user:export"])),
) -> StreamingResponse: ) -> StreamingResponse:
# 获取全量数据 # 获取全量数据
@@ -6,7 +6,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class UserQueryParams: class UserQueryParam:
"""用户管理查询参数""" """用户管理查询参数"""
def __init__( def __init__(
@@ -18,7 +18,7 @@ from ..menu.crud import MenuCRUD
from ..dept.crud import DeptCRUD from ..dept.crud import DeptCRUD
from ..auth.schema import AuthSchema from ..auth.schema import AuthSchema
from ..menu.schema import MenuOutSchema from ..menu.schema import MenuOutSchema
from .param import UserQueryParams from .param import UserQueryParam
from .crud import UserCRUD from .crud import UserCRUD
from .schema import ( from .schema import (
CurrentUserUpdateSchema, CurrentUserUpdateSchema,
@@ -32,7 +32,6 @@ from .schema import (
) )
class UserService: class UserService:
"""用户模块服务层""" """用户模块服务层"""
@@ -53,7 +52,7 @@ class UserService:
return UserOutSchema.model_validate(user).model_dump() return UserOutSchema.model_validate(user).model_dump()
@classmethod @classmethod
async def get_user_list_service(cls, auth: AuthSchema, search: UserQueryParams, order_by: List[Dict]= None) -> List[Dict]: async def get_user_list_service(cls, auth: AuthSchema, search: UserQueryParam, order_by: List[Dict]= None) -> List[Dict]:
if order_by: if order_by:
order_by = eval(order_by) order_by = eval(order_by)
user_list = await UserCRUD(auth).get_list_crud(search=search.__dict__, order_by=order_by) user_list = await UserCRUD(auth).get_list_crud(search=search.__dict__, order_by=order_by)
@@ -6,11 +6,11 @@ from fastapi.responses import JSONResponse
from app.common.response import SuccessResponse from app.common.response import SuccessResponse
from app.common.request import PaginationService from app.common.request import PaginationService
from app.core.router_class import OperationLogRoute from app.core.router_class import OperationLogRoute
from app.core.base_params import PaginationQueryParams from app.core.base_params import PaginationQueryParam
from app.core.dependencies import AuthPermission from app.core.dependencies import AuthPermission
from app.core.logger import logger from app.core.logger import logger
from app.api.v1.module_system.auth.schema import AuthSchema from app.api.v1.module_system.auth.schema import AuthSchema
from .param import VersionQueryParams from .param import VersionQueryParam
from .service import VersionService from .service import VersionService
from .schema import VersionCreateSchema, VersionUpdateSchema from .schema import VersionCreateSchema, VersionUpdateSchema
@@ -30,12 +30,12 @@ async def get_version_detail_controller(
@VersionRouter.get("/list", summary="查询版本列表", description="查询版本列表") @VersionRouter.get("/list", summary="查询版本列表", description="查询版本列表")
async def get_version_list_controller( async def get_version_list_controller(
page: PaginationQueryParams = Depends(), page: PaginationQueryParam = Depends(),
search: VersionQueryParams = Depends(), search: VersionQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["system:version:query"])) auth: AuthSchema = Depends(AuthPermission(permissions=["system:version:query"]))
) -> JSONResponse: ) -> JSONResponse:
result_dict_list = await VersionService.get_version_list_service(auth=auth, search=search, order_by=page.order_by) result_dict_list = await VersionService.get_version_list_service(auth=auth, search=search, order_by=page.order_by)
result_dict = await PaginationService.get_page_obj(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size) result_dict = await PaginationService.paginate(data_list=result_dict_list, page_no=page.page_no, page_size=page.page_size)
logger.info("查询版本列表成功") logger.info("查询版本列表成功")
return SuccessResponse(data=result_dict, msg="查询版本列表成功") return SuccessResponse(data=result_dict, msg="查询版本列表成功")
@@ -7,7 +7,7 @@ from fastapi import Query
from app.core.validator import DateTimeStr from app.core.validator import DateTimeStr
class VersionQueryParams: class VersionQueryParam:
"""版本管理查询参数""" """版本管理查询参数"""
def __init__( def __init__(
@@ -5,7 +5,7 @@ from app.api.v1.module_system.auth.schema import AuthSchema
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
from .crud import VersionCRUD from .crud import VersionCRUD
from .schema import VersionCreateSchema, VersionUpdateSchema, VersionOutSchema from .schema import VersionCreateSchema, VersionUpdateSchema, VersionOutSchema
from .param import VersionQueryParams from .param import VersionQueryParam
class VersionService: class VersionService:
@@ -18,7 +18,7 @@ class VersionService:
return VersionOutSchema.model_validate(obj).model_dump() return VersionOutSchema.model_validate(obj).model_dump()
@classmethod @classmethod
async def get_version_list_service(cls, auth: AuthSchema, search: Optional[VersionQueryParams] = None, order_by: Optional[Union[str, List[Dict[str, str]]]] = None) -> List[Dict]: async def get_version_list_service(cls, auth: AuthSchema, search: Optional[VersionQueryParam] = None, order_by: Optional[Union[str, List[Dict[str, str]]]] = None) -> List[Dict]:
"""获取版本列表""" """获取版本列表"""
# 处理排序参数 # 处理排序参数
processed_order_by = None processed_order_by = None
@@ -70,7 +70,7 @@ class VersionService:
await VersionCRUD(auth).delete_crud(ids=ids) await VersionCRUD(auth).delete_crud(ids=ids)
@classmethod @classmethod
async def get_version_by_status_service(cls, auth: AuthSchema, search: Optional[VersionQueryParams] = None, order_by: Optional[Union[str, List[Dict[str, str]]]] = None) -> List[Dict]: async def get_version_by_status_service(cls, auth: AuthSchema, search: Optional[VersionQueryParam] = None, order_by: Optional[Union[str, List[Dict[str, str]]]] = None) -> List[Dict]:
"""根据状态获取版本列表""" """根据状态获取版本列表"""
# 处理排序参数 # 处理排序参数
processed_order_by = None processed_order_by = None
+1
View File
@@ -67,6 +67,7 @@ class GenConstants:
# 页面不需要查询字段 # 页面不需要查询字段
COLUMN_NAME_NOT_QUERY = ["id", "create_by", "dept_id", "create_time", "del_flag", "update_by", "update_time", "remark"] COLUMN_NAME_NOT_QUERY = ["id", "create_by", "dept_id", "create_time", "del_flag", "update_by", "update_time", "remark"]
# Dao基类字段
DAO_COLUMN_NOT_EDIT = ["create_by", "dept_id", "create_time", "del_flag", "update_time"] DAO_COLUMN_NOT_EDIT = ["create_by", "dept_id", "create_time", "del_flag", "update_time"]
# Entity基类字段 # Entity基类字段
+2 -2
View File
@@ -50,9 +50,9 @@ class RedisInitKeyConfig(Enum):
@property @property
def key(self) -> str: def key(self) -> str:
"""获取Redis键名""" """获取Redis键名"""
return self.value.get('key') return self.value.get('key', '')
@property @property
def remark(self) -> str: def remark(self) -> str:
"""获取Redis键名说明""" """获取Redis键名说明"""
return self.value.get('remark') return self.value.get('remark', '')
+1 -1
View File
@@ -23,7 +23,7 @@ class PaginationService:
"""分页服务类""" """分页服务类"""
@staticmethod @staticmethod
async def get_page_obj(data_list: List[Any], page_no: Optional[int] = None, page_size: Optional[int] = None) -> Dict[str, Any]: async def paginate(data_list: List[Any], page_no: Optional[int] = None, page_size: Optional[int] = None) -> Dict[str, Any]:
""" """
输入数据列表data_list和分页信息,返回分页或非分页数据列表结果。 输入数据列表data_list和分页信息,返回分页或非分页数据列表结果。
如果未传入page_no和page_size,则返回全部数据。 如果未传入page_no和page_size,则返回全部数据。
+4 -4
View File
@@ -1,6 +1,6 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Any, Mapping, Optional, Dict, Union from typing import Any, Mapping, Optional
from fastapi import status from fastapi import status
from fastapi.responses import JSONResponse, StreamingResponse, FileResponse from fastapi.responses import JSONResponse, StreamingResponse, FileResponse
from starlette.background import BackgroundTask from starlette.background import BackgroundTask
@@ -22,7 +22,7 @@ class SuccessResponse(JSONResponse):
def __init__( def __init__(
self, self,
data: Optional[Any] = None, data: Optional[Any] = None,
msg: Optional[str] = RET.OK.msg, msg: str = RET.OK.msg,
code: int = RET.OK.code, code: int = RET.OK.code,
status_code: int = status.HTTP_200_OK, status_code: int = status.HTTP_200_OK,
success: bool = True success: bool = True
@@ -51,7 +51,7 @@ class ErrorResponse(JSONResponse):
def __init__( def __init__(
self, self,
data: Optional[Any] = None, data: Optional[Any] = None,
msg: Optional[str] = RET.ERROR.msg, msg: str = RET.ERROR.msg,
code: int = RET.ERROR.code, code: int = RET.ERROR.code,
status_code: int = status.HTTP_400_BAD_REQUEST, status_code: int = status.HTTP_400_BAD_REQUEST,
success: bool = False success: bool = False
@@ -111,7 +111,7 @@ class UploadFileResponse(FileResponse):
file_path: str, file_path: str,
filename: str, filename: str,
media_type: str = "application/octet-stream", media_type: str = "application/octet-stream",
headers: Optional[Dict[str, Union[str, int]]] = None, headers: Optional[Mapping[str, str]] = None,
background: Optional[BackgroundTask] = None, background: Optional[BackgroundTask] = None,
status_code: int = 200 status_code: int = 200
): ):
+13 -29
View File
@@ -144,7 +144,7 @@ class Settings(BaseSettings):
CAPTCHA_ENABLE: bool = True # 是否启用验证码 CAPTCHA_ENABLE: bool = True # 是否启用验证码
CAPTCHA_EXPIRE_SECONDS: int = 60 * 1 # 验证码过期时间(秒) 1分钟 CAPTCHA_EXPIRE_SECONDS: int = 60 * 1 # 验证码过期时间(秒) 1分钟
CAPTCHA_FONT_SIZE: int = 40 # 字体大小 CAPTCHA_FONT_SIZE: int = 40 # 字体大小
CAPTCHA_FONT_PATH: Path = 'static/assets/font/Arial.ttf' # 字体路径 CAPTCHA_FONT_PATH: str = 'static/assets/font/Arial.ttf' # 字体路径
# ================================================= # # ================================================= #
# ********************* 日志配置 ******************* # # ********************* 日志配置 ******************* #
@@ -230,7 +230,6 @@ class Settings(BaseSettings):
ALI_OSS_END_POINT: str = 'xxxx' ALI_OSS_END_POINT: str = 'xxxx'
ALI_OSS_PRE: str = 'xxxx' ALI_OSS_PRE: str = 'xxxx'
ALI_OSS_BUCKET: str = 'xxxx' ALI_OSS_BUCKET: str = 'xxxx'
UPLOAD_METHOD: str = 'xxxx'
# ================================================= # # ================================================= #
# ***************** Swagger配置 ***************** # # ***************** Swagger配置 ***************** #
@@ -254,11 +253,8 @@ class Settings(BaseSettings):
table_prefix: str = 'sys_' table_prefix: str = 'sys_'
allow_overwrite: bool = False allow_overwrite: bool = False
GEN_PATH: str = 'gen_code/gen_path' GEN_PATH: Path = BASE_DIR.joinpath('app/api/v1/module_generator/gen_backend_code')
# def __init__(self):
# if not os.path.exists(self.GEN_PATH):
# os.makedirs(self.GEN_PATH)
# ================================================= # # ================================================= #
# ******************* AI大模型配置 ****************** # # ******************* AI大模型配置 ****************** #
@@ -279,7 +275,6 @@ class Settings(BaseSettings):
"app.core.middlewares.CustomCORSMiddleware" if self.CORS_ORIGIN_ENABLE else None, "app.core.middlewares.CustomCORSMiddleware" if self.CORS_ORIGIN_ENABLE else None,
"app.core.middlewares.RequestLogMiddleware" if self.OPERATION_LOG_RECORD else None, "app.core.middlewares.RequestLogMiddleware" if self.OPERATION_LOG_RECORD else None,
"app.core.middlewares.CustomGZipMiddleware" if self.GZIP_ENABLE else None, "app.core.middlewares.CustomGZipMiddleware" if self.GZIP_ENABLE else None,
"app.core.middlewares.DemoEnvMiddleware" if self.DEMO_ENABLE else None,
] ]
return MIDDLEWARES return MIDDLEWARES
@@ -297,44 +292,33 @@ class Settings(BaseSettings):
def ASYNC_DB_URI(self) -> str: def ASYNC_DB_URI(self) -> str:
"""获取异步数据库连接""" """获取异步数据库连接"""
if self.DATABASE_TYPE == "mysql": if self.DATABASE_TYPE == "mysql":
uri: MySQLDsn = f"mysql+asyncmy://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}?charset={self.DATABASE_CHARSET}" return f"mysql+asyncmy://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}?charset={self.DATABASE_CHARSET}"
return uri
elif self.DATABASE_TYPE == "postgresql": elif self.DATABASE_TYPE == "postgresql":
uri: PostgresDsn = f"postgresql+asyncpg://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}" return f"postgresql+asyncpg://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}"
return uri
elif self.DATABASE_TYPE == "sqlite": elif self.DATABASE_TYPE == "sqlite":
uri = f"sqlite+aiosqlite:///{self.BASE_DIR.joinpath(self.SQLITE_DB_NAME)}?characterEncoding=UTF-8" return f"sqlite+aiosqlite:///{self.BASE_DIR.joinpath(self.SQLITE_DB_NAME)}?characterEncoding=UTF-8"
return uri
else: else:
supported_db_drivers = ['mysql', 'postgresql', 'sqlite'] raise ValueError(f"数据库驱动不支持: {self.DATABASE_TYPE}, 请选择 请选择 mysql、postgresql、sqlite")
raise ValueError(f"数据库驱动不支持: {self.DATABASE_TYPE}, 请选择 {supported_db_drivers}")
@property @property
def DB_URI(self) -> str: def DB_URI(self) -> str:
"""获取同步数据库连接""" """获取同步数据库连接"""
if self.DATABASE_TYPE == "mysql": if self.DATABASE_TYPE == "mysql":
uri: MySQLDsn = f"mysql+pymysql://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}?charset={self.DATABASE_CHARSET}" return f"mysql+pymysql://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}?charset={self.DATABASE_CHARSET}"
return uri
elif self.DATABASE_TYPE == "postgresql": elif self.DATABASE_TYPE == "postgresql":
uri: PostgresDsn = f"postgresql+psycopg2://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}" return f"postgresql+psycopg2://{self.DATABASE_USER}:{quote_plus(self.DATABASE_PASSWORD)}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}"
return uri
elif self.DATABASE_TYPE == "sqlite": elif self.DATABASE_TYPE == "sqlite":
uri: str = f"sqlite:///{self.BASE_DIR.joinpath(self.SQLITE_DB_NAME)}?characterEncoding=UTF-8" return f"sqlite:///{self.BASE_DIR.joinpath(self.SQLITE_DB_NAME)}?characterEncoding=UTF-8"
return uri
else: else:
supported_db_drivers = ['mysql', 'postgresql', 'sqlite'] raise ValueError(f"数据库驱动不支持: {self.DATABASE_TYPE}, 请选择 mysql、postgresql、sqlite")
raise ValueError(f"数据库驱动不支持: {self.DATABASE_TYPE}, 请选择 {supported_db_drivers}")
@property @property
def MONGO_DB_URI(self) -> MongoDsn: def MONGO_DB_URI(self) -> str:
"""获取MongoDB连接""" """获取MongoDB连接"""
if settings.MONGO_DB_USER and settings.MONGO_DB_PASSWORD: return f"mongodb://{settings.MONGO_DB_USER}:{settings.MONGO_DB_PASSWORD}@{settings.MONGO_DB_HOST}:{settings.MONGO_DB_PORT}/{settings.MONGO_DB_NAME}"
return f"mongodb://{settings.MONGO_DB_USER}:{settings.MONGO_DB_PASSWORD}@{settings.MONGO_DB_HOST}:{settings.MONGO_DB_PORT}/{settings.MONGO_DB_NAME}"
else:
return f"mongodb://{settings.MONGO_DB_HOST}:{settings.MONGO_DB_PORT}/{settings.MONGO_DB_NAME}"
@property @property
def REDIS_URI(self) -> RedisDsn: def REDIS_URI(self) -> str:
"""获取Redis连接""" """获取Redis连接"""
return f"redis://{settings.REDIS_USER}:{self.REDIS_PASSWORD}@{settings.REDIS_HOST}:{settings.REDIS_PORT}/{settings.REDIS_DB_NAME}" return f"redis://{settings.REDIS_USER}:{self.REDIS_PASSWORD}@{settings.REDIS_HOST}:{settings.REDIS_PORT}/{settings.REDIS_DB_NAME}"
+2 -2
View File
@@ -3,7 +3,7 @@
from pydantic import BaseModel from pydantic import BaseModel
from typing import TypeVar, Sequence, Generic, Dict, Any, List, Union, Optional from typing import TypeVar, Sequence, Generic, Dict, Any, List, Union, Optional
from sqlalchemy.sql.elements import ColumnElement from sqlalchemy.sql.elements import ColumnElement
from sqlalchemy.orm import selectinload, DeclarativeBase from sqlalchemy.orm import Session, selectinload, DeclarativeBase
from sqlalchemy.engine import Result from sqlalchemy.engine import Result
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import asc, func, select, delete, Select, desc, update, or_, and_ from sqlalchemy import asc, func, select, delete, Select, desc, update, or_, and_
@@ -32,7 +32,7 @@ class CRUDBase(Generic[ModelType, CreateSchemaType, UpdateSchemaType]):
""" """
self.model = model self.model = model
self.auth = auth self.auth = auth
self.db: AsyncSession = auth.db self.db: AsyncSession | Session | None = auth.db
self.current_user = auth.user self.current_user = auth.user
async def get(self, **kwargs) -> Optional[ModelType]: async def get(self, **kwargs) -> Optional[ModelType]:
+13 -4
View File
@@ -10,7 +10,7 @@ from typing import Optional, Dict, Any
from sqlalchemy import Boolean, String, Integer, DateTime, ForeignKey, Text, BigInteger from sqlalchemy import Boolean, String, Integer, DateTime, ForeignKey, Text, BigInteger
from sqlalchemy.dialects.postgresql import UUID from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.ext.asyncio import AsyncAttrs from sqlalchemy.ext.asyncio import AsyncAttrs
from sqlalchemy.orm import relationship, DeclarativeBase, Mapped, declared_attr, mapped_column from sqlalchemy.orm import relationship, DeclarativeBase, Mapped, declared_attr, mapped_column, MappedAsDataclass
@@ -18,6 +18,12 @@ class MappedBase(AsyncAttrs, DeclarativeBase):
""" """
声明式基类 声明式基类
`AsyncAttrs <https://docs.sqlalchemy.org/en/20/orm/extensions/asyncio.html#sqlalchemy.ext.asyncio.AsyncAttrs>`__
`DeclarativeBase <https://docs.sqlalchemy.org/en/20/orm/declarative_config.html>`__
`mapped_column() <https://docs.sqlalchemy.org/en/20/orm/mapping_api.html#sqlalchemy.orm.mapped_column>`__
兼容 SQLite、MySQL 和 PostgreSQL 兼容 SQLite、MySQL 和 PostgreSQL
""" """
@@ -41,9 +47,12 @@ class ModelMixin(MappedBase):
__abstract__ = True __abstract__ = True
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True, comment='主键ID') id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True, comment='主键ID')
description: Mapped[Optional[str]] = mapped_column(Text, nullable=True, comment="备注说明") description: Mapped[Optional[str]] = mapped_column(Text, nullable=True, default=None, comment="备注/描述")
created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.now, comment='创建时间') created_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True, default=datetime.now, comment='创建时间')
updated_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.now, onupdate=datetime.now, comment='更新时间') create_by: Mapped[Optional[str]] = mapped_column(String(64), default='', comment='创建者')
updated_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True, default=datetime.now, onupdate=datetime.now, comment='更新时间')
update_by: Mapped[Optional[str]] = mapped_column(String(64), default='', comment='更新者')
del_flag: Mapped[str] = mapped_column(String(1), nullable=False, default='0', comment='删除标志(0代表存在 2代表删除)')
class CreatorMixin(ModelMixin): class CreatorMixin(ModelMixin):
+1 -1
View File
@@ -4,7 +4,7 @@ from typing import Optional
from fastapi import Query from fastapi import Query
class PaginationQueryParams: class PaginationQueryParam:
"""分页查询参数基类""" """分页查询参数基类"""
def __init__( def __init__(
+2 -2
View File
@@ -24,6 +24,7 @@ engine: Engine = create_engine(
pool_pre_ping=settings.POOL_PRE_PING, pool_pre_ping=settings.POOL_PRE_PING,
pool_recycle=settings.POOL_RECYCLE, pool_recycle=settings.POOL_RECYCLE,
) )
# 同步数据库会话工厂 # 同步数据库会话工厂
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
@@ -58,7 +59,7 @@ def session_connect() -> AsyncSession:
except Exception as e: except Exception as e:
raise CustomException(msg=f"数据库连接失败: {e}") raise CustomException(msg=f"数据库连接失败: {e}")
async def init_create_table(): async def init_create_table() -> None:
""" """
应用启动时初始化数据库连接 应用启动时初始化数据库连接
@@ -77,7 +78,6 @@ async def redis_connect(app: FastAPI, status: bool) -> aioredis.Redis:
if status: if status:
try: try:
rd = await aioredis.from_url( rd = await aioredis.from_url(
url=settings.REDIS_URI, url=settings.REDIS_URI,
encoding='utf-8', encoding='utf-8',
+3 -2
View File
@@ -8,6 +8,7 @@ from fastapi import Depends, Request
from motor.motor_asyncio import AsyncIOMotorDatabase from motor.motor_asyncio import AsyncIOMotorDatabase
from fastapi import Depends from fastapi import Depends
from app.api.v1.module_system.user.schema import UserOutSchema
from app.common.enums import RedisInitKeyConfig from app.common.enums import RedisInitKeyConfig
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
from app.core.database import session_connect from app.core.database import session_connect
@@ -95,7 +96,7 @@ async def get_current_user(
if hasattr(user, 'positions'): if hasattr(user, 'positions'):
user.positions = [pos for pos in user.positions if pos.status] user.positions = [pos for pos in user.positions if pos.status]
auth.user = user auth.user = UserOutSchema.model_validate(user)
return auth return auth
@@ -133,7 +134,7 @@ class AuthPermission:
auth.check_data_scope = self.check_data_scope auth.check_data_scope = self.check_data_scope
# 超级管理员直接通过 # 超级管理员直接通过
if auth.user.is_superuser: if auth.user and auth.user.is_superuser:
return auth return auth
# 无需验证权限 # 无需验证权限
+12 -78
View File
@@ -1,14 +1,12 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import Any, Optional, List, Tuple, Union from typing import Any, Optional
from fastapi import Request, status from fastapi import Request, status
from fastapi.exceptions import RequestValidationError, ResponseValidationError from fastapi.exceptions import RequestValidationError, ResponseValidationError
from pydantic_validation_decorator import FieldValidationError from pydantic_validation_decorator import FieldValidationError
from starlette.responses import JSONResponse from starlette.responses import JSONResponse
from starlette.exceptions import HTTPException from starlette.exceptions import HTTPException
from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.exc import SQLAlchemyError
from pydantic_core import ErrorDetails
from pydantic import ValidationError
from app.common.constant import RET from app.common.constant import RET
from app.common.response import ErrorResponse from app.common.response import ErrorResponse
@@ -20,7 +18,7 @@ class CustomException(Exception):
def __init__( def __init__(
self, self,
msg: Optional[str] = RET.EXCEPTION.msg, msg: str = RET.EXCEPTION.msg,
code: int = RET.EXCEPTION.code, code: int = RET.EXCEPTION.code,
status_code: int = status.HTTP_500_INTERNAL_SERVER_ERROR, status_code: int = status.HTTP_500_INTERNAL_SERVER_ERROR,
data: Optional[Any] = None, data: Optional[Any] = None,
@@ -59,15 +57,23 @@ async def HttpExceptionHandler(request: Request, exc: HTTPException) -> JSONResp
async def ValidationExceptionHandler(request: Request, exc: RequestValidationError) -> JSONResponse: async def ValidationExceptionHandler(request: Request, exc: RequestValidationError) -> JSONResponse:
"""请求参数验证异常处理器""" """请求参数验证异常处理器"""
msg:List[ErrorDetails] = custom_convert_errors(exc) error_mapping = {
"Field required": "请求失败,缺少必填项!",
"value is not a valid list": "类型错误,提交参数应该为列表!",
"value is not a valid int": "类型错误,提交参数应该为整数!",
"value could not be parsed to a boolean": "类型错误,提交参数应该为布尔值!",
"Input should be a valid list": "类型错误,输入应该是一个有效的列表!"
}
msg = error_mapping.get(exc.errors()[0].get('msg'), exc.errors()[0].get('msg'))
logger.error(f"请求地址: {request.url}, 错误信息: {msg}, 错误详情: {exc}") logger.error(f"请求地址: {request.url}, 错误信息: {msg}, 错误详情: {exc}")
return ErrorResponse(msg=str(msg), status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, data=exc.body) return ErrorResponse(msg=str(msg), status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, data=exc.body)
async def ResponseValidationHandle(request: Request, exc: ResponseValidationError) -> JSONResponse: async def ResponseValidationHandle(request: Request, exc: ResponseValidationError) -> JSONResponse:
logger.error(f"请求地址: {request.url}, 错误详情: {exc}") logger.error(f"请求地址: {request.url}, 错误详情: {exc}")
return ErrorResponse(msg=str(exc), status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, data=exc.body) return ErrorResponse(msg=str(exc), status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, data=exc.body)
async def SQLAlchemyExceptionHandler(request: Request, exc: SQLAlchemyError) -> JSONResponse: async def SQLAlchemyExceptionHandler(request: Request, exc: SQLAlchemyError) -> JSONResponse:
"""数据库异常处理器""" """数据库异常处理器"""
error_msg = f'数据库操作失败: {exc}' error_msg = f'数据库操作失败: {exc}'
@@ -91,75 +97,3 @@ async def AllExceptionHandler(request: Request, exc: Exception) -> JSONResponse:
"""全局异常处理器""" """全局异常处理器"""
logger.error(f"请求地址: {request.url}, 错误详情: {exc}") logger.error(f"请求地址: {request.url}, 错误详情: {exc}")
return ErrorResponse(msg='服务器内部错误', status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, data=str(exc)) return ErrorResponse(msg='服务器内部错误', status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, data=str(exc))
ERROR_MAPPING = {
"missing": "请求失败,缺少必填项!",
# 字符串
"string_pattern_mismatch": "值错误,提交参数不满足正则表达式{pattern}!",
"string_too_long": "值错误,提交参数长度必须小于等于{max_length}!",
"string_too_short": "值错误,提交参数长度必须大于等于{min_length}!",
"string_type": "类型错误,提交参数应该为字符串!",
# 列表
"list_type": "类型错误,提交参数应该为列表!",
# 字典
"dict_type": "类型错误,提交参数应该为字典!",
# 集合
"set_type": "类型错误,提交参数应该为集合!",
# 元组
"tuple_type": "类型错误,提交参数应该为元组!",
# 元素数量
"too_long": "数量错误,提交参数的元素数量必须小于等于{max_length}!",
"too_short": "数量错误,提交参数的元素数量必须大于等于{min_length}!",
# 大小比值
"less_than_equal": "值错误,提交参数必须小于等于{le}!",
"greater_than_equal": "值错误,提交参数必须大于等于{ge}!",
"less_than": "值错误,提交参数必须小于{lt}!",
"greater_than": "值错误,提交参数必须大于{gt}!",
# 布尔值
"bool_type": "类型错误,提交参数应该为布尔值!",
"bool_parsing": "类型错误,提交参数应该为布尔值!",
# 字节
"bytes_type": "类型错误,提交参数应该为字节!",
"bytes_too_long": "值错误,提交参数长度必须小于等于{max_length}!",
"bytes_too_short": "值错误,提交参数长度必须大于等于{min_length}!",
# 整数
"int_parsing": "类型错误,提交参数应该为整数!",
"int_type": "类型错误,提交参数应该为整数!",
# 浮点数
"float_parsing": "类型错误,提交参数应该为浮点数!",
"float_type": "类型错误,提交参数应该为浮点数!",
# 日期时间
"date_parsing": "类型错误,提交参数应该为日期!",
"date_type": "类型错误,提交参数应该为日期!",
"time_parsing": "类型错误,提交参数应该为时间!",
"time_type": "类型错误,提交参数应该为时间!",
# 其他
"literal_error": "值错误,提交参数值在为{expected}中一个!",
"extra_forbidden": "值错误,提交参数值不在允许范围内!",
}
def custom_convert_errors(e: ValidationError | RequestValidationError) -> List[ErrorDetails]:
new_errors: List[ErrorDetails] = []
for error in e.errors():
error['loc'] = loc_to_dot_sep(error['loc'])
custom_message = ERROR_MAPPING.get(error['type'])
if custom_message:
ctx = error.get('ctx')
error['msg'] = (
custom_message.format(**ctx) if ctx else custom_message
)
new_errors.append(error)
return new_errors
def loc_to_dot_sep(loc: Tuple[Union[str, int], ...]) -> str:
path = ''
for i, x in enumerate(loc):
if isinstance(x, str):
if i > 0:
path += '.'
path += x
elif isinstance(x, int):
path += f'[{x}]'
else:
raise TypeError('Unexpected type')
return path
+2 -1
View File
@@ -8,6 +8,7 @@ import logging
from logging.handlers import TimedRotatingFileHandler from logging.handlers import TimedRotatingFileHandler
from typing import Optional, Dict, Any from typing import Optional, Dict, Any
from pathlib import Path from pathlib import Path
import typing
from app.config.setting import settings from app.config.setting import settings
@@ -29,7 +30,7 @@ class CustomTimedRotatingFileHandler(TimedRotatingFileHandler):
# 使用流上下文管理确保资源正确释放 # 使用流上下文管理确保资源正确释放
if self.stream: if self.stream:
self.stream.close() self.stream.close()
self.stream = None self.stream = None # type: ignore
try: try:
# 计算轮换时间(使用缓存避免重复计算) # 计算轮换时间(使用缓存避免重复计算)
+44 -49
View File
@@ -1,12 +1,13 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
import time import time
from typing import Dict, List, Union from typing import Any
from starlette.middleware.cors import CORSMiddleware from starlette.middleware.cors import CORSMiddleware
from starlette.types import ASGIApp from starlette.types import ASGIApp
from starlette.requests import Request from starlette.requests import Request
from starlette.middleware.gzip import GZipMiddleware from starlette.middleware.gzip import GZipMiddleware
from starlette.middleware.base import Response, BaseHTTPMiddleware, RequestResponseEndpoint from starlette.responses import Response
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
from app.common.response import ErrorResponse from app.common.response import ErrorResponse
from app.config.setting import settings from app.config.setting import settings
@@ -17,7 +18,7 @@ from app.core.exceptions import CustomException
class CustomCORSMiddleware(CORSMiddleware): class CustomCORSMiddleware(CORSMiddleware):
"""CORS跨域中间件""" """CORS跨域中间件"""
def __init__(self, app: ASGIApp) -> None: def __init__(self, app: ASGIApp) -> None:
CORSMiddlewareConfig: Dict[str, Union[List[str], bool]] = { CORSMiddlewareConfig: dict[str, Any] = {
"allow_origins": settings.ALLOW_ORIGINS, "allow_origins": settings.ALLOW_ORIGINS,
"allow_methods": settings.ALLOW_METHODS, "allow_methods": settings.ALLOW_METHODS,
"allow_headers": settings.ALLOW_HEADERS, "allow_headers": settings.ALLOW_HEADERS,
@@ -38,65 +39,59 @@ class RequestLogMiddleware(BaseHTTPMiddleware):
) -> Response: ) -> Response:
start_time = time.time() start_time = time.time()
logger.info( # 构建请求日志信息
f"请求来源: {request.client.host}, " request_info = f"请求方法: {request.method}, 请求路径: {request.url.path}"
f"请求方法: {request.method}, " if request.client:
f"请求路径: {request.url.path}, " request_info = f"请求来源: {request.client.host}, {request_info}"
f"客户端IP: {request.client.host}" logger.info(request_info)
)
try: try:
response = await call_next(request)
if settings.DEMO_ENABLE:
# 在演示环境中,只有白名单内的IP或路径才能执行非GET请求
if request.method != "GET":
path = request.scope.get("path")
request_ip = None
x_forwarded_for = request.headers.get('X-Forwarded-For')
if x_forwarded_for:
# 取第一个 IP 地址,通常为客户端真实 IP
request_ip = x_forwarded_for.split(',')[0].strip()
else:
# 若没有 X-Forwarded-For 头,则使用 request.client.host
request_ip = request.client.host if request.client else None
# 检查IP是否在白名单,或路径是否在白名单,或用户是否在白名单
if (request_ip in settings.DEMO_IP_WHITE_LIST) or (path in settings.DEMO_WHITE_LIST_PATH):
response = await call_next(request)
else:
# 非白名单用户,禁止操作
return ErrorResponse(msg="演示环境,禁止操作")
else:
# GET请求在演示环境中总是允许的
response = await call_next(request)
else:
# 非演示环境,正常处理请求
response = await call_next(request)
process_time = round(time.time() - start_time, 5) process_time = round(time.time() - start_time, 5)
response.headers["X-Process-Time"] = str(process_time) response.headers["X-Process-Time"] = str(process_time)
logger.info( # 构建响应日志信息
f"会话ID: {request.scope.get('session_id')}, " session_id = request.scope.get('session_id')
content_length = response.headers.get('content-length', '0')
response_info = (
f"会话ID: {session_id}, "
f"响应状态: {response.status_code}, " f"响应状态: {response.status_code}, "
f"响应内容长度: {response.headers.get('content-length', '0')}, " f"响应内容长度: {content_length}, "
f"处理时间: {process_time}s" f"处理时间: {process_time}s"
) )
logger.info(response_info)
return response return response
except CustomException as e: except CustomException as e:
logger.error(f"系统异常: {str(e)}") logger.error(f"系统异常: {str(e)}")
return ErrorResponse(msg=f"系统异常,请联系管理员: {str(e)}") return ErrorResponse(msg=f"系统异常,请联系管理员", data=str(e))
class DemoEnvMiddleware(BaseHTTPMiddleware):
"""演示环境中间件"""
def __init__(self, app: ASGIApp) -> None:
super().__init__(app)
async def dispatch(
self, request: Request, call_next: RequestResponseEndpoint
) -> Response:
if settings.DEMO_ENABLE and request.method != "GET":
path = request.scope.get("path")
request_ip = None
x_forwarded_for = request.headers.get('X-Forwarded-For')
if x_forwarded_for:
# 取第一个 IP 地址,通常为客户端真实 IP
request_ip = x_forwarded_for.split(',')[0].strip()
else:
# 若没有 X-Forwarded-For 头,则使用 request.client.host
request_ip = request.client.host
user_username = request.scope.get("user_username")
logger.error(f"用户名称: {user_username}")
# 检查IP是否在白名单,或路径是否在白名单,或用户是否在白名单
if (request_ip in settings.DEMO_IP_WHITE_LIST) or (path in settings.DEMO_WHITE_LIST_PATH):
return await call_next(request)
else:
# 非白名单用户,禁止操作
return ErrorResponse(msg="演示环境,禁止操作")
return await call_next(request)
class CustomGZipMiddleware(GZipMiddleware): class CustomGZipMiddleware(GZipMiddleware):
+2 -2
View File
@@ -148,7 +148,7 @@ class RedisCURD:
async def hash_set(self, name: str, key: str, value: Any) -> bool: async def hash_set(self, name: str, key: str, value: Any) -> bool:
"""设置哈希缓存""" """设置哈希缓存"""
try: try:
await self.redis.hset(name=name, key=key, value=value) self.redis.hset(name=name, key=key, value=value)
return True return True
except Exception as e: except Exception as e:
logger.error(f"设置哈希缓存失败: {str(e)}") logger.error(f"设置哈希缓存失败: {str(e)}")
@@ -157,7 +157,7 @@ class RedisCURD:
async def hash_get(self, name: str, keys: list[str]) -> Optional[list[Any]]: async def hash_get(self, name: str, keys: list[str]) -> Optional[list[Any]]:
"""获取哈希缓存""" """获取哈希缓存"""
try: try:
return await self.redis.hmget(name=name, keys=keys) return self.redis.hmget(name=name, keys=keys)
except Exception as e: except Exception as e:
logger.error(f"获取哈希缓存失败: {str(e)}") logger.error(f"获取哈希缓存失败: {str(e)}")
return None return None
+9 -11
View File
@@ -35,7 +35,7 @@ class OperationLogRoute(APIRoute):
return response return response
if request.method not in settings.OPERATION_RECORD_METHOD: if request.method not in settings.OPERATION_RECORD_METHOD:
return response return response
route: APIRoute = request.scope.get("route") route: APIRoute = request.scope.get("route", None)
if route.name in settings.IGNORE_OPERATION_FUNCTION: if route.name in settings.IGNORE_OPERATION_FUNCTION:
return response return response
@@ -66,7 +66,6 @@ class OperationLogRoute(APIRoute):
oper_param['path_params'] = dict(path_params) oper_param['path_params'] = dict(path_params)
payload = json.dumps(oper_param, ensure_ascii=False) payload = json.dumps(oper_param, ensure_ascii=False)
# payload = str(oper_param)
# 日志表请求参数字段长度最大为2000,因此在此处判断长度 # 日志表请求参数字段长度最大为2000,因此在此处判断长度
if len(payload) > 2000: if len(payload) > 2000:
@@ -89,19 +88,18 @@ class OperationLogRoute(APIRoute):
request_ip = x_forwarded_for.split(',')[0].strip() request_ip = x_forwarded_for.split(',')[0].strip()
else: else:
# 若没有 X-Forwarded-For 头,则使用 request.client.host # 若没有 X-Forwarded-For 头,则使用 request.client.host
request_ip = request.client.host if request.client:
request_ip = request.client.host
login_location = await IpLocalUtil.get_ip_location(request_ip) login_location = await IpLocalUtil.get_ip_location(request_ip) if request_ip else None
# 判断请求是否来自api文档 # 判断请求是否来自api文档
request_from_swagger = ( referer = request.headers.get('referer')
request.headers.get('referer').endswith('docs') if request.headers.get('referer') else False request_from_swagger = referer and referer.endswith('docs')
) request_from_redoc = referer and referer.endswith('redoc')
request_from_redoc = (
request.headers.get('referer').endswith('redoc') if request.headers.get('referer') else False
)
if request_from_swagger or request_from_redoc: if request_from_swagger or request_from_redoc:
# 如果请求来自api文档,则不记录日志
pass pass
else: else:
async with session_connect() as session: async with session_connect() as session:
@@ -117,7 +115,7 @@ class OperationLogRoute(APIRoute):
request_os = user_agent.os.family, request_os = user_agent.os.family,
request_browser = user_agent.browser.family, request_browser = user_agent.browser.family,
response_code = response.status_code, response_code = response.status_code,
response_json = response_data.decode(), response_json = response_data.decode() if isinstance(response_data, (bytes, bytearray)) else str(response_data),
process_time = process_time, process_time = process_time,
description = route.summary, description = route.summary,
creator_id = current_user_id creator_id = current_user_id
+1 -1
View File
@@ -37,7 +37,7 @@ class CustomOAuth2PasswordBearer(OAuth2PasswordBearer):
if not authorization or scheme.lower() != settings.TOKEN_TYPE: if not authorization or scheme.lower() != settings.TOKEN_TYPE:
if self.auto_error: if self.auto_error:
raise CustomException(msg="请登录后再试") raise CustomException(msg="认证失败,请登录后再试", code=10401, status_code=401)
return None return None
return token return token
-62
View File
@@ -1,62 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from decimal import Decimal
from typing import Any, Sequence, TypeVar
from fastapi.encoders import decimal_encoder
from sqlalchemy import Row, RowMapping
from sqlalchemy.orm import ColumnProperty, SynonymProperty, class_mapper
RowData = Row | RowMapping | Any
R = TypeVar('R', bound=RowData)
def select_columns_serialize(row: R) -> dict[str, Any]:
"""
序列化 SQLAlchemy 查询表的列,不包含关联列
:param row: SQLAlchemy 查询结果行
:return:
"""
result = {}
for column in row.__table__.columns.keys():
value = getattr(row, column)
if isinstance(value, Decimal):
value = decimal_encoder(value)
result[column] = value
return result
def select_list_serialize(row: Sequence[R]) -> list[dict[str, Any]]:
"""
序列化 SQLAlchemy 查询列表
:param row: SQLAlchemy 查询结果列表
:return:
"""
return [select_columns_serialize(item) for item in row]
def select_as_dict(row: R, use_alias: bool = False) -> dict[str, Any]:
"""
将 SQLAlchemy 查询结果转换为字典,可以包含关联数据
:param row: SQLAlchemy 查询结果行
:param use_alias: 是否使用别名作为列名
:return:
"""
if not use_alias:
result = row.__dict__
if '_sa_instance_state' in result:
del result['_sa_instance_state']
else:
result = {}
mapper = class_mapper(row.__class__) # type: ignore
for prop in mapper.iterate_properties:
if isinstance(prop, (ColumnProperty, SynonymProperty)):
key = prop.key
result[key] = getattr(row, key)
return result
+40
View File
@@ -2256,5 +2256,45 @@
"description": "前端构建" "description": "前端构建"
} }
] ]
},
{
"name": "流程管理",
"type": 1,
"icon": "el-icon-ShoppingBag",
"order": 10,
"permission": null,
"route_name": "Workflow",
"route_path": "/workflow",
"component_path": null,
"status": true,
"keep_alive": false,
"hidden": false,
"always_show": false,
"title": "流程管理",
"params": null,
"affix": false,
"redirect": "/workflow/operator",
"description": "流程管理",
"children": [
{
"name": "我的流程",
"type": 2,
"icon": "el-icon-ShoppingBag",
"order": 1,
"permission": "workflow:operator:query",
"route_name": "Operator",
"route_path": "/workflow/operator",
"component_path": "workflow/operator/index",
"status": true,
"keep_alive": true,
"hidden": false,
"always_show": false,
"title": "我的流程",
"params": null,
"affix": false,
"redirect": null,
"description": "我的流程"
}
]
} }
] ]
@@ -430,5 +430,45 @@
{ {
"role_id": 1, "role_id": 1,
"menu_id": 108 "menu_id": 108
},
{
"role_id": 1,
"menu_id": 109
},
{
"role_id": 1,
"menu_id": 110
},
{
"role_id": 1,
"menu_id": 111
},
{
"role_id": 1,
"menu_id": 112
},
{
"role_id": 1,
"menu_id": 113
},
{
"role_id": 1,
"menu_id": 114
},
{
"role_id": 1,
"menu_id": 115
},
{
"role_id": 1,
"menu_id": 116
},
{
"role_id": 1,
"menu_id": 117
},
{
"role_id": 1,
"menu_id": 118
} }
] ]
-126
View File
@@ -1,126 +0,0 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
from typing import Any, Sequence
from backend.common.enums import BuildTreeType
from backend.utils.serializers import RowData, select_list_serialize
def get_tree_nodes(row: Sequence[RowData], is_sort: bool, sort_key: str) -> list[dict[str, Any]]:
"""
获取所有树形结构节点
:param row: 原始数据行序列
:param is_sort: 是否启用结果排序
:param sort_key: 基于此键对结果进行进行排序
:return:
"""
tree_nodes = select_list_serialize(row)
if is_sort:
tree_nodes.sort(key=lambda x: x[sort_key])
return tree_nodes
def traversal_to_tree(nodes: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""
通过遍历算法构造树形结构
:param nodes: 树节点列表
:return:
"""
tree: list[dict[str, Any]] = []
node_dict = {node['id']: node for node in nodes}
for node in nodes:
parent_id = node['parent_id']
if parent_id is None:
tree.append(node)
else:
parent_node = node_dict.get(parent_id)
if parent_node is not None:
if 'children' not in parent_node:
parent_node['children'] = []
if node not in parent_node['children']:
parent_node['children'].append(node)
else:
if node not in tree:
tree.append(node)
return tree
def recursive_to_tree(nodes: list[dict[str, Any]], *, parent_id: int | None = None) -> list[dict[str, Any]]:
"""
通过递归算法构造树形结构(性能影响较大)
:param nodes: 树节点列表
:param parent_id: 父节点 ID,默认为 None 表示根节点
:return:
"""
tree: list[dict[str, Any]] = []
for node in nodes:
if node['parent_id'] == parent_id:
child_nodes = recursive_to_tree(nodes, parent_id=node['id'])
if child_nodes:
node['children'] = child_nodes
tree.append(node)
return tree
def get_tree_data(
row: Sequence[RowData],
build_type: BuildTreeType = BuildTreeType.traversal,
*,
parent_id: int | None = None,
is_sort: bool = True,
sort_key: str = 'sort',
) -> list[dict[str, Any]]:
"""
获取树形结构数据
:param row: 原始数据行序列
:param build_type: 构建树形结构的算法类型,默认为遍历算法
:param parent_id: 父节点 ID,仅在递归算法中使用
:param is_sort: 是否启用结果排序
:param sort_key: 基于此键对结果进行进行排序
:return:
"""
nodes = get_tree_nodes(row, is_sort, sort_key)
match build_type:
case BuildTreeType.traversal:
tree = traversal_to_tree(nodes)
case BuildTreeType.recursive:
tree = recursive_to_tree(nodes, parent_id=parent_id)
case _:
raise ValueError(f'无效的算法类型:{build_type}')
return tree
def get_vben5_tree_data(row: Sequence[RowData], is_sort: bool = True, sort_key: str = 'sort') -> list[dict[str, Any]]:
"""
获取 vben5 菜单树形结构数据
:param row: 原始数据行序列
:param is_sort: 是否启用结果排序
:param sort_key: 基于此键对结果进行进行排序
:return:
"""
meta_keys = {'title', 'icon', 'link', 'cache', 'display', 'status'}
vben5_nodes = [
{
**{k: v for k, v in node.items() if k not in meta_keys},
'meta': {
'title': node['title'],
'icon': node['icon'],
'iframeSrc': node['link'] if node['type'] == 3 else '',
'link': node['link'] if node['type'] == 4 else '',
'keepAlive': node['cache'],
'hideInMenu': not bool(node['display']),
'menuVisibleWithForbidden': not bool(node['status']),
},
}
for node in get_tree_nodes(row, is_sort, sort_key)
]
return traversal_to_tree(vben5_nodes)
+105 -2
View File
@@ -1,12 +1,17 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from pathlib import Path
import importlib import importlib
import os
import uuid import uuid
import re import re
from sqlalchemy.orm import DeclarativeBase from sqlalchemy.orm import DeclarativeBase
from typing import Any, List, Dict, Sequence, Optional from typing import Any, Generator, List, Dict, Sequence, Optional
from app.config import setting
from app.core.logger import logger from app.core.logger import logger
from app.core.exceptions import CustomException from app.core.exceptions import CustomException
@@ -40,6 +45,7 @@ def import_module(module: str, desc: str) -> Any:
logger.error(f"导入{desc}失败,未找到模块方法:{module}") logger.error(f"导入{desc}失败,未找到模块方法:{module}")
raise AttributeError(f"导入{desc}失败,未找到模块方法:{module}") raise AttributeError(f"导入{desc}失败,未找到模块方法:{module}")
async def import_modules_async(modules: list, desc: str, **kwargs): async def import_modules_async(modules: list, desc: str, **kwargs):
""" """
异步导入模块列表 异步导入模块列表
@@ -62,10 +68,12 @@ async def import_modules_async(modules: list, desc: str, **kwargs):
logger.error(f"导入{desc}失败,未找到模块方法:{module}") logger.error(f"导入{desc}失败,未找到模块方法:{module}")
raise AttributeError(f"导入{desc}失败,未找到模块方法:{module}") raise AttributeError(f"导入{desc}失败,未找到模块方法:{module}")
def get_random_character() -> str: def get_random_character() -> str:
"""生成随机字符串""" """生成随机字符串"""
return uuid.uuid4().hex return uuid.uuid4().hex
def get_parent_id_map(model_list: Sequence[DeclarativeBase]) -> Dict[int, int]: def get_parent_id_map(model_list: Sequence[DeclarativeBase]) -> Dict[int, int]:
""" """
获取父级ID映射字典 获取父级ID映射字典
@@ -74,6 +82,7 @@ def get_parent_id_map(model_list: Sequence[DeclarativeBase]) -> Dict[int, int]:
""" """
return {item.id: item.parent_id for item in model_list} return {item.id: item.parent_id for item in model_list}
def get_parent_recursion( def get_parent_recursion(
id: int, id: int,
id_map: Dict[int, int], id_map: Dict[int, int],
@@ -95,6 +104,7 @@ def get_parent_recursion(
get_parent_recursion(parent_id, id_map, ids) get_parent_recursion(parent_id, id_map, ids)
return ids return ids
def get_child_id_map(model_list: Sequence[DeclarativeBase]) -> Dict[int, List[int]]: def get_child_id_map(model_list: Sequence[DeclarativeBase]) -> Dict[int, List[int]]:
""" """
获取子级ID映射字典 获取子级ID映射字典
@@ -108,6 +118,7 @@ def get_child_id_map(model_list: Sequence[DeclarativeBase]) -> Dict[int, List[in
data_map.setdefault(model.parent_id, []).append(model.id) data_map.setdefault(model.parent_id, []).append(model.id)
return data_map return data_map
def get_child_recursion( def get_child_recursion(
id: int, id: int,
id_map: Dict[int, List[int]], id_map: Dict[int, List[int]],
@@ -197,7 +208,99 @@ def bytes2human(n: int, format_str: str = '%(value).1f%(symbol)s') -> str:
return format_str % locals() return format_str % locals()
return format_str % dict(symbol=symbols[0], value=n) return format_str % dict(symbol=symbols[0], value=n)
def bytes2file_response(bytes_info: bytes):
def bytes2file_response(bytes_info: bytes) -> Generator[bytes, Any, None]:
"""生成文件响应""" """生成文件响应"""
yield bytes_info yield bytes_info
def get_filepath_from_url(url: str) -> Path:
"""
工具方法:根据请求参数获取文件路径
:param url: 请求参数中的url参数
:return: 文件路径
"""
file_info = url.split('?')[1].split('&')
task_id = file_info[0].split('=')[1]
file_name = file_info[1].split('=')[1]
task_path = file_info[2].split('=')[1]
filepath = setting.settings.STATIC_ROOT.joinpath(task_path, task_id, file_name)
return filepath
def export_list2excel(list_data: List) -> Any:
"""
工具方法:将需要导出的list数据转化为对应excel的二进制数据
:param list_data: 数据列表
:return: 字典信息对应excel的二进制数据
"""
df = pd.DataFrame(list_data)
binary_data = io.BytesIO()
df.to_excel(binary_data, index=False, engine='openpyxl')
binary_data = binary_data.getvalue()
return binary_data
def get_excel_template(header_list: List, selector_header_list: List, option_list: List[dict]) -> Any:
"""
工具方法:将需要导出的list数据转化为对应excel的二进制数据
:param header_list: 表头数据列表
:param selector_header_list: 需要设置为选择器格式的表头数据列表
:param option_list: 选择器格式的表头预设的选项列表
:return: 模板excel的二进制数据
"""
# 创建Excel工作簿
wb = Workbook()
# 选择默认的活动工作表
ws = wb.active
# 设置表头文字
headers = header_list
# 设置表头背景样式为灰色,前景色为白色
header_fill = PatternFill(start_color='ababab', end_color='ababab', fill_type='solid')
# 将表头写入第一行
for col_num, header in enumerate(headers, 1):
cell = ws.cell(row=1, column=col_num)
cell.value = header
cell.fill = header_fill
# 设置列宽度为16
ws.column_dimensions[chr(64 + col_num)].width = 12
# 设置水平居中对齐
cell.alignment = Alignment(horizontal='center')
# 设置选择器的预设选项
options = option_list
# 获取selector_header的字母索引
for selector_header in selector_header_list:
column_selector_header_index = headers.index(selector_header) + 1
# 创建数据有效性规则
header_option = []
for option in options:
if option.get(selector_header):
header_option = option.get(selector_header)
dv = DataValidation(type='list', formula1=f'"{",".join(header_option)}"')
# 设置数据有效性规则的起始单元格和结束单元格
dv.add(
f'{get_column_letter(column_selector_header_index)}2:{get_column_letter(column_selector_header_index)}1048576'
)
# 添加数据有效性规则到工作表
ws.add_data_validation(dv)
# 保存Excel文件为字节类型的数据
file = io.BytesIO()
wb.save(file)
file.seek(0)
# 读取字节数据
excel_data = file.getvalue()
return excel_data
+5
View File
@@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
from rich import get_console
console = get_console()
+12 -11
View File
@@ -1,17 +1,18 @@
import re import re
from datetime import datetime from datetime import datetime
from typing import List from typing import List
from config.constant import GenConstant
from config.env import GenConfig from app.common.constant import GenConstant
from module_generator.entity.vo.gen_vo import GenTableColumnModel, GenTableModel from app.config.setting import settings
from utils.string_util import StringUtil from app.api.v1.module_generator.gencode.schema import GenTableColumnSchema, GenTableSchema
from .string_util import StringUtil
class GenUtils: class GenUtils:
"""代码生成器工具类""" """代码生成器工具类"""
@classmethod @classmethod
def init_table(cls, gen_table: GenTableModel, oper_name: str) -> None: def init_table(cls, gen_table: GenTableSchema, oper_name: str) -> None:
""" """
初始化表信息 初始化表信息
@@ -20,18 +21,18 @@ class GenUtils:
:return: :return:
""" """
gen_table.class_name = cls.convert_class_name(gen_table.table_name) gen_table.class_name = cls.convert_class_name(gen_table.table_name)
gen_table.package_name = GenConfig.package_name gen_table.package_name = settings.package_name
gen_table.module_name = cls.get_module_name(GenConfig.package_name) gen_table.module_name = cls.get_module_name(settings.package_name)
gen_table.business_name = cls.get_business_name(gen_table.table_name) gen_table.business_name = cls.get_business_name(gen_table.table_name)
gen_table.function_name = cls.replace_text(gen_table.table_comment) gen_table.function_name = cls.replace_text(gen_table.table_comment)
gen_table.function_author = GenConfig.author gen_table.function_author = settings.author
gen_table.create_by = oper_name gen_table.create_by = oper_name
gen_table.create_time = datetime.now() gen_table.create_time = datetime.now()
gen_table.update_by = oper_name gen_table.update_by = oper_name
gen_table.update_time = datetime.now() gen_table.update_time = datetime.now()
@classmethod @classmethod
def init_column_field(cls, column: GenTableColumnModel, table: GenTableModel) -> None: def init_column_field(cls, column: GenTableColumnSchema, table: GenTableSchema) -> None:
""" """
初始化列属性字段 初始化列属性字段
@@ -143,8 +144,8 @@ class GenUtils:
param table_name: 业务表名 param table_name: 业务表名
:return: Python类名 :return: Python类名
""" """
auto_remove_pre = GenConfig.auto_remove_pre auto_remove_pre = settings.auto_remove_pre
table_prefix = GenConfig.table_prefix table_prefix = settings.table_prefix
if auto_remove_pre and table_prefix: if auto_remove_pre and table_prefix:
search_list = table_prefix.split(',') search_list = table_prefix.split(',')
table_name = cls.replace_first(table_name, search_list) table_name = cls.replace_first(table_name, search_list)
-186
View File
@@ -1,186 +0,0 @@
import os
from typing import List, Dict, Any
from click.types import convert_type
from sqlalchemy import Boolean
from module_gen.constants.gen_constants import GenConstants
from module_gen.entity.do.gen_table_column_do import GenTableColumn
from module_gen.entity.do.gen_table_do import GenTable
from module_gen.entity.vo.gen_table_vo import GenTableModel
from module_gen.entity.vo.gen_table_column_vo import GenTableColumnModel
class GenUtils:
"""代码生成器 工具类"""
@classmethod
def init_table(cls, table: GenTableModel, columns: List[GenTableColumnModel]) -> None:
"""初始化表信息"""
table.class_name = cls.convert_class_name(table.table_name)
table.package_name = cls.get_package_name(table.table_name)
table.module_name = cls.get_module_name(table.table_name)
table.business_name = cls.get_business_name(table.table_name)
table.function_name = table.table_comment
table.function_author = "FluxAdmin"
# 初始化列属性字段
for column in columns:
cls.init_column_field(column, table)
# 设置主键列信息
# for column in columns:
# if column.is_pk == "1":
# table.pk_column = column
# break
@classmethod
def init_column_field(cls, column: GenTableColumnModel, table: GenTableModel):
data_type = cls.get_db_type(column.column_type)
column_name = column.column_name
column.table_id = table.table_id
# 设置python字段名
column.python_field = column_name
# 设置默认类型
column.python_type = GenConstants.MYSQL_TO_PYTHON.get(data_type.upper(), "Any")
column.query_type = GenConstants.QUERY_EQ
if data_type in GenConstants.TYPE_STRING or data_type in GenConstants.TYPE_TEXT:
# 字符串长度超过500设置为文本域
column_length = cls.get_column_length(column.column_type)
html_type = GenConstants.HTML_TEXTAREA if column_length >= 500 or (data_type in GenConstants.TYPE_TEXT) \
else GenConstants.HTML_INPUT
column.html_type = html_type
elif data_type in GenConstants.TYPE_DATE_TIME:
column.html_type = GenConstants.HTML_DATETIME
elif data_type in GenConstants.TYPE_NUMBER:
column.html_type = GenConstants.HTML_INPUT
# 插入字段
if column.column_name not in GenConstants.COLUMN_NAME_NOT_EDIT and not column.is_pk == '1':
column.is_insert = GenConstants.REQUIRE
# 编辑字段
if column.column_name not in GenConstants.COLUMN_NAME_NOT_EDIT and not column.is_pk == '1':
column.is_edit = GenConstants.REQUIRE
# 列表字段
if column.column_name not in GenConstants.COLUMN_NAME_NOT_LIST and not column.is_pk == '1':
column.is_list = GenConstants.REQUIRE
# 查询字段
if column.column_name not in GenConstants.COLUMN_NAME_NOT_QUERY and not column.is_pk == '1':
column.is_query = GenConstants.REQUIRE
@classmethod
def convert_html_type(cls, column_name: str) -> str:
# 状态字段初始化
if column_name.lower().endswith('_status'):
return GenConstants.HTML_RADIO
# 类型字段初始化
elif column_name.lower().endswith('_type'):
return GenConstants.HTML_SELECT
# 内容字段初始化
elif column_name.lower().endswith('_content'):
return GenConstants.HTML_EDITOR
# 文件字段初始化
elif column_name.lower().endswith('_file'):
return GenConstants.HTML_FILE_UPLOAD
# 图片字段初始化
elif column_name.lower().endswith('_image'):
return GenConstants.HTML_IMAGE_UPLOAD
else:
return GenConstants.HTML_INPUT
@classmethod
def get_db_type(cls, column_type):
# 解析数据库类型逻辑,示例返回列的类型
return column_type.split('(')[0]
@classmethod
def get_column_length(cls, column_type):
# 获取列的长度逻辑,这里简化为返回一个默认值
if '(' in column_type:
return int(column_type.split('(')[1].split(')')[0])
return 0
@classmethod
def convert_class_name(cls, table_name: str) -> str:
"""表名转换成Java类名"""
return ''.join(word.title() for word in table_name.lower().split('_'))
@classmethod
def convert_python_field(cls, column_name: str) -> str:
"""列名转换成Python属性名"""
# words = column_name.lower().split('_')
# return words[0] + ''.join(word.title() for word in words[1:])
return column_name.lower()
@classmethod
def get_package_name(cls, table_name: str) -> str:
"""获取包名"""
return "module_admin" # 可配置的包名
@classmethod
def get_module_name(cls, table_name: str) -> str:
"""获取模块名"""
return table_name.split('_')[0]
@classmethod
def get_business_name(cls, table_name: str) -> str:
"""获取业务名"""
words = table_name.split('_')
return words[1] if len(words) > 1 else words[0]
@classmethod
def get_template_path(cls, tpl_category: str) -> Dict[str, str]:
"""获取模板信息"""
templates = {
# Python相关模板
'controller.py': 'python/controller_template.j2',
'do.py': 'python/model_do_template.j2',
'vo.py': 'python/model_vo_template.j2',
'service.py': 'python/service_template.j2',
'dao.py': 'python/dao_template.j2',
# Vue相关模板
'index.vue': 'vue/index.vue.j2',
'api.js': 'vue/api.js.j2',
# SQL脚本模板
'sql': 'sql/sql.j2',
}
# 树表特殊处理
# if tpl_category == "tree":
# templates.update({
# 'entity': 'java/tree_entity.java.vm',
# 'mapper': 'java/tree_mapper.java.vm',
# 'service': 'java/tree_service.java.vm',
# 'service_impl': 'java/tree_service_impl.java.vm',
# 'controller': 'java/tree_controller.java.vm'
# })
return templates
@classmethod
def get_file_name(cls, template_name: str, table) -> str:
"""获取文件名"""
target_file_name = "unknown_file_name"
if template_name.endswith("controller.py"):
target_file_name = f"python/controller/{table.table_name}_{template_name}"
elif template_name.endswith("do.py"):
target_file_name = f"python/entity/do/{table.table_name}_{template_name}"
elif template_name.endswith("vo.py"):
target_file_name = f"python/entity/vo/{table.table_name}_{template_name}"
elif template_name.endswith("service.py"):
target_file_name = f"python/service/{table.table_name}_{template_name}"
elif template_name.endswith("dao.py"):
target_file_name = f"python/dao/{table.table_name}_{template_name}"
elif template_name.endswith('index.vue'):
target_file_name = f'vue/views/{table.module_name}/{table.business_name}/index.vue'
if template_name.endswith('api.js'):
target_file_name = f'vue/api/{table.module_name}/{table.business_name}.js'
if template_name.endswith('sql'):
target_file_name = f'sql/{table.business_name}.sql'
return target_file_name
+1 -1
View File
@@ -1,6 +1,6 @@
import re import re
from module_gen.constants.gen_constants import GenConstants from app.common.constant import GenConstants
def snake_to_pascal_case(value): def snake_to_pascal_case(value):
+101 -26
View File
@@ -1,6 +1,6 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
from typing import List from typing import Dict, List
from app.common.constant import CommonConstant from app.common.constant import CommonConstant
@@ -18,8 +18,15 @@ class StringUtil:
:return: 校验结果 :return: 校验结果
""" """
if string is None: if string is None:
return False
str_len = len(string)
if str_len == 0:
return True
else:
for i in range(str_len):
if string[i] != ' ':
return False
return True return True
return not bool(string.strip())
@classmethod @classmethod
def is_empty(cls, string) -> bool: def is_empty(cls, string) -> bool:
@@ -29,49 +36,82 @@ class StringUtil:
:param string: 需要校验的字符串 :param string: 需要校验的字符串
:return: 校验结果 :return: 校验结果
""" """
return not bool(string) return string is None or len(string) == 0
@classmethod @classmethod
def is_http(cls, link: str) -> bool: def is_not_empty(cls, string: str) -> bool:
"""
校验字符串是否不是''和None
:param string: 需要校验的字符串
:return: 校验结果
"""
return not cls.is_empty(string)
@classmethod
def is_http(cls, link: str):
""" """
判断是否为http(s)://开头 判断是否为http(s)://开头
:param link: 链接 :param link: 链接
:return: 是否为http(s)://开头 :return: 是否为http(s)://开头
""" """
if not link: return link.startswith(CommonConstant.HTTP) or link.startswith(CommonConstant.HTTPS)
return False
return link.lower().startswith((CommonConstant.HTTP.lower(), CommonConstant.HTTPS.lower()))
@classmethod @classmethod
def contains_ignore_case(cls, search_str: str, compare_str: str) -> bool: def contains_ignore_case(cls, search_str: str, compare_str: str):
""" """
查找指定字符串是否包含指定字符串同时串忽略大小写 查找指定字符串是否包含指定字符串同时忽略大小写
:param search_str: 查找的字符串 :param search_str: 查找的字符串
:param compare_str: 比对的字符串 :param compare_str: 比对的字符串
:return: 查找结果 :return: 查找结果
""" """
if not (search_str and compare_str): if compare_str and search_str:
return False return compare_str.lower() in search_str.lower()
return compare_str.lower() in search_str.lower() return False
@classmethod @classmethod
def contains_any_ignore_case(cls, search_str: str, compare_str_list: List[str]) -> bool: def contains_any_ignore_case(cls, search_str: str, compare_str_list: List[str]):
""" """
查找指定字符串是否包含指定字符串列表中的任意一个字符串同时串忽略大小写 查找指定字符串是否包含指定字符串列表中的任意一个字符串同时忽略大小写
:param search_str: 查找的字符串 :param search_str: 查找的字符串
:param compare_str_list: 比对的字符串列表 :param compare_str_list: 比对的字符串列表
:return: 查找结果 :return: 查找结果
""" """
if not (search_str and compare_str_list): if search_str and compare_str_list:
return False return any([cls.contains_ignore_case(search_str, compare_str) for compare_str in compare_str_list])
search_str_lower = search_str.lower() return False
return any(comp_str.lower() in search_str_lower for comp_str in compare_str_list if comp_str)
@classmethod @classmethod
def startswith_case(cls, search_str: str, compare_str: str) -> bool: def equals_ignore_case(cls, search_str: str, compare_str: str):
"""
比较两个字符串是否相等同时忽略大小写
:param search_str: 查找的字符串
:param compare_str: 比对的字符串
:return: 比较结果
"""
if search_str and compare_str:
return search_str.lower() == compare_str.lower()
return False
@classmethod
def equals_any_ignore_case(cls, search_str: str, compare_str_list: List[str]):
"""
比较指定字符串是否与指定字符串列表中的任意一个字符串相等同时忽略大小写
:param search_str: 查找的字符串
:param compare_str_list: 比对的字符串列表
:return: 比较结果
"""
if search_str and compare_str_list:
return any([cls.equals_ignore_case(search_str, compare_str) for compare_str in compare_str_list])
return False
@classmethod
def startswith_case(cls, search_str: str, compare_str: str):
""" """
查找指定字符串是否以指定字符串开头 查找指定字符串是否以指定字符串开头
@@ -79,12 +119,12 @@ class StringUtil:
:param compare_str: 比对的字符串 :param compare_str: 比对的字符串
:return: 查找结果 :return: 查找结果
""" """
if not (search_str and compare_str): if compare_str and search_str:
return False return search_str.startswith(compare_str)
return search_str.startswith(compare_str) return False
@classmethod @classmethod
def startswith_any_case(cls, search_str: str, compare_str_list: List[str]) -> bool: def startswith_any_case(cls, search_str: str, compare_str_list: List[str]):
""" """
查找指定字符串是否以指定字符串列表中的任意一个字符串开头 查找指定字符串是否以指定字符串列表中的任意一个字符串开头
@@ -92,6 +132,41 @@ class StringUtil:
:param compare_str_list: 比对的字符串列表 :param compare_str_list: 比对的字符串列表
:return: 查找结果 :return: 查找结果
""" """
if not (search_str and compare_str_list): if search_str and compare_str_list:
return False return any([cls.startswith_case(search_str, compare_str) for compare_str in compare_str_list])
return any(search_str.startswith(comp_str) for comp_str in compare_str_list if comp_str) return False
@classmethod
def convert_to_camel_case(cls, name: str) -> str:
"""
将下划线大写方式命名的字符串转换为驼峰式。如果转换前的下划线大写方式命名的字符串为空,则返回空字符串
:param name: 转换前的下划线大写方式命名的字符串
:return: 转换后的驼峰式命名的字符串
"""
if not name:
return ''
if '_' not in name:
return name[0].upper() + name[1:]
parts = name.split('_')
result = []
for part in parts:
if not part:
continue
result.append(part[0].upper() + part[1:].lower())
return ''.join(result)
@classmethod
def get_mapping_value_by_key_ignore_case(cls, mapping: Dict[str, str], key: str) -> str:
"""
根据忽略大小写的键获取字典中的对应的值
param mapping: 字典
param key: 字典的键
:return: 字典键对应的值
"""
for k, v in mapping.items():
if key.lower() == k.lower():
return v
return ''
+1 -1
View File
@@ -1,7 +1,7 @@
from jinja2 import Environment, FileSystemLoader, select_autoescape from jinja2 import Environment, FileSystemLoader, select_autoescape
import os import os
from module_gen.utils.jinja2_tools import snake_to_pascal_case, snake_to_camel, snake_2_colon, is_base_column, \ from .jinja2_tools import snake_to_pascal_case, snake_to_camel, snake_2_colon, is_base_column, \
get_sqlalchemy_type, get_column_options get_sqlalchemy_type, get_column_options
@@ -0,0 +1,576 @@
<!-- 演示示例 -->
<template>
<div class="app-container">
<!-- 搜索区域 -->
<div class="search-container">
<el-form ref="queryFormRef" :model="queryFormData" :inline="true" label-suffix=":">
<el-form-item prop="name" label="名称">
<el-input v-model="queryFormData.name" placeholder="请输入名称" clearable />
</el-form-item>
<el-form-item prop="status" label="状态">
<el-select v-model="queryFormData.status" placeholder="请选择状态" style="width: 167.5px" clearable>
<el-option value="true" label="启用" />
<el-option value="false" label="停用" />
</el-select>
</el-form-item>
<!-- 时间范围,收起状态下隐藏 -->
<el-form-item v-if="isExpand" prop="start_time" label="创建时间">
<DatePicker
v-model="dateRange"
@update:model-value="handleDateRangeChange"
/>
</el-form-item>
<!-- 查询、重置、展开/收起按钮 -->
<el-form-item class="search-buttons">
<el-button type="primary" icon="search" @click="handleQuery">
查询
</el-button>
<el-button icon="refresh" @click="handleResetQuery">
重置
</el-button>
<!-- 展开/收起 -->
<template v-if="isExpandable">
<el-link class="ml-3" type="primary" underline="never" @click="isExpand = !isExpand">
{{ isExpand ? "收起" : "展开" }}
<el-icon>
<template v-if="isExpand">
<ArrowUp />
</template>
<template v-else>
<ArrowDown />
</template>
</el-icon>
</el-link>
</template>
</el-form-item>
</el-form>
</div>
<!-- 内容区域 -->
<el-card shadow="hover" class="data-table">
<template #header>
<div class="card-header">
<span>
<el-tooltip content="流程列表">
<QuestionFilled class="w-4 h-4 mx-1" />
</el-tooltip>
演示示例列表
</span>
</div>
</template>
<!-- 功能区域 -->
<div class="data-table__toolbar">
<div class="data-table__toolbar--actions">
<el-button type="success" icon="plus" @click="handleOpenDialog('create')">新增</el-button>
<el-button type="danger" icon="delete" :disabled="selectIds.length === 0"
@click="handleDelete(selectIds)">批量删除</el-button>
<el-dropdown trigger="click">
<el-button type="default" :disabled="selectIds.length === 0" icon="ArrowDown">更多</el-button>
<template #dropdown>
<el-dropdown-menu>
<el-dropdown-item icon="Check" @click="handleMoreClick(true)">批量启用</el-dropdown-item>
<el-dropdown-item icon="CircleClose"
@click="handleMoreClick(false)">批量停用</el-dropdown-item>
</el-dropdown-menu>
</template>
</el-dropdown>
</div>
<div class="data-table__toolbar--tools">
<el-tooltip content="导入">
<el-button type="info" icon="upload" circle @click="handleOpenImportDialog" />
</el-tooltip>
<el-tooltip content="导出">
<el-button type="warning" icon="download" circle @click="handleExport" />
</el-tooltip>
<el-tooltip content="刷新">
<el-button type="primary" icon="refresh" circle @click="handleRefresh" />
</el-tooltip>
<el-tooltip content="列表筛选">
<el-dropdown trigger="click">
<el-button type="default" icon="operation" circle />
<template #dropdown>
<el-dropdown-menu>
<el-dropdown-item v-for="column in tableColumns" :key="column.prop"
:command="column">
<el-checkbox v-model="column.show">
{{ column.label }}
</el-checkbox>
</el-dropdown-item>
</el-dropdown-menu>
</template>
</el-dropdown>
</el-tooltip>
</div>
</div>
<!-- 表格区域:系统配置列表 -->
<el-table ref="dataTableRef" v-loading="loading" :data="pageTableData" highlight-current-row
class="data-table__content" height="450" border stripe @selection-change="handleSelectionChange">
<template #empty>
<el-empty :image-size="80" description="暂无数据" />
</template>
<el-table-column v-if="tableColumns.find(col => col.prop === 'selection')?.show" type="selection"
min-width="55" align="center" />
<el-table-column v-if="tableColumns.find(col => col.prop === 'index')?.show" fixed label="序号"
min-width="60">
<template #default="scope">
{{ (queryFormData.page_no - 1) * queryFormData.page_size + scope.$index + 1 }}
</template>
</el-table-column>
<el-table-column v-if="tableColumns.find(col => col.prop === 'name')?.show" label="名称"
prop="name" min-width="140" />
<el-table-column v-if="tableColumns.find(col => col.prop === 'status')?.show" label="状态" prop="status"
min-width="80">
<template #default="scope">
<el-tag :type="scope.row.status === true ? 'success' : 'danger'">
{{ scope.row.status === true ? "启用" : "停用" }}
</el-tag>
</template>
</el-table-column>
<el-table-column v-if="tableColumns.find(col => col.prop === 'description')?.show" label="描述"
prop="description" min-width="140" />
<el-table-column v-if="tableColumns.find(col => col.prop === 'created_at')?.show" label="创建时间"
prop="created_at" min-width="180" sortable />
<el-table-column v-if="tableColumns.find(col => col.prop === 'updated_at')?.show" label="更新时间"
prop="updated_at" min-width="180" sortable />
<el-table-column v-if="tableColumns.find(col => col.prop === 'creator')?.show" key="creator" label="创建人"
min-width="100">
<template #default="scope">
{{ scope.row.creator?.name }}
</template>
</el-table-column>
<el-table-column v-if="tableColumns.find(col => col.prop === 'operation')?.show" fixed="right"
label="操作" align="center" min-width="200">
<template #default="scope">
<el-button type="info" size="small" link icon="document"
@click="handleOpenDialog('detail', scope.row.id)">详情</el-button>
<el-button type="primary" size="small" link icon="edit"
@click="handleOpenDialog('update', scope.row.id)">编辑</el-button>
<el-button type="danger" size="small" link icon="delete"
@click="handleDelete([scope.row.id])">删除</el-button>
</template>
</el-table-column>
</el-table>
<!-- 分页区域 -->
<template #footer>
<pagination v-model:total="total" v-model:page="queryFormData.page_no"
v-model:limit="queryFormData.page_size" @pagination="loadingData" />
</template>
</el-card>
<!-- 弹窗区域 -->
<el-dialog v-model="dialogVisible.visible" :title="dialogVisible.title" @close="handleCloseDialog">
<!-- 详情 -->
<template v-if="dialogVisible.type === 'detail'">
<el-descriptions :column="4" border>
<el-descriptions-item label="名称" :span="2">
{{ detailFormData.name }}
</el-descriptions-item>
<el-descriptions-item label="状态" :span="2">
<el-tag :type="detailFormData.status ? 'success' : 'danger'">
{{ detailFormData.status ? '启用' : '停用' }}
</el-tag>
</el-descriptions-item>
<el-descriptions-item label="描述" :span="2">
{{ detailFormData.description }}
</el-descriptions-item>
<el-descriptions-item label="创建人" :span="2">
{{ detailFormData.creator?.name }}
</el-descriptions-item>
<el-descriptions-item label="创建时间" :span="2">
{{ detailFormData.created_at }}
</el-descriptions-item>
<el-descriptions-item label="更新时间" :span="2">
{{ detailFormData.updated_at }}
</el-descriptions-item>
</el-descriptions>
</template>
<!-- 新增、编辑表单 -->
<template v-else>
<el-form ref="dataFormRef" :model="formData" :rules="rules" label-suffix=":" label-width="auto"
label-position="right">
<el-form-item label="名称" prop="name">
<el-input v-model="formData.name" placeholder="请输入名称" :maxlength="50" />
</el-form-item>
<el-form-item label="状态" prop="status">
<el-radio-group v-model="formData.status">
<el-radio :value="true">
启用
</el-radio>
<el-radio :value="false">
停用
</el-radio>
</el-radio-group>
</el-form-item>
<el-form-item label="描述" prop="description">
<el-input v-model="formData.description" :rows="4" :maxlength="100" show-word-limit
type="textarea" placeholder="请输入描述" />
</el-form-item>
</el-form>
</template>
<template #footer>
<div class="dialog-footer">
<!-- 详情弹窗不需要确定按钮的提交逻辑 -->
<el-button @click="handleCloseDialog">取消</el-button>
<el-button v-if="dialogVisible.type !== 'detail'" type="primary"
@click="handleSubmit">确定</el-button>
<el-button v-else type="primary" @click="handleCloseDialog">确定</el-button>
</div>
</template>
</el-dialog>
<!-- 用户导入 -->
<ImportModal v-model="importDialogVisible" title="导入数据" @import-success="handleQuery()" @download-template="handleDownloadTemplate" @upload="handleUpload" />
</div>
</template>
<script setup lang="ts">
defineOptions({
name: "Example",
inheritAttrs: false,
});
import { ref, reactive, onMounted } from "vue";
import { ElMessage } from "element-plus";
import { ResultEnum } from "@/enums/api/result.enum";
import ExampleAPI, { ExampleTable, ExampleForm, ExamplePageQuery } from "@/api/demo/example";
import ImportModal from "@/components/Upload/ImportModal.vue";
import DatePicker from "@/components/DatePicker/index.vue";
const emit = defineEmits(['import-success']);
const queryFormRef = ref();
const dataFormRef = ref();
const total = ref(0);
const selectIds = ref<number[]>([]);
const loading = ref(false);
const isExpand = ref(false);
const isExpandable = ref(true);
// 分页表单
const pageTableData = ref<ExampleTable[]>([]);
// 表格列配置
const tableColumns = ref([
{ prop: 'selection', label: '选择框', show: true },
{ prop: 'index', label: '序号', show: true },
{ prop: 'name', label: '名称', show: true },
{ prop: 'status', label: '状态', show: true },
{ prop: 'description', label: '描述', show: true },
{ prop: 'created_at', label: '创建时间', show: true },
{ prop: 'updated_at', label: '更新时间', show: true },
{ prop: 'creator', label: '创建人', show: true },
{ prop: 'operation', label: '操作', show: true }
])
// 详情表单
const detailFormData = ref<ExampleTable>({});
// 日期范围临时变量
const dateRange = ref<[Date, Date] | []>([]);
// 处理日期范围变化
function handleDateRangeChange(range: [Date, Date]) {
dateRange.value = range;
if (range && range.length === 2) {
queryFormData.start_time = range[0].toISOString();
queryFormData.end_time = range[1].toISOString();
} else {
queryFormData.start_time = undefined;
queryFormData.end_time = undefined;
}
}
// 分页查询参数
const queryFormData = reactive<ExamplePageQuery>({
page_no: 1,
page_size: 10,
name: undefined,
status: undefined,
start_time: undefined,
end_time: undefined,
});
// 编辑表单
const formData = reactive<ExampleForm>({
id: undefined,
name: '',
status: true,
description: undefined,
})
// 弹窗状态
const dialogVisible = reactive({
title: "",
visible: false,
type: 'create' as 'create' | 'update' | 'detail',
});
// 表单验证规则
const rules = reactive({
name: [{ required: true, message: "请输入名称", trigger: "blur" }],
status: [{ required: true, message: "请选择状态", trigger: "blur" }],
});
// 导入弹窗显示状态
const importDialogVisible = ref(false);
// 列表刷新
async function handleRefresh() {
await loadingData();
};
// 加载表格数据
async function loadingData() {
loading.value = true;
try {
const response = await ExampleAPI.getExampleList(queryFormData);
pageTableData.value = response.data.data.items;
total.value = response.data.data.total;
}
catch (error: any) {
console.error(error);
}
finally {
loading.value = false;
}
}
// 查询(重置页码后获取数据)
async function handleQuery() {
queryFormData.page_no = 1;
loadingData();
}
// 重置查询
async function handleResetQuery() {
queryFormRef.value.resetFields();
queryFormData.page_no = 1;
// 重置日期范围选择器
dateRange.value = [];
queryFormData.start_time = undefined;
queryFormData.end_time = undefined;
loadingData();
}
// 定义初始表单数据常量
const initialFormData: ExampleForm = {
id: undefined,
name: '',
status: true,
description: '',
}
// 重置表单
async function resetForm() {
if (dataFormRef.value) {
dataFormRef.value.resetFields();
dataFormRef.value.clearValidate();
}
// 完全重置 formData 为初始状态
Object.assign(formData, initialFormData);
}
// 行复选框选中项变化
async function handleSelectionChange(selection: any) {
selectIds.value = selection.map((item: any) => item.id);
}
// 关闭弹窗
async function handleCloseDialog() {
dialogVisible.visible = false;
resetForm();
}
// 打开弹窗
async function handleOpenDialog(type: 'create' | 'update' | 'detail', id?: number) {
dialogVisible.type = type;
if (id) {
const response = await ExampleAPI.getExampleDetail(id);
if (type === 'detail') {
dialogVisible.title = "详情";
Object.assign(detailFormData.value, response.data.data);
} else if (type === 'update') {
dialogVisible.title = "修改";
Object.assign(formData, response.data.data);
}
} else {
dialogVisible.title = "新增公告通知";
formData.id = undefined;
}
dialogVisible.visible = true;
}
// 提交表单(防抖)
async function handleSubmit() {
// 表单校验
dataFormRef.value.validate(async (valid: any) => {
if (valid) {
loading.value = true;
// 根据弹窗传入的参数(deatil\create\update)判断走什么逻辑
const id = formData.id;
if (id) {
try {
await ExampleAPI.updateExample(id, { id, ...formData })
dialogVisible.visible = false;
resetForm();
handleCloseDialog();
handleResetQuery();
} catch (error: any) {
console.error(error);
} finally {
loading.value = false;
}
} else {
try {
await ExampleAPI.createExample(formData)
dialogVisible.visible = false;
resetForm();
handleCloseDialog();
handleResetQuery();
} catch (error: any) {
console.error(error);
} finally {
loading.value = false;
}
}
}
});
}
// 删除、批量删除
async function handleDelete(ids: number[]) {
ElMessageBox.confirm("确认删除该项数据?", "警告", {
confirmButtonText: "确定",
cancelButtonText: "取消",
type: "warning",
}).then(async () => {
try {
loading.value = true;
await ExampleAPI.deleteExample(ids);
handleResetQuery();
} catch (error: any) {
console.error(error);
} finally {
loading.value = false;
}
}).catch(() => {
ElMessageBox.close();
});
}
// 导出
async function handleExport() {
ElMessageBox.confirm('是否确认导出当前系统配置?', '警告', {
confirmButtonText: '确定',
cancelButtonText: '取消',
type: 'warning'
}).then(async () => {
let downloadUrl = '';
try {
loading.value = true;
const response = await ExampleAPI.exportExample(queryFormData);
const fileData = response.data;
const fileName = decodeURI(response.headers["content-disposition"].split(";")[1].split("=")[1]);
const fileType = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet;charset=utf-8";
const blob = new Blob([fileData], { type: fileType });
downloadUrl = window.URL.createObjectURL(blob);
const downloadLink = document.createElement("a");
downloadLink.href = downloadUrl;
downloadLink.download = fileName;
document.body.appendChild(downloadLink);
downloadLink.click();
document.body.removeChild(downloadLink);
} catch (error: any) {
// 错误信息已经在响应拦截器中处理并显示
console.error('导出失败:', error);
} finally {
if (downloadUrl) {
window.URL.revokeObjectURL(downloadUrl);
}
loading.value = false;
}
}).catch(() => {
ElMessageBox.close();
});
}
// 处理上传
const handleUpload = async (formData: FormData, file: File) => {
try {
const response = await ExampleAPI.importExample(formData);
if (response.data.code === ResultEnum.SUCCESS) {
ElMessage.success(`${response.data.msg},${response.data.data}`);
emit('import-success');
}
} catch (error: any) {
console.error(error);
}
};
// 下载导入模板
const handleDownloadTemplate = () => {
ExampleAPI.downloadTemplate().then((response: any) => {
const fileData = response.data;
const fileName = decodeURI(response.headers['content-disposition'].split('; ')[1].split('=')[1]);
const fileType = 'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet;charset=utf-8';
const blob = new Blob([fileData], { type: fileType });
const downloadUrl = window.URL.createObjectURL(blob);
const downloadLink = document.createElement('a');
downloadLink.href = downloadUrl;
downloadLink.download = fileName;
document.body.appendChild(downloadLink);
downloadLink.click();
document.body.removeChild(downloadLink);
window.URL.revokeObjectURL(downloadUrl);
});
};
// 打开导入弹窗
function handleOpenImportDialog() {
importDialogVisible.value = true;
}
// 批量启用/停用
async function handleMoreClick(status: boolean) {
if (selectIds.value.length) {
ElMessageBox.confirm(`确认${status ? '启用' : '停用'}该项数据?`, "警告", {
confirmButtonText: "确定",
cancelButtonText: "取消",
type: "warning",
}).then(async () => {
try {
loading.value = true;
await ExampleAPI.batchAvailableExample({ ids: selectIds.value, status });
handleResetQuery();
} catch (error: any) {
console.error(error);
} finally {
loading.value = false;
}
}).catch(() => {
ElMessageBox.close();
});
}
}
onMounted(() => {
loadingData();
});
</script>
<style lang="scss" scoped></style>