Compare commits
159 Commits
server
...
c3e96ae61b
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c3e96ae61b | ||
|
|
b254fa0985 | ||
|
|
f3625f994c | ||
|
|
78fb894307 | ||
|
|
91b4ae2287 | ||
|
|
2f48934e66 | ||
|
|
fed5694796 | ||
|
|
04c3c1412c | ||
|
|
2adf78a6ed | ||
|
|
4e27b98a84 | ||
|
|
6c1bf61ff7 | ||
|
|
de9c8e4ffe | ||
|
|
5ecca9f0bc | ||
|
|
2f64086177 | ||
|
|
bdaf72d58b | ||
|
|
ca012893f7 | ||
|
|
ae39c45a73 | ||
|
|
b5d2cd7935 | ||
|
|
a12f80492d | ||
|
|
71405def14 | ||
|
|
5f6e7523cf | ||
|
|
f809825a7d | ||
|
|
75cc105ebc | ||
|
|
e397bcc2ca | ||
|
|
c64fa1cd61 | ||
|
|
b5c6557341 | ||
|
|
ec7d8c0a3d | ||
|
|
0292bf5138 | ||
|
|
7f93ed6d09 | ||
|
|
5ec950cc9f | ||
|
|
0271942ba5 | ||
|
|
250e060271 | ||
|
|
7b36bc774e | ||
|
|
4e5c43fad5 | ||
|
|
62a1d03eac | ||
|
|
94230cad16 | ||
|
|
0c601934a0 | ||
|
|
5cc306eb0a | ||
|
|
cc08b164d0 | ||
|
|
24c77a990a | ||
|
|
15c4223f2c | ||
|
|
b975de02da | ||
|
|
46d343fb63 | ||
|
|
0c39f2f5b9 | ||
|
|
c7c9ed925b | ||
|
|
f917a025e1 | ||
|
|
a9ab130d43 | ||
|
|
525fc55cef | ||
|
|
3fd9cf9100 | ||
|
|
a72e2a2520 | ||
|
|
01d2e6c76a | ||
|
|
8a6a6574bb | ||
|
|
9025437a37 | ||
|
|
4623e3fa1c | ||
|
|
f97245b814 | ||
|
|
97cdb5cc68 | ||
|
|
d8a11e4949 | ||
|
|
05c1a5c1e1 | ||
|
|
f3fa5f1a68 | ||
|
|
de2e8952b5 | ||
|
|
ecafb7eb13 | ||
|
|
b82897ca3a | ||
|
|
e6bbb0bb49 | ||
|
|
bce586697b | ||
|
|
e486d36a80 | ||
|
|
d6e325fe9e | ||
|
|
b08a771a61 | ||
|
|
8caaaa5bbc | ||
|
|
3f5fedb9ed | ||
|
|
53844a3a09 | ||
|
|
53014bb381 | ||
|
|
03bd0b6d03 | ||
|
|
b4927a8952 | ||
|
|
b14b2ecf22 | ||
|
|
42c0e4f5c2 | ||
|
|
88a82ed771 | ||
|
|
895983ac20 | ||
|
|
680fa905f8 | ||
|
|
ce0f908d20 | ||
|
|
e4ea1f168c | ||
|
|
915f994c45 | ||
|
|
762f866175 | ||
|
|
9428c6b785 | ||
|
|
25d75f40c7 | ||
|
|
06e7455630 | ||
|
|
396d3f6f47 | ||
|
|
64d7414b04 | ||
|
|
9a5282f39c | ||
|
|
f21a4c954f | ||
|
|
2e2cbb4976 | ||
|
|
2e8278636b | ||
|
|
bbc0df29bf | ||
|
|
b20e7aa595 | ||
|
|
17615aa17d | ||
|
|
07e2999323 | ||
|
|
749c84a62b | ||
|
|
ea0013b99c | ||
|
|
4f8aff5fc4 | ||
|
|
9193f10e3e | ||
|
|
ea08478a37 | ||
|
|
4782981169 | ||
|
|
939a7f8e8f | ||
|
|
e9a121cfeb | ||
|
|
d4b9a76aa5 | ||
|
|
9114f3d4c7 | ||
|
|
9cb77c251a | ||
|
|
e6a5a36bc0 | ||
|
|
b2c570f607 | ||
|
|
994ec6644a | ||
|
|
d6d3d27b2d | ||
|
|
eb6ff93150 | ||
|
|
215b4074e0 | ||
|
|
4544483fc5 | ||
|
|
3266a6fc09 | ||
|
|
a9b06140d0 | ||
|
|
663b73af2e | ||
|
|
4a2f1f5dcd | ||
|
|
8ac39cf007 | ||
|
|
faad88dfcd | ||
|
|
d3a25f3a4b | ||
|
|
b801bd314b | ||
|
|
476502fc0d | ||
|
|
975f55d06c | ||
|
|
93373bc61f | ||
|
|
b79a8e1499 | ||
|
|
0124e28d77 | ||
|
|
ad64e44860 | ||
|
|
33d0ed2e01 | ||
|
|
6d4bf85284 | ||
|
|
2b10c013ce | ||
|
|
b28cfbc6fa | ||
|
|
f04dc479bb | ||
|
|
f453234057 | ||
|
|
a6868ec2e5 | ||
|
|
6cd1e46e86 | ||
|
|
836343b29e | ||
|
|
1e438164c1 | ||
|
|
f4864fafd0 | ||
|
|
9798b34717 | ||
|
|
284995d79c | ||
|
|
e18a367abb | ||
|
|
817d13c8f7 | ||
|
|
a72b8f1e4b | ||
|
|
bccd3bf448 | ||
|
|
a67ca2c19c | ||
|
|
2c1e08a271 | ||
|
|
ba4059fe3b | ||
|
|
4050c120d5 | ||
| 156a952b47 | |||
|
|
4173b53b1b | ||
|
|
ab9e87f948 | ||
|
|
5a040366da | ||
| 8789019db2 | |||
|
|
3cb20a4a28 | ||
|
|
a6085a2612 | ||
|
|
4899bc8779 | ||
|
|
e70538e64d | ||
| 39a5390ecd | |||
| cd354f52e6 |
11
.dockerignore
Normal file
11
.dockerignore
Normal file
@@ -0,0 +1,11 @@
|
|||||||
|
.git
|
||||||
|
.gitignore
|
||||||
|
node_modules
|
||||||
|
frontend/node_modules
|
||||||
|
frontend/dist
|
||||||
|
frontend/.vite
|
||||||
|
npm-debug.log*
|
||||||
|
docker-compose*.yml
|
||||||
|
README.md
|
||||||
|
design-qa.md
|
||||||
|
docs
|
||||||
41
.gitignore
vendored
41
.gitignore
vendored
@@ -12,6 +12,8 @@ __pycache__/
|
|||||||
build/
|
build/
|
||||||
develop-eggs/
|
develop-eggs/
|
||||||
dist/
|
dist/
|
||||||
|
node_modules/
|
||||||
|
*.tsbuildinfo
|
||||||
downloads/
|
downloads/
|
||||||
eggs/
|
eggs/
|
||||||
.eggs/
|
.eggs/
|
||||||
@@ -37,6 +39,16 @@ MANIFEST
|
|||||||
pip-log.txt
|
pip-log.txt
|
||||||
pip-delete-this-directory.txt
|
pip-delete-this-directory.txt
|
||||||
|
|
||||||
|
# Runtime data and logs
|
||||||
|
runtime/
|
||||||
|
backend/runtime/
|
||||||
|
backend/storage/
|
||||||
|
logs/
|
||||||
|
backend/logs/
|
||||||
|
*.db
|
||||||
|
*.sqlite
|
||||||
|
*.sqlite3
|
||||||
|
|
||||||
# Unit test / coverage reports
|
# Unit test / coverage reports
|
||||||
htmlcov/
|
htmlcov/
|
||||||
.tox/
|
.tox/
|
||||||
@@ -130,6 +142,7 @@ celerybeat.pid
|
|||||||
|
|
||||||
# Environments
|
# Environments
|
||||||
.env
|
.env
|
||||||
|
!.env.example
|
||||||
.venv
|
.venv
|
||||||
env/
|
env/
|
||||||
venv/
|
venv/
|
||||||
@@ -137,6 +150,16 @@ ENV/
|
|||||||
env.bak/
|
env.bak/
|
||||||
venv.bak/
|
venv.bak/
|
||||||
|
|
||||||
|
# Local backend config (含数据库账号密码等敏感信息,勿提交)
|
||||||
|
backend/config.yaml
|
||||||
|
|
||||||
|
# Agent / IDE 工具产物,不应进版本库
|
||||||
|
.codex-backups/
|
||||||
|
.pnpm-store/
|
||||||
|
.zcode/
|
||||||
|
.claude/
|
||||||
|
CLAUDE.md
|
||||||
|
|
||||||
# Spyder project settings
|
# Spyder project settings
|
||||||
.spyderproject
|
.spyderproject
|
||||||
.spyproject
|
.spyproject
|
||||||
@@ -174,3 +197,21 @@ cython_debug/
|
|||||||
# PyPI configuration file
|
# PyPI configuration file
|
||||||
.pypirc
|
.pypirc
|
||||||
|
|
||||||
|
docker/llamafactory-latest.tar.gz
|
||||||
|
# Compute data - 保留目录结构和 README,忽略子目录内容(日志、模型、数据集等)
|
||||||
|
!docker/compute/data/yg-ft/logs/
|
||||||
|
docker/compute/data/yg-ft/datasets/*
|
||||||
|
docker/compute/data/yg-ft/models/*
|
||||||
|
docker/compute/data/yg-ft/outputs/*
|
||||||
|
docker/compute/data/yg-ft/logs/**
|
||||||
|
!docker/compute/data/yg-ft/logs/compute/
|
||||||
|
!docker/compute/data/yg-ft/logs/training/
|
||||||
|
!docker/compute/data/yg-ft/**/.gitkeep
|
||||||
|
!docker/compute/data/yg-ft/**/README.md
|
||||||
|
|
||||||
|
# Offline deployment bundle - 离线部署包(镜像、运行时等大文件,不提交)
|
||||||
|
docker/offline/
|
||||||
|
|
||||||
|
# MinIO object storage data - 对象存储运行时数据,勿提交,保留目录结构
|
||||||
|
docker/minio/data/*
|
||||||
|
!docker/minio/data/.gitkeep
|
||||||
|
|||||||
164
README.md
164
README.md
@@ -1,133 +1,59 @@
|
|||||||
# YG_FT
|
# YG_FT 模型微调平台
|
||||||
|
|
||||||
远光微调平台 - 面向大语言模型的微调、评测、推理与对比一体化前端。
|
YG_FT 是面向多用户、多租户和多算力节点的模型训练与推理平台,提供数据集、模型、训练、权重合并、推理、评测、算力节点、项目隔离、权限和审计能力。
|
||||||
|
|
||||||
## 技术栈
|
## 架构
|
||||||
|
|
||||||
| 类别 | 技术 | 版本 |
|
```text
|
||||||
|------|------|------|
|
浏览器 -> Frontend Nginx:16801 -> Backend API:17861
|
||||||
| 框架 | Vue 3 | ^3.5.13 |
|
|-> PostgreSQL(元数据、权限、审计、任务状态)
|
||||||
| 语言 | TypeScript | ~5.7.2 |
|
|-> Redis(缓存及任务辅助状态)
|
||||||
| 构建工具 | Vite | ^6.0.7 |
|
|-> MinIO:19000(模型和数据唯一对象源)
|
||||||
| 路由 | Vue Router | ^4.5.0 |
|
|-> Compute API:19100
|
||||||
| 状态管理 | Pinia | ^2.3.0 |
|
|-> Compute Agent/GPU/LLaMA-Factory
|
||||||
| UI 组件库 | Element Plus | ^2.9.1 |
|
`-> File Gateway:19101
|
||||||
| HTTP 客户端 | axios | ^1.7.9 |
|
```
|
||||||
| 图表 | ECharts / vue-echarts | ^6.1.0 / ^8.0.1 |
|
|
||||||
| Markdown | marked + DOMPurify | ^15.0.5 / ^3.2.3 |
|
|
||||||
| 编辑器 | md-editor-v3 | ^5.1.4 |
|
|
||||||
| 工具集 | @vueuse/core | ^11.3.0 |
|
|
||||||
| 样式 | Sass | ^1.83.0 |
|
|
||||||
|
|
||||||
**项目版本**:1.0.0
|
Backend、MinIO 和 Compute 节点可以部署在不同服务器,不依赖跨服务器 Docker 网络,通过 IP、DNS 或负载均衡地址通信。Compute 节点只保存按需准备的本地缓存,MinIO 是模型、数据集、权重、评测报告和训练产物的唯一数据源。
|
||||||
|
|
||||||
## 环境要求
|
## 目录
|
||||||
|
|
||||||
- **Node.js** >= 18(推荐 20 LTS)
|
| 目录 | 作用 |
|
||||||
- **npm** >= 9
|
| --- | --- |
|
||||||
- 后端服务运行于 `http://localhost:7861`(前端通过代理转发,见下文)
|
| `frontend/` | Vue 3、TypeScript、Element Plus 控制台 |
|
||||||
|
| `backend/app/api/v1/` | 平台 REST API |
|
||||||
|
| `backend/app/core/` | 配置、认证、权限和日志 |
|
||||||
|
| `backend/app/db/` | PostgreSQL 访问和初始化 SQL |
|
||||||
|
| `backend/app/modules/` | 系统、资源、审批、数据处理和存储模块 |
|
||||||
|
| `backend/app/workers/` | 节点轮询、任务对账和资源同步 |
|
||||||
|
| `compute/` | Compute API、Agent、GPU 和训练引擎 |
|
||||||
|
| `docker/` | 应用、MinIO、算力服务部署文件 |
|
||||||
|
| `docs/` | 架构、权限、部署和测试文档 |
|
||||||
|
|
||||||
## 快速开始
|
## 端口
|
||||||
|
|
||||||
### 1. 安装依赖
|
| 服务 | 主机端口 | 容器端口 |
|
||||||
|
| --- | ---: | ---: |
|
||||||
|
| Frontend | 16801 | 80 |
|
||||||
|
| Backend API | 17861 | 8000 |
|
||||||
|
| Redis | 16379 | 6379 |
|
||||||
|
| MinIO API/Console | 19000/19001 | 9000/9001 |
|
||||||
|
| Compute API/File Gateway | 19100/19101 | 9100 |
|
||||||
|
|
||||||
|
## 启动
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
cd frontend
|
cd frontend && npm ci && npm run build && cd ..
|
||||||
npm install
|
docker build -f docker/app/Dockerfile.backend -t yg-ft-backend-api:latest .
|
||||||
|
docker build -f docker/app/Dockerfile.frontend -t yg-ft-frontend-runtime:latest .
|
||||||
|
docker build -f docker/compute/Dockerfile.compute -t yg-ft-compute-api:latest .
|
||||||
|
cd docker/minio && docker compose up -d
|
||||||
|
cd ../app && docker compose up -d
|
||||||
|
cd ../compute && docker compose up -d
|
||||||
```
|
```
|
||||||
|
|
||||||
### 2. 启动开发服务器
|
## 权限与性能
|
||||||
|
|
||||||
```bash
|
系统使用角色权限、资源 ACL、用户/项目/租户归属联合校验;删除为软删除,关键操作写入审计。模型合并前准备 base model 和 adapter,结果归档 MinIO;推理前按选择节点准备缓存。远程 PostgreSQL 延迟会影响全量列表和看板,页面慢时应检查浏览器 Network、Nginx、Backend 日志、连接池和节点可达性。
|
||||||
npm run dev
|
|
||||||
```
|
|
||||||
|
|
||||||
开发服务器默认运行在 `http://localhost:6801`。
|
详细部署见 `docker/README.md`,测试见 `测试用例.md`,本次快照见 `docs/20260812/`。
|
||||||
|
|
||||||
### 3. 构建生产包
|
|
||||||
|
|
||||||
```bash
|
|
||||||
npm run build # 类型检查 + 生产构建,产物输出到 dist/
|
|
||||||
npm run preview # 本地预览构建产物
|
|
||||||
```
|
|
||||||
|
|
||||||
### 4. 类型检查
|
|
||||||
|
|
||||||
```bash
|
|
||||||
npm run type-check
|
|
||||||
```
|
|
||||||
|
|
||||||
## 测试
|
|
||||||
|
|
||||||
内置基于 Playwright 的 UI 回归脚本,首次运行前需安装浏览器:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
npx playwright install chromium
|
|
||||||
```
|
|
||||||
|
|
||||||
执行已注册的回归脚本:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
npm run test:data-process-wizard # 数据处理向导
|
|
||||||
npm run test:model-manage # 模型管理
|
|
||||||
npm run test:training-log-layout # 训练日志布局
|
|
||||||
npm run test:page-surface # 页面表层级
|
|
||||||
```
|
|
||||||
|
|
||||||
其余脚本可直接运行:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
node scripts/regression-back-navigation.mjs # 返回导航
|
|
||||||
node scripts/regression-fine-tune-create-ui.mjs # 调优创建 UI
|
|
||||||
```
|
|
||||||
|
|
||||||
> 回归脚本默认连接 `http://localhost:6801`,需先启动开发服务器。
|
|
||||||
|
|
||||||
## 目录结构
|
|
||||||
|
|
||||||
```
|
|
||||||
YG-FT/
|
|
||||||
├── frontend/ # 前端工程(Vue 3 SPA)
|
|
||||||
│ ├── src/
|
|
||||||
│ │ ├── api/ # axios 封装 + 各业务模块 API
|
|
||||||
│ │ ├── components/ # 公共组件
|
|
||||||
│ │ ├── composables/ # 组合式函数
|
|
||||||
│ │ ├── constants/ # 常量与映射表
|
|
||||||
│ │ ├── layouts/ # 主布局
|
|
||||||
│ │ ├── mock/ # Mock 数据与适配器
|
|
||||||
│ │ ├── plugins/ # 第三方插件注册
|
|
||||||
│ │ ├── router/ # 路由配置 + 登录守卫
|
|
||||||
│ │ ├── stores/ # Pinia 状态
|
|
||||||
│ │ ├── styles/ # 全局样式
|
|
||||||
│ │ ├── types/ # TypeScript 类型定义
|
|
||||||
│ │ └── views/ # 业务页面
|
|
||||||
│ ├── scripts/ # UI 回归测试脚本
|
|
||||||
│ ├── public/ # 静态资源
|
|
||||||
│ └── vite.config.ts # Vite 构建与代理配置
|
|
||||||
├── docs/ # 设计文档与视觉走查记录
|
|
||||||
└── design-qa.md # 视觉走查汇总
|
|
||||||
```
|
|
||||||
|
|
||||||
## 端口与代理
|
|
||||||
|
|
||||||
| 服务 | 地址 |
|
|
||||||
|------|------|
|
|
||||||
| 前端开发服务器 | `http://localhost:6801` |
|
|
||||||
| 后端 API | `http://localhost:7861` |
|
|
||||||
|
|
||||||
前端统一使用 `/api` 相对路径发请求,由 Vite 开发代理转发到后端 `http://localhost:7861`(配置见 `frontend/vite.config.ts`)。
|
|
||||||
|
|
||||||
## 业务模块
|
|
||||||
|
|
||||||
| 模块 | 说明 |
|
|
||||||
|------|------|
|
|
||||||
| 登录 | 用户登录鉴权 |
|
|
||||||
| 模型调优 | 微调任务创建与管理 |
|
|
||||||
| 模型评测 | 评测任务与评测维度配置 |
|
|
||||||
| 模型推理 | 在线推理对话 |
|
|
||||||
| 模型对比 | 多模型对话与结果对比 |
|
|
||||||
| 模型管理 | 模型 CRUD 与权重合并 |
|
|
||||||
| 数据集 | 数据集管理与预览 |
|
|
||||||
| 数据处理 | 数据处理任务向导 |
|
|
||||||
| 工具 | 辅助工具集 |
|
|
||||||
| 系统 | 硬件监控、日志、训练日志 |
|
|
||||||
|
|||||||
65
backend/README.md
Normal file
65
backend/README.md
Normal file
@@ -0,0 +1,65 @@
|
|||||||
|
# Backend Service
|
||||||
|
|
||||||
|
后端工程使用 FastAPI,定位为模型微调平台的应用平台服务,负责用户中心、多租户、权限隔离、项目、数据集、模型、训练任务、审批、审计和算力平台编排。
|
||||||
|
|
||||||
|
## 目录结构
|
||||||
|
|
||||||
|
```text
|
||||||
|
backend/
|
||||||
|
app/
|
||||||
|
main.py # FastAPI 应用入口
|
||||||
|
api/v1/ # 对前端暴露的 接口路由
|
||||||
|
core/ # 配置、日志、中间件、权限等基础能力
|
||||||
|
db/ # 数据库连接、迁移集成、事务工具
|
||||||
|
modules/ # 业务模块
|
||||||
|
auth/
|
||||||
|
tenant/
|
||||||
|
project/
|
||||||
|
model/
|
||||||
|
dataset/
|
||||||
|
data_process/
|
||||||
|
fine_tune/
|
||||||
|
eval/
|
||||||
|
inference/
|
||||||
|
approval/
|
||||||
|
audit/
|
||||||
|
compute_gateway/
|
||||||
|
file_gateway/
|
||||||
|
engine_registry/
|
||||||
|
retention/
|
||||||
|
system/
|
||||||
|
schemas/ # Pydantic 入参/出参模型
|
||||||
|
services/ # 跨模块应用服务
|
||||||
|
workers/ # 后台任务入口
|
||||||
|
requirements.txt # 后端第三方依赖
|
||||||
|
logs/ # 本地开发日志目录,生产环境建议挂载到独立日志盘
|
||||||
|
```
|
||||||
|
|
||||||
|
## 本地启动
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd backend
|
||||||
|
python -m venv .venv
|
||||||
|
.venv\Scripts\activate
|
||||||
|
pip install -r requirements.txt
|
||||||
|
uvicorn app.main:app --reload
|
||||||
|
```
|
||||||
|
|
||||||
|
健康检查:
|
||||||
|
|
||||||
|
```text
|
||||||
|
GET /modelTF/health
|
||||||
|
```
|
||||||
|
|
||||||
|
## 日志
|
||||||
|
|
||||||
|
日志模块位于 `app/core/logging.py`,使用说明见 `../docs/backend-logging.md`。
|
||||||
|
|
||||||
|
默认日志文件:
|
||||||
|
|
||||||
|
```text
|
||||||
|
logs/backend-YYYY-MM-DD.log
|
||||||
|
logs/error-YYYY-MM-DD.log
|
||||||
|
```
|
||||||
|
|
||||||
|
文件日志为 JSON Lines 格式,单个文件不超过 20MB,只保存最近 10 天,错误日志按 `ERROR` 级别独立拆分,便于 ELK/日志平台采集。
|
||||||
10
backend/_check_sessions.py
Normal file
10
backend/_check_sessions.py
Normal file
@@ -0,0 +1,10 @@
|
|||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT id, user_id, login_at, logout_at, duration_seconds FROM sessions ORDER BY login_at DESC LIMIT 10"
|
||||||
|
).fetchall()
|
||||||
|
print(f"sessions count: {len(rows)}")
|
||||||
|
for r in rows:
|
||||||
|
print(f" user={r['user_id'][:25]}... login={r['login_at']} logout={r['logout_at']} dur={r['duration_seconds']}")
|
||||||
1
backend/app/__init__.py
Normal file
1
backend/app/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Application package."""
|
||||||
1
backend/app/api/__init__.py
Normal file
1
backend/app/api/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""API package."""
|
||||||
1
backend/app/api/v1/__init__.py
Normal file
1
backend/app/api/v1/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Versioned API package."""
|
||||||
1
backend/app/api/v1/endpoints/__init__.py
Normal file
1
backend/app/api/v1/endpoints/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""API endpoint modules."""
|
||||||
2503
backend/app/api/v1/endpoints/data_process.py
Normal file
2503
backend/app/api/v1/endpoints/data_process.py
Normal file
File diff suppressed because it is too large
Load Diff
40
backend/app/api/v1/endpoints/health.py
Normal file
40
backend/app/api/v1/endpoints/health.py
Normal file
@@ -0,0 +1,40 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
from app.core.logging import get_logger
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
from app.core.config import get_settings
|
||||||
|
from app.modules.storage.minio_store import ObjectStorageError, get_object_storage
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/health")
|
||||||
|
async def health_check() -> dict[str, object]:
|
||||||
|
# Health endpoints are called frequently by Docker and the frontend.
|
||||||
|
# Keep failures visible without emitting one INFO line per probe.
|
||||||
|
logger.debug("health check requested")
|
||||||
|
storage_status: dict[str, object] = {"enabled": get_settings().minio_enabled, "status": "disabled"}
|
||||||
|
if get_settings().minio_enabled:
|
||||||
|
try:
|
||||||
|
get_object_storage().ensure_bucket()
|
||||||
|
storage_status = {
|
||||||
|
"enabled": True,
|
||||||
|
"status": "ready",
|
||||||
|
"endpoint": get_settings().minio_endpoint,
|
||||||
|
"bucket": get_settings().minio_bucket,
|
||||||
|
}
|
||||||
|
except (ObjectStorageError, OSError) as exc:
|
||||||
|
logger.warning("MinIO health check failed: %s", exc)
|
||||||
|
storage_status = {
|
||||||
|
"enabled": True,
|
||||||
|
"status": "unavailable",
|
||||||
|
"endpoint": get_settings().minio_endpoint,
|
||||||
|
"error": str(exc),
|
||||||
|
}
|
||||||
|
return {
|
||||||
|
"code": 0,
|
||||||
|
"message": "ok",
|
||||||
|
"data": {**get_platform_store().health_metrics(), "storage": storage_status},
|
||||||
|
}
|
||||||
|
|
||||||
3040
backend/app/api/v1/endpoints/platform.py
Normal file
3040
backend/app/api/v1/endpoints/platform.py
Normal file
File diff suppressed because it is too large
Load Diff
26
backend/app/api/v1/router.py
Normal file
26
backend/app/api/v1/router.py
Normal file
@@ -0,0 +1,26 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.data_process import router as data_process_router
|
||||||
|
from app.api.v1.endpoints.platform import router as platform_router
|
||||||
|
from app.api.v1.endpoints.health import router as health_router
|
||||||
|
from app.modules.tenant.router import router as tenant_router
|
||||||
|
from app.modules.project.router import router as project_router
|
||||||
|
from app.modules.approval.router import router as approval_router
|
||||||
|
from app.modules.system.router import router as system_router
|
||||||
|
from app.modules.retention.router import router as retention_router
|
||||||
|
from app.modules.resource.router import router as resource_router
|
||||||
|
from app.modules.gpu.router import router as gpu_router
|
||||||
|
from app.modules.data_convert.router import router as data_convert_router
|
||||||
|
|
||||||
|
api_router = APIRouter()
|
||||||
|
api_router.include_router(health_router, tags=["health"])
|
||||||
|
api_router.include_router(data_process_router, tags=["data-process"])
|
||||||
|
api_router.include_router(platform_router, tags=["platform"])
|
||||||
|
api_router.include_router(system_router, tags=["system"])
|
||||||
|
api_router.include_router(tenant_router, tags=["tenant"])
|
||||||
|
api_router.include_router(project_router, tags=["project"])
|
||||||
|
api_router.include_router(approval_router, tags=["approval"])
|
||||||
|
api_router.include_router(retention_router, tags=["retention"])
|
||||||
|
api_router.include_router(resource_router, tags=["resource"])
|
||||||
|
api_router.include_router(gpu_router, tags=["gpu-assignment"])
|
||||||
|
api_router.include_router(data_convert_router, tags=["data-convert"])
|
||||||
1
backend/app/core/__init__.py
Normal file
1
backend/app/core/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Core infrastructure modules."""
|
||||||
199
backend/app/core/audit.py
Normal file
199
backend/app/core/audit.py
Normal file
@@ -0,0 +1,199 @@
|
|||||||
|
"""
|
||||||
|
审计日志装饰器模块
|
||||||
|
|
||||||
|
提供 @audit_log 装饰器,用于自动记录关键业务操作的审计日志。
|
||||||
|
|
||||||
|
使用示例:
|
||||||
|
from app.core.audit import audit_log
|
||||||
|
|
||||||
|
@audit_log(action="create_dataset", target_type="dataset")
|
||||||
|
async def create_dataset(request: Request, ...):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import functools
|
||||||
|
import time
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any, Callable, Optional, TypeVar
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
|
||||||
|
from app.core.logging import get_logger, request_id_var
|
||||||
|
|
||||||
|
logger = get_logger("app.audit")
|
||||||
|
|
||||||
|
F = TypeVar("F", bound=Callable[..., Any])
|
||||||
|
|
||||||
|
|
||||||
|
def audit_log(
|
||||||
|
action: str,
|
||||||
|
target_type: str = "",
|
||||||
|
*,
|
||||||
|
detail_template: str = "",
|
||||||
|
extract_target_id: Optional[Callable[[Any], str]] = None,
|
||||||
|
) -> Callable[[F], F]:
|
||||||
|
"""
|
||||||
|
审计日志装饰器
|
||||||
|
|
||||||
|
Args:
|
||||||
|
action: 操作类型,如 create_dataset、update_model 等
|
||||||
|
target_type: 目标资源类型,如 dataset、model 等
|
||||||
|
detail_template: 日志详情模板(支持 format 参数)
|
||||||
|
extract_target_id: 从返回值中提取目标 ID 的函数
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
装饰后的函数
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(func: F) -> F:
|
||||||
|
if asyncio.iscoroutinefunction(func):
|
||||||
|
@functools.wraps(func)
|
||||||
|
async def async_wrapper(*args, **kwargs):
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
trace_id = request_id_var.get("-")
|
||||||
|
try:
|
||||||
|
result = await func(*args, **kwargs)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
target_id = _extract_target_id(result, kwargs, extract_target_id)
|
||||||
|
detail = _build_detail(detail_template, kwargs)
|
||||||
|
_record_audit(
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=target_id,
|
||||||
|
detail=detail,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
except Exception:
|
||||||
|
logger.error(
|
||||||
|
"审计日志记录失败 action=%s", action, exc_info=True
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
return async_wrapper # type: ignore
|
||||||
|
else:
|
||||||
|
@functools.wraps(func)
|
||||||
|
def sync_wrapper(*args, **kwargs):
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
trace_id = request_id_var.get("-")
|
||||||
|
try:
|
||||||
|
result = func(*args, **kwargs)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
target_id = _extract_target_id(result, kwargs, extract_target_id)
|
||||||
|
detail = _build_detail(detail_template, kwargs)
|
||||||
|
_record_audit(
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=target_id,
|
||||||
|
detail=detail,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
except Exception:
|
||||||
|
logger.error(
|
||||||
|
"审计日志记录失败 action=%s", action, exc_info=True
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
return sync_wrapper # type: ignore
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_target_id(
|
||||||
|
result: Any, kwargs: dict, extractor: Optional[Callable[[Any], str]]
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""从返回值或 kwargs 中提取目标 ID"""
|
||||||
|
if extractor:
|
||||||
|
try:
|
||||||
|
return extractor(result)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if isinstance(result, dict):
|
||||||
|
return result.get("id")
|
||||||
|
# 尝试从路径参数中提取
|
||||||
|
for key in ("dataset_id", "model_id", "task_id", "resource_id"):
|
||||||
|
val = kwargs.get(key)
|
||||||
|
if val:
|
||||||
|
return str(val)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _build_detail(template: str, kwargs: dict) -> str:
|
||||||
|
"""构建审计详情"""
|
||||||
|
if not template:
|
||||||
|
return ""
|
||||||
|
try:
|
||||||
|
return template.format(**kwargs)
|
||||||
|
except (KeyError, IndexError):
|
||||||
|
return template
|
||||||
|
|
||||||
|
|
||||||
|
def _record_audit(
|
||||||
|
action: str,
|
||||||
|
target_type: str,
|
||||||
|
target_id: Optional[str],
|
||||||
|
detail: str,
|
||||||
|
trace_id: str,
|
||||||
|
duration_ms: float,
|
||||||
|
) -> None:
|
||||||
|
"""通过已有的 record_audit 方法写入审计日志"""
|
||||||
|
try:
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
store = get_platform_store()
|
||||||
|
store.record_audit(
|
||||||
|
action=action,
|
||||||
|
target_type=target_type or None,
|
||||||
|
target_id=target_id,
|
||||||
|
detail=f"{detail} trace_id={trace_id} duration_ms={duration_ms:.1f}" if detail else f"trace_id={trace_id} duration_ms={duration_ms:.1f}",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.error("写入审计日志失败 action=%s", action, exc_info=True)
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== 预定义的审计操作常量 ====================
|
||||||
|
|
||||||
|
class AuditActions:
|
||||||
|
"""预定义的审计操作类型"""
|
||||||
|
# 数据集操作
|
||||||
|
CREATE_DATASET = "create_dataset"
|
||||||
|
UPDATE_DATASET = "update_dataset"
|
||||||
|
DELETE_DATASET = "delete_dataset"
|
||||||
|
|
||||||
|
# 模型操作
|
||||||
|
CREATE_MODEL = "create_model"
|
||||||
|
UPDATE_MODEL = "update_model"
|
||||||
|
DELETE_MODEL = "delete_model"
|
||||||
|
|
||||||
|
# 微调任务
|
||||||
|
CREATE_FINE_TUNE = "create_fine_tune"
|
||||||
|
UPDATE_FINE_TUNE = "update_fine_tune"
|
||||||
|
DELETE_FINE_TUNE = "delete_fine_tune"
|
||||||
|
|
||||||
|
# 推理任务
|
||||||
|
CREATE_INFERENCE = "create_inference"
|
||||||
|
UPDATE_INFERENCE = "update_inference"
|
||||||
|
DELETE_INFERENCE = "delete_inference"
|
||||||
|
|
||||||
|
# 用户管理
|
||||||
|
CREATE_USER = "create_user"
|
||||||
|
UPDATE_USER = "update_user"
|
||||||
|
DELETE_USER = "delete_user"
|
||||||
|
|
||||||
|
# 租户管理
|
||||||
|
CREATE_TENANT = "create_tenant"
|
||||||
|
UPDATE_TENANT = "update_tenant"
|
||||||
|
DELETE_TENANT = "delete_tenant"
|
||||||
|
|
||||||
|
# 权限授权
|
||||||
|
GRANT_ACL = "grant_acl"
|
||||||
|
REVOKE_ACL = "revoke_acl"
|
||||||
|
|
||||||
|
# 系统配置
|
||||||
|
UPDATE_CONFIG = "update_config"
|
||||||
240
backend/app/core/auth.py
Normal file
240
backend/app/core/auth.py
Normal file
@@ -0,0 +1,240 @@
|
|||||||
|
"""鉴权依赖:从 Authorization header 解析当前用户,提供权限校验。"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from fastapi import Depends, HTTPException, Query, Request, status
|
||||||
|
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
# 无需鉴权的路径前缀(健康检查、登录等)
|
||||||
|
PUBLIC_PATHS = ("/health", "/login", "/system-info")
|
||||||
|
OWNER_TABLES = {
|
||||||
|
"dataset": ("datasets", "created_by"), "model": ("models", "created_by"),
|
||||||
|
"trained_model": ("trained_models", "created_by"), "eval": ("eval_tasks", "created_by"),
|
||||||
|
"fine-tune": ("fine_tune_tasks", "payload"), "fine_tune_task": ("fine_tune_tasks", "payload"),
|
||||||
|
"compare": ("compare_tasks", "payload"), "inference": ("compare_tasks", "payload"),
|
||||||
|
"project": ("projects", "created_by"), "data_process": ("data_process_tasks", "created_by"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_token(request: Request) -> str | None:
|
||||||
|
"""从 Authorization header 提取 token(格式: Bearer platform-token-{user_id})。"""
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
if token.startswith("platform-token-"):
|
||||||
|
return token[len("platform-token-"):]
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _session_token(user_id: str, session_id: str) -> str:
|
||||||
|
return f"platform-token-{user_id}.{session_id}"
|
||||||
|
|
||||||
|
|
||||||
|
def get_current_user(request: Request) -> dict[str, Any]:
|
||||||
|
"""
|
||||||
|
FastAPI 依赖:解析当前登录用户。
|
||||||
|
- 公开路径(/health, /login 等)直接放行,返回匿名用户。
|
||||||
|
- 无 token 或 token 无效时抛 401。
|
||||||
|
- admin 用户标记为超级管理员,拥有全部权限。
|
||||||
|
"""
|
||||||
|
path = request.url.path
|
||||||
|
# 去掉路由前缀后判断
|
||||||
|
for prefix in PUBLIC_PATHS:
|
||||||
|
if path.endswith(prefix):
|
||||||
|
return {"id": None, "username": "anonymous", "role": "viewer", "permissions": [], "protected": False}
|
||||||
|
|
||||||
|
token_value = _extract_token(request)
|
||||||
|
if not token_value:
|
||||||
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing or invalid token")
|
||||||
|
|
||||||
|
store = get_platform_store()
|
||||||
|
user_id, _, session_id = token_value.partition(".")
|
||||||
|
if session_id:
|
||||||
|
with store.connect() as conn:
|
||||||
|
session = conn.execute(
|
||||||
|
"SELECT user_id, logout_at, expires_at FROM sessions WHERE id=?", (session_id,)
|
||||||
|
).fetchone()
|
||||||
|
if not session or session["user_id"] != user_id or session["logout_at"]:
|
||||||
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="session expired")
|
||||||
|
if session["expires_at"]:
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
try:
|
||||||
|
if datetime.fromisoformat(str(session["expires_at"]).replace("Z", "+00:00")) <= datetime.now(timezone.utc):
|
||||||
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="session expired")
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
with store.connect() as conn:
|
||||||
|
user_row = conn.execute("SELECT * FROM users WHERE id=?", (user_id,)).fetchone()
|
||||||
|
if user_row:
|
||||||
|
return store._user(user_row)
|
||||||
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||||
|
|
||||||
|
|
||||||
|
def require_admin(current_user: dict[str, Any] = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
"""FastAPI 依赖:要求当前用户是管理员(role=admin 或 protected)。"""
|
||||||
|
if current_user.get("role") == "admin" or current_user.get("protected"):
|
||||||
|
return current_user
|
||||||
|
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="admin permission required")
|
||||||
|
|
||||||
|
|
||||||
|
def is_admin(user: dict[str, Any]) -> bool:
|
||||||
|
"""判断用户是否为管理员(admin 角色或 protected 标记)。"""
|
||||||
|
return user.get("role") == "admin" or user.get("protected", False)
|
||||||
|
|
||||||
|
|
||||||
|
def has_resource_access(
|
||||||
|
resource_type: str,
|
||||||
|
resource_id: str,
|
||||||
|
user: dict[str, Any],
|
||||||
|
permission: str = "read",
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
检查用户对某资源是否有指定权限。
|
||||||
|
- admin/protected 用户直接放行(旁路)。
|
||||||
|
- 其他用户检查 acls 表中是否有对应授权。
|
||||||
|
"""
|
||||||
|
if user.get("role") == "admin" or user.get("protected"):
|
||||||
|
return True
|
||||||
|
|
||||||
|
store = get_platform_store()
|
||||||
|
acls = store.get_acl(resource_type, resource_id)
|
||||||
|
user_id = user.get("id")
|
||||||
|
user_role = user.get("role")
|
||||||
|
|
||||||
|
owner_tables = OWNER_TABLES
|
||||||
|
table_info = owner_tables.get(resource_type)
|
||||||
|
if table_info and user_id:
|
||||||
|
table, column = table_info
|
||||||
|
with store.connect() as conn:
|
||||||
|
row = conn.execute(f"SELECT {column} FROM {table} WHERE id=?", (resource_id,)).fetchone()
|
||||||
|
if row:
|
||||||
|
owner = row[column]
|
||||||
|
if column == "payload":
|
||||||
|
try:
|
||||||
|
import json
|
||||||
|
owner = json.loads(owner or "{}").get("created_by")
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
owner = None
|
||||||
|
if owner == user_id:
|
||||||
|
return True
|
||||||
|
|
||||||
|
for entry in acls:
|
||||||
|
# 按 user 授权
|
||||||
|
if entry.get("principal_type") == "user" and entry.get("principal_id") == user_id:
|
||||||
|
if _permission_covers(entry.get("permission"), permission):
|
||||||
|
return True
|
||||||
|
# 按 role 授权
|
||||||
|
if entry.get("principal_type") == "role" and entry.get("principal_id") == user_role:
|
||||||
|
if _permission_covers(entry.get("permission"), permission):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _permission_covers(granted: str | None, required: str) -> bool:
|
||||||
|
"""权限覆盖判断:write/execute 覆盖 read;admin 覆盖一切。"""
|
||||||
|
if not granted:
|
||||||
|
return False
|
||||||
|
if granted == "admin":
|
||||||
|
return True
|
||||||
|
if granted == required:
|
||||||
|
return True
|
||||||
|
# write 覆盖 read
|
||||||
|
if required == "read" and granted in ("write", "execute"):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def filter_accessible_resource_ids(
|
||||||
|
resource_type: str,
|
||||||
|
all_ids: list[str],
|
||||||
|
user: dict[str, Any],
|
||||||
|
) -> list[str]:
|
||||||
|
"""
|
||||||
|
从全部资源 ID 中过滤出当前用户可访问的 ID 列表。
|
||||||
|
- admin 直接返回全部。
|
||||||
|
- 普通用户查 acls 表取交集。
|
||||||
|
"""
|
||||||
|
if user.get("role") == "admin" or user.get("protected"):
|
||||||
|
return all_ids
|
||||||
|
|
||||||
|
if not all_ids:
|
||||||
|
return []
|
||||||
|
|
||||||
|
store = get_platform_store()
|
||||||
|
user_id = user.get("id")
|
||||||
|
user_role = user.get("role")
|
||||||
|
|
||||||
|
# 查询该用户在该资源类型下有 read 权限的所有 resource_id
|
||||||
|
with store.connect() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT DISTINCT resource_id FROM acls
|
||||||
|
WHERE resource_type=? AND (
|
||||||
|
(principal_type='user' AND principal_id=?)
|
||||||
|
OR (principal_type='role' AND principal_id=?)
|
||||||
|
)
|
||||||
|
""",
|
||||||
|
(resource_type, user_id, user_role),
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
accessible = {r["resource_id"] for r in rows}
|
||||||
|
if resource_type in OWNER_TABLES:
|
||||||
|
table, column = OWNER_TABLES[resource_type]
|
||||||
|
if column == "payload":
|
||||||
|
# payload 是 JSON 字符串,需要查出后解析 created_by
|
||||||
|
with store.connect() as conn:
|
||||||
|
owned = conn.execute(f"SELECT id, {column} FROM {table}").fetchall()
|
||||||
|
for row in owned:
|
||||||
|
try:
|
||||||
|
import json
|
||||||
|
payload = json.loads(row[column] or "{}")
|
||||||
|
if payload.get("created_by") == user_id:
|
||||||
|
accessible.add(row["id"])
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
with store.connect() as conn:
|
||||||
|
owned = conn.execute(f"SELECT id FROM {table} WHERE {column}=?", (user_id,)).fetchall()
|
||||||
|
accessible.update(row["id"] for row in owned)
|
||||||
|
return [rid for rid in all_ids if rid in accessible]
|
||||||
|
|
||||||
|
|
||||||
|
def filter_accessible_resource_ids_batch(
|
||||||
|
resource_type: str,
|
||||||
|
resource_ids: list[str],
|
||||||
|
user: dict[str, Any],
|
||||||
|
) -> set[str]:
|
||||||
|
"""Filter a list endpoint with one ACL query instead of one query per row."""
|
||||||
|
if is_admin(user):
|
||||||
|
return set(resource_ids)
|
||||||
|
if not resource_ids:
|
||||||
|
return set()
|
||||||
|
store = get_platform_store()
|
||||||
|
placeholders = ",".join("?" for _ in resource_ids)
|
||||||
|
with store.connect() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
f"SELECT DISTINCT resource_id FROM acls WHERE resource_type=? AND resource_id IN ({placeholders}) "
|
||||||
|
"AND ((principal_type='user' AND principal_id=?) OR (principal_type='role' AND principal_id=?))",
|
||||||
|
(resource_type, *resource_ids, user.get("id"), user.get("role")),
|
||||||
|
).fetchall()
|
||||||
|
accessible = {row["resource_id"] for row in rows}
|
||||||
|
table_info = OWNER_TABLES.get(resource_type)
|
||||||
|
if table_info and user.get("id"):
|
||||||
|
table, column = table_info
|
||||||
|
with store.connect() as conn:
|
||||||
|
owned = conn.execute(
|
||||||
|
f"SELECT id, {column} FROM {table} WHERE id IN ({placeholders})",
|
||||||
|
(*resource_ids,),
|
||||||
|
).fetchall()
|
||||||
|
for row in owned:
|
||||||
|
owner = row[column]
|
||||||
|
# 如果列是 payload(JSON),需要解析后提取 created_by
|
||||||
|
if column == "payload":
|
||||||
|
try:
|
||||||
|
import json
|
||||||
|
owner = json.loads(owner or "{}").get("created_by")
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
owner = None
|
||||||
|
if owner == user["id"]:
|
||||||
|
accessible.add(row["id"])
|
||||||
|
return accessible
|
||||||
110
backend/app/core/config.py
Normal file
110
backend/app/core/config.py
Normal file
@@ -0,0 +1,110 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
from functools import lru_cache
|
||||||
|
import os
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
try:
|
||||||
|
from pathlib import Path as _Path
|
||||||
|
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
|
# 显式指定 backend 目录下的 .env,并强制覆盖已有环境变量,
|
||||||
|
# 确保远程数据库配置生效,不被本地默认值或残留环境变量影响。
|
||||||
|
_env_path = _Path(__file__).resolve().parent.parent.parent / ".env"
|
||||||
|
load_dotenv(dotenv_path=_env_path, override=True)
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _int_env(name: str, default: int) -> int:
|
||||||
|
raw = os.getenv(name)
|
||||||
|
if raw is None or raw == "":
|
||||||
|
return default
|
||||||
|
return int(raw)
|
||||||
|
|
||||||
|
|
||||||
|
def _list_env(name: str, default: list[str]) -> list[str]:
|
||||||
|
raw = os.getenv(name)
|
||||||
|
if raw is None or raw.strip() == "":
|
||||||
|
return default
|
||||||
|
return [item.strip() for item in raw.split(",") if item.strip()]
|
||||||
|
|
||||||
|
|
||||||
|
def _bool_env(name: str, default: bool) -> bool:
|
||||||
|
raw = os.getenv(name)
|
||||||
|
if raw is None or raw.strip() == "":
|
||||||
|
return default
|
||||||
|
return raw.strip().lower() in {"1", "true", "yes", "on"}
|
||||||
|
|
||||||
|
|
||||||
|
def docs_kwargs(enabled: bool) -> dict[str, Any]:
|
||||||
|
"""Swagger UI / ReDoc / OpenAPI schema 路由开关。
|
||||||
|
|
||||||
|
关闭时 FastAPI 不注册 /docs、/redoc、/openapi.json,访问一律返回 404,
|
||||||
|
避免未授权访问泄露 API 结构。
|
||||||
|
"""
|
||||||
|
if enabled:
|
||||||
|
return {}
|
||||||
|
return {"docs_url": None, "redoc_url": None, "openapi_url": None}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class Settings:
|
||||||
|
app_name: str = os.getenv("APP_NAME", "YG Fine-Tune Platform API")
|
||||||
|
app_env: str = os.getenv("APP_ENV", "local")
|
||||||
|
route_prefix: str = os.getenv("MODELTF_ROUTE_PREFIX", "/modelTF")
|
||||||
|
app_mode: str = os.getenv("APP_MODE", "local")
|
||||||
|
database_url: str = os.getenv("DATABASE_URL", "postgresql+psycopg://yg_ft:change_me@localhost:15432/yg_ft")
|
||||||
|
cors_allow_origins: list[str] = None # type: ignore[assignment]
|
||||||
|
compute_mode: str = os.getenv("COMPUTE_MODE", "real")
|
||||||
|
compute_status_sync_mode: str = os.getenv("COMPUTE_STATUS_SYNC_MODE", "polling")
|
||||||
|
compute_poll_interval_seconds: int = _int_env("COMPUTE_POLL_INTERVAL_SECONDS", 3)
|
||||||
|
compute_request_timeout_seconds: int = _int_env("COMPUTE_REQUEST_TIMEOUT_SECONDS", 5)
|
||||||
|
minio_enabled: bool = _bool_env("MINIO_ENABLED", False)
|
||||||
|
# MinIO is an independent service and may run on another host. The
|
||||||
|
# endpoint must therefore be reachable from the Backend container.
|
||||||
|
minio_endpoint: str = os.getenv("MINIO_ENDPOINT", "http://host.docker.internal:19000")
|
||||||
|
minio_access_key: str = os.getenv("MINIO_ACCESS_KEY", "minioadmin")
|
||||||
|
minio_secret_key: str = os.getenv("MINIO_SECRET_KEY", "minioadmin")
|
||||||
|
minio_bucket: str = os.getenv("MINIO_BUCKET", "yg-ft-resources")
|
||||||
|
minio_secure: bool = _bool_env("MINIO_SECURE", False)
|
||||||
|
storage_wait_seconds: int = _int_env("STORAGE_WAIT_SECONDS", 300)
|
||||||
|
storage_check_interval_seconds: int = _int_env("STORAGE_CHECK_INTERVAL_SECONDS", 10)
|
||||||
|
compute_service_token: str = os.getenv("COMPUTE_SERVICE_TOKEN", "")
|
||||||
|
log_level: str = os.getenv("LOG_LEVEL", "INFO")
|
||||||
|
log_dir: str = os.getenv("LOG_DIR", "./logs")
|
||||||
|
log_file_prefix: str = os.getenv("LOG_FILE_PREFIX", "backend")
|
||||||
|
log_error_file_prefix: str = os.getenv("LOG_ERROR_FILE_PREFIX", "error")
|
||||||
|
log_max_bytes: int = _int_env("LOG_MAX_BYTES", 20 * 1024 * 1024)
|
||||||
|
log_retention_days: int = _int_env("LOG_RETENTION_DAYS", 10)
|
||||||
|
enable_docs: bool = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
def __post_init__(self) -> None:
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"cors_allow_origins",
|
||||||
|
_list_env(
|
||||||
|
"CORS_ALLOW_ORIGINS",
|
||||||
|
[
|
||||||
|
"http://localhost:16801",
|
||||||
|
"http://127.0.0.1:16801",
|
||||||
|
"http://localhost:17861",
|
||||||
|
"http://127.0.0.1:17861",
|
||||||
|
],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
# Swagger UI / ReDoc / OpenAPI 文档路由开关:
|
||||||
|
# 未显式配置 ENABLE_DOCS 时,仅本地/开发环境开放,生产环境默认关闭,
|
||||||
|
# 避免未授权访问泄露 API 结构。从运行时环境读取 APP_ENV,而非类定义时
|
||||||
|
# 缓存的默认值,保证生产默认关闭始终生效且便于测试。
|
||||||
|
object.__setattr__(
|
||||||
|
self,
|
||||||
|
"enable_docs",
|
||||||
|
_bool_env("ENABLE_DOCS", os.getenv("APP_ENV", "local") != "prod"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache
|
||||||
|
def get_settings() -> Settings:
|
||||||
|
return Settings()
|
||||||
|
|
||||||
468
backend/app/core/logging.py
Normal file
468
backend/app/core/logging.py
Normal file
@@ -0,0 +1,468 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from contextvars import ContextVar
|
||||||
|
from datetime import date, datetime, timedelta
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import sys
|
||||||
|
from logging import Handler, LogRecord
|
||||||
|
from pathlib import Path
|
||||||
|
import re
|
||||||
|
import time
|
||||||
|
from typing import Any, Callable, Optional
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
from fastapi import FastAPI, Request
|
||||||
|
|
||||||
|
from app.core.config import Settings, get_settings
|
||||||
|
|
||||||
|
request_id_var: ContextVar[str] = ContextVar("request_id", default="-")
|
||||||
|
|
||||||
|
# ==================== 敏感数据脱敏规则 ====================
|
||||||
|
|
||||||
|
SENSITIVE_PATTERNS: dict[str, Callable | str] = {
|
||||||
|
"token": "***",
|
||||||
|
"password": "***",
|
||||||
|
"access_token": "***",
|
||||||
|
"refresh_token": "***",
|
||||||
|
"secret_key": "***",
|
||||||
|
"authorization": "***",
|
||||||
|
"bearer": "***",
|
||||||
|
"api_key": "***",
|
||||||
|
"private_key": "***",
|
||||||
|
}
|
||||||
|
|
||||||
|
def mask_value(key: str, value: Any) -> str:
|
||||||
|
"""对单个值进行脱敏处理"""
|
||||||
|
if value is None:
|
||||||
|
return ""
|
||||||
|
str_val = str(value)
|
||||||
|
|
||||||
|
handler = SENSITIVE_PATTERNS.get(key)
|
||||||
|
if callable(handler):
|
||||||
|
return handler(str_val)
|
||||||
|
elif isinstance(handler, str):
|
||||||
|
# 支持正则替换模式,如 r"1\d{3}\d{4}"
|
||||||
|
try:
|
||||||
|
return re.sub(handler, "***", str_val)
|
||||||
|
except re.error:
|
||||||
|
return "***"
|
||||||
|
return handler
|
||||||
|
|
||||||
|
|
||||||
|
def mask_sensitive_dict(data: dict) -> dict:
|
||||||
|
"""递归脱敏字典中的敏感字段"""
|
||||||
|
if not data or not isinstance(data, dict):
|
||||||
|
return data
|
||||||
|
|
||||||
|
result = {}
|
||||||
|
for key, value in data.items():
|
||||||
|
result[key] = mask_value(key, value)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def mask_sensitive_string(text: str) -> str:
|
||||||
|
"""从文本中脱敏常见敏感信息"""
|
||||||
|
if not text:
|
||||||
|
return text
|
||||||
|
|
||||||
|
patterns = [
|
||||||
|
(r'Bearer\s+[A-Za-z0-9\-._]+', '***'),
|
||||||
|
(r'token\s*[:=]\s*', '***'),
|
||||||
|
(r'password\s*[:=]\s*', '***'),
|
||||||
|
(r'secret[_-]?key\s*[:=]', '***'),
|
||||||
|
(r'api[-_]?key\s*[:=]', '***'),
|
||||||
|
(r'private[_-]?key\s*[:=]', '***'),
|
||||||
|
(r'\d{11}', r'\d{3}\*\d{4}'), # 手机号/身份证
|
||||||
|
(r'1[3-9]\d{9}', r'1\*{3}\*{4}'), # 手机号
|
||||||
|
]
|
||||||
|
|
||||||
|
for pattern, replacement in patterns:
|
||||||
|
try:
|
||||||
|
text = re.sub(pattern, replacement, text, flags=re.IGNORECASE)
|
||||||
|
except re.error:
|
||||||
|
pass
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== RequestId Filter ====================
|
||||||
|
|
||||||
|
class RequestIdFilter(logging.Filter):
|
||||||
|
def filter(self, record: LogRecord) -> bool:
|
||||||
|
record.request_id = request_id_var.get()
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== Enhanced JSON Formatter ====================
|
||||||
|
|
||||||
|
class JsonLogFormatter(logging.Formatter):
|
||||||
|
"""
|
||||||
|
增强的 JSON 日志格式化器,支持结构化字段输出。
|
||||||
|
|
||||||
|
输出示例:
|
||||||
|
{
|
||||||
|
"@timestamp": "2026-08-17T18:30:00.123Z",
|
||||||
|
"level": "INFO",
|
||||||
|
"logger": "dataset.router",
|
||||||
|
"message": "数据集创建成功",
|
||||||
|
"module": "dataset.router",
|
||||||
|
"function": "create_dataset",
|
||||||
|
"file": "dataset/router.py",
|
||||||
|
"line": 45,
|
||||||
|
"process": 12345,
|
||||||
|
"thread": "MainThread",
|
||||||
|
"request_id": "req-abc123",
|
||||||
|
"user_id": "u_admin",
|
||||||
|
"client_ip": "192.168.1.100",
|
||||||
|
"extra": {...}
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
def format(self, record: LogRecord) -> str:
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"@timestamp": datetime.fromtimestamp(record.created).astimezone().isoformat(
|
||||||
|
timespec="milliseconds"
|
||||||
|
),
|
||||||
|
"level": record.levelname,
|
||||||
|
"logger": record.name,
|
||||||
|
"message": record.getMessage(),
|
||||||
|
"module": record.module,
|
||||||
|
"function": record.funcName,
|
||||||
|
"file": record.pathname,
|
||||||
|
"line": record.lineno,
|
||||||
|
"process": record.process,
|
||||||
|
"thread": record.thread,
|
||||||
|
"thread_name": record.threadName,
|
||||||
|
"request_id": getattr(record, "request_id", "-"),
|
||||||
|
}
|
||||||
|
|
||||||
|
# 从 record 中提取额外字段(通过 extra 参数传入)
|
||||||
|
for attr in ("user_id", "client_ip", "target_type", "target_id",
|
||||||
|
"duration_ms", "status_code", "error"):
|
||||||
|
val = getattr(record, attr, None)
|
||||||
|
if val is not None:
|
||||||
|
payload[attr] = val
|
||||||
|
|
||||||
|
# 处理异常信息
|
||||||
|
if record.exc_info:
|
||||||
|
payload["exception"] = self.formatException(record.exc_info)
|
||||||
|
if record.stack_info:
|
||||||
|
payload["stack"] = self.formatStack(record.stack_info)
|
||||||
|
|
||||||
|
return json.dumps(payload, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== DateSizeRotatingFileHandler ====================
|
||||||
|
# (保持不变,已有实现)
|
||||||
|
|
||||||
|
class DateSizeRotatingFileHandler(Handler):
|
||||||
|
"""Rotate log files by date and size while keeping date in every file name."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
log_dir: str | Path,
|
||||||
|
file_prefix: str,
|
||||||
|
max_bytes: int,
|
||||||
|
retention_days: int,
|
||||||
|
encoding: str = "utf-8",
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.log_dir = Path(log_dir)
|
||||||
|
self.file_prefix = file_prefix
|
||||||
|
self.max_bytes = max_bytes
|
||||||
|
self.retention_days = retention_days
|
||||||
|
self.encoding = encoding
|
||||||
|
self._current_date: date | None = None
|
||||||
|
self._stream: Any | None = None
|
||||||
|
self._current_path: Path | None = None
|
||||||
|
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
def emit(self, record: LogRecord) -> None:
|
||||||
|
try:
|
||||||
|
message = self.format(record) + self.terminator
|
||||||
|
encoded_size = len(message.encode(self.encoding))
|
||||||
|
self._ensure_stream()
|
||||||
|
if self._should_rotate(encoded_size):
|
||||||
|
self._rotate_by_size()
|
||||||
|
self._ensure_stream(force=True)
|
||||||
|
self._stream.write(message)
|
||||||
|
self.flush()
|
||||||
|
self._cleanup_expired_files()
|
||||||
|
except Exception:
|
||||||
|
self.handleError(record)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def terminator(self) -> str:
|
||||||
|
return "\n"
|
||||||
|
|
||||||
|
def flush(self) -> None:
|
||||||
|
if self._stream and not self._stream.closed:
|
||||||
|
self._stream.flush()
|
||||||
|
|
||||||
|
def close(self) -> None:
|
||||||
|
try:
|
||||||
|
if self._stream and not self._stream.closed:
|
||||||
|
self._stream.close()
|
||||||
|
finally:
|
||||||
|
self._stream = None
|
||||||
|
super().close()
|
||||||
|
|
||||||
|
def _dated_path(self, target_date: date) -> Path:
|
||||||
|
return self.log_dir / f"{self.file_prefix}-{target_date.isoformat()}.log"
|
||||||
|
|
||||||
|
def _ensure_stream(self, force: bool = False) -> None:
|
||||||
|
today = date.today()
|
||||||
|
if not force and self._stream and self._current_date == today:
|
||||||
|
return
|
||||||
|
|
||||||
|
if self._stream and not self._stream.closed:
|
||||||
|
self._stream.close()
|
||||||
|
|
||||||
|
self._current_date = today
|
||||||
|
self._current_path = self._dated_path(today)
|
||||||
|
self._stream = self._current_path.open("a", encoding=self.encoding)
|
||||||
|
|
||||||
|
def _should_rotate(self, incoming_size: int) -> bool:
|
||||||
|
if not self._current_path or self.max_bytes <= 0:
|
||||||
|
return False
|
||||||
|
if not self._current_path.exists():
|
||||||
|
return False
|
||||||
|
return self._current_path.stat().st_size + incoming_size > self.max_bytes
|
||||||
|
|
||||||
|
def _rotate_by_size(self) -> None:
|
||||||
|
if not self._current_path or not self._current_path.exists():
|
||||||
|
return
|
||||||
|
|
||||||
|
if self._stream and not self._stream.closed:
|
||||||
|
self._stream.close()
|
||||||
|
self._stream = None
|
||||||
|
|
||||||
|
stem = self._current_path.stem
|
||||||
|
suffix = self._current_path.suffix
|
||||||
|
index = 1
|
||||||
|
while True:
|
||||||
|
rotated_path = self.log_dir / f"{stem}.{index}{suffix}"
|
||||||
|
if not rotated_path.exists():
|
||||||
|
self._current_path.rename(rotated_path)
|
||||||
|
return
|
||||||
|
index += 1
|
||||||
|
|
||||||
|
def _cleanup_expired_files(self) -> None:
|
||||||
|
if self.retention_days <= 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
cutoff = date.today() - timedelta(days=self.retention_days - 1)
|
||||||
|
pattern = re.compile(
|
||||||
|
rf"^{re.escape(self.file_prefix)}-(\d{{4}}-\d{{2}}-\d{{2}})(?:\.\d+)?\.log$"
|
||||||
|
)
|
||||||
|
for path in self.log_dir.glob(f"{self.file_prefix}-*.log"):
|
||||||
|
match = pattern.match(path.name)
|
||||||
|
if not match:
|
||||||
|
continue
|
||||||
|
file_date = datetime.strptime(match.group(1), "%Y-%m-%d").date()
|
||||||
|
if file_date < cutoff:
|
||||||
|
path.unlink(missing_ok=True)
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== Structured Logger 封装 ====================
|
||||||
|
|
||||||
|
class StructuredLogger:
|
||||||
|
"""
|
||||||
|
结构化日志记录器,提供统一的日志接口。
|
||||||
|
|
||||||
|
使用方式:
|
||||||
|
logger = get_structured_logger('dataset.router')
|
||||||
|
logger.info('创建数据集', dataset_id='ds_123')
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, name: str, module: str = ""):
|
||||||
|
self.logger = logging.getLogger(name)
|
||||||
|
self.name = name
|
||||||
|
self.module = module
|
||||||
|
|
||||||
|
@property
|
||||||
|
def trace_id(self) -> str:
|
||||||
|
return request_id_var.get("-")
|
||||||
|
|
||||||
|
def info(self, message: str, **extra: Any) -> None:
|
||||||
|
self._log("INFO", message, **extra)
|
||||||
|
|
||||||
|
def warning(self, message: str, **extra: Any) -> None:
|
||||||
|
self._log("WARNING", message, **extra)
|
||||||
|
|
||||||
|
def error(self, message: str, **extra: Any) -> None:
|
||||||
|
self._log("ERROR", message, **extra)
|
||||||
|
|
||||||
|
def debug(self, message: str, **extra: Any) -> None:
|
||||||
|
self._log("DEBUG", message, **extra)
|
||||||
|
|
||||||
|
def _log(self, level: str, message: str, **extra: Any) -> None:
|
||||||
|
"""统一日志记录方法"""
|
||||||
|
log_entry: dict[str, Any] = {
|
||||||
|
"timestamp": datetime.utcnow().isoformat(),
|
||||||
|
"level": level,
|
||||||
|
"logger": self.name,
|
||||||
|
"module": self.module,
|
||||||
|
"message": message,
|
||||||
|
"trace_id": self.trace_id,
|
||||||
|
"extra": extra,
|
||||||
|
}
|
||||||
|
self.logger.log(getattr(logging, level, logging.INFO), json.dumps(log_entry, ensure_ascii=False, default=str))
|
||||||
|
|
||||||
|
|
||||||
|
def get_structured_logger(name: str, module: str = "") -> StructuredLogger:
|
||||||
|
"""获取结构化日志记录器"""
|
||||||
|
return StructuredLogger(name, module)
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== 快捷函数 ====================
|
||||||
|
|
||||||
|
def get_logger(name: str) -> logging.Logger:
|
||||||
|
"""获取标准 Python logger"""
|
||||||
|
return logging.getLogger(name)
|
||||||
|
|
||||||
|
|
||||||
|
def set_request_id(request_id: str) -> None:
|
||||||
|
"""设置当前请求的追踪 ID"""
|
||||||
|
request_id_var.set(request_id)
|
||||||
|
|
||||||
|
|
||||||
|
def setup_request_logging(app: FastAPI) -> None:
|
||||||
|
"""配置 FastAPI 请求日志中间件"""
|
||||||
|
logger = get_logger("app.access")
|
||||||
|
|
||||||
|
@app.middleware("http")
|
||||||
|
async def request_logging_middleware(request: Request, call_next): # type: ignore[no-untyped-def]
|
||||||
|
request_id = request.headers.get("X-Request-ID") or str(uuid4())
|
||||||
|
token = request_id_var.set(request_id)
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
try:
|
||||||
|
response = await call_next(request)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
|
||||||
|
noisy_paths = ("/health", "/system-info", "/compute/jobs/", "/model-eval/", "/model-compare/")
|
||||||
|
log_method = logger.debug if request.method == "GET" and response.status_code < 400 else logger.info
|
||||||
|
if any(request.url.path.endswith(path) or path in request.url.path for path in noisy_paths) and response.status_code < 400:
|
||||||
|
log_method = logger.debug
|
||||||
|
if response.status_code >= 400:
|
||||||
|
log_method = logger.warning
|
||||||
|
log_method(
|
||||||
|
"request completed method=%s path=%s status_code=%s duration_ms=%.2f client=%s",
|
||||||
|
request.method,
|
||||||
|
request.url.path,
|
||||||
|
response.status_code,
|
||||||
|
elapsed_ms,
|
||||||
|
request.client.host if request.client else "-",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 5xx 系统错误自动写入操作日志(未被 @op_log 覆盖的系统级异常)
|
||||||
|
if response.status_code >= 500:
|
||||||
|
try:
|
||||||
|
from app.core.op_log import log_operation, OpModule, OpStatus
|
||||||
|
log_operation(
|
||||||
|
module=OpModule.SYSTEM,
|
||||||
|
action="request",
|
||||||
|
target_type="api",
|
||||||
|
target_name=request.url.path,
|
||||||
|
status=OpStatus.FAILURE,
|
||||||
|
error_message=f"HTTP {response.status_code} - 系统内部错误",
|
||||||
|
error_type="HTTPError",
|
||||||
|
detail=f'{{"method":"{request.method}","path":"{request.url.path}","status":{response.status_code}}}',
|
||||||
|
func_name="request_logging_middleware",
|
||||||
|
request=request,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass # 日志写入失败不影响主流程
|
||||||
|
|
||||||
|
response.headers["X-Request-ID"] = request_id
|
||||||
|
return response
|
||||||
|
except Exception:
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
logger.exception(
|
||||||
|
"request failed method=%s path=%s duration_ms=%.2f client=%s",
|
||||||
|
request.method,
|
||||||
|
request.url.path,
|
||||||
|
elapsed_ms,
|
||||||
|
request.client.host if request.client else "-",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 未被捕获的异常,写入操作日志
|
||||||
|
try:
|
||||||
|
import traceback as _tb
|
||||||
|
from app.core.op_log import log_operation, OpModule, OpStatus
|
||||||
|
log_operation(
|
||||||
|
module=OpModule.SYSTEM,
|
||||||
|
action="request",
|
||||||
|
target_type="api",
|
||||||
|
target_name=request.url.path,
|
||||||
|
status=OpStatus.FAILURE,
|
||||||
|
error_message=str(sys.exc_info()[1])[:1000] if sys.exc_info()[1] else "未知异常",
|
||||||
|
error_type=type(sys.exc_info()[1]).__name__ if sys.exc_info()[1] else "UnknownError",
|
||||||
|
error_traceback="".join(_tb.format_exception(*sys.exc_info()))[:5000],
|
||||||
|
func_name="request_logging_middleware",
|
||||||
|
detail=f'{{"method":"{request.method}","path":"{request.url.path}"}}',
|
||||||
|
request=request,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass # 日志写入失败不影响主流程
|
||||||
|
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
request_id_var.reset(token)
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== 配置函数 ====================
|
||||||
|
|
||||||
|
def configure_logging(settings: Settings | None = None) -> None:
|
||||||
|
settings = settings or get_settings()
|
||||||
|
|
||||||
|
root_logger = logging.getLogger()
|
||||||
|
root_logger.handlers.clear()
|
||||||
|
root_logger.setLevel(settings.log_level.upper())
|
||||||
|
|
||||||
|
console_formatter = logging.Formatter(
|
||||||
|
fmt=(
|
||||||
|
"%(asctime)s | %(levelname)s | pid=%(process)d | %(threadName)s | "
|
||||||
|
"request_id=%(request_id)s | %(name)s | %(pathname)s:%(lineno)d | %(message)s"
|
||||||
|
),
|
||||||
|
datefmt="%Y-%m-%d %H:%M:%S",
|
||||||
|
)
|
||||||
|
json_formatter = JsonLogFormatter()
|
||||||
|
request_filter = RequestIdFilter()
|
||||||
|
|
||||||
|
console_handler = logging.StreamHandler()
|
||||||
|
console_handler.setFormatter(console_formatter)
|
||||||
|
console_handler.addFilter(request_filter)
|
||||||
|
|
||||||
|
file_handler = DateSizeRotatingFileHandler(
|
||||||
|
log_dir=settings.log_dir,
|
||||||
|
file_prefix=settings.log_file_prefix,
|
||||||
|
max_bytes=settings.log_max_bytes,
|
||||||
|
retention_days=settings.log_retention_days,
|
||||||
|
)
|
||||||
|
file_handler.setFormatter(json_formatter)
|
||||||
|
file_handler.addFilter(request_filter)
|
||||||
|
|
||||||
|
error_file_handler = DateSizeRotatingFileHandler(
|
||||||
|
log_dir=settings.log_dir,
|
||||||
|
file_prefix=settings.log_error_file_prefix,
|
||||||
|
max_bytes=settings.log_max_bytes,
|
||||||
|
retention_days=settings.log_retention_days,
|
||||||
|
)
|
||||||
|
error_file_handler.setLevel(logging.ERROR)
|
||||||
|
error_file_handler.setFormatter(json_formatter)
|
||||||
|
error_file_handler.addFilter(request_filter)
|
||||||
|
|
||||||
|
root_logger.addHandler(console_handler)
|
||||||
|
root_logger.addHandler(file_handler)
|
||||||
|
root_logger.addHandler(error_file_handler)
|
||||||
|
|
||||||
|
for logger_name in ("uvicorn", "uvicorn.error", "uvicorn.access"):
|
||||||
|
logger = logging.getLogger(logger_name)
|
||||||
|
logger.handlers.clear()
|
||||||
|
logger.propagate = True
|
||||||
|
|
||||||
|
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
|
||||||
|
logging.getLogger("psycopg.pool").setLevel(logging.ERROR)
|
||||||
382
backend/app/core/op_log.py
Normal file
382
backend/app/core/op_log.py
Normal file
@@ -0,0 +1,382 @@
|
|||||||
|
"""
|
||||||
|
操作日志工具模块
|
||||||
|
|
||||||
|
提供 @op_log 装饰器和 log_operation 函数,用于记录用户在各业务模块的详细操作。
|
||||||
|
自动捕获成功/失败状态、完整报错堆栈、操作耗时等。
|
||||||
|
|
||||||
|
核心设计:
|
||||||
|
- 失败操作必须清晰记录完整异常堆栈(traceback)
|
||||||
|
- 记录异常类型(如 RuntimeError / ValueError / ConnectionError)
|
||||||
|
- 记录具体出错的函数名和文件位置,方便定位 bug
|
||||||
|
- 记录 HTTP 状态码,方便区分用户错误(4xx)和系统错误(5xx)
|
||||||
|
|
||||||
|
使用示例:
|
||||||
|
from app.core.op_log import op_log, OpModule, OpAction
|
||||||
|
|
||||||
|
@router.post("/inference/start")
|
||||||
|
@op_log(module=OpModule.INFERENCE, action=OpAction.START, target_type="inference")
|
||||||
|
async def start_inference(...):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import functools
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
import traceback
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any, Callable, Optional, TypeVar
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
|
||||||
|
from app.core.logging import get_logger, request_id_var
|
||||||
|
from app.db.platform_store import get_platform_store, new_id, utcnow
|
||||||
|
|
||||||
|
logger = get_logger("app.op_log")
|
||||||
|
|
||||||
|
F = TypeVar("F", bound=Callable[..., Any])
|
||||||
|
|
||||||
|
|
||||||
|
class OpModule:
|
||||||
|
"""业务模块常量"""
|
||||||
|
FINE_TUNE = "fine-tune" # 模型训练
|
||||||
|
MODEL_EVAL = "model-eval" # 模型评测
|
||||||
|
INFERENCE = "model-inference" # 模型推理
|
||||||
|
MODEL_MANAGE = "model-manage" # 模型管理
|
||||||
|
DATASET = "dataset" # 数据集
|
||||||
|
DATA_PROCESS = "data-process" # 数据处理
|
||||||
|
DATA_CONVERT = "data-convert" # 数据类型转换
|
||||||
|
COMPUTE = "compute" # 算力节点
|
||||||
|
SYSTEM = "system" # 系统
|
||||||
|
|
||||||
|
|
||||||
|
class OpAction:
|
||||||
|
"""操作动作常量"""
|
||||||
|
CREATE = "create"
|
||||||
|
UPDATE = "update"
|
||||||
|
DELETE = "delete"
|
||||||
|
START = "start"
|
||||||
|
STOP = "stop"
|
||||||
|
UPLOAD = "upload"
|
||||||
|
DOWNLOAD = "download"
|
||||||
|
CONVERT = "convert"
|
||||||
|
MERGE = "merge"
|
||||||
|
IMPORT = "import"
|
||||||
|
LOGIN = "login"
|
||||||
|
LOGOUT = "logout"
|
||||||
|
PUBLISH = "publish"
|
||||||
|
RETRY = "retry"
|
||||||
|
|
||||||
|
|
||||||
|
class OpStatus:
|
||||||
|
"""操作状态常量"""
|
||||||
|
SUCCESS = "success"
|
||||||
|
FAILURE = "failure"
|
||||||
|
|
||||||
|
|
||||||
|
def op_log(
|
||||||
|
module: str,
|
||||||
|
action: str,
|
||||||
|
target_type: str = "",
|
||||||
|
*,
|
||||||
|
target_name_param: str = "name",
|
||||||
|
detail_params: Optional[list[str]] = None,
|
||||||
|
) -> Callable[[F], F]:
|
||||||
|
"""
|
||||||
|
操作日志装饰器
|
||||||
|
|
||||||
|
自动记录:
|
||||||
|
- 谁在什么时间操作了什么
|
||||||
|
- 成功还是失败
|
||||||
|
- 失败时记录完整异常堆栈(traceback)、异常类型、异常消息
|
||||||
|
- 出错的函数名和文件位置,方便定位 bug
|
||||||
|
- 操作耗时(ms)
|
||||||
|
- 客户端 IP、请求路径
|
||||||
|
|
||||||
|
Args:
|
||||||
|
module: 业务模块(OpModule 常量)
|
||||||
|
action: 操作动作(OpAction 常量)
|
||||||
|
target_type: 资源类型
|
||||||
|
target_name_param: 从 kwargs 中提取目标名称的参数名
|
||||||
|
detail_params: 需要记录到 detail 的参数名列表
|
||||||
|
"""
|
||||||
|
def decorator(func: F) -> F:
|
||||||
|
func_name = f"{func.__module__}.{func.__qualname__}"
|
||||||
|
|
||||||
|
if asyncio.iscoroutinefunction(func):
|
||||||
|
@functools.wraps(func)
|
||||||
|
async def async_wrapper(*args, **kwargs):
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
trace_id = request_id_var.get("-")
|
||||||
|
|
||||||
|
user = _extract_user(args, kwargs)
|
||||||
|
request = _extract_request(args)
|
||||||
|
|
||||||
|
target_name = _get_param(kwargs, target_name_param, "")
|
||||||
|
target_id = _get_param(kwargs, "task_id", "") or _get_param(kwargs, "dataset_id", "") or _get_param(kwargs, "model_id", "")
|
||||||
|
|
||||||
|
detail_dict = _build_detail(detail_params, kwargs)
|
||||||
|
detail_str = json.dumps(detail_dict, ensure_ascii=False) if detail_dict else ""
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await func(*args, **kwargs)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
|
||||||
|
if not target_id and isinstance(result, dict):
|
||||||
|
target_id = str(result.get("id", ""))
|
||||||
|
|
||||||
|
_write_log(
|
||||||
|
module=module,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=str(target_id) if target_id else None,
|
||||||
|
target_name=str(target_name) if target_name else None,
|
||||||
|
status=OpStatus.SUCCESS,
|
||||||
|
error_message="",
|
||||||
|
error_type="",
|
||||||
|
error_traceback="",
|
||||||
|
func_name=func_name,
|
||||||
|
detail=detail_str,
|
||||||
|
user=user,
|
||||||
|
request=request,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
|
||||||
|
# 捕获完整异常堆栈
|
||||||
|
tb_lines = traceback.format_exception(type(exc), exc, exc.__traceback__)
|
||||||
|
full_traceback = "".join(tb_lines)
|
||||||
|
error_msg = str(exc)[:1000]
|
||||||
|
error_type = type(exc).__name__
|
||||||
|
|
||||||
|
_write_log(
|
||||||
|
module=module,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=str(target_id) if target_id else None,
|
||||||
|
target_name=str(target_name) if target_name else None,
|
||||||
|
status=OpStatus.FAILURE,
|
||||||
|
error_message=error_msg,
|
||||||
|
error_type=error_type,
|
||||||
|
error_traceback=full_traceback,
|
||||||
|
func_name=func_name,
|
||||||
|
detail=detail_str,
|
||||||
|
user=user,
|
||||||
|
request=request,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
return async_wrapper # type: ignore
|
||||||
|
else:
|
||||||
|
@functools.wraps(func)
|
||||||
|
def sync_wrapper(*args, **kwargs):
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
trace_id = request_id_var.get("-")
|
||||||
|
|
||||||
|
user = _extract_user(args, kwargs)
|
||||||
|
request = _extract_request(args)
|
||||||
|
|
||||||
|
target_name = _get_param(kwargs, target_name_param, "")
|
||||||
|
target_id = _get_param(kwargs, "task_id", "") or _get_param(kwargs, "dataset_id", "") or _get_param(kwargs, "model_id", "")
|
||||||
|
|
||||||
|
detail_dict = _build_detail(detail_params, kwargs)
|
||||||
|
detail_str = json.dumps(detail_dict, ensure_ascii=False) if detail_dict else ""
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = func(*args, **kwargs)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
|
||||||
|
if not target_id and isinstance(result, dict):
|
||||||
|
target_id = str(result.get("id", ""))
|
||||||
|
|
||||||
|
_write_log(
|
||||||
|
module=module,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=str(target_id) if target_id else None,
|
||||||
|
target_name=str(target_name) if target_name else None,
|
||||||
|
status=OpStatus.SUCCESS,
|
||||||
|
error_message="",
|
||||||
|
error_type="",
|
||||||
|
error_traceback="",
|
||||||
|
func_name=func_name,
|
||||||
|
detail=detail_str,
|
||||||
|
user=user,
|
||||||
|
request=request,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
|
||||||
|
tb_lines = traceback.format_exception(type(exc), exc, exc.__traceback__)
|
||||||
|
full_traceback = "".join(tb_lines)
|
||||||
|
error_msg = str(exc)[:1000]
|
||||||
|
error_type = type(exc).__name__
|
||||||
|
|
||||||
|
_write_log(
|
||||||
|
module=module,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=str(target_id) if target_id else None,
|
||||||
|
target_name=str(target_name) if target_name else None,
|
||||||
|
status=OpStatus.FAILURE,
|
||||||
|
error_message=error_msg,
|
||||||
|
error_type=error_type,
|
||||||
|
error_traceback=full_traceback,
|
||||||
|
func_name=func_name,
|
||||||
|
detail=detail_str,
|
||||||
|
user=user,
|
||||||
|
request=request,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=elapsed_ms,
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
return sync_wrapper # type: ignore
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
def log_operation(
|
||||||
|
*,
|
||||||
|
module: str,
|
||||||
|
action: str,
|
||||||
|
target_type: str = "",
|
||||||
|
target_id: str = "",
|
||||||
|
target_name: str = "",
|
||||||
|
status: str = OpStatus.SUCCESS,
|
||||||
|
error_message: str = "",
|
||||||
|
error_type: str = "",
|
||||||
|
error_traceback: str = "",
|
||||||
|
detail: str = "",
|
||||||
|
func_name: str = "",
|
||||||
|
user: Optional[dict] = None,
|
||||||
|
request: Optional[Request] = None,
|
||||||
|
duration_ms: float = 0,
|
||||||
|
) -> None:
|
||||||
|
"""手动记录操作日志(不方便用装饰器时使用)"""
|
||||||
|
trace_id = request_id_var.get("-")
|
||||||
|
_write_log(
|
||||||
|
module=module,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
target_id=target_id or None,
|
||||||
|
target_name=target_name or None,
|
||||||
|
status=status,
|
||||||
|
error_message=error_message,
|
||||||
|
error_type=error_type,
|
||||||
|
error_traceback=error_traceback,
|
||||||
|
func_name=func_name,
|
||||||
|
detail=detail,
|
||||||
|
user=user,
|
||||||
|
request=request,
|
||||||
|
trace_id=trace_id,
|
||||||
|
duration_ms=duration_ms,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_detail(detail_params: Optional[list[str]], kwargs: dict) -> dict:
|
||||||
|
"""从 kwargs 中提取需要记录的参数"""
|
||||||
|
detail_dict = {}
|
||||||
|
if detail_params:
|
||||||
|
for p in detail_params:
|
||||||
|
val = kwargs.get(p)
|
||||||
|
if val is not None:
|
||||||
|
detail_dict[p] = str(val)[:200]
|
||||||
|
return detail_dict
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_user(args: tuple, kwargs: dict) -> Optional[dict]:
|
||||||
|
"""从函数参数中提取 current_user dict"""
|
||||||
|
for arg in args:
|
||||||
|
if isinstance(arg, dict) and "id" in arg and "username" in arg:
|
||||||
|
return arg
|
||||||
|
for v in kwargs.values():
|
||||||
|
if isinstance(v, dict) and "id" in v and "username" in v:
|
||||||
|
return v
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_request(args: tuple) -> Optional[Request]:
|
||||||
|
"""从函数参数中提取 Request 对象"""
|
||||||
|
for arg in args:
|
||||||
|
if isinstance(arg, Request):
|
||||||
|
return arg
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_param(kwargs: dict, key: str, default: str = "") -> str:
|
||||||
|
"""安全获取参数值"""
|
||||||
|
val = kwargs.get(key, default)
|
||||||
|
if val is None:
|
||||||
|
return default
|
||||||
|
return str(val)
|
||||||
|
|
||||||
|
|
||||||
|
def _write_log(
|
||||||
|
module: str,
|
||||||
|
action: str,
|
||||||
|
target_type: str,
|
||||||
|
target_id: Optional[str],
|
||||||
|
target_name: Optional[str],
|
||||||
|
status: str,
|
||||||
|
error_message: str,
|
||||||
|
error_type: str,
|
||||||
|
error_traceback: str,
|
||||||
|
func_name: str,
|
||||||
|
detail: str,
|
||||||
|
user: Optional[dict],
|
||||||
|
request: Optional[Request],
|
||||||
|
trace_id: str,
|
||||||
|
duration_ms: float,
|
||||||
|
) -> None:
|
||||||
|
"""写入操作日志到数据库"""
|
||||||
|
try:
|
||||||
|
store = get_platform_store()
|
||||||
|
log_id = new_id("op")
|
||||||
|
user_id = user.get("id") if user else None
|
||||||
|
username = user.get("username") if user else None
|
||||||
|
client_ip = None
|
||||||
|
req_method = None
|
||||||
|
req_path = None
|
||||||
|
if request:
|
||||||
|
client_ip = request.client.host if request.client else None
|
||||||
|
req_method = request.method
|
||||||
|
req_path = request.url.path
|
||||||
|
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO operation_logs
|
||||||
|
(id, user_id, username, module, action, target_type, target_id,
|
||||||
|
target_name, status, error_message, error_type, error_traceback,
|
||||||
|
func_name, detail, client_ip,
|
||||||
|
request_method, request_path, trace_id, duration_ms, create_time)
|
||||||
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
log_id, user_id, username, module, action,
|
||||||
|
target_type or None, target_id, target_name,
|
||||||
|
status,
|
||||||
|
error_message[:1000] if error_message else None,
|
||||||
|
error_type or None,
|
||||||
|
error_traceback[:5000] if error_traceback else None,
|
||||||
|
func_name or None,
|
||||||
|
detail[:2000] if detail else None,
|
||||||
|
client_ip, req_method, req_path, trace_id,
|
||||||
|
round(duration_ms, 2), utcnow(),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.error("写入操作日志失败 module=%s action=%s", module, action, exc_info=True)
|
||||||
1
backend/app/db/__init__.py
Normal file
1
backend/app/db/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Database infrastructure package."""
|
||||||
4518
backend/app/db/platform_store.py
Normal file
4518
backend/app/db/platform_store.py
Normal file
File diff suppressed because it is too large
Load Diff
40
backend/app/db/session.py
Normal file
40
backend/app/db/session.py
Normal file
@@ -0,0 +1,40 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from collections.abc import Generator
|
||||||
|
from contextlib import contextmanager
|
||||||
|
|
||||||
|
from sqlalchemy import create_engine
|
||||||
|
from sqlalchemy.orm import Session, sessionmaker
|
||||||
|
|
||||||
|
|
||||||
|
DATABASE_URL = os.getenv("DATABASE_URL", "postgresql+psycopg://yg_ft:change_me@localhost:15432/yg_ft")
|
||||||
|
|
||||||
|
engine = create_engine(
|
||||||
|
DATABASE_URL,
|
||||||
|
pool_pre_ping=True,
|
||||||
|
future=True,
|
||||||
|
)
|
||||||
|
SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False, expire_on_commit=False, future=True)
|
||||||
|
|
||||||
|
|
||||||
|
def get_db() -> Generator[Session, None, None]:
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
yield db
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def session_scope() -> Generator[Session, None, None]:
|
||||||
|
db = SessionLocal()
|
||||||
|
try:
|
||||||
|
yield db
|
||||||
|
db.commit()
|
||||||
|
except Exception:
|
||||||
|
db.rollback()
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
893
backend/app/db/sql/000_full_init.sql
Normal file
893
backend/app/db/sql/000_full_init.sql
Normal file
@@ -0,0 +1,893 @@
|
|||||||
|
-- ============================================================================
|
||||||
|
-- YG Fine-Tune Platform — PostgreSQL 完整初始化脚本(一键建库建表)
|
||||||
|
-- ============================================================================
|
||||||
|
-- 用途:切换到新的 PG 数据集时,一次性创建平台运行所需的全部数据库对象与
|
||||||
|
-- 基础种子数据(幂等,可重复执行)。
|
||||||
|
--
|
||||||
|
-- 覆盖范围(与运行时代码实际使用的表一致):
|
||||||
|
-- 001_platform_runtime.sql 平台核心表
|
||||||
|
-- 002_governance.sql 治理表(租户 / 审批 / 审计 / 留存)
|
||||||
|
-- 003_tenant_quota.sql 租户配额列
|
||||||
|
-- 003_model_path_governance.sql 模型可训练标识列
|
||||||
|
-- 002_data_process.sql 数据处理表 + 数据集扩展列
|
||||||
|
-- 本文件补充:data_convert_tasks(数据转换任务,运行时代码引用但原脚本缺失)
|
||||||
|
-- 种子数据:admin / operator 两个初始用户
|
||||||
|
--
|
||||||
|
-- 说明:
|
||||||
|
-- * 本脚本通过 psql 执行,包含 DO $$ ... $$ 块与事务,不能用应用的
|
||||||
|
-- executescript()(按分号切分)执行。
|
||||||
|
-- * 应用启动时 PlatformStore.ensure_schema() 只会自动执行
|
||||||
|
-- 001 / 002_governance / 003_tenant_quota;数据处理表需另跑
|
||||||
|
-- 002_data_process.sql(本脚本已包含)。应用首次启动还会自动补充
|
||||||
|
-- admin/operator 种子用户(本脚本已包含,二选一即可)。
|
||||||
|
-- * 脚本内所有 DDL 均使用 IF NOT EXISTS / ADD COLUMN IF NOT EXISTS,
|
||||||
|
-- 可在已初始化的库上安全重复执行。
|
||||||
|
--
|
||||||
|
-- 执行步骤(详见 docs/database-config.md):
|
||||||
|
-- 1. 以超级用户创建角色与数据库(必须单独执行,不能放进事务):
|
||||||
|
-- CREATE ROLE yg_ft LOGIN PASSWORD '请改为强密码';
|
||||||
|
-- CREATE DATABASE yg_ft OWNER yg_ft;
|
||||||
|
-- 2. 连接目标库执行本脚本:
|
||||||
|
-- psql "postgresql://yg_ft:密码@<host>:5432/yg_ft" -f backend/app/db/sql/000_full_init.sql
|
||||||
|
-- 3. 可选:为 superuser 授权
|
||||||
|
-- ALTER ROLE yg_ft SUPERUSER; -- 仅当需要执行 CREATE EXTENSION 等
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
BEGIN;
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 一、平台核心表(来源:001_platform_runtime.sql)
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS users (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
username TEXT NOT NULL UNIQUE,
|
||||||
|
password_hash TEXT NOT NULL,
|
||||||
|
display_name TEXT NOT NULL,
|
||||||
|
role TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
permissions TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
last_login TEXT,
|
||||||
|
protected INTEGER NOT NULL DEFAULT 0
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS models (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
type TEXT NOT NULL,
|
||||||
|
purpose TEXT NOT NULL,
|
||||||
|
model_source TEXT NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
path TEXT,
|
||||||
|
api_url TEXT,
|
||||||
|
api_key TEXT,
|
||||||
|
online_model_name TEXT,
|
||||||
|
can_train INTEGER NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
created_by TEXT,
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
deleted_at TEXT,
|
||||||
|
deleted_by TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS trained_models (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
train_methods TEXT NOT NULL,
|
||||||
|
base_model_path TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
merged INTEGER NOT NULL DEFAULT 0,
|
||||||
|
merging INTEGER NOT NULL DEFAULT 0,
|
||||||
|
merged_path TEXT,
|
||||||
|
artifact_dir TEXT,
|
||||||
|
compute_node_id TEXT,
|
||||||
|
compute_node_name TEXT,
|
||||||
|
created_by TEXT,
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
deleted_at TEXT,
|
||||||
|
deleted_by TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS model_lineage (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
child_resource_type TEXT NOT NULL,
|
||||||
|
child_resource_id TEXT NOT NULL,
|
||||||
|
parent_resource_type TEXT NOT NULL,
|
||||||
|
parent_resource_id TEXT NOT NULL,
|
||||||
|
relation_type TEXT NOT NULL,
|
||||||
|
compute_job_id TEXT,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS model_artifacts (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
model_id TEXT NOT NULL,
|
||||||
|
model_kind TEXT NOT NULL,
|
||||||
|
artifact_type TEXT NOT NULL,
|
||||||
|
path TEXT NOT NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0,
|
||||||
|
checksum_sha256 TEXT,
|
||||||
|
metadata TEXT NOT NULL,
|
||||||
|
compute_job_id TEXT,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS model_export_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
trained_model_id TEXT,
|
||||||
|
compute_job_id TEXT NOT NULL,
|
||||||
|
node_id TEXT,
|
||||||
|
export_type TEXT NOT NULL,
|
||||||
|
quantization_bit INTEGER NOT NULL DEFAULT 0,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
output_dir TEXT,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS datasets (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
type TEXT NOT NULL,
|
||||||
|
storage_type TEXT NOT NULL,
|
||||||
|
source TEXT NOT NULL,
|
||||||
|
task_id TEXT,
|
||||||
|
size TEXT,
|
||||||
|
count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
description TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
created_by TEXT,
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
deleted_at TEXT,
|
||||||
|
deleted_by TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS dataset_files (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
dataset_id TEXT NOT NULL REFERENCES datasets(id) ON DELETE CASCADE,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
size TEXT,
|
||||||
|
content TEXT NOT NULL,
|
||||||
|
active_version_id TEXT NOT NULL,
|
||||||
|
versions TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS compute_nodes (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
code TEXT NOT NULL UNIQUE,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
api_base_url TEXT NOT NULL,
|
||||||
|
file_gateway_url TEXT NOT NULL,
|
||||||
|
enabled INTEGER NOT NULL DEFAULT 1,
|
||||||
|
scheduler_status TEXT NOT NULL,
|
||||||
|
scheduler_weight INTEGER NOT NULL DEFAULT 100,
|
||||||
|
tags TEXT NOT NULL,
|
||||||
|
gpu_count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
current_running_jobs INTEGER NOT NULL DEFAULT 0,
|
||||||
|
max_parallel_jobs INTEGER NOT NULL DEFAULT 2,
|
||||||
|
data_root TEXT NOT NULL,
|
||||||
|
model_root TEXT NOT NULL,
|
||||||
|
log_root TEXT NOT NULL,
|
||||||
|
api_version TEXT NOT NULL DEFAULT 'v1',
|
||||||
|
capabilities TEXT NOT NULL DEFAULT '[]',
|
||||||
|
description TEXT,
|
||||||
|
last_health_check_at TEXT,
|
||||||
|
health_detail TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gpus (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
gpu_index INTEGER NOT NULL,
|
||||||
|
uuid TEXT NOT NULL,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
memory_total_gb DOUBLE PRECISION NOT NULL,
|
||||||
|
power_limit_w DOUBLE PRECISION NOT NULL,
|
||||||
|
base_temperature INTEGER NOT NULL,
|
||||||
|
last_seen_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS fine_tune_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
progress INTEGER NOT NULL DEFAULT 0,
|
||||||
|
process_id INTEGER,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
start_time TEXT,
|
||||||
|
completed_at TEXT,
|
||||||
|
compute_node_id TEXT REFERENCES compute_nodes(id) ON DELETE SET NULL,
|
||||||
|
gpus TEXT NOT NULL,
|
||||||
|
sync_job_id TEXT,
|
||||||
|
compute_job_id TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS fine_tune_metrics (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||||
|
step INTEGER NOT NULL,
|
||||||
|
epoch DOUBLE PRECISION,
|
||||||
|
loss DOUBLE PRECISION,
|
||||||
|
grad_norm DOUBLE PRECISION,
|
||||||
|
learning_rate DOUBLE PRECISION,
|
||||||
|
raw TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS fine_tune_checkpoints (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||||
|
step INTEGER NOT NULL,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
path TEXT NOT NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS compute_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT REFERENCES fine_tune_tasks(id) ON DELETE SET NULL,
|
||||||
|
node_id TEXT REFERENCES compute_nodes(id) ON DELETE SET NULL,
|
||||||
|
engine TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
command TEXT NOT NULL,
|
||||||
|
output_dir TEXT,
|
||||||
|
log_file TEXT,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
update_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gpu_allocations (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||||
|
compute_job_id TEXT,
|
||||||
|
node_id TEXT REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
gpu_index INTEGER NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
released_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS scheduler_locks (
|
||||||
|
lock_key TEXT PRIMARY KEY,
|
||||||
|
owner TEXT NOT NULL,
|
||||||
|
expires_at TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
update_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS resource_replicas (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
resource_type TEXT NOT NULL,
|
||||||
|
resource_id TEXT NOT NULL,
|
||||||
|
local_path TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
sync_status TEXT NOT NULL,
|
||||||
|
checksum_sha256 TEXT,
|
||||||
|
byte_size BIGINT NOT NULL DEFAULT 0,
|
||||||
|
last_checked_at TEXT,
|
||||||
|
last_error TEXT,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS resource_sync_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
target_node_id TEXT NOT NULL,
|
||||||
|
resources TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
progress INTEGER NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS storage_objects (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
resource_type TEXT NOT NULL,
|
||||||
|
resource_id TEXT NOT NULL,
|
||||||
|
version_id TEXT NOT NULL,
|
||||||
|
bucket TEXT NOT NULL,
|
||||||
|
object_key TEXT NOT NULL,
|
||||||
|
file_name TEXT,
|
||||||
|
content_type TEXT,
|
||||||
|
checksum_sha256 TEXT,
|
||||||
|
byte_size BIGINT NOT NULL DEFAULT 0,
|
||||||
|
status TEXT NOT NULL DEFAULT 'pending',
|
||||||
|
created_by TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
UNIQUE (resource_type, resource_id, version_id, object_key)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS storage_cache_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
storage_object_id TEXT NOT NULL REFERENCES storage_objects(id) ON DELETE CASCADE,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
direction TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
progress INTEGER NOT NULL DEFAULT 0,
|
||||||
|
local_path TEXT,
|
||||||
|
error TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_storage_objects_resource ON storage_objects(resource_type, resource_id, version_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_storage_cache_jobs_node_status ON storage_cache_jobs(node_id, status);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS eval_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS eval_dimensions (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
is_active INTEGER NOT NULL DEFAULT 1,
|
||||||
|
is_default INTEGER NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS compare_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_status ON fine_tune_tasks(status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_compute_job ON fine_tune_tasks(compute_job_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_compute_node_status ON fine_tune_tasks(compute_node_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_lineage_child ON model_lineage(child_resource_type, child_resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_lineage_parent ON model_lineage(parent_resource_type, parent_resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_artifacts_model ON model_artifacts(model_kind, model_id, artifact_type);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_export_jobs_model ON model_export_jobs(trained_model_id, create_time DESC);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_export_jobs_compute ON model_export_jobs(compute_job_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_metrics_task_step ON fine_tune_metrics(task_id, step);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_fine_tune_metrics_task_step_epoch ON fine_tune_metrics(task_id, step, epoch);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_checkpoints_task_step ON fine_tune_checkpoints(task_id, step);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_fine_tune_checkpoints_task_path ON fine_tune_checkpoints(task_id, path);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_compute_jobs_task ON compute_jobs(task_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_compute_jobs_node_status ON compute_jobs(node_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpu_allocations_node_status ON gpu_allocations(node_id, status);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_gpu_allocations_active ON gpu_allocations(node_id, gpu_index) WHERE status IN ('allocated','running');
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_scheduler_locks_expires ON scheduler_locks(expires_at);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_files_dataset ON dataset_files(dataset_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpus_node ON gpus(node_id);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_gpus_node_index ON gpus(node_id, gpu_index);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_replicas_resource ON resource_replicas(resource_type, resource_id);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_replicas_node_resource ON resource_replicas(node_id, resource_type, resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_sync_jobs_node_status ON resource_sync_jobs(target_node_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_eval_tasks_status ON eval_tasks(status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_eval_dimensions_active ON eval_dimensions(is_active);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_compare_tasks_status ON compare_tasks(status);
|
||||||
|
|
||||||
|
-- ---- 项目 / 租户 ----
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS projects (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
tenant_id TEXT NOT NULL DEFAULT 'default',
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
code TEXT NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
quota TEXT,
|
||||||
|
status TEXT NOT NULL DEFAULT 'active',
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
create_by TEXT,
|
||||||
|
updated_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS project_members (
|
||||||
|
project_id TEXT NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
|
||||||
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
role TEXT NOT NULL DEFAULT 'member',
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
PRIMARY KEY (project_id, user_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS roles (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
permissions TEXT NOT NULL DEFAULT '[]',
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS sessions (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
user_id TEXT NOT NULL,
|
||||||
|
username TEXT,
|
||||||
|
login_at TEXT,
|
||||||
|
logout_at TEXT,
|
||||||
|
duration_seconds INTEGER,
|
||||||
|
issued_at TEXT,
|
||||||
|
expires_at TEXT,
|
||||||
|
ip TEXT,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS acls (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
resource_type TEXT NOT NULL,
|
||||||
|
resource_id TEXT NOT NULL,
|
||||||
|
principal_type TEXT NOT NULL,
|
||||||
|
principal_id TEXT NOT NULL,
|
||||||
|
permission TEXT NOT NULL,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
ALTER TABLE models ADD COLUMN IF NOT EXISTS deleted_at TEXT;
|
||||||
|
ALTER TABLE models ADD COLUMN IF NOT EXISTS deleted_by TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS deleted_at TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS deleted_by TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS deleted_at TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS deleted_by TEXT;
|
||||||
|
ALTER TABLE eval_tasks ADD COLUMN IF NOT EXISTS deleted_at TEXT;
|
||||||
|
ALTER TABLE eval_tasks ADD COLUMN IF NOT EXISTS deleted_by TEXT;
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_acls_subject_permission
|
||||||
|
ON acls(resource_type, resource_id, principal_type, principal_id, permission);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_acls_resource ON acls(resource_type, resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_acls_principal ON acls(principal_type, principal_id);
|
||||||
|
|
||||||
|
-- ---- 权限扩展(来源:004_permissions.sql) ----
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gpu_assignments (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
gpu_index INTEGER NOT NULL,
|
||||||
|
user_id TEXT NOT NULL,
|
||||||
|
assigned_by TEXT,
|
||||||
|
assigned_at TEXT NOT NULL,
|
||||||
|
UNIQUE (node_id, gpu_index, user_id)
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpu_assignments_user ON gpu_assignments(user_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpu_assignments_gpu ON gpu_assignments(node_id, gpu_index);
|
||||||
|
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE models ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE eval_tasks ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 二、治理表(来源:002_governance.sql)
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS tenants (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
code TEXT,
|
||||||
|
status TEXT DEFAULT 'active',
|
||||||
|
owner_user_id TEXT,
|
||||||
|
quota TEXT,
|
||||||
|
retention_policy_id TEXT,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS approval_templates (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
steps TEXT,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS approval_instances (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
template_id TEXT,
|
||||||
|
resource_type TEXT,
|
||||||
|
resource_id TEXT,
|
||||||
|
applicant_id TEXT,
|
||||||
|
status TEXT DEFAULT 'pending',
|
||||||
|
current_step INTEGER DEFAULT 0,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS approval_steps (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
instance_id TEXT,
|
||||||
|
step_index INTEGER,
|
||||||
|
approver_id TEXT,
|
||||||
|
status TEXT DEFAULT 'pending',
|
||||||
|
comment TEXT,
|
||||||
|
time TEXT
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_approval_instances_applicant ON approval_instances(applicant_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_approval_instances_resource ON approval_instances(resource_type, resource_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_approval_steps_approver ON approval_steps(approver_id, status);
|
||||||
|
|
||||||
|
-- Compatibility for databases created from an older approval_steps definition.
|
||||||
|
ALTER TABLE approval_steps ADD COLUMN IF NOT EXISTS id TEXT;
|
||||||
|
UPDATE approval_steps
|
||||||
|
SET id = 'astep_' || md5(concat_ws(':', instance_id, step_index, coalesce(approver_id, ''), coalesce(time, '')))
|
||||||
|
WHERE id IS NULL;
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS idx_approval_steps_id ON approval_steps(id);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS audit_logs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
actor_id TEXT,
|
||||||
|
action TEXT,
|
||||||
|
target_type TEXT,
|
||||||
|
target_id TEXT,
|
||||||
|
detail TEXT,
|
||||||
|
client_ip TEXT,
|
||||||
|
time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_tenant ON audit_logs(tenant_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_project ON audit_logs(project_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_action ON audit_logs(action);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_time ON audit_logs(time);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS retention_policies (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
scope TEXT,
|
||||||
|
rule TEXT,
|
||||||
|
status TEXT DEFAULT 'active',
|
||||||
|
create_time TEXT,
|
||||||
|
create_by TEXT,
|
||||||
|
updated_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Compatibility for the earlier retention policy schema
|
||||||
|
-- (resource_type/retention_days). Keep legacy columns if they exist.
|
||||||
|
ALTER TABLE retention_policies ADD COLUMN IF NOT EXISTS scope TEXT;
|
||||||
|
ALTER TABLE retention_policies ADD COLUMN IF NOT EXISTS rule TEXT;
|
||||||
|
ALTER TABLE retention_policies ADD COLUMN IF NOT EXISTS status TEXT DEFAULT 'active';
|
||||||
|
ALTER TABLE retention_policies ADD COLUMN IF NOT EXISTS create_by TEXT;
|
||||||
|
ALTER TABLE retention_policies ADD COLUMN IF NOT EXISTS updated_at TEXT;
|
||||||
|
DO $$
|
||||||
|
BEGIN
|
||||||
|
IF EXISTS (
|
||||||
|
SELECT 1 FROM information_schema.columns
|
||||||
|
WHERE table_schema='public' AND table_name='retention_policies' AND column_name='resource_type'
|
||||||
|
) AND EXISTS (
|
||||||
|
SELECT 1 FROM information_schema.columns
|
||||||
|
WHERE table_schema='public' AND table_name='retention_policies' AND column_name='retention_days'
|
||||||
|
) THEN
|
||||||
|
EXECUTE $migration$
|
||||||
|
UPDATE retention_policies
|
||||||
|
SET scope = COALESCE(scope, resource_type),
|
||||||
|
rule = COALESCE(rule, retention_days::text),
|
||||||
|
status = COALESCE(status, 'active'),
|
||||||
|
updated_at = COALESCE(updated_at, create_time)
|
||||||
|
WHERE scope IS NULL OR rule IS NULL OR status IS NULL OR updated_at IS NULL
|
||||||
|
$migration$;
|
||||||
|
ELSE
|
||||||
|
UPDATE retention_policies
|
||||||
|
SET status = COALESCE(status, 'active'),
|
||||||
|
updated_at = COALESCE(updated_at, create_time)
|
||||||
|
WHERE status IS NULL OR updated_at IS NULL;
|
||||||
|
END IF;
|
||||||
|
END $$;
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 三、租户配额扩展(来源:003_tenant_quota.sql)
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
ALTER TABLE tenants ADD COLUMN IF NOT EXISTS gpu_quota TEXT;
|
||||||
|
ALTER TABLE tenants ADD COLUMN IF NOT EXISTS storage_quota TEXT;
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 四、模型路径治理(来源:003_model_path_governance.sql)
|
||||||
|
-- models.can_train 已在建表语句中声明;以下为兼容旧库的幂等语句。
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
ALTER TABLE models ADD COLUMN IF NOT EXISTS can_train INTEGER NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS artifact_dir TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS compute_node_id TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS compute_node_name TEXT;
|
||||||
|
|
||||||
|
-- 按规则推定已有模型的 can_train(新库为空表,此语句为 no-op)
|
||||||
|
UPDATE models
|
||||||
|
SET can_train = CASE
|
||||||
|
WHEN path IS NOT NULL AND path != '' AND model_source IS NOT NULL AND model_source != 'api' THEN 1
|
||||||
|
ELSE 0
|
||||||
|
END;
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 五、数据处理(来源:002_data_process.sql,去掉其外层 BEGIN/COMMIT)
|
||||||
|
-- 数据集扩展列
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS size_bytes BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS record_count BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}';
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS tenant_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS project_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS owner_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ;
|
||||||
|
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS storage_object_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS current_version_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS size_bytes BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS record_count BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS file_format VARCHAR(40);
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS checksum_sha256 CHAR(64);
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS version_no INTEGER NOT NULL DEFAULT 1;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS tenant_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS project_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}';
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ;
|
||||||
|
|
||||||
|
-- ---- 数据处理任务 / 源文件 / 预览 / 结果 ----
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name VARCHAR(150) NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'pending'
|
||||||
|
CHECK (status IN ('pending', 'running', 'completed', 'failed', 'stopped')),
|
||||||
|
process_type VARCHAR(20) NOT NULL
|
||||||
|
CHECK (process_type IN ('structured', 'unstructured', 'external')),
|
||||||
|
source_dataset_id TEXT REFERENCES datasets(id) ON DELETE SET NULL,
|
||||||
|
output_dataset_id TEXT REFERENCES datasets(id) ON DELETE SET NULL,
|
||||||
|
config TEXT NOT NULL DEFAULT '{}',
|
||||||
|
progress NUMERIC(5,2) NOT NULL DEFAULT 0 CHECK (progress >= 0 AND progress <= 100),
|
||||||
|
input_count BIGINT NOT NULL DEFAULT 0 CHECK (input_count >= 0),
|
||||||
|
output_count BIGINT NOT NULL DEFAULT 0 CHECK (output_count >= 0),
|
||||||
|
filtered_count BIGINT NOT NULL DEFAULT 0 CHECK (filtered_count >= 0),
|
||||||
|
duplicate_count BIGINT NOT NULL DEFAULT 0 CHECK (duplicate_count >= 0),
|
||||||
|
error_count BIGINT NOT NULL DEFAULT 0 CHECK (error_count >= 0),
|
||||||
|
failure_reason TEXT,
|
||||||
|
generation_run_id TEXT,
|
||||||
|
results_confirmed BOOLEAN NOT NULL DEFAULT TRUE,
|
||||||
|
workflow_step VARCHAR(20) NOT NULL DEFAULT 'create'
|
||||||
|
CHECK (workflow_step IN ('create', 'model', 'upload', 'preview', 'generate', 'results')),
|
||||||
|
preview_status VARCHAR(20) NOT NULL DEFAULT 'idle'
|
||||||
|
CHECK (preview_status IN ('idle', 'queued', 'running', 'completed', 'failed', 'cancelled')),
|
||||||
|
preview_progress NUMERIC(5,2) NOT NULL DEFAULT 0
|
||||||
|
CHECK (preview_progress >= 0 AND preview_progress <= 100),
|
||||||
|
preview_run_id TEXT,
|
||||||
|
preview_failure_reason TEXT,
|
||||||
|
preview_total_files INTEGER NOT NULL DEFAULT 0 CHECK (preview_total_files >= 0),
|
||||||
|
preview_completed_files INTEGER NOT NULL DEFAULT 0 CHECK (preview_completed_files >= 0),
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
owner_id TEXT,
|
||||||
|
approval_status VARCHAR(30) NOT NULL DEFAULT 'not_required',
|
||||||
|
created_by TEXT,
|
||||||
|
updated_by TEXT,
|
||||||
|
deleted_by TEXT,
|
||||||
|
started_at TIMESTAMPTZ,
|
||||||
|
completed_at TIMESTAMPTZ,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS generation_run_id TEXT;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS results_confirmed BOOLEAN NOT NULL DEFAULT TRUE;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS workflow_step VARCHAR(20);
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_status VARCHAR(20);
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_progress NUMERIC(5,2);
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_run_id TEXT;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_failure_reason TEXT;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_total_files INTEGER;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_completed_files INTEGER;
|
||||||
|
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN workflow_step SET DEFAULT 'create';
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN workflow_step SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_status SET DEFAULT 'idle';
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_status SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_progress SET DEFAULT 0;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_progress SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_total_files SET DEFAULT 0;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_total_files SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_completed_files SET DEFAULT 0;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_completed_files SET NOT NULL;
|
||||||
|
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_data_process_tasks_name_alive
|
||||||
|
ON data_process_tasks(name) WHERE deleted_at IS NULL;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_tasks_scope_status
|
||||||
|
ON data_process_tasks(tenant_id, project_id, status, created_at DESC)
|
||||||
|
WHERE deleted_at IS NULL;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_tasks_creator_created
|
||||||
|
ON data_process_tasks(created_by, created_at DESC) WHERE deleted_at IS NULL;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_source_files (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE,
|
||||||
|
storage_object_id TEXT,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0 CHECK (size_bytes >= 0),
|
||||||
|
record_count BIGINT NOT NULL DEFAULT 0 CHECK (record_count >= 0),
|
||||||
|
file_format VARCHAR(40),
|
||||||
|
checksum_sha256 CHAR(64) NOT NULL,
|
||||||
|
version_no INTEGER NOT NULL DEFAULT 1 CHECK (version_no > 0),
|
||||||
|
content TEXT NOT NULL,
|
||||||
|
content_preview TEXT,
|
||||||
|
metadata TEXT NOT NULL DEFAULT '{}',
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
created_by TEXT,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_source_files_task
|
||||||
|
ON data_process_source_files(task_id, created_at) WHERE deleted_at IS NULL;
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_data_process_source_checksum_alive
|
||||||
|
ON data_process_source_files(task_id, checksum_sha256) WHERE deleted_at IS NULL;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_preview_items (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE,
|
||||||
|
source_file_id TEXT REFERENCES data_process_source_files(id) ON DELETE CASCADE,
|
||||||
|
original_content TEXT NOT NULL DEFAULT '',
|
||||||
|
edited_content TEXT NOT NULL DEFAULT '',
|
||||||
|
source_start INTEGER CHECK (source_start IS NULL OR source_start >= 0),
|
||||||
|
source_end INTEGER CHECK (source_end IS NULL OR source_end >= 0),
|
||||||
|
source_start_line INTEGER CHECK (source_start_line IS NULL OR source_start_line > 0),
|
||||||
|
source_end_line INTEGER CHECK (source_end_line IS NULL OR source_end_line > 0),
|
||||||
|
token_count INTEGER NOT NULL DEFAULT 0 CHECK (token_count >= 0),
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'original'
|
||||||
|
CHECK (status IN ('original', 'modified', 'manual', 'invalid')),
|
||||||
|
quality_score TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
CHECK (source_start IS NULL OR source_end IS NULL OR source_end >= source_start),
|
||||||
|
CHECK (source_start_line IS NULL OR source_end_line IS NULL OR source_end_line >= source_start_line)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_preview_task_file
|
||||||
|
ON data_process_preview_items(task_id, source_file_id, created_at);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_results (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE,
|
||||||
|
preview_item_id TEXT REFERENCES data_process_preview_items(id) ON DELETE SET NULL,
|
||||||
|
instruction TEXT NOT NULL,
|
||||||
|
input TEXT NOT NULL DEFAULT '',
|
||||||
|
output TEXT NOT NULL,
|
||||||
|
chosen TEXT NOT NULL DEFAULT '',
|
||||||
|
rejected TEXT NOT NULL DEFAULT '',
|
||||||
|
original_instruction TEXT,
|
||||||
|
original_input TEXT,
|
||||||
|
original_output TEXT,
|
||||||
|
original_chosen TEXT,
|
||||||
|
original_rejected TEXT,
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'valid'
|
||||||
|
CHECK (status IN ('valid', 'modified', 'invalid')),
|
||||||
|
error TEXT,
|
||||||
|
split VARCHAR(20) CHECK (split IS NULL OR split IN ('train', 'validation', 'test')),
|
||||||
|
quality_score TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Keep existing databases compatible with the current generation result model.
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS chosen TEXT NOT NULL DEFAULT '';
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS rejected TEXT NOT NULL DEFAULT '';
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS original_chosen TEXT;
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS original_rejected TEXT;
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_results_task_status
|
||||||
|
ON data_process_results(task_id, status, id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_results_task_split
|
||||||
|
ON data_process_results(task_id, split);
|
||||||
|
|
||||||
|
-- ---- 数据集版本 / 记录 ----
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS dataset_file_versions (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
dataset_file_id TEXT NOT NULL REFERENCES dataset_files(id) ON DELETE CASCADE,
|
||||||
|
version_no INTEGER NOT NULL CHECK (version_no > 0),
|
||||||
|
storage_object_id TEXT NOT NULL,
|
||||||
|
content_preview TEXT,
|
||||||
|
description TEXT,
|
||||||
|
base_version_id TEXT REFERENCES dataset_file_versions(id) ON DELETE SET NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0 CHECK (size_bytes >= 0),
|
||||||
|
record_count BIGINT NOT NULL DEFAULT 0 CHECK (record_count >= 0),
|
||||||
|
checksum_sha256 CHAR(64) NOT NULL,
|
||||||
|
source_task_id TEXT REFERENCES data_process_tasks(id) ON DELETE SET NULL,
|
||||||
|
metadata TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_by TEXT,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE dataset_file_versions ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE dataset_file_versions ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}';
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_dataset_file_versions_no_002
|
||||||
|
ON dataset_file_versions(dataset_file_id, version_no);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_file_versions_source_task_002
|
||||||
|
ON dataset_file_versions(source_task_id) WHERE source_task_id IS NOT NULL;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS dataset_records (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
dataset_id TEXT NOT NULL REFERENCES datasets(id) ON DELETE CASCADE,
|
||||||
|
dataset_file_id TEXT REFERENCES dataset_files(id) ON DELETE CASCADE,
|
||||||
|
version_id TEXT REFERENCES dataset_file_versions(id) ON DELETE CASCADE,
|
||||||
|
line_no INTEGER,
|
||||||
|
split VARCHAR(20) CHECK (split IS NULL OR split IN ('train', 'validation', 'test')),
|
||||||
|
instruction TEXT,
|
||||||
|
input TEXT,
|
||||||
|
output TEXT,
|
||||||
|
raw TEXT NOT NULL DEFAULT '{}',
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'valid'
|
||||||
|
CHECK (status IN ('valid', 'modified', 'invalid')),
|
||||||
|
source_task_id TEXT REFERENCES data_process_tasks(id) ON DELETE SET NULL,
|
||||||
|
source_result_id TEXT REFERENCES data_process_results(id) ON DELETE SET NULL,
|
||||||
|
preview_item_id TEXT REFERENCES data_process_preview_items(id) ON DELETE SET NULL,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS source_result_id TEXT;
|
||||||
|
ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS preview_item_id TEXT;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_records_dataset_002
|
||||||
|
ON dataset_records(dataset_id, id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_records_source_task_002
|
||||||
|
ON dataset_records(source_task_id, source_result_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_datasets_source_task_002
|
||||||
|
ON datasets(source_task_id) WHERE source_task_id IS NOT NULL;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_files_source_task_002
|
||||||
|
ON dataset_files(source_task_id) WHERE source_task_id IS NOT NULL;
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 六、数据转换任务(data_convert_tasks)
|
||||||
|
-- 运行时 router(app/modules/data_convert/router.py)引用但原脚本缺失,本文件补齐。
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_convert_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
output_filename TEXT DEFAULT 'converted-data.jsonl',
|
||||||
|
status TEXT NOT NULL DEFAULT 'pending',
|
||||||
|
input_count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
output_count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
error_message TEXT,
|
||||||
|
create_time TEXT NOT NULL DEFAULT (to_char(now(), 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"')),
|
||||||
|
update_time TEXT NOT NULL DEFAULT (to_char(now(), 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"')),
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_convert_tasks_status ON data_convert_tasks(status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_convert_tasks_create_time ON data_convert_tasks(create_time DESC);
|
||||||
|
|
||||||
|
-- ============================================================================
|
||||||
|
-- 七、种子数据:初始管理员 / 操作员
|
||||||
|
-- 应用首次启动(ensure_seed_data)也会自动创建;此处提供以便脱离应用直接初始化。
|
||||||
|
-- 密码:admin / admin123,operator / operator123(上线前请改密)。
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
INSERT INTO users
|
||||||
|
(id, username, password_hash, display_name, role, status, permissions, create_time, protected)
|
||||||
|
VALUES
|
||||||
|
(
|
||||||
|
'u_admin', 'admin', 'pbkdf2_sha256$390000$ygft_init_salt_admin$2b6f31f22968c4f5a30bcf0acf066b7a0f58d4773d15c5ab898ba715ea87b5bd',
|
||||||
|
'Platform Admin', 'admin', 'active',
|
||||||
|
'["dashboard","fine-tune","model-eval","model-inference","model-manage","dataset","data-process","data-convert","compute","hardware","logs","user-settings"]',
|
||||||
|
to_char(now(), 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), 1
|
||||||
|
),
|
||||||
|
(
|
||||||
|
'u_operator', 'operator', 'pbkdf2_sha256$390000$ygft_init_salt_op$525bf35d02ed26f37952cbd6862b0ae358b9d1a7fa0cbbf0217aa2b5dd544125',
|
||||||
|
'Platform Operator', 'operator', 'active',
|
||||||
|
'["dashboard","fine-tune","model-eval","model-inference","model-manage","dataset","data-process","data-convert","compute","hardware","logs"]',
|
||||||
|
to_char(now(), 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), 0
|
||||||
|
)
|
||||||
|
ON CONFLICT (username) DO NOTHING;
|
||||||
|
|
||||||
|
COMMIT;
|
||||||
373
backend/app/db/sql/001_platform_runtime.sql
Normal file
373
backend/app/db/sql/001_platform_runtime.sql
Normal file
@@ -0,0 +1,373 @@
|
|||||||
|
CREATE TABLE IF NOT EXISTS users (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
username TEXT NOT NULL UNIQUE,
|
||||||
|
password_hash TEXT NOT NULL,
|
||||||
|
display_name TEXT NOT NULL,
|
||||||
|
role TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
permissions TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
last_login TEXT,
|
||||||
|
protected INTEGER NOT NULL DEFAULT 0
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS models (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
type TEXT NOT NULL,
|
||||||
|
purpose TEXT NOT NULL,
|
||||||
|
model_source TEXT NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
path TEXT,
|
||||||
|
api_url TEXT,
|
||||||
|
api_key TEXT,
|
||||||
|
online_model_name TEXT,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS trained_models (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
train_methods TEXT NOT NULL,
|
||||||
|
base_model_path TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
merged INTEGER NOT NULL DEFAULT 0,
|
||||||
|
merging INTEGER NOT NULL DEFAULT 0,
|
||||||
|
merged_path TEXT,
|
||||||
|
artifact_dir TEXT,
|
||||||
|
compute_node_id TEXT,
|
||||||
|
compute_node_name TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS model_lineage (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
child_resource_type TEXT NOT NULL,
|
||||||
|
child_resource_id TEXT NOT NULL,
|
||||||
|
parent_resource_type TEXT NOT NULL,
|
||||||
|
parent_resource_id TEXT NOT NULL,
|
||||||
|
relation_type TEXT NOT NULL,
|
||||||
|
compute_job_id TEXT,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS model_artifacts (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
model_id TEXT NOT NULL,
|
||||||
|
model_kind TEXT NOT NULL,
|
||||||
|
artifact_type TEXT NOT NULL,
|
||||||
|
path TEXT NOT NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0,
|
||||||
|
checksum_sha256 TEXT,
|
||||||
|
metadata TEXT NOT NULL,
|
||||||
|
compute_job_id TEXT,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS model_export_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
trained_model_id TEXT,
|
||||||
|
compute_job_id TEXT NOT NULL,
|
||||||
|
node_id TEXT,
|
||||||
|
export_type TEXT NOT NULL,
|
||||||
|
quantization_bit INTEGER NOT NULL DEFAULT 0,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
output_dir TEXT,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS datasets (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
type TEXT NOT NULL,
|
||||||
|
storage_type TEXT NOT NULL,
|
||||||
|
source TEXT NOT NULL,
|
||||||
|
task_id TEXT,
|
||||||
|
size TEXT,
|
||||||
|
count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
description TEXT,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS dataset_files (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
dataset_id TEXT NOT NULL REFERENCES datasets(id) ON DELETE CASCADE,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
size TEXT,
|
||||||
|
content TEXT NOT NULL,
|
||||||
|
active_version_id TEXT NOT NULL,
|
||||||
|
versions TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS compute_nodes (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
code TEXT NOT NULL UNIQUE,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
api_base_url TEXT NOT NULL,
|
||||||
|
file_gateway_url TEXT NOT NULL,
|
||||||
|
enabled INTEGER NOT NULL DEFAULT 1,
|
||||||
|
scheduler_status TEXT NOT NULL,
|
||||||
|
scheduler_weight INTEGER NOT NULL DEFAULT 100,
|
||||||
|
tags TEXT NOT NULL,
|
||||||
|
gpu_count INTEGER NOT NULL DEFAULT 0,
|
||||||
|
current_running_jobs INTEGER NOT NULL DEFAULT 0,
|
||||||
|
max_parallel_jobs INTEGER NOT NULL DEFAULT 2,
|
||||||
|
data_root TEXT NOT NULL,
|
||||||
|
model_root TEXT NOT NULL,
|
||||||
|
log_root TEXT NOT NULL,
|
||||||
|
api_version TEXT NOT NULL DEFAULT 'v1',
|
||||||
|
capabilities TEXT NOT NULL DEFAULT '[]',
|
||||||
|
description TEXT,
|
||||||
|
last_health_check_at TEXT,
|
||||||
|
health_detail TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gpus (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
gpu_index INTEGER NOT NULL,
|
||||||
|
uuid TEXT NOT NULL,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
memory_total_gb DOUBLE PRECISION NOT NULL,
|
||||||
|
power_limit_w DOUBLE PRECISION NOT NULL,
|
||||||
|
base_temperature INTEGER NOT NULL,
|
||||||
|
last_seen_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS fine_tune_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL UNIQUE,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
progress INTEGER NOT NULL DEFAULT 0,
|
||||||
|
process_id INTEGER,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
start_time TEXT,
|
||||||
|
completed_at TEXT,
|
||||||
|
compute_node_id TEXT REFERENCES compute_nodes(id) ON DELETE SET NULL,
|
||||||
|
gpus TEXT NOT NULL,
|
||||||
|
sync_job_id TEXT,
|
||||||
|
compute_job_id TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS fine_tune_metrics (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||||
|
step INTEGER NOT NULL,
|
||||||
|
epoch DOUBLE PRECISION,
|
||||||
|
loss DOUBLE PRECISION,
|
||||||
|
grad_norm DOUBLE PRECISION,
|
||||||
|
learning_rate DOUBLE PRECISION,
|
||||||
|
raw TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS fine_tune_checkpoints (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||||
|
step INTEGER NOT NULL,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
path TEXT NOT NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS compute_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT REFERENCES fine_tune_tasks(id) ON DELETE SET NULL,
|
||||||
|
node_id TEXT REFERENCES compute_nodes(id) ON DELETE SET NULL,
|
||||||
|
engine TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
command TEXT NOT NULL,
|
||||||
|
output_dir TEXT,
|
||||||
|
log_file TEXT,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
update_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gpu_allocations (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||||
|
compute_job_id TEXT,
|
||||||
|
node_id TEXT REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
gpu_index INTEGER NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
released_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS scheduler_locks (
|
||||||
|
lock_key TEXT PRIMARY KEY,
|
||||||
|
owner TEXT NOT NULL,
|
||||||
|
expires_at TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
update_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS resource_replicas (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
resource_type TEXT NOT NULL,
|
||||||
|
resource_id TEXT NOT NULL,
|
||||||
|
local_path TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
sync_status TEXT NOT NULL,
|
||||||
|
checksum_sha256 TEXT,
|
||||||
|
byte_size BIGINT NOT NULL DEFAULT 0,
|
||||||
|
last_checked_at TEXT,
|
||||||
|
last_error TEXT,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS resource_sync_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
target_node_id TEXT NOT NULL,
|
||||||
|
resources TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
progress INTEGER NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS storage_objects (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
resource_type TEXT NOT NULL,
|
||||||
|
resource_id TEXT NOT NULL,
|
||||||
|
version_id TEXT NOT NULL,
|
||||||
|
bucket TEXT NOT NULL,
|
||||||
|
object_key TEXT NOT NULL,
|
||||||
|
file_name TEXT,
|
||||||
|
content_type TEXT,
|
||||||
|
checksum_sha256 TEXT,
|
||||||
|
byte_size BIGINT NOT NULL DEFAULT 0,
|
||||||
|
status TEXT NOT NULL DEFAULT 'pending',
|
||||||
|
created_by TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
UNIQUE (resource_type, resource_id, version_id, object_key)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS storage_cache_jobs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
storage_object_id TEXT NOT NULL REFERENCES storage_objects(id) ON DELETE CASCADE,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
direction TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
progress INTEGER NOT NULL DEFAULT 0,
|
||||||
|
local_path TEXT,
|
||||||
|
error TEXT,
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
completed_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_storage_objects_resource ON storage_objects(resource_type, resource_id, version_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_storage_cache_jobs_node_status ON storage_cache_jobs(node_id, status);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS eval_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS eval_dimensions (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
is_active INTEGER NOT NULL DEFAULT 1,
|
||||||
|
is_default INTEGER NOT NULL DEFAULT 0,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS compare_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
payload TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL,
|
||||||
|
create_time TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_status ON fine_tune_tasks(status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_compute_job ON fine_tune_tasks(compute_job_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_compute_node_status ON fine_tune_tasks(compute_node_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_lineage_child ON model_lineage(child_resource_type, child_resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_lineage_parent ON model_lineage(parent_resource_type, parent_resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_artifacts_model ON model_artifacts(model_kind, model_id, artifact_type);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_export_jobs_model ON model_export_jobs(trained_model_id, create_time DESC);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_model_export_jobs_compute ON model_export_jobs(compute_job_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_metrics_task_step ON fine_tune_metrics(task_id, step);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_fine_tune_metrics_task_step_epoch ON fine_tune_metrics(task_id, step, epoch);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_fine_tune_checkpoints_task_step ON fine_tune_checkpoints(task_id, step);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_fine_tune_checkpoints_task_path ON fine_tune_checkpoints(task_id, path);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_compute_jobs_task ON compute_jobs(task_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_compute_jobs_node_status ON compute_jobs(node_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpu_allocations_node_status ON gpu_allocations(node_id, status);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_gpu_allocations_active ON gpu_allocations(node_id, gpu_index) WHERE status IN ('allocated','running');
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_scheduler_locks_expires ON scheduler_locks(expires_at);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_files_dataset ON dataset_files(dataset_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpus_node ON gpus(node_id);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_gpus_node_index ON gpus(node_id, gpu_index);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_replicas_resource ON resource_replicas(resource_type, resource_id);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_replicas_node_resource ON resource_replicas(node_id, resource_type, resource_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_sync_jobs_node_status ON resource_sync_jobs(target_node_id, status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_eval_tasks_status ON eval_tasks(status);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_eval_dimensions_active ON eval_dimensions(is_active);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_compare_tasks_status ON compare_tasks(status);
|
||||||
|
|
||||||
|
-- ===================== Project / Tenant =====================
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS projects (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
tenant_id TEXT NOT NULL DEFAULT 'default',
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
code TEXT NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
quota TEXT,
|
||||||
|
status TEXT NOT NULL DEFAULT 'active',
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
create_by TEXT,
|
||||||
|
updated_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS project_members (
|
||||||
|
project_id TEXT NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
|
||||||
|
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||||
|
role TEXT NOT NULL DEFAULT 'member',
|
||||||
|
create_time TEXT NOT NULL,
|
||||||
|
PRIMARY KEY (project_id, user_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS roles (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
permissions TEXT NOT NULL DEFAULT '[]',
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS sessions (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
user_id TEXT NOT NULL,
|
||||||
|
username TEXT,
|
||||||
|
login_at TEXT,
|
||||||
|
logout_at TEXT,
|
||||||
|
duration_seconds INTEGER,
|
||||||
|
issued_at TEXT,
|
||||||
|
expires_at TEXT,
|
||||||
|
ip TEXT,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS acls (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
resource_type TEXT NOT NULL,
|
||||||
|
resource_id TEXT NOT NULL,
|
||||||
|
principal_type TEXT NOT NULL,
|
||||||
|
principal_id TEXT NOT NULL,
|
||||||
|
permission TEXT NOT NULL,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
352
backend/app/db/sql/002_data_process.sql
Normal file
352
backend/app/db/sql/002_data_process.sql
Normal file
@@ -0,0 +1,352 @@
|
|||||||
|
-- Data processing migration.
|
||||||
|
--
|
||||||
|
-- IMPORTANT: This file is intentionally NOT wired into application startup.
|
||||||
|
-- Apply it explicitly in a controlled deployment, or call
|
||||||
|
-- DataProcessStore.ensure_schema() from an administrative command.
|
||||||
|
|
||||||
|
BEGIN;
|
||||||
|
|
||||||
|
-- This migration targets the current runtime schema created by
|
||||||
|
-- 001_platform_runtime.sql. Refuse the UUID/JSONB target-design schema instead
|
||||||
|
-- of partially altering it with incompatible TEXT foreign keys.
|
||||||
|
DO $$
|
||||||
|
DECLARE
|
||||||
|
datasets_id_type TEXT;
|
||||||
|
BEGIN
|
||||||
|
SELECT format_type(a.atttypid, a.atttypmod)
|
||||||
|
INTO datasets_id_type
|
||||||
|
FROM pg_attribute a
|
||||||
|
JOIN pg_class c ON c.oid = a.attrelid
|
||||||
|
JOIN pg_namespace n ON n.oid = c.relnamespace
|
||||||
|
WHERE n.nspname = current_schema()
|
||||||
|
AND c.relname = 'datasets'
|
||||||
|
AND a.attname = 'id'
|
||||||
|
AND a.attnum > 0
|
||||||
|
AND NOT a.attisdropped;
|
||||||
|
IF datasets_id_type IS NULL THEN
|
||||||
|
RAISE EXCEPTION '002_data_process.sql requires 001_platform_runtime.sql first';
|
||||||
|
END IF;
|
||||||
|
IF datasets_id_type <> 'text' THEN
|
||||||
|
RAISE EXCEPTION
|
||||||
|
'002_data_process.sql supports only the current TEXT runtime schema; found datasets.id type %',
|
||||||
|
datasets_id_type;
|
||||||
|
END IF;
|
||||||
|
END $$;
|
||||||
|
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS size_bytes BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS record_count BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}';
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS tenant_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS project_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS owner_id TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ;
|
||||||
|
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS storage_object_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS current_version_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS size_bytes BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS record_count BIGINT NOT NULL DEFAULT 0;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS file_format VARCHAR(40);
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS checksum_sha256 CHAR(64);
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS version_no INTEGER NOT NULL DEFAULT 1;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS tenant_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS project_id TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}';
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS created_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||||
|
ALTER TABLE dataset_files ADD COLUMN IF NOT EXISTS deleted_at TIMESTAMPTZ;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_tasks (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name VARCHAR(150) NOT NULL,
|
||||||
|
description TEXT,
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'pending'
|
||||||
|
CHECK (status IN ('pending', 'running', 'completed', 'failed', 'stopped')),
|
||||||
|
process_type VARCHAR(20) NOT NULL
|
||||||
|
CHECK (process_type IN ('structured', 'unstructured', 'external')),
|
||||||
|
source_dataset_id TEXT REFERENCES datasets(id) ON DELETE SET NULL,
|
||||||
|
output_dataset_id TEXT REFERENCES datasets(id) ON DELETE SET NULL,
|
||||||
|
config TEXT NOT NULL DEFAULT '{}',
|
||||||
|
progress NUMERIC(5,2) NOT NULL DEFAULT 0 CHECK (progress >= 0 AND progress <= 100),
|
||||||
|
input_count BIGINT NOT NULL DEFAULT 0 CHECK (input_count >= 0),
|
||||||
|
output_count BIGINT NOT NULL DEFAULT 0 CHECK (output_count >= 0),
|
||||||
|
filtered_count BIGINT NOT NULL DEFAULT 0 CHECK (filtered_count >= 0),
|
||||||
|
duplicate_count BIGINT NOT NULL DEFAULT 0 CHECK (duplicate_count >= 0),
|
||||||
|
error_count BIGINT NOT NULL DEFAULT 0 CHECK (error_count >= 0),
|
||||||
|
failure_reason TEXT,
|
||||||
|
generation_run_id TEXT,
|
||||||
|
results_confirmed BOOLEAN NOT NULL DEFAULT TRUE,
|
||||||
|
workflow_step VARCHAR(20) NOT NULL DEFAULT 'create'
|
||||||
|
CHECK (workflow_step IN ('create', 'model', 'upload', 'preview', 'generate', 'results')),
|
||||||
|
preview_status VARCHAR(20) NOT NULL DEFAULT 'idle'
|
||||||
|
CHECK (preview_status IN ('idle', 'queued', 'running', 'completed', 'failed', 'cancelled')),
|
||||||
|
preview_progress NUMERIC(5,2) NOT NULL DEFAULT 0
|
||||||
|
CHECK (preview_progress >= 0 AND preview_progress <= 100),
|
||||||
|
preview_run_id TEXT,
|
||||||
|
preview_failure_reason TEXT,
|
||||||
|
preview_total_files INTEGER NOT NULL DEFAULT 0 CHECK (preview_total_files >= 0),
|
||||||
|
preview_completed_files INTEGER NOT NULL DEFAULT 0 CHECK (preview_completed_files >= 0),
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
owner_id TEXT,
|
||||||
|
approval_status VARCHAR(30) NOT NULL DEFAULT 'not_required',
|
||||||
|
created_by TEXT,
|
||||||
|
updated_by TEXT,
|
||||||
|
deleted_by TEXT,
|
||||||
|
started_at TIMESTAMPTZ,
|
||||||
|
completed_at TIMESTAMPTZ,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS generation_run_id TEXT;
|
||||||
|
-- 历史任务在引入六步确认流程前已经完成审核,默认保留为已确认;
|
||||||
|
-- 新任务由创建接口显式写入 FALSE,并在第六步确认后转为 TRUE。
|
||||||
|
ALTER TABLE data_process_tasks
|
||||||
|
ADD COLUMN IF NOT EXISTS results_confirmed BOOLEAN NOT NULL DEFAULT TRUE;
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET results_confirmed=FALSE
|
||||||
|
WHERE status <> 'completed' AND results_confirmed=TRUE;
|
||||||
|
|
||||||
|
-- 先以可空列接入旧库,才能只回填历史行;随后再收紧默认值与约束。
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS workflow_step VARCHAR(20);
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_status VARCHAR(20);
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_progress NUMERIC(5,2);
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_run_id TEXT;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_failure_reason TEXT;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_total_files INTEGER;
|
||||||
|
ALTER TABLE data_process_tasks ADD COLUMN IF NOT EXISTS preview_completed_files INTEGER;
|
||||||
|
|
||||||
|
CREATE TEMP TABLE data_process_workflow_backfill_ids ON COMMIT DROP AS
|
||||||
|
SELECT id FROM data_process_tasks WHERE workflow_step IS NULL;
|
||||||
|
|
||||||
|
UPDATE data_process_tasks task
|
||||||
|
SET workflow_step = CASE
|
||||||
|
WHEN task.status IN ('running', 'failed', 'stopped') THEN 'generate'
|
||||||
|
WHEN task.status = 'completed' AND task.results_confirmed=FALSE THEN 'generate'
|
||||||
|
WHEN task.status = 'completed' THEN 'results'
|
||||||
|
ELSE 'create'
|
||||||
|
END
|
||||||
|
WHERE task.workflow_step IS NULL;
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET preview_status='idle', preview_progress=0,
|
||||||
|
preview_total_files=0, preview_completed_files=0
|
||||||
|
WHERE preview_status IS NULL OR preview_progress IS NULL
|
||||||
|
OR preview_total_files IS NULL OR preview_completed_files IS NULL;
|
||||||
|
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN workflow_step SET DEFAULT 'create';
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN workflow_step SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_status SET DEFAULT 'idle';
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_status SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_progress SET DEFAULT 0;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_progress SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_total_files SET DEFAULT 0;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_total_files SET NOT NULL;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_completed_files SET DEFAULT 0;
|
||||||
|
ALTER TABLE data_process_tasks ALTER COLUMN preview_completed_files SET NOT NULL;
|
||||||
|
|
||||||
|
DO $$
|
||||||
|
BEGIN
|
||||||
|
IF NOT EXISTS (
|
||||||
|
SELECT 1 FROM pg_constraint
|
||||||
|
WHERE conrelid='data_process_tasks'::regclass
|
||||||
|
AND conname='ck_data_process_tasks_workflow_step'
|
||||||
|
) THEN
|
||||||
|
ALTER TABLE data_process_tasks ADD CONSTRAINT ck_data_process_tasks_workflow_step
|
||||||
|
CHECK (workflow_step IN ('create', 'model', 'upload', 'preview', 'generate', 'results'));
|
||||||
|
END IF;
|
||||||
|
IF NOT EXISTS (
|
||||||
|
SELECT 1 FROM pg_constraint
|
||||||
|
WHERE conrelid='data_process_tasks'::regclass
|
||||||
|
AND conname='ck_data_process_tasks_preview_status'
|
||||||
|
) THEN
|
||||||
|
ALTER TABLE data_process_tasks ADD CONSTRAINT ck_data_process_tasks_preview_status
|
||||||
|
CHECK (preview_status IN ('idle', 'queued', 'running', 'completed', 'failed', 'cancelled'));
|
||||||
|
END IF;
|
||||||
|
IF NOT EXISTS (
|
||||||
|
SELECT 1 FROM pg_constraint
|
||||||
|
WHERE conrelid='data_process_tasks'::regclass
|
||||||
|
AND conname='ck_data_process_tasks_preview_progress'
|
||||||
|
) THEN
|
||||||
|
ALTER TABLE data_process_tasks ADD CONSTRAINT ck_data_process_tasks_preview_progress
|
||||||
|
CHECK (preview_progress >= 0 AND preview_progress <= 100);
|
||||||
|
END IF;
|
||||||
|
IF NOT EXISTS (
|
||||||
|
SELECT 1 FROM pg_constraint
|
||||||
|
WHERE conrelid='data_process_tasks'::regclass
|
||||||
|
AND conname='ck_data_process_tasks_preview_file_counts'
|
||||||
|
) THEN
|
||||||
|
ALTER TABLE data_process_tasks ADD CONSTRAINT ck_data_process_tasks_preview_file_counts
|
||||||
|
CHECK (preview_total_files >= 0 AND preview_completed_files >= 0
|
||||||
|
AND preview_completed_files <= preview_total_files);
|
||||||
|
END IF;
|
||||||
|
END $$;
|
||||||
|
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_data_process_tasks_name_alive
|
||||||
|
ON data_process_tasks(name) WHERE deleted_at IS NULL;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_tasks_scope_status
|
||||||
|
ON data_process_tasks(tenant_id, project_id, status, created_at DESC)
|
||||||
|
WHERE deleted_at IS NULL;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_tasks_creator_created
|
||||||
|
ON data_process_tasks(created_by, created_at DESC) WHERE deleted_at IS NULL;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_source_files (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE,
|
||||||
|
storage_object_id TEXT,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0 CHECK (size_bytes >= 0),
|
||||||
|
record_count BIGINT NOT NULL DEFAULT 0 CHECK (record_count >= 0),
|
||||||
|
file_format VARCHAR(40),
|
||||||
|
checksum_sha256 CHAR(64) NOT NULL,
|
||||||
|
version_no INTEGER NOT NULL DEFAULT 1 CHECK (version_no > 0),
|
||||||
|
content TEXT NOT NULL,
|
||||||
|
content_preview TEXT,
|
||||||
|
metadata TEXT NOT NULL DEFAULT '{}',
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
created_by TEXT,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
deleted_at TIMESTAMPTZ
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_source_files_task
|
||||||
|
ON data_process_source_files(task_id, created_at) WHERE deleted_at IS NULL;
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_data_process_source_checksum_alive
|
||||||
|
ON data_process_source_files(task_id, checksum_sha256) WHERE deleted_at IS NULL;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_preview_items (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE,
|
||||||
|
source_file_id TEXT REFERENCES data_process_source_files(id) ON DELETE CASCADE,
|
||||||
|
original_content TEXT NOT NULL DEFAULT '',
|
||||||
|
edited_content TEXT NOT NULL DEFAULT '',
|
||||||
|
source_start INTEGER CHECK (source_start IS NULL OR source_start >= 0),
|
||||||
|
source_end INTEGER CHECK (source_end IS NULL OR source_end >= 0),
|
||||||
|
source_start_line INTEGER CHECK (source_start_line IS NULL OR source_start_line > 0),
|
||||||
|
source_end_line INTEGER CHECK (source_end_line IS NULL OR source_end_line > 0),
|
||||||
|
token_count INTEGER NOT NULL DEFAULT 0 CHECK (token_count >= 0),
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'original'
|
||||||
|
CHECK (status IN ('original', 'modified', 'manual', 'invalid')),
|
||||||
|
quality_score TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
CHECK (source_start IS NULL OR source_end IS NULL OR source_end >= source_start),
|
||||||
|
CHECK (source_start_line IS NULL OR source_end_line IS NULL OR source_end_line >= source_start_line)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_preview_task_file
|
||||||
|
ON data_process_preview_items(task_id, source_file_id, created_at);
|
||||||
|
|
||||||
|
-- 子表在新库中到这里才存在;只修复本次新增 workflow_step 前的历史任务。
|
||||||
|
UPDATE data_process_tasks task
|
||||||
|
SET workflow_step = CASE
|
||||||
|
WHEN EXISTS (
|
||||||
|
SELECT 1 FROM data_process_preview_items preview
|
||||||
|
WHERE preview.task_id=task.id
|
||||||
|
) THEN 'preview'
|
||||||
|
WHEN EXISTS (
|
||||||
|
SELECT 1 FROM data_process_source_files source_file
|
||||||
|
WHERE source_file.task_id=task.id AND source_file.deleted_at IS NULL
|
||||||
|
) THEN 'upload'
|
||||||
|
ELSE task.workflow_step
|
||||||
|
END
|
||||||
|
WHERE task.id IN (SELECT id FROM data_process_workflow_backfill_ids)
|
||||||
|
AND task.status='pending';
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS data_process_results (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
task_id TEXT NOT NULL REFERENCES data_process_tasks(id) ON DELETE CASCADE,
|
||||||
|
preview_item_id TEXT REFERENCES data_process_preview_items(id) ON DELETE SET NULL,
|
||||||
|
instruction TEXT NOT NULL,
|
||||||
|
input TEXT NOT NULL DEFAULT '',
|
||||||
|
output TEXT NOT NULL,
|
||||||
|
chosen TEXT NOT NULL DEFAULT '',
|
||||||
|
rejected TEXT NOT NULL DEFAULT '',
|
||||||
|
original_instruction TEXT,
|
||||||
|
original_input TEXT,
|
||||||
|
original_output TEXT,
|
||||||
|
original_chosen TEXT,
|
||||||
|
original_rejected TEXT,
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'valid'
|
||||||
|
CHECK (status IN ('valid', 'modified', 'invalid')),
|
||||||
|
error TEXT,
|
||||||
|
split VARCHAR(20) CHECK (split IS NULL OR split IN ('train', 'validation', 'test')),
|
||||||
|
quality_score TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||||
|
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS chosen TEXT NOT NULL DEFAULT '';
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS rejected TEXT NOT NULL DEFAULT '';
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS original_chosen TEXT;
|
||||||
|
ALTER TABLE data_process_results ADD COLUMN IF NOT EXISTS original_rejected TEXT;
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_results_task_status
|
||||||
|
ON data_process_results(task_id, status, id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_data_process_results_task_split
|
||||||
|
ON data_process_results(task_id, split);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS dataset_file_versions (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
dataset_file_id TEXT NOT NULL REFERENCES dataset_files(id) ON DELETE CASCADE,
|
||||||
|
version_no INTEGER NOT NULL CHECK (version_no > 0),
|
||||||
|
storage_object_id TEXT NOT NULL,
|
||||||
|
content_preview TEXT,
|
||||||
|
description TEXT,
|
||||||
|
base_version_id TEXT REFERENCES dataset_file_versions(id) ON DELETE SET NULL,
|
||||||
|
size_bytes BIGINT NOT NULL DEFAULT 0 CHECK (size_bytes >= 0),
|
||||||
|
record_count BIGINT NOT NULL DEFAULT 0 CHECK (record_count >= 0),
|
||||||
|
checksum_sha256 CHAR(64) NOT NULL,
|
||||||
|
source_task_id TEXT REFERENCES data_process_tasks(id) ON DELETE SET NULL,
|
||||||
|
metadata TEXT NOT NULL DEFAULT '{}',
|
||||||
|
created_by TEXT,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE dataset_file_versions ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE dataset_file_versions ADD COLUMN IF NOT EXISTS metadata TEXT NOT NULL DEFAULT '{}';
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_dataset_file_versions_no_002
|
||||||
|
ON dataset_file_versions(dataset_file_id, version_no);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_file_versions_source_task_002
|
||||||
|
ON dataset_file_versions(source_task_id) WHERE source_task_id IS NOT NULL;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS dataset_records (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
dataset_id TEXT NOT NULL REFERENCES datasets(id) ON DELETE CASCADE,
|
||||||
|
dataset_file_id TEXT REFERENCES dataset_files(id) ON DELETE CASCADE,
|
||||||
|
version_id TEXT REFERENCES dataset_file_versions(id) ON DELETE CASCADE,
|
||||||
|
line_no INTEGER,
|
||||||
|
split VARCHAR(20) CHECK (split IS NULL OR split IN ('train', 'validation', 'test')),
|
||||||
|
instruction TEXT,
|
||||||
|
input TEXT,
|
||||||
|
output TEXT,
|
||||||
|
raw TEXT NOT NULL DEFAULT '{}',
|
||||||
|
status VARCHAR(20) NOT NULL DEFAULT 'valid'
|
||||||
|
CHECK (status IN ('valid', 'modified', 'invalid')),
|
||||||
|
source_task_id TEXT REFERENCES data_process_tasks(id) ON DELETE SET NULL,
|
||||||
|
source_result_id TEXT REFERENCES data_process_results(id) ON DELETE SET NULL,
|
||||||
|
preview_item_id TEXT REFERENCES data_process_preview_items(id) ON DELETE SET NULL,
|
||||||
|
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS source_task_id TEXT;
|
||||||
|
ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS source_result_id TEXT;
|
||||||
|
ALTER TABLE dataset_records ADD COLUMN IF NOT EXISTS preview_item_id TEXT;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_records_dataset_002
|
||||||
|
ON dataset_records(dataset_id, id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_records_source_task_002
|
||||||
|
ON dataset_records(source_task_id, source_result_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_datasets_source_task_002
|
||||||
|
ON datasets(source_task_id) WHERE source_task_id IS NOT NULL;
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_dataset_files_source_task_002
|
||||||
|
ON dataset_files(source_task_id) WHERE source_task_id IS NOT NULL;
|
||||||
|
|
||||||
|
COMMIT;
|
||||||
111
backend/app/db/sql/002_governance.sql
Normal file
111
backend/app/db/sql/002_governance.sql
Normal file
@@ -0,0 +1,111 @@
|
|||||||
|
-- 平台治理:租户 / 审批 / 审计(字段以 platform_store 实际写入为准)
|
||||||
|
CREATE TABLE IF NOT EXISTS tenants (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
code TEXT,
|
||||||
|
status TEXT DEFAULT 'active',
|
||||||
|
owner_user_id TEXT,
|
||||||
|
quota TEXT,
|
||||||
|
retention_policy_id TEXT,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS approval_templates (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
steps TEXT,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS approval_instances (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
template_id TEXT,
|
||||||
|
resource_type TEXT,
|
||||||
|
resource_id TEXT,
|
||||||
|
applicant_id TEXT,
|
||||||
|
status TEXT DEFAULT 'pending',
|
||||||
|
current_step INTEGER DEFAULT 0,
|
||||||
|
create_time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS approval_steps (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
instance_id TEXT,
|
||||||
|
step_index INTEGER,
|
||||||
|
approver_id TEXT,
|
||||||
|
status TEXT DEFAULT 'pending',
|
||||||
|
comment TEXT,
|
||||||
|
time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS audit_logs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
tenant_id TEXT,
|
||||||
|
project_id TEXT,
|
||||||
|
actor_id TEXT,
|
||||||
|
action TEXT,
|
||||||
|
target_type TEXT,
|
||||||
|
target_id TEXT,
|
||||||
|
detail TEXT,
|
||||||
|
client_ip TEXT,
|
||||||
|
time TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
-- 幂等升级 audit_logs 表:新增字段(已存在则跳过)
|
||||||
|
ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS trace_id TEXT;
|
||||||
|
ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS request_method TEXT;
|
||||||
|
ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS request_path TEXT;
|
||||||
|
ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS status_code INTEGER;
|
||||||
|
ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS duration_ms REAL;
|
||||||
|
ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS extra JSONB;
|
||||||
|
|
||||||
|
-- 索引(已存在则跳过)
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_trace ON audit_logs(trace_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_actor_time ON audit_logs(actor_id, time);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_audit_target ON audit_logs(target_type, target_id);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS retention_policies (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
name TEXT NOT NULL,
|
||||||
|
scope TEXT,
|
||||||
|
rule TEXT,
|
||||||
|
status TEXT DEFAULT 'active',
|
||||||
|
create_time TEXT,
|
||||||
|
create_by TEXT,
|
||||||
|
updated_at TEXT
|
||||||
|
);
|
||||||
|
|
||||||
|
-- ===================== 操作日志表 =====================
|
||||||
|
-- 记录用户在每个业务模块的详细操作(成功/失败、报错信息等)
|
||||||
|
CREATE TABLE IF NOT EXISTS operation_logs (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
user_id TEXT, -- 操作用户 ID
|
||||||
|
username TEXT, -- 操作用户名(冗余,便于查询)
|
||||||
|
module TEXT, -- 业务模块:fine-tune/model-eval/model-inference/dataset/data-process/data-convert/model-manage
|
||||||
|
action TEXT, -- 具体动作:create/start/stop/delete/upload/convert 等
|
||||||
|
target_type TEXT, -- 资源类型
|
||||||
|
target_id TEXT, -- 资源 ID
|
||||||
|
target_name TEXT, -- 资源名称(便于阅读)
|
||||||
|
status TEXT NOT NULL, -- success / failure
|
||||||
|
error_message TEXT, -- 失败时的报错信息
|
||||||
|
detail TEXT, -- 操作详情 JSON(参数摘要)
|
||||||
|
client_ip TEXT, -- 客户端 IP
|
||||||
|
request_method TEXT, -- HTTP 方法
|
||||||
|
request_path TEXT, -- 请求路径
|
||||||
|
trace_id TEXT, -- 链路追踪 ID
|
||||||
|
duration_ms REAL, -- 耗时(ms)
|
||||||
|
create_time TEXT -- 操作时间
|
||||||
|
);
|
||||||
|
|
||||||
|
-- 幂等升级 operation_logs 表:新增字段(已存在则跳过)
|
||||||
|
ALTER TABLE operation_logs ADD COLUMN IF NOT EXISTS error_type TEXT; -- 异常类型:RuntimeError / ValueError / ConnectionError
|
||||||
|
ALTER TABLE operation_logs ADD COLUMN IF NOT EXISTS error_traceback TEXT; -- 完整异常堆栈
|
||||||
|
ALTER TABLE operation_logs ADD COLUMN IF NOT EXISTS func_name TEXT; -- 出错的函数名
|
||||||
|
|
||||||
|
-- 索引
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_op_user_time ON operation_logs(user_id, create_time DESC);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_op_module_time ON operation_logs(module, create_time DESC);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_op_status ON operation_logs(status, create_time DESC);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_op_action ON operation_logs(action);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_op_create_time ON operation_logs(create_time DESC);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_op_error_type ON operation_logs(error_type);
|
||||||
20
backend/app/db/sql/003_model_path_governance.sql
Normal file
20
backend/app/db/sql/003_model_path_governance.sql
Normal file
@@ -0,0 +1,20 @@
|
|||||||
|
-- 003_model_path_governance
|
||||||
|
-- 模型路径治理:增加 can_train 标识,区分本地可训练模型与 API / 远程模型。
|
||||||
|
-- 训练预检阶段依赖该字段拦截不适合 LLaMA-Factory 本地训练的基座模型。
|
||||||
|
|
||||||
|
-- 1. models 表增加 can_train(默认 0,后设搬迁为 1 的规则如下)
|
||||||
|
ALTER TABLE models ADD COLUMN IF NOT EXISTS can_train INTEGER NOT NULL DEFAULT 0;
|
||||||
|
|
||||||
|
-- 2. 将已有模型按规则推定 can_train:
|
||||||
|
-- - path 非空 且 model_source != 'api' → 可训练 (1)
|
||||||
|
-- - 其余 → 不可训练 (0)
|
||||||
|
UPDATE models
|
||||||
|
SET can_train = CASE
|
||||||
|
WHEN path IS NOT NULL AND path != '' AND model_source IS NOT NULL AND model_source != 'api' THEN 1
|
||||||
|
ELSE 0
|
||||||
|
END;
|
||||||
|
|
||||||
|
-- 3. 给 trained_models 增加 artifact_dir(训练产物目录扫描结果目录)
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS artifact_dir TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS compute_node_id TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS compute_node_name TEXT;
|
||||||
3
backend/app/db/sql/003_tenant_quota.sql
Normal file
3
backend/app/db/sql/003_tenant_quota.sql
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
-- 租户配额与保留策略扩展(如后续治理表需补列,可在此追加)
|
||||||
|
ALTER TABLE tenants ADD COLUMN IF NOT EXISTS gpu_quota TEXT;
|
||||||
|
ALTER TABLE tenants ADD COLUMN IF NOT EXISTS storage_quota TEXT;
|
||||||
22
backend/app/db/sql/004_permissions.sql
Normal file
22
backend/app/db/sql/004_permissions.sql
Normal file
@@ -0,0 +1,22 @@
|
|||||||
|
-- ============================================================
|
||||||
|
-- 权限体系扩展:GPU 分配表 + 资源所有权字段
|
||||||
|
-- ============================================================
|
||||||
|
|
||||||
|
-- GPU 分配表:管理员指定哪些用户可以使用哪些 GPU 卡
|
||||||
|
CREATE TABLE IF NOT EXISTS gpu_assignments (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
node_id TEXT NOT NULL REFERENCES compute_nodes(id) ON DELETE CASCADE,
|
||||||
|
gpu_index INTEGER NOT NULL,
|
||||||
|
user_id TEXT NOT NULL,
|
||||||
|
assigned_by TEXT,
|
||||||
|
assigned_at TEXT NOT NULL,
|
||||||
|
UNIQUE (node_id, gpu_index, user_id)
|
||||||
|
);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpu_assignments_user ON gpu_assignments(user_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gpu_assignments_gpu ON gpu_assignments(node_id, gpu_index);
|
||||||
|
|
||||||
|
-- 资源所有权字段:用户创建的数据集/模型/训练产物/评测任务
|
||||||
|
ALTER TABLE datasets ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE models ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE trained_models ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
|
ALTER TABLE eval_tasks ADD COLUMN IF NOT EXISTS created_by TEXT;
|
||||||
43
backend/app/main.py
Normal file
43
backend/app/main.py
Normal file
@@ -0,0 +1,43 @@
|
|||||||
|
import asyncio
|
||||||
|
from contextlib import suppress
|
||||||
|
|
||||||
|
from fastapi import FastAPI
|
||||||
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
|
|
||||||
|
from app.api.v1.router import api_router
|
||||||
|
from app.core.config import docs_kwargs, get_settings
|
||||||
|
from app.core.logging import configure_logging, setup_request_logging
|
||||||
|
from app.workers.compute_poller import run_compute_poller
|
||||||
|
|
||||||
|
|
||||||
|
def create_app() -> FastAPI:
|
||||||
|
settings = get_settings()
|
||||||
|
configure_logging(settings)
|
||||||
|
|
||||||
|
app = FastAPI(title=settings.app_name, **docs_kwargs(settings.enable_docs))
|
||||||
|
app.add_middleware(
|
||||||
|
CORSMiddleware,
|
||||||
|
allow_origins=settings.cors_allow_origins,
|
||||||
|
allow_credentials=True,
|
||||||
|
allow_methods=["*"],
|
||||||
|
allow_headers=["*"],
|
||||||
|
)
|
||||||
|
setup_request_logging(app)
|
||||||
|
app.include_router(api_router, prefix=settings.route_prefix)
|
||||||
|
|
||||||
|
@app.on_event("startup")
|
||||||
|
async def start_workers() -> None:
|
||||||
|
app.state.compute_poller_task = asyncio.create_task(run_compute_poller())
|
||||||
|
|
||||||
|
@app.on_event("shutdown")
|
||||||
|
async def stop_workers() -> None:
|
||||||
|
task = getattr(app.state, "compute_poller_task", None)
|
||||||
|
if task:
|
||||||
|
task.cancel()
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
|
||||||
|
return app
|
||||||
|
|
||||||
|
|
||||||
|
app = create_app()
|
||||||
15
backend/app/modules/README.md
Normal file
15
backend/app/modules/README.md
Normal file
@@ -0,0 +1,15 @@
|
|||||||
|
# Backend Module Convention
|
||||||
|
|
||||||
|
每个业务模块建议保持一致结构:
|
||||||
|
|
||||||
|
```text
|
||||||
|
module_name/
|
||||||
|
__init__.py
|
||||||
|
router.py # FastAPI router
|
||||||
|
schemas.py # Pydantic request/response models
|
||||||
|
service.py # Business orchestration
|
||||||
|
repository.py # Database access
|
||||||
|
permissions.py # Optional resource permission checks
|
||||||
|
```
|
||||||
|
|
||||||
|
模块边界以 `docs/system-development-plan.md` 的页面模块开发工作包为准。
|
||||||
1
backend/app/modules/approval/__init__.py
Normal file
1
backend/app/modules/approval/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Approval workflow module."""
|
||||||
100
backend/app/modules/approval/router.py
Normal file
100
backend/app/modules/approval/router.py
Normal file
@@ -0,0 +1,100 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Depends
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
from app.core.auth import get_current_user, is_admin
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/approvals", tags=["approval"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/templates")
|
||||||
|
def list_templates(current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
return ok(get_platform_store().approval_templates())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/templates")
|
||||||
|
def create_template(payload: dict[str, Any] = Body(...), current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
if not is_admin(current_user): raise fail(403, "admin permission required")
|
||||||
|
if not payload.get("name"):
|
||||||
|
raise fail(400, "name 必填")
|
||||||
|
return ok(get_platform_store().create_approval_template(payload))
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/templates/{template_id}")
|
||||||
|
def get_template(template_id: str) -> dict[str, Any]:
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().approval_template(template_id))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "template not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/templates/{template_id}")
|
||||||
|
def update_template(template_id: str, payload: dict[str, Any] = Body(...), current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
if not is_admin(current_user): raise fail(403, "admin permission required")
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().update_approval_template(template_id, payload))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "template not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/templates/{template_id}")
|
||||||
|
def delete_template(template_id: str, current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
if not is_admin(current_user): raise fail(403, "admin permission required")
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().delete_approval_template(template_id))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "template not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("")
|
||||||
|
def list_instances(status: str | None = None, current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
items = get_platform_store().approval_instances(status=status)
|
||||||
|
return ok(items if is_admin(current_user) else [item for item in items if item.get("applicant_id") == current_user.get("id")])
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("")
|
||||||
|
def create_instance(payload: dict[str, Any] = Body(...), current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
payload["applicant_id"] = current_user.get("id")
|
||||||
|
for field in ("resource_type", "resource_id"):
|
||||||
|
if not payload.get(field):
|
||||||
|
raise fail(400, f"{field} 必填")
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().create_approval_instance(payload))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "template not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{instance_id}")
|
||||||
|
def get_instance(instance_id: str, current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
try:
|
||||||
|
item = get_platform_store().approval_instance(instance_id)
|
||||||
|
if not is_admin(current_user) and item.get("applicant_id") != current_user.get("id"):
|
||||||
|
raise fail(403, "no permission to access approval")
|
||||||
|
return ok(item)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "instance not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{instance_id}/steps/{step_index}/decision")
|
||||||
|
def decide(
|
||||||
|
instance_id: str,
|
||||||
|
step_index: int,
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if not payload.get("approver_id"):
|
||||||
|
raise fail(400, "approver_id 必填")
|
||||||
|
try:
|
||||||
|
return ok(
|
||||||
|
get_platform_store().decide_approval_step(
|
||||||
|
instance_id,
|
||||||
|
step_index,
|
||||||
|
approver_id=payload["approver_id"],
|
||||||
|
approved=bool(payload.get("approved", False)),
|
||||||
|
comment=payload.get("comment"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
except (KeyError, ValueError) as e:
|
||||||
|
raise fail(400, str(e))
|
||||||
1
backend/app/modules/audit/__init__.py
Normal file
1
backend/app/modules/audit/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Audit log module."""
|
||||||
1
backend/app/modules/auth/__init__.py
Normal file
1
backend/app/modules/auth/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Authentication and user session module."""
|
||||||
1
backend/app/modules/compute_gateway/__init__.py
Normal file
1
backend/app/modules/compute_gateway/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Application-side compute platform gateway module."""
|
||||||
273
backend/app/modules/compute_gateway/client.py
Normal file
273
backend/app/modules/compute_gateway/client.py
Normal file
@@ -0,0 +1,273 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import urljoin
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
|
||||||
|
|
||||||
|
def _join_url(base_url: str, path: str) -> str:
|
||||||
|
return urljoin(base_url.rstrip("/") + "/", path.lstrip("/"))
|
||||||
|
|
||||||
|
|
||||||
|
def _unwrap_items(payload: Any) -> list[dict[str, Any]]:
|
||||||
|
if isinstance(payload, list):
|
||||||
|
return [item for item in payload if isinstance(item, dict)]
|
||||||
|
if isinstance(payload, dict):
|
||||||
|
data = payload.get("data")
|
||||||
|
if isinstance(data, dict) and isinstance(data.get("items"), list):
|
||||||
|
return [item for item in data["items"] if isinstance(item, dict)]
|
||||||
|
if isinstance(payload.get("items"), list):
|
||||||
|
return [item for item in payload["items"] if isinstance(item, dict)]
|
||||||
|
if isinstance(data, list):
|
||||||
|
return [item for item in data if isinstance(item, dict)]
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def _unwrap_dict(payload: Any) -> dict[str, Any]:
|
||||||
|
if isinstance(payload, dict) and isinstance(payload.get("data"), dict):
|
||||||
|
return payload["data"]
|
||||||
|
return payload if isinstance(payload, dict) else {}
|
||||||
|
|
||||||
|
|
||||||
|
# Inference calls are intentionally short-timeout:
|
||||||
|
# - load dispatch only confirms the compute node accepted the request
|
||||||
|
# (the actual model load now runs asynchronously on the node).
|
||||||
|
# - status/unload must never block the platform for long when a node is
|
||||||
|
# unreachable but still marked online.
|
||||||
|
INFERENCE_LOAD_TIMEOUT = httpx.Timeout(30, connect=10)
|
||||||
|
INFERENCE_STATUS_TIMEOUT = httpx.Timeout(30, connect=5)
|
||||||
|
INFERENCE_UNLOAD_TIMEOUT = httpx.Timeout(30, connect=5)
|
||||||
|
|
||||||
|
|
||||||
|
class ComputeNodeClient:
|
||||||
|
"""Application-side client for one compute node.
|
||||||
|
|
||||||
|
The client accepts both current YG Compute API responses and common
|
||||||
|
wrapper shapes such as `{code,message,data}` to make future engine/node
|
||||||
|
adapters less brittle.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, api_base_url: str, token: str | None = None, timeout: float | None = None) -> None:
|
||||||
|
settings = get_settings()
|
||||||
|
self.api_base_url = api_base_url.rstrip("/")
|
||||||
|
self.token = token or settings.compute_service_token
|
||||||
|
self.timeout = timeout or settings.compute_request_timeout_seconds
|
||||||
|
self.route_prefix = settings.route_prefix.rstrip("/") or "/modelTF"
|
||||||
|
|
||||||
|
def headers(self) -> dict[str, str]:
|
||||||
|
if not self.token:
|
||||||
|
return {}
|
||||||
|
return {"X-Compute-Token": self.token}
|
||||||
|
|
||||||
|
async def test_connection(self) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
health = await self.health()
|
||||||
|
gpus = await self.gpus()
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"latency_ms": int((time.perf_counter() - started) * 1000),
|
||||||
|
"health": health,
|
||||||
|
"gpus": gpus,
|
||||||
|
}
|
||||||
|
|
||||||
|
async def health(self) -> dict[str, Any]:
|
||||||
|
paths = [f"{self.route_prefix}/v1/compute/health", f"{self.route_prefix}/health", "/health"]
|
||||||
|
last_error = ""
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
for path in paths:
|
||||||
|
try:
|
||||||
|
response = await client.get(_join_url(self.api_base_url, path))
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
except Exception as exc: # noqa: BLE001 - keep endpoint compatibility fallback broad
|
||||||
|
last_error = str(exc)
|
||||||
|
raise RuntimeError(last_error or "compute health check failed")
|
||||||
|
|
||||||
|
async def gpus(self) -> list[dict[str, Any]]:
|
||||||
|
paths = [
|
||||||
|
f"{self.route_prefix}/compute/resources/gpus",
|
||||||
|
f"{self.route_prefix}/v1/compute/resources/gpus",
|
||||||
|
"/compute/resources/gpus",
|
||||||
|
]
|
||||||
|
last_error = ""
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
for path in paths:
|
||||||
|
try:
|
||||||
|
response = await client.get(_join_url(self.api_base_url, path))
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_items(response.json())
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
last_error = str(exc)
|
||||||
|
raise RuntimeError(last_error or "compute gpu discovery failed")
|
||||||
|
|
||||||
|
async def create_job(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(_join_url(self.api_base_url, f"{self.route_prefix}/compute/jobs"), json=payload)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def preview_job(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/jobs/preview"),
|
||||||
|
json=payload,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def validate_job(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/jobs/validate"),
|
||||||
|
json=payload,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def check_paths(self, paths: list[dict[str, Any]]) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/files/check-paths"),
|
||||||
|
json={"paths": paths},
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def list_files(
|
||||||
|
self,
|
||||||
|
root: str = "data",
|
||||||
|
relative_path: str = "",
|
||||||
|
directories_only: bool = False,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.get(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/files/list"),
|
||||||
|
params={"root": root, "relative_path": relative_path, "directories_only": directories_only},
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def get_job(self, job_id: str) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.get(_join_url(self.api_base_url, f"{self.route_prefix}/compute/jobs/{job_id}"))
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def stop_job(self, job_id: str) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(_join_url(self.api_base_url, f"{self.route_prefix}/compute/jobs/{job_id}/stop"))
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def job_logs(
|
||||||
|
self,
|
||||||
|
job_id: str,
|
||||||
|
tail_lines: int | None = None,
|
||||||
|
offset: int | None = None,
|
||||||
|
limit: int | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
params = {
|
||||||
|
key: value
|
||||||
|
for key, value in {"tail_lines": tail_lines, "offset": offset, "limit": limit}.items()
|
||||||
|
if value is not None
|
||||||
|
}
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.get(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/jobs/{job_id}/logs"),
|
||||||
|
params=params,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def import_local_file(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/files/import-local"),
|
||||||
|
json=payload,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def prepare_cache(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
async with httpx.AsyncClient(timeout=httpx.Timeout(900, connect=30), headers=self.headers()) as client:
|
||||||
|
response = await client.post(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/cache/prepare"),
|
||||||
|
json=payload,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def cache_status(self, resource_id: str, version_id: str | None = None) -> dict[str, Any]:
|
||||||
|
params = {"resource_id": resource_id}
|
||||||
|
if version_id:
|
||||||
|
params["version_id"] = version_id
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.get(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/cache/status"),
|
||||||
|
params=params,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def _request(
|
||||||
|
self,
|
||||||
|
method: str,
|
||||||
|
path: str,
|
||||||
|
json_data: dict[str, Any] | None = None,
|
||||||
|
timeout: float | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Generic request method for compute API endpoints."""
|
||||||
|
url = _join_url(self.api_base_url, f"{self.route_prefix}{path}")
|
||||||
|
async with httpx.AsyncClient(timeout=timeout or 300, headers=self.headers()) as client:
|
||||||
|
if method.upper() == "GET":
|
||||||
|
response = await client.get(url)
|
||||||
|
else:
|
||||||
|
response = await client.post(url, json=json_data)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
# ── Inference helpers (short timeouts — see module constants) ──────────
|
||||||
|
|
||||||
|
async def inference_load(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
"""Dispatch a model load. Returns as soon as the node accepts the
|
||||||
|
request; the node now loads asynchronously (status goes 'loading')."""
|
||||||
|
return await self._request("POST", "/inference/load", json_data=payload, timeout=INFERENCE_LOAD_TIMEOUT)
|
||||||
|
|
||||||
|
async def inference_status(self) -> dict[str, Any]:
|
||||||
|
return await self._request("GET", "/inference/status", timeout=INFERENCE_STATUS_TIMEOUT)
|
||||||
|
|
||||||
|
async def inference_unload(self) -> dict[str, Any]:
|
||||||
|
return await self._request("POST", "/inference/unload", json_data={}, timeout=INFERENCE_UNLOAD_TIMEOUT)
|
||||||
|
|
||||||
|
async def upload_file(
|
||||||
|
self,
|
||||||
|
filename: str,
|
||||||
|
content: bytes,
|
||||||
|
target_relative_path: str,
|
||||||
|
resource_type: str | None = None,
|
||||||
|
resource_id: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
data = {
|
||||||
|
"target_relative_path": target_relative_path,
|
||||||
|
"resource_type": resource_type or "",
|
||||||
|
"resource_id": resource_id or "",
|
||||||
|
}
|
||||||
|
files = {"file": (filename, content)}
|
||||||
|
timeout = httpx.Timeout(max(self.timeout, 60), connect=self.timeout)
|
||||||
|
async with httpx.AsyncClient(timeout=timeout, headers=self.headers()) as client:
|
||||||
|
response = await client.post(
|
||||||
|
_join_url(self.api_base_url, f"{self.route_prefix}/compute/files/upload"),
|
||||||
|
data=data,
|
||||||
|
files=files,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return _unwrap_dict(response.json())
|
||||||
|
|
||||||
|
async def upload_file_to_url(self, path: str, upload_url: str, object_key: str = "", content_type: str = "application/octet-stream") -> dict[str, Any]:
|
||||||
|
return await self._request("POST", "/compute/files/upload-to-url", json_data={
|
||||||
|
"path": path, "upload_url": upload_url, "object_key": object_key, "content_type": content_type,
|
||||||
|
}, timeout=900)
|
||||||
191
backend/app/modules/compute_gateway/sync.py
Normal file
191
backend/app/modules/compute_gateway/sync.py
Normal file
@@ -0,0 +1,191 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
from app.modules.compute_gateway.client import ComputeNodeClient
|
||||||
|
|
||||||
|
# starting 状态允许的最大轮询次数(约 40 * 3s ≈ 2 分钟),超过即判定节点不可达
|
||||||
|
MAX_STARTING_ATTEMPTS = 40
|
||||||
|
|
||||||
|
|
||||||
|
def _node_for_task(task: dict[str, Any]) -> dict[str, Any] | None:
|
||||||
|
return next((node for node in get_platform_store().compute_nodes() if node["id"] == task.get("compute_node_id")), None)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_inference_load_status(task: dict[str, Any]) -> tuple[list[dict[str, Any]], dict[str, Any]]:
|
||||||
|
load_status = task.get("load_status") or {}
|
||||||
|
if isinstance(load_status, str):
|
||||||
|
try:
|
||||||
|
load_status = json.loads(load_status)
|
||||||
|
except (json.JSONDecodeError, TypeError):
|
||||||
|
load_status = {}
|
||||||
|
return load_status.get("loaded_models") or [], load_status
|
||||||
|
|
||||||
|
|
||||||
|
async def reconcile_inference_loads(store: Any) -> list[dict[str, Any]]:
|
||||||
|
"""推进处于 starting 状态的推理加载。
|
||||||
|
|
||||||
|
模型加载已改为异步派发:/model-compare/{id}/load 立即返回,这里在每次
|
||||||
|
轮询时查询对应计算节点的 /inference/status,把任务从 starting 推进到
|
||||||
|
ready/error。使用短超时,单节点不可达不会阻塞整轮轮询。
|
||||||
|
"""
|
||||||
|
reconciled: list[dict[str, Any]] = []
|
||||||
|
now = time.time()
|
||||||
|
for task in store.compare_tasks():
|
||||||
|
items, _ = _parse_inference_load_status(task)
|
||||||
|
if not any(item.get("status") == "starting" for item in items):
|
||||||
|
continue
|
||||||
|
# dirty 只要处理过任一 starting 项就置位:load_attempts / last_polled_at
|
||||||
|
# 必须落库,否则节点不可达时计数不会累积,封顶逻辑永远触发不了
|
||||||
|
dirty = False
|
||||||
|
for item in items:
|
||||||
|
if item.get("status") != "starting":
|
||||||
|
continue
|
||||||
|
# 节流:同一 item 每 3s 只查询一次
|
||||||
|
if now - float(item.get("last_polled_at") or 0) < 3:
|
||||||
|
continue
|
||||||
|
item["last_polled_at"] = now
|
||||||
|
item["load_attempts"] = int(item.get("load_attempts") or 0) + 1
|
||||||
|
dirty = True
|
||||||
|
node = next((n for n in store.compute_nodes() if n["id"] == item.get("node_id")), None)
|
||||||
|
if not node:
|
||||||
|
item["status"] = "error"
|
||||||
|
item["error"] = "compute node deleted"
|
||||||
|
store.mark_inference_unloaded(item.get("node_id") or "")
|
||||||
|
continue
|
||||||
|
if not node.get("enabled") or node.get("scheduler_status") != "online":
|
||||||
|
item["status"] = "error"
|
||||||
|
item["error"] = "compute node offline"
|
||||||
|
store.mark_inference_unloaded(node["id"])
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
status = await ComputeNodeClient(node["api_base_url"]).inference_status()
|
||||||
|
except Exception as exc: # noqa: BLE001 - node unreachable; keep retrying until cap
|
||||||
|
if int(item.get("load_attempts") or 0) >= MAX_STARTING_ATTEMPTS:
|
||||||
|
item["status"] = "error"
|
||||||
|
item["error"] = f"compute node unreachable: {exc}"
|
||||||
|
store.mark_inference_unloaded(node["id"])
|
||||||
|
continue
|
||||||
|
node_status = status.get("status")
|
||||||
|
if node_status == "ready":
|
||||||
|
item["status"] = "ready"
|
||||||
|
item.pop("error", None)
|
||||||
|
store.mark_inference_loaded(node["id"])
|
||||||
|
elif node_status == "error":
|
||||||
|
item["status"] = "error"
|
||||||
|
item["error"] = status.get("error") or "model load failed on compute node"
|
||||||
|
store.mark_inference_unloaded(node["id"])
|
||||||
|
elif node_status == "idle":
|
||||||
|
# 节点重启导致已加载模型丢失
|
||||||
|
item["status"] = "error"
|
||||||
|
item["error"] = "model disappeared from compute node (node may have restarted)"
|
||||||
|
store.mark_inference_unloaded(node["id"])
|
||||||
|
# node_status == "loading" -> 保持 starting,下轮再查
|
||||||
|
if dirty:
|
||||||
|
if any(i.get("status") in {"ready", "running"} for i in items):
|
||||||
|
new_status = "loaded"
|
||||||
|
elif any(i.get("status") == "starting" for i in items):
|
||||||
|
new_status = "starting" # 仍在加载中,保持 starting
|
||||||
|
else:
|
||||||
|
new_status = "failed"
|
||||||
|
store.update_compare_task(task["id"], {"status": new_status, "load_status": {"loaded_models": items}})
|
||||||
|
reconciled.append({"task_id": task["id"], "status": new_status})
|
||||||
|
return reconciled
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_eval_result_content(client: ComputeNodeClient, node: dict[str, Any], job: dict[str, Any]) -> dict[str, Any] | None:
|
||||||
|
output_dir = job.get("output_dir")
|
||||||
|
if not output_dir:
|
||||||
|
return None
|
||||||
|
full_path = f"{str(output_dir).rstrip('/')}/eval_results.json"
|
||||||
|
data_root = "/data/yg-ft/"
|
||||||
|
if full_path.startswith(data_root):
|
||||||
|
full_path = full_path[len(data_root):]
|
||||||
|
rel_path = full_path.lstrip("/")
|
||||||
|
import httpx
|
||||||
|
url = f"{node['api_base_url'].rstrip('/')}/modelTF/compute/files/read"
|
||||||
|
async with httpx.AsyncClient(timeout=30, headers=client.headers()) as http:
|
||||||
|
response = await http.get(url, params={"path": rel_path})
|
||||||
|
response.raise_for_status()
|
||||||
|
payload = response.json()
|
||||||
|
return payload if isinstance(payload, dict) else None
|
||||||
|
|
||||||
|
|
||||||
|
async def poll_compute_jobs_once() -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
synced: list[dict[str, Any]] = []
|
||||||
|
failed: list[dict[str, str]] = []
|
||||||
|
for task in store.running_compute_tasks():
|
||||||
|
node = _node_for_task(task)
|
||||||
|
if not node:
|
||||||
|
failed.append({"task_id": task["id"], "error": "compute node not found"})
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
client = ComputeNodeClient(node["api_base_url"])
|
||||||
|
job = await client.get_job(task["compute_job_id"])
|
||||||
|
try:
|
||||||
|
logs = await client.job_logs(task["compute_job_id"], tail_lines=5000)
|
||||||
|
store.record_training_log_metrics(task["id"], str(logs.get("content") or ""))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
# P0-4: Force-fetch last log snippet when job reaches terminal state
|
||||||
|
if job.get("status") in {"failed", "stopped"}:
|
||||||
|
try:
|
||||||
|
last_logs = await client.job_logs(task["compute_job_id"], tail_lines=200)
|
||||||
|
job["log_snippet"] = str(last_logs.get("content") or "")[:8192]
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
synced.append(store.apply_compute_job(task["id"], job))
|
||||||
|
except Exception as exc: # noqa: BLE001 - keep polling other jobs
|
||||||
|
failed.append({"task_id": task["id"], "error": str(exc)})
|
||||||
|
standalone_synced: list[dict[str, Any]] = []
|
||||||
|
for record in store.active_standalone_compute_jobs():
|
||||||
|
node = next((item for item in store.compute_nodes() if item["id"] == record.get("node_id")), None)
|
||||||
|
if not node:
|
||||||
|
failed.append({"job_id": record["id"], "error": "compute node not found"})
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
job = await ComputeNodeClient(node["api_base_url"]).get_job(record["id"])
|
||||||
|
standalone_synced.append(store.sync_model_merge_job(record["id"], job))
|
||||||
|
except Exception as exc: # noqa: BLE001 - keep polling other jobs
|
||||||
|
failed.append({"job_id": record["id"], "error": str(exc)})
|
||||||
|
|
||||||
|
# ── Eval job sync ────────────────────────────────────────────────
|
||||||
|
eval_synced = 0
|
||||||
|
for eval_task in store.running_eval_tasks():
|
||||||
|
node = next(
|
||||||
|
(item for item in store.compute_nodes() if item["id"] == eval_task.get("compute_node_id")),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not node:
|
||||||
|
failed.append({"eval_task_id": eval_task["id"], "error": "compute node not found"})
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
client = ComputeNodeClient(node["api_base_url"])
|
||||||
|
job = await client.get_job(eval_task["compute_job_id"])
|
||||||
|
result_content = None
|
||||||
|
# Try to read eval_results.json from the job output directory
|
||||||
|
if job.get("status") == "completed" and job.get("output_dir"):
|
||||||
|
try:
|
||||||
|
result_content = await fetch_eval_result_content(client, node, job)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
store.apply_eval_job_result(eval_task["id"], job, result_content)
|
||||||
|
# 评测 GPU 占用由 eval_tasks 状态派生,无需维护推理内存标记
|
||||||
|
eval_synced += 1
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
failed.append({"eval_task_id": eval_task["id"], "error": str(exc)})
|
||||||
|
|
||||||
|
# ── Inference load reconciliation ─────────────────────────────────────
|
||||||
|
try:
|
||||||
|
inference_reconciled = await reconcile_inference_loads(store)
|
||||||
|
except Exception as exc: # noqa: BLE001 - keep polling alive
|
||||||
|
failed.append({"inference_reconcile": str(exc)})
|
||||||
|
inference_reconciled = []
|
||||||
|
|
||||||
|
return {"synced": len(synced) + len(standalone_synced) + eval_synced, "failed": failed,
|
||||||
|
"items": synced, "standalone": standalone_synced, "eval_synced": eval_synced,
|
||||||
|
"inference_reconciled": inference_reconciled}
|
||||||
3
backend/app/modules/data_convert/__init__.py
Normal file
3
backend/app/modules/data_convert/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
from .router import router
|
||||||
|
|
||||||
|
__all__ = ["router"]
|
||||||
378
backend/app/modules/data_convert/router.py
Normal file
378
backend/app/modules/data_convert/router.py
Normal file
@@ -0,0 +1,378 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Depends, File, UploadFile
|
||||||
|
from fastapi.responses import FileResponse
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.core.auth import get_current_user, is_admin
|
||||||
|
from app.core.op_log import op_log, OpModule, OpAction
|
||||||
|
from app.db.platform_store import get_platform_store, new_id
|
||||||
|
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/data-convert", tags=["data-convert"])
|
||||||
|
|
||||||
|
# 存储根目录
|
||||||
|
STORAGE_ROOT = Path(__file__).resolve().parents[3] / "storage" / "data-convert"
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_output_filename(value: Any) -> str:
|
||||||
|
"""输出文件名白名单校验:仅允许普通文件名,阻断 ``../``、``/``、``\\`` 等路径穿越。
|
||||||
|
|
||||||
|
转换结果始终写入 ``STORAGE_ROOT/<task_id>/output/<output_filename>``,
|
||||||
|
若文件名可被注入路径分隔符,将导致任意文件读写/删除。
|
||||||
|
"""
|
||||||
|
name = str(value or "converted-data.jsonl").strip()
|
||||||
|
if (
|
||||||
|
not name
|
||||||
|
or name in {".", ".."}
|
||||||
|
or name != Path(name).name
|
||||||
|
or "/" in name
|
||||||
|
or "\\" in name
|
||||||
|
or any(ord(character) < 32 or ord(character) == 127 for character in name)
|
||||||
|
):
|
||||||
|
raise fail(400, "output filename must be a plain file name")
|
||||||
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
def _task_output_path(task: dict[str, Any]) -> Path:
|
||||||
|
"""返回经过白名单校验的转换输出文件路径(始终位于任务 output 目录内)。"""
|
||||||
|
return _output_dir(task["id"]) / _safe_output_filename(task.get("output_filename"))
|
||||||
|
|
||||||
|
|
||||||
|
def _task_dir(task_id: str) -> Path:
|
||||||
|
return STORAGE_ROOT / task_id
|
||||||
|
|
||||||
|
|
||||||
|
def _input_dir(task_id: str) -> Path:
|
||||||
|
return _task_dir(task_id) / "input"
|
||||||
|
|
||||||
|
|
||||||
|
def _output_dir(task_id: str) -> Path:
|
||||||
|
return _task_dir(task_id) / "output"
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("")
|
||||||
|
def list_tasks(
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 20,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
if is_admin(current_user):
|
||||||
|
# 管理员可见全部
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT * FROM data_convert_tasks WHERE deleted_at IS NULL "
|
||||||
|
"ORDER BY create_time DESC LIMIT %s OFFSET %s",
|
||||||
|
(page_size, (page - 1) * page_size),
|
||||||
|
).fetchall()
|
||||||
|
total = conn.execute(
|
||||||
|
"SELECT COUNT(*) FROM data_convert_tasks WHERE deleted_at IS NULL"
|
||||||
|
).fetchone()[0]
|
||||||
|
else:
|
||||||
|
# 普通用户只能看到自己创建的
|
||||||
|
user_id = current_user.get("id")
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT * FROM data_convert_tasks WHERE deleted_at IS NULL AND created_by=%s "
|
||||||
|
"ORDER BY create_time DESC LIMIT %s OFFSET %s",
|
||||||
|
(user_id, page_size, (page - 1) * page_size),
|
||||||
|
).fetchall()
|
||||||
|
total = conn.execute(
|
||||||
|
"SELECT COUNT(*) FROM data_convert_tasks WHERE deleted_at IS NULL AND created_by=%s",
|
||||||
|
(user_id,)
|
||||||
|
).fetchone()[0]
|
||||||
|
return ok({"items": [dict(r) for r in rows], "total": total})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("")
|
||||||
|
@op_log(module=OpModule.DATA_CONVERT, action=OpAction.CREATE, target_type="convert_task", target_name_param="name")
|
||||||
|
def create_task(
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
name = str(payload.get("name") or "").strip()
|
||||||
|
if not name:
|
||||||
|
raise fail(400, "name is required")
|
||||||
|
task_id = new_id("dct")
|
||||||
|
output_filename = _safe_output_filename(payload.get("output_filename"))
|
||||||
|
description = str(payload.get("description") or "").strip()
|
||||||
|
user_id = current_user.get("id")
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"INSERT INTO data_convert_tasks (id, name, description, output_filename, created_by) "
|
||||||
|
"VALUES (%s, %s, %s, %s, %s)",
|
||||||
|
(task_id, name, description, output_filename, user_id),
|
||||||
|
)
|
||||||
|
# 创建目录
|
||||||
|
_input_dir(task_id).mkdir(parents=True, exist_ok=True)
|
||||||
|
_output_dir(task_id).mkdir(parents=True, exist_ok=True)
|
||||||
|
return ok(_get_task(task_id))
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{task_id}")
|
||||||
|
def get_task(
|
||||||
|
task_id: str,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
task = _get_task(task_id)
|
||||||
|
if not task:
|
||||||
|
raise fail(404, "task not found")
|
||||||
|
# 附加输入文件列表
|
||||||
|
input_dir = _input_dir(task_id)
|
||||||
|
files = []
|
||||||
|
if input_dir.exists():
|
||||||
|
for f in sorted(input_dir.iterdir()):
|
||||||
|
if f.is_file():
|
||||||
|
files.append({"name": f.name, "size": f.stat().st_size})
|
||||||
|
task["input_files"] = files
|
||||||
|
return ok(task)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{task_id}/source-files")
|
||||||
|
@op_log(module=OpModule.DATA_CONVERT, action=OpAction.UPLOAD, target_type="convert_task", target_name_param="task_id")
|
||||||
|
async def upload_source_files(
|
||||||
|
task_id: str,
|
||||||
|
files: list[UploadFile] = File(...),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
task = _get_task(task_id)
|
||||||
|
if not task:
|
||||||
|
raise fail(404, "task not found")
|
||||||
|
if task["status"] not in ("pending", "uploaded"):
|
||||||
|
raise fail(400, "task is not editable")
|
||||||
|
input_dir = _input_dir(task_id)
|
||||||
|
input_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
staged = []
|
||||||
|
for upload in files:
|
||||||
|
name = Path(upload.filename or "input.json").name
|
||||||
|
if not name.lower().endswith(".json"):
|
||||||
|
raise fail(415, f"only JSON files are supported: {name}")
|
||||||
|
target = input_dir / name
|
||||||
|
content = await upload.read()
|
||||||
|
target.write_bytes(content)
|
||||||
|
staged.append({"name": name, "size": len(content)})
|
||||||
|
store = get_platform_store()
|
||||||
|
# 标记上传完成
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET status='uploaded', update_time=NOW() WHERE id=%s",
|
||||||
|
(task_id,),
|
||||||
|
)
|
||||||
|
# 自动转换并导入数据集
|
||||||
|
try:
|
||||||
|
output_dir = _output_dir(task_id)
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
output_path = _task_output_path(task)
|
||||||
|
# 清空旧输出(如果重新上传)
|
||||||
|
if output_path.exists():
|
||||||
|
output_path.unlink()
|
||||||
|
input_count = 0
|
||||||
|
output_count = 0
|
||||||
|
for json_file in sorted(input_dir.iterdir()):
|
||||||
|
if not json_file.is_file() or not json_file.name.lower().endswith(".json"):
|
||||||
|
continue
|
||||||
|
input_count += 1
|
||||||
|
with open(json_file, "r", encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
if isinstance(data, list):
|
||||||
|
records = data
|
||||||
|
elif isinstance(data, dict):
|
||||||
|
records = [data]
|
||||||
|
else:
|
||||||
|
raise ValueError(f"JSON must be object or array: {json_file.name}")
|
||||||
|
with open(output_path, "a", encoding="utf-8") as f:
|
||||||
|
for record in records:
|
||||||
|
f.write(json.dumps(record, ensure_ascii=False) + "\n")
|
||||||
|
output_count += 1
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET status='completed', "
|
||||||
|
"input_count=%s, output_count=%s, update_time=NOW() WHERE id=%s",
|
||||||
|
(input_count, output_count, task_id),
|
||||||
|
)
|
||||||
|
# 自动导入数据集
|
||||||
|
content = output_path.read_text(encoding="utf-8")
|
||||||
|
size_bytes = len(content.encode("utf-8"))
|
||||||
|
dataset = store.create_dataset({
|
||||||
|
"name": task["name"],
|
||||||
|
"type": "train",
|
||||||
|
"storage_type": "local",
|
||||||
|
"source": "upload",
|
||||||
|
"task_id": task_id,
|
||||||
|
"size": f"{size_bytes} B",
|
||||||
|
"count": output_count,
|
||||||
|
"description": f"由数据类型转换任务 {task_id} 自动导入",
|
||||||
|
"created_by": task.get("created_by") or current_user.get("id"),
|
||||||
|
})
|
||||||
|
dataset_id = dataset["id"]
|
||||||
|
with store.connect() as conn:
|
||||||
|
store.add_dataset_file(conn, dataset_id, _safe_output_filename(task.get("output_filename")), content)
|
||||||
|
return ok({
|
||||||
|
"staged_files": staged,
|
||||||
|
"auto_converted": True,
|
||||||
|
"dataset_id": dataset_id,
|
||||||
|
"input_count": input_count,
|
||||||
|
"output_count": output_count,
|
||||||
|
})
|
||||||
|
except Exception as exc:
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET status='failed', error_message=%s, update_time=NOW() WHERE id=%s",
|
||||||
|
(str(exc)[:500], task_id),
|
||||||
|
)
|
||||||
|
return ok({"staged_files": staged, "auto_converted": False, "error": str(exc)[:500]})
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{task_id}/run")
|
||||||
|
@op_log(module=OpModule.DATA_CONVERT, action=OpAction.CONVERT, target_type="convert_task", target_name_param="task_id")
|
||||||
|
def run_convert(
|
||||||
|
task_id: str,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
task = _get_task(task_id)
|
||||||
|
if not task:
|
||||||
|
raise fail(404, "task not found")
|
||||||
|
if task["status"] not in ("uploaded", "completed", "failed"):
|
||||||
|
raise fail(400, "please upload source files first")
|
||||||
|
# 标记运行中
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET status='running', error_message='', update_time=NOW() WHERE id=%s",
|
||||||
|
(task_id,),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
input_dir = _input_dir(task_id)
|
||||||
|
output_dir = _output_dir(task_id)
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
output_path = _task_output_path(task)
|
||||||
|
input_count = 0
|
||||||
|
output_count = 0
|
||||||
|
for json_file in sorted(input_dir.iterdir()):
|
||||||
|
if not json_file.is_file() or not json_file.name.lower().endswith(".json"):
|
||||||
|
continue
|
||||||
|
input_count += 1
|
||||||
|
with open(json_file, "r", encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
if isinstance(data, list):
|
||||||
|
records = data
|
||||||
|
elif isinstance(data, dict):
|
||||||
|
records = [data]
|
||||||
|
else:
|
||||||
|
raise ValueError(f"JSON must be object or array: {json_file.name}")
|
||||||
|
with open(output_path, "a", encoding="utf-8") as f:
|
||||||
|
for record in records:
|
||||||
|
f.write(json.dumps(record, ensure_ascii=False) + "\n")
|
||||||
|
output_count += 1
|
||||||
|
# 更新任务状态
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET status='completed', "
|
||||||
|
"input_count=%s, output_count=%s, update_time=NOW() WHERE id=%s",
|
||||||
|
(input_count, output_count, task_id),
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET status='failed', error_message=%s, update_time=NOW() WHERE id=%s",
|
||||||
|
(str(exc)[:500], task_id),
|
||||||
|
)
|
||||||
|
raise fail(500, f"convert failed: {exc}")
|
||||||
|
return ok(_get_task(task_id))
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{task_id}/download")
|
||||||
|
def download_result(
|
||||||
|
task_id: str,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
):
|
||||||
|
task = _get_task(task_id)
|
||||||
|
if not task:
|
||||||
|
raise fail(404, "task not found")
|
||||||
|
if task["status"] != "completed":
|
||||||
|
raise fail(400, "task is not completed")
|
||||||
|
output_path = _task_output_path(task)
|
||||||
|
if not output_path.exists():
|
||||||
|
raise fail(404, "output file not found")
|
||||||
|
return FileResponse(
|
||||||
|
str(output_path),
|
||||||
|
media_type="application/octet-stream",
|
||||||
|
filename=_safe_output_filename(task.get("output_filename")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{task_id}/import-as-dataset")
|
||||||
|
def import_as_dataset(
|
||||||
|
task_id: str,
|
||||||
|
payload: dict[str, Any] = Body(default={}),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""把已转换的 JSONL 文件导入为数据集管理中的上传任务记录(source='task')。"""
|
||||||
|
task = _get_task(task_id)
|
||||||
|
if not task:
|
||||||
|
raise fail(404, "task not found")
|
||||||
|
if task["status"] != "completed":
|
||||||
|
raise fail(400, "task is not completed")
|
||||||
|
output_path = _task_output_path(task)
|
||||||
|
if not output_path.exists():
|
||||||
|
raise fail(404, "output file not found")
|
||||||
|
content = output_path.read_text(encoding="utf-8")
|
||||||
|
dataset_name = str(payload.get("name") or task["name"]).strip()
|
||||||
|
description = str(payload.get("description") or f"由数据类型转换任务 {task_id} 导入").strip()
|
||||||
|
size_bytes = len(content.encode("utf-8"))
|
||||||
|
store = get_platform_store()
|
||||||
|
# 用 store 提供的接口创建数据集与文件
|
||||||
|
dataset = store.create_dataset({
|
||||||
|
"name": dataset_name,
|
||||||
|
"type": "train",
|
||||||
|
"storage_type": "local",
|
||||||
|
"source": "upload",
|
||||||
|
"task_id": task_id,
|
||||||
|
"size": f"{size_bytes} B",
|
||||||
|
"count": task["output_count"],
|
||||||
|
"description": description,
|
||||||
|
"created_by": task.get("created_by") or (current_user.get("id") if current_user else None),
|
||||||
|
})
|
||||||
|
dataset_id = dataset["id"]
|
||||||
|
with store.connect() as conn:
|
||||||
|
store.add_dataset_file(conn, dataset_id, _safe_output_filename(task.get("output_filename")), content)
|
||||||
|
return ok({"dataset_id": dataset_id, "name": dataset_name})
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{task_id}")
|
||||||
|
@op_log(module=OpModule.DATA_CONVERT, action=OpAction.DELETE, target_type="convert_task", target_name_param="task_id")
|
||||||
|
def delete_task(
|
||||||
|
task_id: str,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
task = _get_task(task_id)
|
||||||
|
if not task:
|
||||||
|
raise fail(404, "task not found")
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_convert_tasks SET deleted_at=NOW() WHERE id=%s",
|
||||||
|
(task_id,),
|
||||||
|
)
|
||||||
|
# 清理文件
|
||||||
|
import shutil
|
||||||
|
task_dir = _task_dir(task_id)
|
||||||
|
if task_dir.exists():
|
||||||
|
shutil.rmtree(task_dir, ignore_errors=True)
|
||||||
|
return ok({"deleted": task_id})
|
||||||
|
|
||||||
|
|
||||||
|
def _get_task(task_id: str) -> dict[str, Any] | None:
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT * FROM data_convert_tasks WHERE id=%s AND deleted_at IS NULL",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
return dict(row) if row else None
|
||||||
46
backend/app/modules/data_process/__init__.py
Normal file
46
backend/app/modules/data_process/__init__.py
Normal file
@@ -0,0 +1,46 @@
|
|||||||
|
"""数据处理模块:从原始文件接入到生成标准训练数据的全流程。
|
||||||
|
|
||||||
|
整体分层
|
||||||
|
--------
|
||||||
|
- ``algorithms/`` 纯算法层(无副作用:不访问 DB / 文件系统 / 网络)。
|
||||||
|
解析、格式检测、质量评分、去重、数据集切分、结构化预处理。
|
||||||
|
API、后台任务与测试共同复用。
|
||||||
|
- ``store/`` 持久层(PostgreSQL)。按 Mixin 拆分:
|
||||||
|
tasks / source_files / preview / generation / results / datasets。
|
||||||
|
构造时不连库、不迁移;部署方须显式执行 002 迁移(见 schema_cli)。
|
||||||
|
- 其余顶层文件 围绕上述两层的“服务 / 适配器”模块,由 endpoints 编排。
|
||||||
|
|
||||||
|
顶层模块速查
|
||||||
|
------------
|
||||||
|
constants.py 共享常量(MAX_QA_PAIRS_PER_ITEM、MODEL_GENERATION_BATCH_SIZE)。
|
||||||
|
storage.py 原始源文件的本地对象存储(受控暂存于 storage/data-process)。
|
||||||
|
office_preview.py Word/Excel 原文件的安全受限预览(仅返回绘制所需的结构化数据)。
|
||||||
|
document_chunking.py 基于 Docling / LlamaIndex 的文档切分。
|
||||||
|
dataset_format.py Alpaca/ShareGPT/DPO/CPT 数据集格式校验(训练提交前预检)。
|
||||||
|
generation.py 大模型生成适配器,把预览内容转为标准 instruction/output 记录。
|
||||||
|
schema_cli.py data_process 运行表的显式检查 / 安装命令(运维工具,应用启动不自动调用)。
|
||||||
|
|
||||||
|
调用关系
|
||||||
|
--------
|
||||||
|
``app/api/v1/endpoints/data_process.py`` 是编排入口,组合调用以上模块与两个子包;
|
||||||
|
``dataset_format.py`` 另被 ``db/platform_store.py`` 用于训练预检。
|
||||||
|
|
||||||
|
典型链路
|
||||||
|
--------
|
||||||
|
上传源文件 → storage 暂存 → algorithms.parse_text_content 解析
|
||||||
|
→ office_preview / document_chunking 处理 → store 落库预览
|
||||||
|
→ generation 调模型生成 → store 写 results → dataset_format 校验后发布。
|
||||||
|
|
||||||
|
导入约定
|
||||||
|
--------
|
||||||
|
为避免循环导入,本 ``__init__`` 不统一再导出;请按需从子模块直接导入:
|
||||||
|
|
||||||
|
from app.modules.data_process.store import DataProcessStore
|
||||||
|
from app.modules.data_process.algorithms import estimate_token_count
|
||||||
|
from app.modules.data_process.generation import generate_model_records
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 注意:为了避免循环导入,不在此处导入所有内容
|
||||||
|
# 请直接从子模块导入所需功能
|
||||||
|
|
||||||
|
__all__ = ["store", "algorithms"]
|
||||||
145
backend/app/modules/data_process/algorithms/__init__.py
Normal file
145
backend/app/modules/data_process/algorithms/__init__.py
Normal file
@@ -0,0 +1,145 @@
|
|||||||
|
"""数据处理算法模块。
|
||||||
|
|
||||||
|
重构为多个按职责拆分的子模块:
|
||||||
|
|
||||||
|
- types.py: 类型定义、常量与 dataclass
|
||||||
|
- text_utils.py: 文本解码、归一化与格式检测
|
||||||
|
- format_detection.py: 文档结构检测
|
||||||
|
- parsers/: PDF / Office / JSON / CSV 解析器
|
||||||
|
- quality.py: 质量评分与去重
|
||||||
|
- transforms.py: 数据集分割
|
||||||
|
- structured_processing.py: 结构化数据预处理
|
||||||
|
|
||||||
|
使用方式:
|
||||||
|
from app.modules.data_process.algorithms import estimate_token_count
|
||||||
|
|
||||||
|
子模块之间存在导入分层,顶层按依赖顺序导入以避免循环导入。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
# Layer 0: 类型与常量(叶子,不依赖内部模块)
|
||||||
|
from .types import (
|
||||||
|
DatasetSplit,
|
||||||
|
DocumentHeading,
|
||||||
|
DocumentNoiseSpan,
|
||||||
|
DocumentStructure,
|
||||||
|
MAX_QA_PAIRS_PER_ITEM,
|
||||||
|
ParsedText,
|
||||||
|
PdfPageText,
|
||||||
|
ProcessedStructuredRecord,
|
||||||
|
QualityScore,
|
||||||
|
SUPPORTED_TEXT_FORMATS,
|
||||||
|
StructuredPreprocessOption,
|
||||||
|
TextFormat,
|
||||||
|
_MAX_WORKBOOK_COLUMNS,
|
||||||
|
_MAX_WORKBOOK_HEADER_SCAN_ROWS,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Layer 2: 文本工具(仅依赖 types;对 parsers 的依赖在函数体内延迟导入)
|
||||||
|
from .text_utils import (
|
||||||
|
_normalize_spreadsheet_value,
|
||||||
|
decode_utf8,
|
||||||
|
detect_text_format,
|
||||||
|
normalize_text,
|
||||||
|
parse_text_content,
|
||||||
|
parse_utf8_text,
|
||||||
|
structured_json_dumps,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Layer 1: 解析器(依赖 types 与 text_utils)
|
||||||
|
from .parsers import (
|
||||||
|
_infer_xlsx_header_region,
|
||||||
|
_rewrite_xlsx_workbook_relationships,
|
||||||
|
_validate_office_archive,
|
||||||
|
_xlsx_sheet_merge_ranges,
|
||||||
|
detect_pdf_document_noise,
|
||||||
|
extract_pdf_page_texts,
|
||||||
|
remove_document_noise,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Layer 3: 数据转换与质量评分
|
||||||
|
from .transforms import stable_split, stable_split_assignments
|
||||||
|
from .quality import (
|
||||||
|
content_quality_flags,
|
||||||
|
deduplicate_structured_records,
|
||||||
|
estimate_token_count,
|
||||||
|
fingerprints_are_near_duplicate,
|
||||||
|
is_low_quality_content,
|
||||||
|
is_near_duplicate,
|
||||||
|
near_duplicate_fingerprint,
|
||||||
|
record_fingerprint,
|
||||||
|
score_quality,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Layer 4: 结构化数据处理(依赖 text_utils / quality / transforms / parsers)
|
||||||
|
from .structured_processing import (
|
||||||
|
canonical_record_json,
|
||||||
|
desensitize_pii,
|
||||||
|
desensitize_structured_record,
|
||||||
|
expand_to_context_boundaries,
|
||||||
|
extract_structured_records,
|
||||||
|
filter_anomalous_structured_records,
|
||||||
|
flatten_structured_record,
|
||||||
|
generate_standard_records,
|
||||||
|
merge_short_blocks,
|
||||||
|
normalize_structured_record,
|
||||||
|
preprocess_structured_records,
|
||||||
|
preprocess_structured_records_with_lineage,
|
||||||
|
protected_context_ranges,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Layer 5: 文档结构检测(依赖 structured_processing)
|
||||||
|
from .format_detection import detect_document_structure
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"DatasetSplit",
|
||||||
|
"DocumentHeading",
|
||||||
|
"DocumentNoiseSpan",
|
||||||
|
"DocumentStructure",
|
||||||
|
"MAX_QA_PAIRS_PER_ITEM",
|
||||||
|
"ParsedText",
|
||||||
|
"PdfPageText",
|
||||||
|
"ProcessedStructuredRecord",
|
||||||
|
"QualityScore",
|
||||||
|
"SUPPORTED_TEXT_FORMATS",
|
||||||
|
"StructuredPreprocessOption",
|
||||||
|
"TextFormat",
|
||||||
|
"_MAX_WORKBOOK_COLUMNS",
|
||||||
|
"_MAX_WORKBOOK_HEADER_SCAN_ROWS",
|
||||||
|
"_normalize_spreadsheet_value",
|
||||||
|
"canonical_record_json",
|
||||||
|
"content_quality_flags",
|
||||||
|
"decode_utf8",
|
||||||
|
"deduplicate_structured_records",
|
||||||
|
"desensitize_pii",
|
||||||
|
"desensitize_structured_record",
|
||||||
|
"detect_document_structure",
|
||||||
|
"detect_pdf_document_noise",
|
||||||
|
"detect_text_format",
|
||||||
|
"estimate_token_count",
|
||||||
|
"expand_to_context_boundaries",
|
||||||
|
"extract_pdf_page_texts",
|
||||||
|
"extract_structured_records",
|
||||||
|
"filter_anomalous_structured_records",
|
||||||
|
"fingerprints_are_near_duplicate",
|
||||||
|
"flatten_structured_record",
|
||||||
|
"generate_standard_records",
|
||||||
|
"is_low_quality_content",
|
||||||
|
"is_near_duplicate",
|
||||||
|
"merge_short_blocks",
|
||||||
|
"near_duplicate_fingerprint",
|
||||||
|
"normalize_structured_record",
|
||||||
|
"normalize_text",
|
||||||
|
"parse_text_content",
|
||||||
|
"parse_utf8_text",
|
||||||
|
"preprocess_structured_records",
|
||||||
|
"preprocess_structured_records_with_lineage",
|
||||||
|
"protected_context_ranges",
|
||||||
|
"record_fingerprint",
|
||||||
|
"remove_document_noise",
|
||||||
|
"score_quality",
|
||||||
|
"stable_split",
|
||||||
|
"stable_split_assignments",
|
||||||
|
"structured_json_dumps",
|
||||||
|
]
|
||||||
119
backend/app/modules/data_process/algorithms/format_detection.py
Normal file
119
backend/app/modules/data_process/algorithms/format_detection.py
Normal file
@@ -0,0 +1,119 @@
|
|||||||
|
"""数据处理算法 - 格式检测。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
|
||||||
|
from .structured_processing import protected_context_ranges
|
||||||
|
from .text_utils import normalize_text
|
||||||
|
from .types import DocumentHeading, DocumentStructure
|
||||||
|
|
||||||
|
|
||||||
|
def detect_document_structure(text: str) -> DocumentStructure:
|
||||||
|
"""识别 Markdown、中文章节和数字编号标题及不可拆分块。"""
|
||||||
|
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
if not normalized:
|
||||||
|
return DocumentStructure(
|
||||||
|
line_count=0,
|
||||||
|
paragraph_count=0,
|
||||||
|
headings=(),
|
||||||
|
code_block_count=0,
|
||||||
|
table_block_count=0,
|
||||||
|
list_block_count=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
code_ranges = protected_context_ranges(
|
||||||
|
normalized,
|
||||||
|
preserve_code_blocks=True,
|
||||||
|
preserve_tables=False,
|
||||||
|
preserve_lists=False,
|
||||||
|
)
|
||||||
|
table_candidates = protected_context_ranges(
|
||||||
|
normalized,
|
||||||
|
preserve_code_blocks=False,
|
||||||
|
preserve_tables=True,
|
||||||
|
preserve_lists=False,
|
||||||
|
)
|
||||||
|
list_candidates = protected_context_ranges(
|
||||||
|
normalized,
|
||||||
|
preserve_code_blocks=False,
|
||||||
|
preserve_tables=False,
|
||||||
|
preserve_lists=True,
|
||||||
|
)
|
||||||
|
table_ranges = tuple(
|
||||||
|
item
|
||||||
|
for item in table_candidates
|
||||||
|
if not any(
|
||||||
|
item[0] < code_end and item[1] > code_start
|
||||||
|
for code_start, code_end in code_ranges
|
||||||
|
)
|
||||||
|
)
|
||||||
|
list_ranges = tuple(
|
||||||
|
item
|
||||||
|
for item in list_candidates
|
||||||
|
if not any(
|
||||||
|
item[0] < code_end and item[1] > code_start
|
||||||
|
for code_start, code_end in code_ranges
|
||||||
|
)
|
||||||
|
)
|
||||||
|
markdown_heading = re.compile(r"^\s*(?P<marks>#{1,6})\s+(?P<title>.+?)\s*#*\s*$")
|
||||||
|
chinese_heading = re.compile(
|
||||||
|
r"^\s*(?P<title>第[一二三四五六七八九十百千万0-9]+[章节篇部分].*)$"
|
||||||
|
)
|
||||||
|
numbered_heading = re.compile(
|
||||||
|
r"^\s*(?P<number>\d+(?:\.\d+)*)(?:[、.]|\s+)\s*(?P<title>\S.*)$"
|
||||||
|
)
|
||||||
|
|
||||||
|
headings: list[DocumentHeading] = []
|
||||||
|
cursor = 0
|
||||||
|
for line_number, raw_line in enumerate(normalized.splitlines(keepends=True), start=1):
|
||||||
|
line = raw_line.rstrip("\n")
|
||||||
|
line_end = cursor + len(line)
|
||||||
|
if not any(range_start <= cursor < range_end for range_start, range_end in code_ranges):
|
||||||
|
match = markdown_heading.match(line)
|
||||||
|
if match:
|
||||||
|
headings.append(
|
||||||
|
DocumentHeading(
|
||||||
|
level=len(match.group("marks")),
|
||||||
|
title=normalize_text(match.group("title")),
|
||||||
|
line_number=line_number,
|
||||||
|
start=cursor,
|
||||||
|
end=line_end,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
match = chinese_heading.match(line)
|
||||||
|
if match:
|
||||||
|
headings.append(
|
||||||
|
DocumentHeading(
|
||||||
|
level=1,
|
||||||
|
title=normalize_text(match.group("title")),
|
||||||
|
line_number=line_number,
|
||||||
|
start=cursor,
|
||||||
|
end=line_end,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
match = numbered_heading.match(line)
|
||||||
|
if match:
|
||||||
|
headings.append(
|
||||||
|
DocumentHeading(
|
||||||
|
level=min(6, match.group("number").count(".") + 1),
|
||||||
|
title=normalize_text(match.group("title")),
|
||||||
|
line_number=line_number,
|
||||||
|
start=cursor,
|
||||||
|
end=line_end,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
cursor += len(raw_line)
|
||||||
|
|
||||||
|
paragraphs = [part for part in re.split(r"\n\s*\n", normalized) if part.strip()]
|
||||||
|
return DocumentStructure(
|
||||||
|
line_count=len(normalized.splitlines()),
|
||||||
|
paragraph_count=len(paragraphs),
|
||||||
|
headings=tuple(headings),
|
||||||
|
code_block_count=len(code_ranges),
|
||||||
|
table_block_count=len(table_ranges),
|
||||||
|
list_block_count=len(list_ranges),
|
||||||
|
)
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
"""文档解析器模块。"""
|
||||||
|
|
||||||
|
from .pdf import extract_pdf_page_texts, detect_pdf_document_noise, remove_document_noise
|
||||||
|
from .office import (
|
||||||
|
_validate_office_archive,
|
||||||
|
_rewrite_xlsx_workbook_relationships,
|
||||||
|
_xlsx_sheet_merge_ranges,
|
||||||
|
_infer_xlsx_header_region,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'extract_pdf_page_texts',
|
||||||
|
'detect_pdf_document_noise',
|
||||||
|
'remove_document_noise',
|
||||||
|
'_validate_office_archive',
|
||||||
|
'_rewrite_xlsx_workbook_relationships',
|
||||||
|
'_xlsx_sheet_merge_ranges',
|
||||||
|
'_infer_xlsx_header_region',
|
||||||
|
]
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
"""数据处理算法 - CSV 解析。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import csv
|
||||||
|
import io
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from ..text_utils import normalize_text
|
||||||
|
|
||||||
|
|
||||||
|
def _source_line_offsets(
|
||||||
|
text: str,
|
||||||
|
start_line: int,
|
||||||
|
end_line: int,
|
||||||
|
) -> tuple[int, int]:
|
||||||
|
"""把 1-based 物理行范围转换为左闭右开的字符范围。"""
|
||||||
|
|
||||||
|
line_starts = [0]
|
||||||
|
line_starts.extend(match.end() for match in re.finditer("\n", text))
|
||||||
|
if start_line < 1 or end_line < start_line or end_line > len(line_starts):
|
||||||
|
raise ValueError("source line range is outside normalized text")
|
||||||
|
source_start = line_starts[start_line - 1]
|
||||||
|
source_end = (
|
||||||
|
line_starts[end_line] - 1
|
||||||
|
if end_line < len(line_starts)
|
||||||
|
else len(text)
|
||||||
|
)
|
||||||
|
return source_start, source_end
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_csv_records_with_locators(
|
||||||
|
text: str,
|
||||||
|
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||||
|
"""从 CSV 文本中提取记录及其行级定位信息。"""
|
||||||
|
|
||||||
|
normalized_text = normalize_text(text)
|
||||||
|
if not normalized_text:
|
||||||
|
return [], []
|
||||||
|
|
||||||
|
try:
|
||||||
|
dialect = csv.Sniffer().sniff(normalized_text[:8192], delimiters=",\t;")
|
||||||
|
except csv.Error:
|
||||||
|
dialect = csv.excel
|
||||||
|
reader = csv.DictReader(io.StringIO(normalized_text), dialect=dialect)
|
||||||
|
if not reader.fieldnames:
|
||||||
|
raise ValueError("CSV header is required")
|
||||||
|
headers = [normalize_text(header or "") for header in reader.fieldnames]
|
||||||
|
if any(not header for header in headers):
|
||||||
|
raise ValueError("CSV header cannot be empty")
|
||||||
|
if len(set(headers)) != len(headers):
|
||||||
|
raise ValueError("CSV headers must be unique")
|
||||||
|
reader.fieldnames = headers
|
||||||
|
|
||||||
|
records: list[dict[str, Any]] = []
|
||||||
|
locators: list[dict[str, Any]] = []
|
||||||
|
source_lines = normalized_text.splitlines()
|
||||||
|
previous_end_line = reader.line_num
|
||||||
|
for row in reader:
|
||||||
|
end_line = reader.line_num
|
||||||
|
start_line = previous_end_line + 1
|
||||||
|
previous_end_line = end_line
|
||||||
|
while start_line < end_line and not source_lines[start_line - 1].strip():
|
||||||
|
start_line += 1
|
||||||
|
if None in row:
|
||||||
|
raise ValueError("CSV row has more fields than the header")
|
||||||
|
normalized_row = {
|
||||||
|
key: normalize_text(value or "")
|
||||||
|
for key, value in row.items()
|
||||||
|
}
|
||||||
|
if any(value for value in normalized_row.values()):
|
||||||
|
records.append(normalized_row)
|
||||||
|
source_start, source_end = _source_line_offsets(
|
||||||
|
normalized_text,
|
||||||
|
start_line,
|
||||||
|
end_line,
|
||||||
|
)
|
||||||
|
locators.append(
|
||||||
|
{
|
||||||
|
"kind": "csv",
|
||||||
|
"record_index": len(records),
|
||||||
|
"start_line": start_line,
|
||||||
|
"end_line": end_line,
|
||||||
|
"source_start": source_start,
|
||||||
|
"source_end": source_end,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return records, locators
|
||||||
@@ -0,0 +1,395 @@
|
|||||||
|
"""数据处理算法 - JSON/JSONL 解析。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
from collections.abc import Mapping, Sequence
|
||||||
|
from decimal import Decimal
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from ..text_utils import _normalize_format, _normalize_value, normalize_text
|
||||||
|
from ..types import (
|
||||||
|
_DuplicateJsonKeyError,
|
||||||
|
_JSON_ENVELOPE_KEYS,
|
||||||
|
_JSON_RECORD_ARRAY_KEYS,
|
||||||
|
_JSON_RESPONSE_METADATA_KEYS,
|
||||||
|
_JSON_WRAPPER_METADATA_KEYS,
|
||||||
|
_MAX_JSON_DEPTH,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _record_from_value(value: Any, *, normalize: bool = True) -> dict[str, Any]:
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return dict(_normalize_value(value)) if normalize else dict(value)
|
||||||
|
return {"value": _normalize_value(value) if normalize else value}
|
||||||
|
|
||||||
|
def _json_pointer_segment(value: Any) -> str:
|
||||||
|
return str(value).replace("~", "~0").replace("/", "~1")
|
||||||
|
|
||||||
|
def _reject_duplicate_json_keys(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
|
||||||
|
result: dict[str, Any] = {}
|
||||||
|
for key, value in pairs:
|
||||||
|
if key in result:
|
||||||
|
raise _DuplicateJsonKeyError(f"duplicate JSON object key: {key!r}")
|
||||||
|
result[key] = value
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _reject_json_constant(value: str) -> Any:
|
||||||
|
raise ValueError(f"non-finite JSON number is not allowed: {value}")
|
||||||
|
|
||||||
|
def _skip_json_whitespace(text: str, offset: int) -> int:
|
||||||
|
while offset < len(text) and text[offset] in " \t\r\n":
|
||||||
|
offset += 1
|
||||||
|
return offset
|
||||||
|
|
||||||
|
def _validate_json_nesting(text: str) -> None:
|
||||||
|
"""在构造 Python 对象前限制容器深度,避免依赖解释器递归阈值。"""
|
||||||
|
|
||||||
|
depth = 0
|
||||||
|
in_string = False
|
||||||
|
escaped = False
|
||||||
|
for char in text:
|
||||||
|
if in_string:
|
||||||
|
if escaped:
|
||||||
|
escaped = False
|
||||||
|
elif char == "\\":
|
||||||
|
escaped = True
|
||||||
|
elif char == '"':
|
||||||
|
in_string = False
|
||||||
|
continue
|
||||||
|
if char == '"':
|
||||||
|
in_string = True
|
||||||
|
elif char in "[{":
|
||||||
|
depth += 1
|
||||||
|
if depth > _MAX_JSON_DEPTH:
|
||||||
|
raise ValueError(
|
||||||
|
f"JSON nesting exceeds the supported depth of {_MAX_JSON_DEPTH}"
|
||||||
|
)
|
||||||
|
elif char in "]}":
|
||||||
|
depth = max(0, depth - 1)
|
||||||
|
|
||||||
|
def _strict_json_loads(text: str) -> tuple[Any, int, int]:
|
||||||
|
"""严格解析单个 JSON 值并返回其左闭右开源码区间。"""
|
||||||
|
|
||||||
|
start = _skip_json_whitespace(text, 0)
|
||||||
|
if start >= len(text):
|
||||||
|
raise ValueError("JSON content is empty")
|
||||||
|
_validate_json_nesting(text)
|
||||||
|
decoder = json.JSONDecoder(
|
||||||
|
object_pairs_hook=_reject_duplicate_json_keys,
|
||||||
|
parse_float=Decimal,
|
||||||
|
parse_int=int,
|
||||||
|
parse_constant=_reject_json_constant,
|
||||||
|
strict=True,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
payload, end = decoder.raw_decode(text, start)
|
||||||
|
except json.JSONDecodeError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"invalid JSON at line {exc.lineno}, column {exc.colno}: {exc.msg}"
|
||||||
|
) from exc
|
||||||
|
except RecursionError as exc:
|
||||||
|
raise ValueError("JSON nesting exceeds the supported depth") from exc
|
||||||
|
except _DuplicateJsonKeyError as exc:
|
||||||
|
raise ValueError(str(exc)) from exc
|
||||||
|
except ValueError as exc:
|
||||||
|
# parse_int/parse_float/parse_constant 的异常也必须稳定映射为客户端错误。
|
||||||
|
raise ValueError(f"invalid JSON number: {exc}") from exc
|
||||||
|
trailing = _skip_json_whitespace(text, end)
|
||||||
|
if trailing != len(text):
|
||||||
|
line = text.count("\n", 0, trailing) + 1
|
||||||
|
line_start = text.rfind("\n", 0, trailing) + 1
|
||||||
|
column = trailing - line_start + 1
|
||||||
|
raise ValueError(
|
||||||
|
f"invalid JSON at line {line}, column {column}: extra data"
|
||||||
|
)
|
||||||
|
return payload, start, end
|
||||||
|
|
||||||
|
def _json_value_end(text: str, start: int) -> int:
|
||||||
|
"""在已验证 JSON 中定位一个值的结束偏移,不对数值做二次解析。"""
|
||||||
|
|
||||||
|
if start >= len(text):
|
||||||
|
raise ValueError("invalid JSON source span")
|
||||||
|
first = text[start]
|
||||||
|
if first == '"':
|
||||||
|
escaped = False
|
||||||
|
for offset in range(start + 1, len(text)):
|
||||||
|
char = text[offset]
|
||||||
|
if escaped:
|
||||||
|
escaped = False
|
||||||
|
elif char == "\\":
|
||||||
|
escaped = True
|
||||||
|
elif char == '"':
|
||||||
|
return offset + 1
|
||||||
|
raise ValueError("invalid JSON source span")
|
||||||
|
if first in "[{":
|
||||||
|
stack = [first]
|
||||||
|
in_string = False
|
||||||
|
escaped = False
|
||||||
|
for offset in range(start + 1, len(text)):
|
||||||
|
char = text[offset]
|
||||||
|
if in_string:
|
||||||
|
if escaped:
|
||||||
|
escaped = False
|
||||||
|
elif char == "\\":
|
||||||
|
escaped = True
|
||||||
|
elif char == '"':
|
||||||
|
in_string = False
|
||||||
|
continue
|
||||||
|
if char == '"':
|
||||||
|
in_string = True
|
||||||
|
elif char in "[{":
|
||||||
|
stack.append(char)
|
||||||
|
elif char in "]}":
|
||||||
|
expected = "[" if char == "]" else "{"
|
||||||
|
if not stack or stack[-1] != expected:
|
||||||
|
raise ValueError("invalid JSON source span")
|
||||||
|
stack.pop()
|
||||||
|
if not stack:
|
||||||
|
return offset + 1
|
||||||
|
raise ValueError("invalid JSON source span")
|
||||||
|
end = start
|
||||||
|
while end < len(text) and text[end] not in " \t\r\n,]}":
|
||||||
|
end += 1
|
||||||
|
if end == start:
|
||||||
|
raise ValueError("invalid JSON source span")
|
||||||
|
return end
|
||||||
|
|
||||||
|
def _json_object_value_spans(
|
||||||
|
text: str,
|
||||||
|
start: int,
|
||||||
|
end: int,
|
||||||
|
) -> dict[str, tuple[int, int]]:
|
||||||
|
"""返回已验证 JSON 对象直接子字段的值区间。"""
|
||||||
|
|
||||||
|
if text[start] != "{" or text[end - 1] != "}":
|
||||||
|
raise ValueError("JSON source value is not an object")
|
||||||
|
result: dict[str, tuple[int, int]] = {}
|
||||||
|
offset = _skip_json_whitespace(text, start + 1)
|
||||||
|
key_decoder = json.JSONDecoder()
|
||||||
|
while offset < end - 1:
|
||||||
|
key, key_end = key_decoder.raw_decode(text, offset)
|
||||||
|
if not isinstance(key, str):
|
||||||
|
raise ValueError("invalid JSON object key")
|
||||||
|
offset = _skip_json_whitespace(text, key_end)
|
||||||
|
if offset >= end or text[offset] != ":":
|
||||||
|
raise ValueError("invalid JSON object member")
|
||||||
|
value_start = _skip_json_whitespace(text, offset + 1)
|
||||||
|
value_end = _json_value_end(text, value_start)
|
||||||
|
result[key] = (value_start, value_end)
|
||||||
|
offset = _skip_json_whitespace(text, value_end)
|
||||||
|
if offset >= end - 1:
|
||||||
|
break
|
||||||
|
if text[offset] != ",":
|
||||||
|
raise ValueError("invalid JSON object member")
|
||||||
|
offset = _skip_json_whitespace(text, offset + 1)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _json_array_item_spans(
|
||||||
|
text: str,
|
||||||
|
start: int,
|
||||||
|
end: int,
|
||||||
|
) -> list[tuple[int, int]]:
|
||||||
|
"""返回已验证 JSON 数组中每个直接元素的源码区间。"""
|
||||||
|
|
||||||
|
if text[start] != "[" or text[end - 1] != "]":
|
||||||
|
raise ValueError("JSON source value is not an array")
|
||||||
|
result: list[tuple[int, int]] = []
|
||||||
|
offset = _skip_json_whitespace(text, start + 1)
|
||||||
|
while offset < end - 1:
|
||||||
|
item_end = _json_value_end(text, offset)
|
||||||
|
result.append((offset, item_end))
|
||||||
|
offset = _skip_json_whitespace(text, item_end)
|
||||||
|
if offset >= end - 1:
|
||||||
|
break
|
||||||
|
if text[offset] != ",":
|
||||||
|
raise ValueError("invalid JSON array item")
|
||||||
|
offset = _skip_json_whitespace(text, offset + 1)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _json_span_at_path(
|
||||||
|
text: str,
|
||||||
|
root_span: tuple[int, int],
|
||||||
|
path: Sequence[str],
|
||||||
|
) -> tuple[int, int]:
|
||||||
|
span = root_span
|
||||||
|
for key in path:
|
||||||
|
try:
|
||||||
|
span = _json_object_value_spans(text, *span)[key]
|
||||||
|
except KeyError as exc:
|
||||||
|
raise ValueError(f"JSON source path cannot be located: {key}") from exc
|
||||||
|
return span
|
||||||
|
|
||||||
|
|
||||||
|
def _pure_json_record_wrapper(
|
||||||
|
payload: Any,
|
||||||
|
) -> tuple[list[Any], tuple[str, ...]] | None:
|
||||||
|
"""识别不会与业务字段冲突的纯记录包装对象。"""
|
||||||
|
|
||||||
|
if not isinstance(payload, Mapping):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def direct_wrapper(
|
||||||
|
value: Mapping[str, Any],
|
||||||
|
metadata_keys: frozenset[str] | set[str] = _JSON_WRAPPER_METADATA_KEYS,
|
||||||
|
) -> tuple[list[Any], tuple[str, ...]] | None:
|
||||||
|
candidates = [
|
||||||
|
key
|
||||||
|
for key in _JSON_RECORD_ARRAY_KEYS
|
||||||
|
if isinstance(value.get(key), list)
|
||||||
|
]
|
||||||
|
if len(candidates) != 1:
|
||||||
|
return None
|
||||||
|
record_key = candidates[0]
|
||||||
|
records = value[record_key]
|
||||||
|
if any(not isinstance(record, Mapping) for record in records):
|
||||||
|
return None
|
||||||
|
if any(
|
||||||
|
key != record_key and key not in metadata_keys
|
||||||
|
for key in value
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
return records, (record_key,)
|
||||||
|
|
||||||
|
direct = direct_wrapper(payload)
|
||||||
|
if direct is not None:
|
||||||
|
return direct
|
||||||
|
|
||||||
|
envelope_keys = [
|
||||||
|
key
|
||||||
|
for key in _JSON_ENVELOPE_KEYS
|
||||||
|
if isinstance(payload.get(key), Mapping)
|
||||||
|
]
|
||||||
|
if len(envelope_keys) != 1:
|
||||||
|
return None
|
||||||
|
envelope_key = envelope_keys[0]
|
||||||
|
if any(
|
||||||
|
key != envelope_key and key not in _JSON_RESPONSE_METADATA_KEYS
|
||||||
|
for key in payload
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
nested = direct_wrapper(
|
||||||
|
payload[envelope_key],
|
||||||
|
_JSON_RESPONSE_METADATA_KEYS,
|
||||||
|
)
|
||||||
|
if nested is None:
|
||||||
|
return None
|
||||||
|
records, nested_path = nested
|
||||||
|
return records, (envelope_key, *nested_path)
|
||||||
|
|
||||||
|
|
||||||
|
def _json_record_locator(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
record_index: int,
|
||||||
|
json_pointer: str,
|
||||||
|
span: tuple[int, int],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
source_start, source_end = span
|
||||||
|
start_line = text.count("\n", 0, source_start) + 1
|
||||||
|
last_character = max(source_start, source_end - 1)
|
||||||
|
end_line = text.count("\n", 0, last_character) + 1
|
||||||
|
return {
|
||||||
|
"kind": "json",
|
||||||
|
"record_index": record_index,
|
||||||
|
"json_pointer": json_pointer,
|
||||||
|
"source_start": source_start,
|
||||||
|
"source_end": source_end,
|
||||||
|
"start_line": start_line,
|
||||||
|
"end_line": end_line,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_structured_records_with_locators(
|
||||||
|
text: str,
|
||||||
|
file_format: str,
|
||||||
|
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||||
|
"""从 JSON、JSONL 或 CSV 中提取记录。
|
||||||
|
|
||||||
|
JSON 根数组始终表示多条记录;对象仅在满足纯包装契约时展开,其他
|
||||||
|
对象均视为一条业务记录。JSON 字段和值在解析阶段保持原样,只有用户
|
||||||
|
明确选择 ``normalize_format`` 后才会规范化。
|
||||||
|
"""
|
||||||
|
|
||||||
|
normalized_format = _normalize_format(file_format)
|
||||||
|
if normalized_format not in {"json", "jsonl", "csv"}:
|
||||||
|
raise ValueError("structured record extraction only supports JSON, JSONL and CSV")
|
||||||
|
|
||||||
|
if normalized_format == "json":
|
||||||
|
if _skip_json_whitespace(text, 0) == len(text):
|
||||||
|
return [], []
|
||||||
|
payload, root_start, root_end = _strict_json_loads(text)
|
||||||
|
values: Sequence[Any]
|
||||||
|
pointer_path: tuple[str, ...] = ()
|
||||||
|
record_spans: list[tuple[int, int]]
|
||||||
|
if isinstance(payload, list):
|
||||||
|
values = payload
|
||||||
|
record_spans = _json_array_item_spans(text, root_start, root_end)
|
||||||
|
else:
|
||||||
|
wrapper = _pure_json_record_wrapper(payload)
|
||||||
|
if wrapper is None:
|
||||||
|
values = [payload]
|
||||||
|
record_spans = [(root_start, root_end)]
|
||||||
|
else:
|
||||||
|
values, pointer_path = wrapper
|
||||||
|
array_span = _json_span_at_path(
|
||||||
|
text,
|
||||||
|
(root_start, root_end),
|
||||||
|
pointer_path,
|
||||||
|
)
|
||||||
|
record_spans = _json_array_item_spans(text, *array_span)
|
||||||
|
if len(record_spans) != len(values):
|
||||||
|
raise ValueError("JSON record source spans do not match parsed records")
|
||||||
|
records = [_record_from_value(value, normalize=False) for value in values]
|
||||||
|
pointer_prefix = "".join(
|
||||||
|
f"/{_json_pointer_segment(segment)}" for segment in pointer_path
|
||||||
|
)
|
||||||
|
locators = [
|
||||||
|
_json_record_locator(
|
||||||
|
text,
|
||||||
|
record_index=index + 1,
|
||||||
|
json_pointer=(
|
||||||
|
f"{pointer_prefix}/{index}"
|
||||||
|
if pointer_path or isinstance(payload, list)
|
||||||
|
else ""
|
||||||
|
),
|
||||||
|
span=record_spans[index],
|
||||||
|
)
|
||||||
|
for index in range(len(records))
|
||||||
|
]
|
||||||
|
return records, locators
|
||||||
|
|
||||||
|
if normalized_format == "jsonl":
|
||||||
|
records: list[dict[str, Any]] = []
|
||||||
|
locators: list[dict[str, Any]] = []
|
||||||
|
source_offset = 0
|
||||||
|
for line_number, line in enumerate(text.split("\n"), start=1):
|
||||||
|
line_content_end = len(line)
|
||||||
|
if _skip_json_whitespace(line, 0) == line_content_end:
|
||||||
|
source_offset += len(line) + 1
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
value, value_start, value_end = _strict_json_loads(line)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"invalid JSONL at line {line_number}: {exc}"
|
||||||
|
) from exc
|
||||||
|
records.append(_record_from_value(value, normalize=False))
|
||||||
|
locators.append(
|
||||||
|
{
|
||||||
|
"kind": "jsonl",
|
||||||
|
"record_index": len(records),
|
||||||
|
"start_line": line_number,
|
||||||
|
"end_line": line_number,
|
||||||
|
"source_start": source_offset + value_start,
|
||||||
|
"source_end": source_offset + value_end,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
source_offset += len(line) + 1
|
||||||
|
return records, locators
|
||||||
|
|
||||||
|
# CSV 分支独立放在 csv_parser 中,避免与 JSON 机制耦合。
|
||||||
|
from .csv_parser import _extract_csv_records_with_locators
|
||||||
|
|
||||||
|
return _extract_csv_records_with_locators(text)
|
||||||
700
backend/app/modules/data_process/algorithms/parsers/office.py
Normal file
700
backend/app/modules/data_process/algorithms/parsers/office.py
Normal file
@@ -0,0 +1,700 @@
|
|||||||
|
"""数据处理算法 - Office 文档解析。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import io
|
||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
import zipfile
|
||||||
|
import xml.etree.ElementTree as ET
|
||||||
|
from collections.abc import Mapping, Sequence
|
||||||
|
from pathlib import PurePosixPath
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import unquote, urlsplit
|
||||||
|
|
||||||
|
from docx import Document
|
||||||
|
from docx.oxml.table import CT_Tbl
|
||||||
|
from docx.oxml.text.paragraph import CT_P
|
||||||
|
from docx.table import Table
|
||||||
|
from docx.text.paragraph import Paragraph
|
||||||
|
from openpyxl import load_workbook
|
||||||
|
from openpyxl.utils.cell import range_boundaries
|
||||||
|
from pptx import Presentation
|
||||||
|
|
||||||
|
from ..text_utils import _append_bounded_text, _normalize_spreadsheet_value, normalize_text
|
||||||
|
from ..types import (
|
||||||
|
_MAX_ARCHIVE_COMPRESSION_RATIO,
|
||||||
|
_MAX_ARCHIVE_ENTRIES,
|
||||||
|
_MAX_ARCHIVE_ENTRY_BYTES,
|
||||||
|
_MAX_ARCHIVE_UNCOMPRESSED_BYTES,
|
||||||
|
_MAX_PRESENTATION_SLIDES,
|
||||||
|
_MAX_WORKBOOK_CELLS,
|
||||||
|
_MAX_WORKBOOK_COLUMNS,
|
||||||
|
_MAX_WORKBOOK_HEADER_ROWS,
|
||||||
|
_MAX_WORKBOOK_HEADER_SCAN_ROWS,
|
||||||
|
_MAX_WORKBOOK_MERGED_RANGES,
|
||||||
|
_MAX_WORKBOOK_ROWS,
|
||||||
|
_MAX_WORKBOOK_SCANNED_ROWS,
|
||||||
|
_MAX_WORKBOOK_SHEETS,
|
||||||
|
TextFormat,
|
||||||
|
)
|
||||||
|
|
||||||
|
_XLSX_REPORT_METADATA_PATTERN = re.compile(
|
||||||
|
r"^(?:报表|报告|标题|说明|备注|制表|统计|日期|时间|期间|"
|
||||||
|
r"report|title|note|remark|date|time|period)\b",
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_office_archive(raw: bytes, file_format: TextFormat) -> None:
|
||||||
|
"""在交给 Office 解析库前限制 ZIP 包规模并拒绝活动 XML。"""
|
||||||
|
|
||||||
|
required_members = {
|
||||||
|
"docx": {"[Content_Types].xml", "word/document.xml"},
|
||||||
|
"xlsx": {"[Content_Types].xml", "xl/workbook.xml"},
|
||||||
|
"pptx": {"[Content_Types].xml", "ppt/presentation.xml"},
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
with zipfile.ZipFile(io.BytesIO(raw)) as archive:
|
||||||
|
members = archive.infolist()
|
||||||
|
if len(members) > _MAX_ARCHIVE_ENTRIES:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive contains too many entries "
|
||||||
|
f"(limit {_MAX_ARCHIVE_ENTRIES})"
|
||||||
|
)
|
||||||
|
|
||||||
|
names: set[str] = set()
|
||||||
|
total_size = 0
|
||||||
|
for member in members:
|
||||||
|
path = PurePosixPath(member.filename.replace("\\", "/"))
|
||||||
|
if path.is_absolute() or ".." in path.parts:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive contains an unsafe member path"
|
||||||
|
)
|
||||||
|
if member.filename in names:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive contains duplicate member names"
|
||||||
|
)
|
||||||
|
names.add(member.filename)
|
||||||
|
if member.flag_bits & 0x1:
|
||||||
|
raise ValueError(f"encrypted {file_format.upper()} files are not supported")
|
||||||
|
if member.is_dir():
|
||||||
|
continue
|
||||||
|
if member.file_size > _MAX_ARCHIVE_ENTRY_BYTES:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive entry exceeds "
|
||||||
|
f"{_MAX_ARCHIVE_ENTRY_BYTES} bytes"
|
||||||
|
)
|
||||||
|
total_size += member.file_size
|
||||||
|
if total_size > _MAX_ARCHIVE_UNCOMPRESSED_BYTES:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive expands beyond "
|
||||||
|
f"{_MAX_ARCHIVE_UNCOMPRESSED_BYTES} bytes"
|
||||||
|
)
|
||||||
|
if member.file_size >= 1024 * 1024:
|
||||||
|
if member.compress_size <= 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive has an unsafe compression ratio"
|
||||||
|
)
|
||||||
|
ratio = member.file_size / member.compress_size
|
||||||
|
if ratio > _MAX_ARCHIVE_COMPRESSION_RATIO:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive has an unsafe compression ratio"
|
||||||
|
)
|
||||||
|
|
||||||
|
missing = required_members[file_format] - names
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"invalid {file_format.upper()} package: missing "
|
||||||
|
f"{', '.join(sorted(missing))}"
|
||||||
|
)
|
||||||
|
|
||||||
|
for member in members:
|
||||||
|
if member.is_dir() or not member.filename.lower().endswith((".xml", ".rels")):
|
||||||
|
continue
|
||||||
|
with archive.open(member) as stream:
|
||||||
|
prefix = stream.read(min(member.file_size, 1024 * 1024)).upper()
|
||||||
|
if b"<!DOCTYPE" in prefix or b"<!ENTITY" in prefix:
|
||||||
|
raise ValueError(
|
||||||
|
f"{file_format.upper()} archive contains unsupported active XML"
|
||||||
|
)
|
||||||
|
except zipfile.BadZipFile as exc:
|
||||||
|
raise ValueError(f"invalid {file_format.upper()} file: not an Office ZIP package") from exc
|
||||||
|
|
||||||
|
def _extract_docx_text(raw: bytes) -> str:
|
||||||
|
_validate_office_archive(raw, "docx")
|
||||||
|
try:
|
||||||
|
document = Document(io.BytesIO(raw))
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"invalid DOCX file: {exc}") from exc
|
||||||
|
|
||||||
|
parts: list[str] = []
|
||||||
|
total = 0
|
||||||
|
for child in document.element.body.iterchildren():
|
||||||
|
if isinstance(child, CT_P):
|
||||||
|
total = _append_bounded_text(parts, Paragraph(child, document).text, total)
|
||||||
|
continue
|
||||||
|
if isinstance(child, CT_Tbl):
|
||||||
|
table = Table(child, document)
|
||||||
|
for row in table.rows:
|
||||||
|
cells = [normalize_text(cell.text) for cell in row.cells]
|
||||||
|
total = _append_bounded_text(parts, "\t".join(cells), total)
|
||||||
|
return normalize_text("\n\n".join(parts))
|
||||||
|
|
||||||
|
def _presentation_shape_text(shape: Any) -> list[str]:
|
||||||
|
if getattr(shape, "has_table", False):
|
||||||
|
return [
|
||||||
|
"\t".join(normalize_text(cell.text) for cell in row.cells)
|
||||||
|
for row in shape.table.rows
|
||||||
|
]
|
||||||
|
if getattr(shape, "has_text_frame", False):
|
||||||
|
return [shape.text]
|
||||||
|
child_shapes = getattr(shape, "shapes", None)
|
||||||
|
if child_shapes is not None:
|
||||||
|
values: list[str] = []
|
||||||
|
for child in child_shapes:
|
||||||
|
values.extend(_presentation_shape_text(child))
|
||||||
|
return values
|
||||||
|
return []
|
||||||
|
|
||||||
|
def _extract_pptx_text(raw: bytes) -> str:
|
||||||
|
_validate_office_archive(raw, "pptx")
|
||||||
|
try:
|
||||||
|
presentation = Presentation(io.BytesIO(raw))
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"invalid PPTX file: {exc}") from exc
|
||||||
|
if len(presentation.slides) > _MAX_PRESENTATION_SLIDES:
|
||||||
|
raise ValueError(
|
||||||
|
f"PPTX contains too many slides (limit {_MAX_PRESENTATION_SLIDES})"
|
||||||
|
)
|
||||||
|
|
||||||
|
parts: list[str] = []
|
||||||
|
total = 0
|
||||||
|
for slide in presentation.slides:
|
||||||
|
slide_parts: list[str] = []
|
||||||
|
for shape in slide.shapes:
|
||||||
|
slide_parts.extend(_presentation_shape_text(shape))
|
||||||
|
total = _append_bounded_text(parts, "\n".join(slide_parts), total)
|
||||||
|
return normalize_text("\n\n".join(parts))
|
||||||
|
|
||||||
|
def _xml_local_name(tag: str) -> str:
|
||||||
|
return tag.rsplit("}", 1)[-1]
|
||||||
|
|
||||||
|
def _resolve_xlsx_relationship_target(
|
||||||
|
archive: zipfile.ZipFile,
|
||||||
|
target: str,
|
||||||
|
*,
|
||||||
|
source_part: str = "xl/workbook.xml",
|
||||||
|
) -> str:
|
||||||
|
"""按 OPC URI 规则解析内部关系目标,并保证结果仍位于 ZIP 根内。"""
|
||||||
|
|
||||||
|
raw_target = target.strip()
|
||||||
|
if (
|
||||||
|
raw_target != target
|
||||||
|
or not raw_target
|
||||||
|
or "\\" in raw_target
|
||||||
|
or any(unicodedata.category(char).startswith("C") for char in raw_target)
|
||||||
|
or re.search(r"%(?![0-9A-Fa-f]{2})", raw_target)
|
||||||
|
):
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
try:
|
||||||
|
parsed = urlsplit(raw_target)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path") from exc
|
||||||
|
if parsed.scheme or parsed.netloc or parsed.query or parsed.fragment:
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
if re.search(r"%(?:2[fF]|5[cC]|0{2})", parsed.path):
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
try:
|
||||||
|
decoded_path = unquote(parsed.path, encoding="utf-8", errors="strict")
|
||||||
|
except UnicodeDecodeError as exc:
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path") from exc
|
||||||
|
if not decoded_path or "\\" in decoded_path or "\x00" in decoded_path or "%" in decoded_path:
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
|
||||||
|
parts = [] if decoded_path.startswith("/") else list(PurePosixPath(source_part).parent.parts)
|
||||||
|
for part in decoded_path.lstrip("/").split("/"):
|
||||||
|
if part in {"", "."}:
|
||||||
|
continue
|
||||||
|
if part == "..":
|
||||||
|
if not parts:
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
parts.pop()
|
||||||
|
continue
|
||||||
|
if unicodedata.category(part[0]).startswith("C") or any(
|
||||||
|
unicodedata.category(char).startswith("C") for char in part
|
||||||
|
):
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
parts.append(part)
|
||||||
|
if not parts:
|
||||||
|
raise ValueError("XLSX workbook contains an unsafe worksheet path")
|
||||||
|
|
||||||
|
member_name = "/".join(parts)
|
||||||
|
try:
|
||||||
|
member = archive.getinfo(member_name)
|
||||||
|
except KeyError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet relationship target does not exist: {member_name}"
|
||||||
|
) from exc
|
||||||
|
if member.is_dir():
|
||||||
|
raise ValueError("XLSX worksheet relationship target must be a file")
|
||||||
|
return member_name
|
||||||
|
|
||||||
|
def _rewrite_xlsx_workbook_relationships(
|
||||||
|
raw: bytes,
|
||||||
|
replacements: Mapping[str, str],
|
||||||
|
) -> bytes:
|
||||||
|
"""把已验证的 worksheet Target 改为解析库稳定支持的包内绝对路径。"""
|
||||||
|
|
||||||
|
relationships_member = "xl/_rels/workbook.xml.rels"
|
||||||
|
output = io.BytesIO()
|
||||||
|
with zipfile.ZipFile(io.BytesIO(raw)) as source, zipfile.ZipFile(output, "w") as target:
|
||||||
|
target.comment = source.comment
|
||||||
|
for member in source.infolist():
|
||||||
|
if member.filename == relationships_member:
|
||||||
|
root = ET.fromstring(source.read(member))
|
||||||
|
pending = dict(replacements)
|
||||||
|
for element in root:
|
||||||
|
if _xml_local_name(element.tag) != "Relationship":
|
||||||
|
continue
|
||||||
|
relationship_id = element.attrib.get("Id")
|
||||||
|
if relationship_id in pending:
|
||||||
|
element.set("Target", pending.pop(relationship_id))
|
||||||
|
if pending:
|
||||||
|
raise ValueError(
|
||||||
|
"XLSX workbook relationship changed during normalization"
|
||||||
|
)
|
||||||
|
content = ET.tostring(root, encoding="utf-8", xml_declaration=True)
|
||||||
|
target.writestr(member, content)
|
||||||
|
continue
|
||||||
|
if member.is_dir():
|
||||||
|
target.writestr(member, b"")
|
||||||
|
continue
|
||||||
|
with source.open(member) as source_stream, target.open(
|
||||||
|
member,
|
||||||
|
"w",
|
||||||
|
force_zip64=True,
|
||||||
|
) as target_stream:
|
||||||
|
while chunk := source_stream.read(1024 * 1024):
|
||||||
|
target_stream.write(chunk)
|
||||||
|
return output.getvalue()
|
||||||
|
|
||||||
|
def _xlsx_sheet_merge_ranges(
|
||||||
|
raw: bytes,
|
||||||
|
) -> tuple[
|
||||||
|
dict[str, tuple[tuple[int, int, int, int], ...]],
|
||||||
|
dict[str, str],
|
||||||
|
]:
|
||||||
|
"""流式读取 XLSX 合并单元格,不把工作表 XML 整体载入内存。"""
|
||||||
|
|
||||||
|
relationship_namespace = (
|
||||||
|
"http://schemas.openxmlformats.org/officeDocument/2006/relationships"
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
with zipfile.ZipFile(io.BytesIO(raw)) as archive:
|
||||||
|
relationships: dict[str, tuple[str, str, str]] = {}
|
||||||
|
with archive.open("xl/_rels/workbook.xml.rels") as stream:
|
||||||
|
for _, element in ET.iterparse(stream, events=("end",)):
|
||||||
|
if _xml_local_name(element.tag) != "Relationship":
|
||||||
|
element.clear()
|
||||||
|
continue
|
||||||
|
relationship_id = element.attrib.get("Id")
|
||||||
|
target = element.attrib.get("Target")
|
||||||
|
target_mode = element.attrib.get("TargetMode", "Internal")
|
||||||
|
relationship_type = element.attrib.get("Type", "")
|
||||||
|
if relationship_id:
|
||||||
|
if relationship_id in relationships:
|
||||||
|
raise ValueError(
|
||||||
|
"XLSX workbook contains duplicate relationship identifiers"
|
||||||
|
)
|
||||||
|
relationships[relationship_id] = (
|
||||||
|
target or "",
|
||||||
|
target_mode,
|
||||||
|
relationship_type,
|
||||||
|
)
|
||||||
|
element.clear()
|
||||||
|
|
||||||
|
sheet_paths: dict[str, str] = {}
|
||||||
|
normalized_targets: dict[str, str] = {}
|
||||||
|
with archive.open("xl/workbook.xml") as stream:
|
||||||
|
for _, element in ET.iterparse(stream, events=("end",)):
|
||||||
|
if _xml_local_name(element.tag) != "sheet":
|
||||||
|
element.clear()
|
||||||
|
continue
|
||||||
|
title = element.attrib.get("name")
|
||||||
|
relationship_id = element.attrib.get(
|
||||||
|
f"{{{relationship_namespace}}}id"
|
||||||
|
)
|
||||||
|
if title and relationship_id:
|
||||||
|
relationship = relationships.get(relationship_id)
|
||||||
|
if relationship is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {title!r} references a missing relationship"
|
||||||
|
)
|
||||||
|
target, target_mode, relationship_type = relationship
|
||||||
|
if target_mode.strip().lower() != "internal":
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {title!r} uses an external relationship"
|
||||||
|
)
|
||||||
|
if not relationship_type.endswith("/worksheet"):
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {title!r} has an invalid relationship type"
|
||||||
|
)
|
||||||
|
member_name = _resolve_xlsx_relationship_target(
|
||||||
|
archive,
|
||||||
|
target,
|
||||||
|
)
|
||||||
|
sheet_paths[title] = member_name
|
||||||
|
canonical_target = f"/{member_name}"
|
||||||
|
if target != canonical_target:
|
||||||
|
normalized_targets[relationship_id] = canonical_target
|
||||||
|
element.clear()
|
||||||
|
|
||||||
|
result: dict[str, tuple[tuple[int, int, int, int], ...]] = {}
|
||||||
|
total_ranges = 0
|
||||||
|
for title, member_name in sheet_paths.items():
|
||||||
|
ranges: list[tuple[int, int, int, int]] = []
|
||||||
|
with archive.open(member_name) as stream:
|
||||||
|
for _, element in ET.iterparse(stream, events=("end",)):
|
||||||
|
if _xml_local_name(element.tag) != "mergeCell":
|
||||||
|
element.clear()
|
||||||
|
continue
|
||||||
|
reference = element.attrib.get("ref")
|
||||||
|
if reference:
|
||||||
|
try:
|
||||||
|
boundaries = range_boundaries(reference)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {title!r} has an invalid merged range"
|
||||||
|
) from exc
|
||||||
|
ranges.append(boundaries)
|
||||||
|
total_ranges += 1
|
||||||
|
if total_ranges > _MAX_WORKBOOK_MERGED_RANGES:
|
||||||
|
raise ValueError(
|
||||||
|
"XLSX workbook contains too many merged ranges "
|
||||||
|
f"(limit {_MAX_WORKBOOK_MERGED_RANGES})"
|
||||||
|
)
|
||||||
|
element.clear()
|
||||||
|
result[title] = tuple(ranges)
|
||||||
|
return result, normalized_targets
|
||||||
|
except (KeyError, ET.ParseError, zipfile.BadZipFile) as exc:
|
||||||
|
raise ValueError(f"invalid XLSX workbook structure: {exc}") from exc
|
||||||
|
|
||||||
|
def _xlsx_header_end_row(
|
||||||
|
first_row: int,
|
||||||
|
rows: Mapping[int, Sequence[Any]],
|
||||||
|
merged_ranges: Sequence[tuple[int, int, int, int]],
|
||||||
|
) -> int:
|
||||||
|
"""在固定深度内闭包表头合并关系,忽略越界或跨空行的可疑级联。"""
|
||||||
|
|
||||||
|
header_end = first_row
|
||||||
|
maximum_end = first_row + _MAX_WORKBOOK_HEADER_ROWS - 1
|
||||||
|
has_header_hierarchy = any(
|
||||||
|
max_column > min_column and min_row == first_row
|
||||||
|
for min_column, min_row, max_column, _ in merged_ranges
|
||||||
|
)
|
||||||
|
changed = True
|
||||||
|
while changed:
|
||||||
|
changed = False
|
||||||
|
for min_column, min_row, max_column, max_row in merged_ranges:
|
||||||
|
if min_row < first_row or min_row > header_end or max_row < first_row:
|
||||||
|
continue
|
||||||
|
if max_column == min_column and not has_header_hierarchy:
|
||||||
|
continue
|
||||||
|
candidate = max_row + 1 if max_column > min_column else max_row
|
||||||
|
if candidate <= header_end or candidate > maximum_end:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
max_column > min_column
|
||||||
|
and len(_xlsx_nonempty_values(rows.get(min_row, ()))) < 2
|
||||||
|
and len(_xlsx_nonempty_values(rows.get(candidate, ()))) < 2
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
if any(
|
||||||
|
not rows.get(row_number)
|
||||||
|
for row_number in range(header_end + 1, candidate + 1)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
header_end = candidate
|
||||||
|
changed = True
|
||||||
|
return header_end
|
||||||
|
|
||||||
|
def _xlsx_headers(
|
||||||
|
title: str,
|
||||||
|
rows: Mapping[int, Sequence[Any]],
|
||||||
|
first_row: int,
|
||||||
|
header_end: int,
|
||||||
|
merged_ranges: Sequence[tuple[int, int, int, int]],
|
||||||
|
) -> list[str]:
|
||||||
|
width = max((len(row) for row in rows.values()), default=0)
|
||||||
|
for min_column, min_row, max_column, max_row in merged_ranges:
|
||||||
|
if min_row <= header_end and max_row >= first_row:
|
||||||
|
width = max(width, max_column)
|
||||||
|
if width > _MAX_WORKBOOK_COLUMNS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {title!r} exceeds {_MAX_WORKBOOK_COLUMNS} columns"
|
||||||
|
)
|
||||||
|
|
||||||
|
matrix = [
|
||||||
|
list(rows.get(row_number, ())) + [None] * (width - len(rows.get(row_number, ())))
|
||||||
|
for row_number in range(first_row, header_end + 1)
|
||||||
|
]
|
||||||
|
for min_column, min_row, max_column, max_row in merged_ranges:
|
||||||
|
if min_row > header_end or max_row < first_row:
|
||||||
|
continue
|
||||||
|
source_row = max(first_row, min_row) - first_row
|
||||||
|
source_column = min_column - 1
|
||||||
|
source = matrix[source_row][source_column]
|
||||||
|
for row_number in range(max(first_row, min_row), min(header_end, max_row) + 1):
|
||||||
|
for column_number in range(min_column, max_column + 1):
|
||||||
|
matrix[row_number - first_row][column_number - 1] = source
|
||||||
|
|
||||||
|
headers: list[str] = []
|
||||||
|
for column in range(width):
|
||||||
|
components: list[str] = []
|
||||||
|
for row in matrix:
|
||||||
|
component = normalize_text(str(row[column] or ""))
|
||||||
|
if component and (not components or component != components[-1]):
|
||||||
|
components.append(component)
|
||||||
|
header = ".".join(components)
|
||||||
|
if not header:
|
||||||
|
raise ValueError(f"XLSX worksheet {title!r} contains an empty header")
|
||||||
|
headers.append(header)
|
||||||
|
if len(set(headers)) != len(headers):
|
||||||
|
raise ValueError(f"XLSX worksheet {title!r} contains duplicate headers")
|
||||||
|
return headers
|
||||||
|
|
||||||
|
def _xlsx_nonempty_values(row: Sequence[Any]) -> list[Any]:
|
||||||
|
return [value for value in row if value not in {None, ""}]
|
||||||
|
|
||||||
|
def _infer_xlsx_header_region(
|
||||||
|
title: str,
|
||||||
|
rows: Mapping[int, Sequence[Any]],
|
||||||
|
merged_ranges: Sequence[tuple[int, int, int, int]],
|
||||||
|
) -> tuple[int, int, list[str]]:
|
||||||
|
"""从有界前缀中推断表头,区分报表说明、多级表头和正文合并。"""
|
||||||
|
|
||||||
|
ordered_rows = sorted(rows)
|
||||||
|
if not ordered_rows:
|
||||||
|
return 0, 0, []
|
||||||
|
first_nonempty_row = ordered_rows[0]
|
||||||
|
horizontal_merge_rows = {
|
||||||
|
min_row
|
||||||
|
for min_column, min_row, max_column, _ in merged_ranges
|
||||||
|
if max_column > min_column
|
||||||
|
}
|
||||||
|
candidates: list[tuple[float, int, int, list[str]]] = []
|
||||||
|
for first_row in ordered_rows:
|
||||||
|
raw_values = _xlsx_nonempty_values(rows[first_row])
|
||||||
|
if not raw_values:
|
||||||
|
continue
|
||||||
|
if len(raw_values) < 2 and first_row in horizontal_merge_rows:
|
||||||
|
continue
|
||||||
|
|
||||||
|
header_end = _xlsx_header_end_row(first_row, rows, merged_ranges)
|
||||||
|
header_rows = {
|
||||||
|
row_number: rows[row_number]
|
||||||
|
for row_number in range(first_row, header_end + 1)
|
||||||
|
if row_number in rows
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
headers = _xlsx_headers(
|
||||||
|
title,
|
||||||
|
header_rows,
|
||||||
|
first_row,
|
||||||
|
header_end,
|
||||||
|
merged_ranges,
|
||||||
|
)
|
||||||
|
except ValueError:
|
||||||
|
if header_end == first_row:
|
||||||
|
continue
|
||||||
|
header_end = first_row
|
||||||
|
try:
|
||||||
|
headers = _xlsx_headers(
|
||||||
|
title,
|
||||||
|
{first_row: rows[first_row]},
|
||||||
|
first_row,
|
||||||
|
first_row,
|
||||||
|
merged_ranges,
|
||||||
|
)
|
||||||
|
except ValueError:
|
||||||
|
continue
|
||||||
|
|
||||||
|
data_rows = [rows[row_number] for row_number in ordered_rows if row_number > header_end]
|
||||||
|
if any(len(row) > len(headers) for row in data_rows):
|
||||||
|
continue
|
||||||
|
|
||||||
|
text_ratio = sum(isinstance(value, str) for value in raw_values) / len(raw_values)
|
||||||
|
score = text_ratio * 6 + min(len(raw_values), 4)
|
||||||
|
if data_rows:
|
||||||
|
first_data_values = _xlsx_nonempty_values(data_rows[0])
|
||||||
|
score += 4 * min(len(first_data_values), len(headers)) / len(headers)
|
||||||
|
if first_data_values:
|
||||||
|
score += (
|
||||||
|
2
|
||||||
|
* sum(
|
||||||
|
not isinstance(value, str)
|
||||||
|
for value in first_data_values
|
||||||
|
)
|
||||||
|
/ len(first_data_values)
|
||||||
|
)
|
||||||
|
first_label = normalize_text(str(raw_values[0]))
|
||||||
|
if len(raw_values) <= 2 and _XLSX_REPORT_METADATA_PATTERN.match(first_label):
|
||||||
|
score -= 8
|
||||||
|
candidates.append((score, first_row, header_end, headers))
|
||||||
|
|
||||||
|
if not candidates:
|
||||||
|
first_row = first_nonempty_row
|
||||||
|
headers = _xlsx_headers(
|
||||||
|
title,
|
||||||
|
{first_row: rows[first_row]},
|
||||||
|
first_row,
|
||||||
|
first_row,
|
||||||
|
(),
|
||||||
|
)
|
||||||
|
return first_row, first_row, headers
|
||||||
|
_, first_row, header_end, headers = max(
|
||||||
|
candidates,
|
||||||
|
key=lambda candidate: (candidate[0], -candidate[1]),
|
||||||
|
)
|
||||||
|
return first_row, header_end, headers
|
||||||
|
|
||||||
|
def _extract_xlsx_records(
|
||||||
|
raw: bytes,
|
||||||
|
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||||
|
_validate_office_archive(raw, "xlsx")
|
||||||
|
merged_by_sheet, normalized_targets = _xlsx_sheet_merge_ranges(raw)
|
||||||
|
workbook_raw = (
|
||||||
|
_rewrite_xlsx_workbook_relationships(raw, normalized_targets)
|
||||||
|
if normalized_targets
|
||||||
|
else raw
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
workbook = load_workbook(
|
||||||
|
io.BytesIO(workbook_raw),
|
||||||
|
read_only=True,
|
||||||
|
data_only=True,
|
||||||
|
keep_links=False,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"invalid XLSX file: {exc}") from exc
|
||||||
|
|
||||||
|
records: list[dict[str, Any]] = []
|
||||||
|
locators: list[dict[str, Any]] = []
|
||||||
|
total_cells = 0
|
||||||
|
try:
|
||||||
|
if len(workbook.worksheets) > _MAX_WORKBOOK_SHEETS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX contains too many worksheets (limit {_MAX_WORKBOOK_SHEETS})"
|
||||||
|
)
|
||||||
|
for sheet_index, worksheet in enumerate(workbook.worksheets):
|
||||||
|
reset_dimensions = getattr(worksheet, "reset_dimensions", None)
|
||||||
|
if callable(reset_dimensions):
|
||||||
|
reset_dimensions()
|
||||||
|
merged_ranges = merged_by_sheet.get(worksheet.title, ())
|
||||||
|
sheet_rows = 0
|
||||||
|
scanned_rows = 0
|
||||||
|
row_iterator = enumerate(
|
||||||
|
worksheet.iter_rows(values_only=True),
|
||||||
|
start=1,
|
||||||
|
)
|
||||||
|
buffered_rows: dict[int, Sequence[Any]] = {}
|
||||||
|
|
||||||
|
def normalized_row_values(
|
||||||
|
row: Sequence[Any],
|
||||||
|
sheet_title: str = worksheet.title,
|
||||||
|
) -> list[Any]:
|
||||||
|
values = list(row)
|
||||||
|
while values and values[-1] in {None, ""}:
|
||||||
|
values.pop()
|
||||||
|
if len(values) > _MAX_WORKBOOK_COLUMNS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {sheet_title!r} exceeds "
|
||||||
|
f"{_MAX_WORKBOOK_COLUMNS} columns"
|
||||||
|
)
|
||||||
|
return values
|
||||||
|
|
||||||
|
for row_number, row in row_iterator:
|
||||||
|
scanned_rows += 1
|
||||||
|
if scanned_rows > _MAX_WORKBOOK_SCANNED_ROWS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {worksheet.title!r} exceeds "
|
||||||
|
f"{_MAX_WORKBOOK_SCANNED_ROWS} scanned rows"
|
||||||
|
)
|
||||||
|
values = normalized_row_values(row)
|
||||||
|
if not values or all(value in {None, ""} for value in values):
|
||||||
|
continue
|
||||||
|
buffered_rows[row_number] = tuple(values)
|
||||||
|
if len(buffered_rows) >= _MAX_WORKBOOK_HEADER_SCAN_ROWS:
|
||||||
|
break
|
||||||
|
|
||||||
|
if not buffered_rows:
|
||||||
|
continue
|
||||||
|
_, header_end_row, headers = _infer_xlsx_header_region(
|
||||||
|
worksheet.title,
|
||||||
|
buffered_rows,
|
||||||
|
merged_ranges,
|
||||||
|
)
|
||||||
|
|
||||||
|
def append_record(
|
||||||
|
row_number: int,
|
||||||
|
values: Sequence[Any],
|
||||||
|
record_headers: Sequence[str] = tuple(headers),
|
||||||
|
locator_sheet_index: int = sheet_index,
|
||||||
|
sheet_title: str = worksheet.title,
|
||||||
|
) -> None:
|
||||||
|
nonlocal total_cells, sheet_rows
|
||||||
|
row_values = list(values)
|
||||||
|
if len(row_values) > len(record_headers):
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {sheet_title!r} has a row wider than its header"
|
||||||
|
)
|
||||||
|
row_values.extend([None] * (len(record_headers) - len(row_values)))
|
||||||
|
record = {
|
||||||
|
header: _normalize_spreadsheet_value(value)
|
||||||
|
for header, value in zip(record_headers, row_values, strict=True)
|
||||||
|
}
|
||||||
|
if not any(value not in {"", None} for value in record.values()):
|
||||||
|
return
|
||||||
|
sheet_record_index = sheet_rows
|
||||||
|
sheet_rows += 1
|
||||||
|
total_cells += len(record_headers)
|
||||||
|
if sheet_rows > _MAX_WORKBOOK_ROWS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {sheet_title!r} exceeds "
|
||||||
|
f"{_MAX_WORKBOOK_ROWS} data rows"
|
||||||
|
)
|
||||||
|
if total_cells > _MAX_WORKBOOK_CELLS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX workbook exceeds {_MAX_WORKBOOK_CELLS} populated cells"
|
||||||
|
)
|
||||||
|
records.append(record)
|
||||||
|
locators.append(
|
||||||
|
{
|
||||||
|
"kind": "xlsx",
|
||||||
|
"record_index": len(records),
|
||||||
|
"sheet_index": locator_sheet_index,
|
||||||
|
"sheet_name": sheet_title,
|
||||||
|
"row_number": row_number,
|
||||||
|
"sheet_record_index": sheet_record_index,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
for row_number, values in buffered_rows.items():
|
||||||
|
if row_number > header_end_row:
|
||||||
|
append_record(row_number, values)
|
||||||
|
|
||||||
|
for row_number, row in row_iterator:
|
||||||
|
scanned_rows += 1
|
||||||
|
if scanned_rows > _MAX_WORKBOOK_SCANNED_ROWS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {worksheet.title!r} exceeds "
|
||||||
|
f"{_MAX_WORKBOOK_SCANNED_ROWS} scanned rows"
|
||||||
|
)
|
||||||
|
values = normalized_row_values(row)
|
||||||
|
if not values or all(value in {None, ""} for value in values):
|
||||||
|
continue
|
||||||
|
append_record(row_number, values)
|
||||||
|
finally:
|
||||||
|
workbook.close()
|
||||||
|
return records, locators
|
||||||
299
backend/app/modules/data_process/algorithms/parsers/pdf.py
Normal file
299
backend/app/modules/data_process/algorithms/parsers/pdf.py
Normal file
@@ -0,0 +1,299 @@
|
|||||||
|
"""数据处理算法 - PDF 文档解析。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import io
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
from collections.abc import Sequence
|
||||||
|
from typing import Literal
|
||||||
|
|
||||||
|
from pypdf import PdfReader
|
||||||
|
|
||||||
|
from ..text_utils import normalize_text
|
||||||
|
from ..types import (
|
||||||
|
DocumentNoiseSpan,
|
||||||
|
PdfPageText,
|
||||||
|
_MAX_EXTRACTED_TEXT_CHARS,
|
||||||
|
_MAX_PDF_PAGES,
|
||||||
|
_PdfLine,
|
||||||
|
)
|
||||||
|
|
||||||
|
_PDF_PAGE_NUMBER_LINE_PATTERN = re.compile(
|
||||||
|
r"^(?:页次\s*)?(?:第\s*)?(?P<page>\d+)\s*页\s*"
|
||||||
|
r"(?:(?:[//]\s*)?共\s*(?P<total>\d+)\s*页)?$"
|
||||||
|
)
|
||||||
|
_PDF_FRACTION_PAGE_LINE_PATTERN = re.compile(
|
||||||
|
r"^[—–-]?\s*(?P<page>\d+)\s*[//]\s*(?P<total>\d+)\s*[—–-]?$"
|
||||||
|
)
|
||||||
|
_PDF_CLASSIFICATION_LABEL_PATTERN = re.compile(
|
||||||
|
r"^(?:(?:秘密等级|密级)\s*)?(?:商密|秘密|机密|绝密)"
|
||||||
|
r"\s*(?:[【\[((][^】\]))]{1,8}[】\]))])?$"
|
||||||
|
)
|
||||||
|
_TOC_TITLE_PATTERN = re.compile(r"^(?:目\s*录|contents)$", re.IGNORECASE)
|
||||||
|
_TOC_LEADER_ENTRY_PATTERN = re.compile(
|
||||||
|
r"(?:[..…·•]\s*){3,}\s*\d{1,4}\s*$"
|
||||||
|
)
|
||||||
|
_TOC_NUMBERED_ENTRY_PATTERN = re.compile(
|
||||||
|
r"^(?:第[\u3400-\u4dbf\u4e00-\u9fff]{1,12}章|附表\s*\d+|\d+(?:\.\d+)+)"
|
||||||
|
r"\s+.+\s+\d{1,4}\s*$",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
_MARGIN_TEMPLATE_KEYWORDS = (
|
||||||
|
"页",
|
||||||
|
"页次",
|
||||||
|
"版本",
|
||||||
|
"文件编码",
|
||||||
|
"秘密等级",
|
||||||
|
"密级",
|
||||||
|
"商密",
|
||||||
|
"confidential",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def extract_pdf_page_texts(raw: bytes) -> tuple[PdfPageText, ...]:
|
||||||
|
"""提取 PDF 各页文本,并保留与切片字符偏移一致的页范围。"""
|
||||||
|
|
||||||
|
if b"%PDF-" not in raw[:1024]:
|
||||||
|
raise ValueError("invalid PDF file: missing PDF header")
|
||||||
|
try:
|
||||||
|
reader = PdfReader(io.BytesIO(raw), strict=True)
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"invalid PDF file: {exc}") from exc
|
||||||
|
if reader.is_encrypted and not reader.decrypt(""):
|
||||||
|
raise ValueError("password-protected PDF files are not supported")
|
||||||
|
if len(reader.pages) > _MAX_PDF_PAGES:
|
||||||
|
raise ValueError(f"PDF contains too many pages (limit {_MAX_PDF_PAGES})")
|
||||||
|
|
||||||
|
pages: list[PdfPageText] = []
|
||||||
|
total = 0
|
||||||
|
has_text = False
|
||||||
|
for page_number, page in enumerate(reader.pages, start=1):
|
||||||
|
try:
|
||||||
|
text = normalize_text(page.extract_text() or "")
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(
|
||||||
|
f"failed to extract text from PDF page {page_number}: {exc}"
|
||||||
|
) from exc
|
||||||
|
if not text:
|
||||||
|
pages.append(
|
||||||
|
PdfPageText(
|
||||||
|
page_number=page_number,
|
||||||
|
text="",
|
||||||
|
source_start=total,
|
||||||
|
source_end=total,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
if has_text:
|
||||||
|
total += 2
|
||||||
|
start = total
|
||||||
|
total += len(text)
|
||||||
|
if total > _MAX_EXTRACTED_TEXT_CHARS:
|
||||||
|
raise ValueError(
|
||||||
|
f"extracted document text exceeds {_MAX_EXTRACTED_TEXT_CHARS} characters"
|
||||||
|
)
|
||||||
|
pages.append(
|
||||||
|
PdfPageText(
|
||||||
|
page_number=page_number,
|
||||||
|
text=text,
|
||||||
|
source_start=start,
|
||||||
|
source_end=total,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
has_text = True
|
||||||
|
if not has_text:
|
||||||
|
raise ValueError(
|
||||||
|
"PDF contains no extractable text; scanned or image-only PDF files are not supported"
|
||||||
|
)
|
||||||
|
return tuple(pages)
|
||||||
|
|
||||||
|
def _pdf_page_lines(page: PdfPageText) -> tuple[_PdfLine, ...]:
|
||||||
|
lines: list[_PdfLine] = []
|
||||||
|
local_offset = 0
|
||||||
|
for raw_line in page.text.splitlines(keepends=True):
|
||||||
|
content = raw_line.rstrip("\r\n")
|
||||||
|
leading = len(content) - len(content.lstrip())
|
||||||
|
trailing = len(content.rstrip())
|
||||||
|
text = content.strip()
|
||||||
|
if text:
|
||||||
|
lines.append(
|
||||||
|
_PdfLine(
|
||||||
|
text=text,
|
||||||
|
start=page.source_start + local_offset + leading,
|
||||||
|
end=page.source_start + local_offset + trailing,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
local_offset += len(raw_line)
|
||||||
|
return tuple(lines)
|
||||||
|
|
||||||
|
def _is_standalone_page_number(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
physical_page: int,
|
||||||
|
page_count: int,
|
||||||
|
) -> bool:
|
||||||
|
normalized = unicodedata.normalize("NFKC", text).strip()
|
||||||
|
match = _PDF_PAGE_NUMBER_LINE_PATTERN.fullmatch(normalized)
|
||||||
|
if match is None:
|
||||||
|
match = _PDF_FRACTION_PAGE_LINE_PATTERN.fullmatch(normalized)
|
||||||
|
if match is None or int(match.group("page")) != physical_page:
|
||||||
|
return False
|
||||||
|
total = match.groupdict().get("total")
|
||||||
|
return total is None or int(total) == page_count
|
||||||
|
|
||||||
|
def _margin_signature(text: str) -> str:
|
||||||
|
normalized = unicodedata.normalize("NFKC", text).casefold()
|
||||||
|
normalized = re.sub(r"\s+", " ", normalized).strip()
|
||||||
|
if any(keyword in normalized for keyword in _MARGIN_TEMPLATE_KEYWORDS):
|
||||||
|
normalized = re.sub(r"\d+", "#", normalized)
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
def _has_margin_metadata_keyword(text: str) -> bool:
|
||||||
|
normalized = unicodedata.normalize("NFKC", text).casefold()
|
||||||
|
return any(keyword in normalized for keyword in _MARGIN_TEMPLATE_KEYWORDS)
|
||||||
|
|
||||||
|
def _has_meaningful_margin_signature(signature: str) -> bool:
|
||||||
|
return len(re.sub(r"[#\W_]+", "", signature, flags=re.UNICODE)) >= 2
|
||||||
|
|
||||||
|
def _is_toc_leader_entry(text: str) -> bool:
|
||||||
|
return bool(_TOC_LEADER_ENTRY_PATTERN.search(text))
|
||||||
|
|
||||||
|
def _is_toc_numbered_entry(text: str) -> bool:
|
||||||
|
return bool(_TOC_NUMBERED_ENTRY_PATTERN.fullmatch(text))
|
||||||
|
|
||||||
|
def detect_pdf_document_noise(
|
||||||
|
pages: Sequence[PdfPageText],
|
||||||
|
) -> tuple[DocumentNoiseSpan, ...]:
|
||||||
|
"""识别 PDF 中的独立页码、重复页边内容和高置信目录。
|
||||||
|
|
||||||
|
规则只查看每页顶部 5 行和底部 3 行来推断页眉页脚;目录必须有
|
||||||
|
明显的点引导线密度,避免仅因正文中出现“目录”或章节标题而误删。
|
||||||
|
"""
|
||||||
|
|
||||||
|
page_lines = tuple(_pdf_page_lines(page) for page in pages)
|
||||||
|
detected: dict[tuple[int, int], DocumentNoiseSpan] = {}
|
||||||
|
|
||||||
|
def mark(
|
||||||
|
line: _PdfLine,
|
||||||
|
kind: Literal["page_number", "repeated_margin", "table_of_contents"],
|
||||||
|
) -> None:
|
||||||
|
detected.setdefault(
|
||||||
|
(line.start, line.end),
|
||||||
|
DocumentNoiseSpan(
|
||||||
|
start=line.start,
|
||||||
|
end=line.end,
|
||||||
|
kind=kind,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
for page, lines in zip(pages, page_lines, strict=True):
|
||||||
|
for line in lines:
|
||||||
|
if _is_standalone_page_number(
|
||||||
|
line.text,
|
||||||
|
physical_page=page.page_number,
|
||||||
|
page_count=len(pages),
|
||||||
|
):
|
||||||
|
mark(line, "page_number")
|
||||||
|
outer_margin_lines = (*lines[:2], *lines[-2:])
|
||||||
|
for line in outer_margin_lines:
|
||||||
|
if _PDF_CLASSIFICATION_LABEL_PATTERN.fullmatch(line.text):
|
||||||
|
mark(line, "repeated_margin")
|
||||||
|
|
||||||
|
# 只在三页及以上文档中推断通用页眉页脚,避免短文档误删。
|
||||||
|
if len(pages) >= 3:
|
||||||
|
signature_pages: dict[str, set[int]] = {}
|
||||||
|
candidate_lines: list[tuple[int, _PdfLine, str]] = []
|
||||||
|
for page_index, lines in enumerate(page_lines):
|
||||||
|
boundary_lines = (
|
||||||
|
*((line, index < 2) for index, line in enumerate(lines[:5])),
|
||||||
|
*((line, index < 2) for index, line in enumerate(reversed(lines[-3:]))),
|
||||||
|
)
|
||||||
|
seen_ranges: set[tuple[int, int]] = set()
|
||||||
|
for line, is_outer_margin in boundary_lines:
|
||||||
|
line_range = (line.start, line.end)
|
||||||
|
if (
|
||||||
|
line_range in seen_ranges
|
||||||
|
or line_range in detected
|
||||||
|
or len(line.text) > 160
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
seen_ranges.add(line_range)
|
||||||
|
if not is_outer_margin and not _has_margin_metadata_keyword(line.text):
|
||||||
|
continue
|
||||||
|
signature = _margin_signature(line.text)
|
||||||
|
if not _has_meaningful_margin_signature(signature):
|
||||||
|
continue
|
||||||
|
signature_pages.setdefault(signature, set()).add(page_index)
|
||||||
|
candidate_lines.append((page_index, line, signature))
|
||||||
|
minimum_pages = max(3, math.ceil(len(pages) * 0.3))
|
||||||
|
repeated_signatures = {
|
||||||
|
signature
|
||||||
|
for signature, matching_pages in signature_pages.items()
|
||||||
|
if len(matching_pages) >= minimum_pages
|
||||||
|
}
|
||||||
|
for _, line, signature in candidate_lines:
|
||||||
|
if signature in repeated_signatures:
|
||||||
|
mark(line, "repeated_margin")
|
||||||
|
|
||||||
|
# 先依据强证据判定目录页,再补充删除少量不带点引导线的编号目录项。
|
||||||
|
toc_active = False
|
||||||
|
for lines in page_lines:
|
||||||
|
content_lines = [
|
||||||
|
line for line in lines if (line.start, line.end) not in detected
|
||||||
|
]
|
||||||
|
leader_entries = [line for line in content_lines if _is_toc_leader_entry(line.text)]
|
||||||
|
titles = [line for line in content_lines if _TOC_TITLE_PATTERN.fullmatch(line.text)]
|
||||||
|
starts_toc = bool(titles and len(leader_entries) >= 2)
|
||||||
|
is_toc_dense = bool(
|
||||||
|
len(leader_entries) >= 3
|
||||||
|
and len(leader_entries) / max(1, len(content_lines)) >= 0.5
|
||||||
|
)
|
||||||
|
if not (starts_toc or (toc_active and is_toc_dense)):
|
||||||
|
toc_active = False
|
||||||
|
continue
|
||||||
|
toc_active = True
|
||||||
|
for line in content_lines:
|
||||||
|
if (
|
||||||
|
line in titles
|
||||||
|
or _is_toc_leader_entry(line.text)
|
||||||
|
or _is_toc_numbered_entry(line.text)
|
||||||
|
):
|
||||||
|
mark(line, "table_of_contents")
|
||||||
|
|
||||||
|
return tuple(sorted(detected.values(), key=lambda span: (span.start, span.end)))
|
||||||
|
|
||||||
|
def remove_document_noise(
|
||||||
|
text: str,
|
||||||
|
spans: Sequence[DocumentNoiseSpan],
|
||||||
|
*,
|
||||||
|
source_offset: int = 0,
|
||||||
|
) -> str:
|
||||||
|
"""按原文绝对偏移移除噪声,不改动调用方保留的原文及偏移。"""
|
||||||
|
|
||||||
|
text_end = source_offset + len(text)
|
||||||
|
intersections = sorted(
|
||||||
|
(
|
||||||
|
max(0, span.start - source_offset),
|
||||||
|
min(len(text), span.end - source_offset),
|
||||||
|
)
|
||||||
|
for span in spans
|
||||||
|
if span.start < text_end and span.end > source_offset
|
||||||
|
)
|
||||||
|
if not intersections:
|
||||||
|
return text
|
||||||
|
parts: list[str] = []
|
||||||
|
cursor = 0
|
||||||
|
for start, end in intersections:
|
||||||
|
if end <= cursor:
|
||||||
|
continue
|
||||||
|
if start > cursor:
|
||||||
|
parts.append(text[cursor:start])
|
||||||
|
cursor = end
|
||||||
|
parts.append(text[cursor:])
|
||||||
|
cleaned = normalize_text("".join(parts))
|
||||||
|
return re.sub(r"\n{3,}", "\n\n", cleaned)
|
||||||
|
|
||||||
|
def _extract_pdf_text(raw: bytes) -> str:
|
||||||
|
return "\n\n".join(page.text for page in extract_pdf_page_texts(raw) if page.text)
|
||||||
343
backend/app/modules/data_process/algorithms/quality.py
Normal file
343
backend/app/modules/data_process/algorithms/quality.py
Normal file
@@ -0,0 +1,343 @@
|
|||||||
|
"""数据处理算法 - 质量评分和去重。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
from collections import Counter
|
||||||
|
from collections.abc import Iterable, Mapping, Sequence
|
||||||
|
from copy import deepcopy
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from .text_utils import normalize_text
|
||||||
|
from .types import (
|
||||||
|
_MAX_ANOMALY_TEXT_CHARS,
|
||||||
|
_MOJIBAKE_MARKERS,
|
||||||
|
_TOKEN_PATTERN,
|
||||||
|
ProcessedStructuredRecord,
|
||||||
|
QualityScore,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def estimate_token_count(text: str) -> int:
|
||||||
|
"""粗略估计文本的 token 数量。"""
|
||||||
|
return len(_TOKEN_PATTERN.findall(text))
|
||||||
|
|
||||||
|
|
||||||
|
def content_quality_flags(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
min_chars: int = 20,
|
||||||
|
min_tokens: int = 5,
|
||||||
|
max_chars: int = _MAX_ANOMALY_TEXT_CHARS,
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
"""返回非结构化内容的确定性低质量原因。"""
|
||||||
|
|
||||||
|
if min_chars < 0 or min_tokens < 0 or max_chars <= 0:
|
||||||
|
raise ValueError("content quality limits must be non-negative")
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
if not normalized:
|
||||||
|
return ("empty_content",)
|
||||||
|
flags: list[str] = []
|
||||||
|
if len(normalized) < min_chars or estimate_token_count(normalized) < min_tokens:
|
||||||
|
flags.append("content_too_short")
|
||||||
|
if len(normalized) > max_chars:
|
||||||
|
flags.append("content_too_long")
|
||||||
|
if any(marker in normalized for marker in _MOJIBAKE_MARKERS):
|
||||||
|
flags.append("mojibake")
|
||||||
|
nonspace = [char for char in normalized if not char.isspace()]
|
||||||
|
if nonspace:
|
||||||
|
readable_ratio = sum(
|
||||||
|
char.isprintable()
|
||||||
|
and unicodedata.category(char) not in {"Co", "Cs", "Cn"}
|
||||||
|
for char in nonspace
|
||||||
|
) / len(nonspace)
|
||||||
|
if readable_ratio < 0.85:
|
||||||
|
flags.append("low_printable_ratio")
|
||||||
|
if len(nonspace) >= 100:
|
||||||
|
most_common = Counter(nonspace).most_common(1)[0][1]
|
||||||
|
if most_common / len(nonspace) > 0.9:
|
||||||
|
flags.append("repetitive_content")
|
||||||
|
return tuple(dict.fromkeys(flags))
|
||||||
|
|
||||||
|
def is_low_quality_content(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
min_chars: int = 20,
|
||||||
|
min_tokens: int = 5,
|
||||||
|
max_chars: int = _MAX_ANOMALY_TEXT_CHARS,
|
||||||
|
) -> bool:
|
||||||
|
"""判断内容是否命中任一低质量规则。"""
|
||||||
|
|
||||||
|
return bool(
|
||||||
|
content_quality_flags(
|
||||||
|
text,
|
||||||
|
min_chars=min_chars,
|
||||||
|
min_tokens=min_tokens,
|
||||||
|
max_chars=max_chars,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _deduplicate_structured_entries(
|
||||||
|
entries: Sequence[ProcessedStructuredRecord],
|
||||||
|
) -> list[ProcessedStructuredRecord]:
|
||||||
|
"""仅按整条 canonical JSON 稳定去重,避免误删同 ID 的更新记录。"""
|
||||||
|
|
||||||
|
# canonical_record_json 位于 structured_processing,延迟导入以断开循环依赖。
|
||||||
|
from .structured_processing import canonical_record_json
|
||||||
|
|
||||||
|
exact_seen: set[str] = set()
|
||||||
|
unique: list[ProcessedStructuredRecord] = []
|
||||||
|
for entry in entries:
|
||||||
|
record = entry.record
|
||||||
|
fingerprint = hashlib.sha256(canonical_record_json(record).encode("utf-8")).hexdigest()
|
||||||
|
if fingerprint in exact_seen:
|
||||||
|
continue
|
||||||
|
exact_seen.add(fingerprint)
|
||||||
|
unique.append(
|
||||||
|
ProcessedStructuredRecord(
|
||||||
|
entry.source_index,
|
||||||
|
deepcopy(dict(record)),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return unique
|
||||||
|
|
||||||
|
def deduplicate_structured_records(
|
||||||
|
records: Sequence[Mapping[str, Any]],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""仅按整条 canonical JSON 稳定去重。"""
|
||||||
|
|
||||||
|
entries = [
|
||||||
|
ProcessedStructuredRecord(index, deepcopy(dict(record)))
|
||||||
|
for index, record in enumerate(records)
|
||||||
|
]
|
||||||
|
return [entry.record for entry in _deduplicate_structured_entries(entries)]
|
||||||
|
|
||||||
|
def _near_duplicate_features(text: str, shingle_size: int) -> tuple[str, ...]:
|
||||||
|
if isinstance(shingle_size, bool) or not isinstance(shingle_size, int):
|
||||||
|
raise TypeError("shingle_size must be an integer")
|
||||||
|
if shingle_size <= 0:
|
||||||
|
raise ValueError("shingle_size must be greater than 0")
|
||||||
|
tokens = re.findall(
|
||||||
|
r"[\u3400-\u4dbf\u4e00-\u9fff]|[A-Za-z0-9_]+",
|
||||||
|
normalize_text(text).casefold(),
|
||||||
|
)
|
||||||
|
if not tokens:
|
||||||
|
return ()
|
||||||
|
if len(tokens) < shingle_size:
|
||||||
|
return ("\x1f".join(tokens),)
|
||||||
|
return tuple(
|
||||||
|
"\x1f".join(tokens[index : index + shingle_size])
|
||||||
|
for index in range(len(tokens) - shingle_size + 1)
|
||||||
|
)
|
||||||
|
|
||||||
|
def near_duplicate_fingerprint(text: str, *, shingle_size: int = 3) -> str:
|
||||||
|
"""生成 64 位 SimHash 指纹,用于低成本近重复候选筛选。"""
|
||||||
|
|
||||||
|
if isinstance(shingle_size, bool) or not isinstance(shingle_size, int):
|
||||||
|
raise TypeError("shingle_size must be an integer")
|
||||||
|
if shingle_size <= 0:
|
||||||
|
raise ValueError("shingle_size must be greater than 0")
|
||||||
|
features = Counter(_near_duplicate_features(text, shingle_size))
|
||||||
|
if not features:
|
||||||
|
return "0" * 16
|
||||||
|
vector = [0] * 64
|
||||||
|
for feature, weight in features.items():
|
||||||
|
digest = int.from_bytes(hashlib.sha256(feature.encode("utf-8")).digest()[:8], "big")
|
||||||
|
for bit in range(64):
|
||||||
|
vector[bit] += weight if digest & (1 << bit) else -weight
|
||||||
|
fingerprint = sum(1 << bit for bit, value in enumerate(vector) if value >= 0)
|
||||||
|
return f"{fingerprint:016x}"
|
||||||
|
|
||||||
|
def fingerprints_are_near_duplicate(
|
||||||
|
left: str,
|
||||||
|
right: str,
|
||||||
|
*,
|
||||||
|
max_hamming_distance: int = 3,
|
||||||
|
) -> bool:
|
||||||
|
"""比较两个 64 位十六进制 SimHash 指纹。"""
|
||||||
|
|
||||||
|
if isinstance(max_hamming_distance, bool) or not isinstance(max_hamming_distance, int):
|
||||||
|
raise TypeError("max_hamming_distance must be an integer")
|
||||||
|
if not 0 <= max_hamming_distance <= 64:
|
||||||
|
raise ValueError("max_hamming_distance must be in [0, 64]")
|
||||||
|
if not re.fullmatch(r"[0-9a-fA-F]{16}", left) or not re.fullmatch(
|
||||||
|
r"[0-9a-fA-F]{16}", right
|
||||||
|
):
|
||||||
|
raise ValueError("fingerprints must be 16-character hexadecimal strings")
|
||||||
|
distance = (int(left, 16) ^ int(right, 16)).bit_count()
|
||||||
|
return distance <= max_hamming_distance
|
||||||
|
|
||||||
|
def is_near_duplicate(
|
||||||
|
left: str,
|
||||||
|
right: str,
|
||||||
|
*,
|
||||||
|
shingle_size: int = 3,
|
||||||
|
similarity_threshold: float = 0.9,
|
||||||
|
max_hamming_distance: int = 3,
|
||||||
|
) -> bool:
|
||||||
|
"""结合词片 Jaccard 和 SimHash 判断两段内容是否近重复。"""
|
||||||
|
|
||||||
|
if isinstance(similarity_threshold, bool) or not isinstance(
|
||||||
|
similarity_threshold, (int, float)
|
||||||
|
):
|
||||||
|
raise TypeError("similarity_threshold must be a number")
|
||||||
|
if not 0 <= similarity_threshold <= 1:
|
||||||
|
raise ValueError("similarity_threshold must be in [0, 1]")
|
||||||
|
if isinstance(max_hamming_distance, bool) or not isinstance(max_hamming_distance, int):
|
||||||
|
raise TypeError("max_hamming_distance must be an integer")
|
||||||
|
if not 0 <= max_hamming_distance <= 64:
|
||||||
|
raise ValueError("max_hamming_distance must be in [0, 64]")
|
||||||
|
left_normalized = normalize_text(left)
|
||||||
|
right_normalized = normalize_text(right)
|
||||||
|
if not left_normalized or not right_normalized:
|
||||||
|
return left_normalized == right_normalized
|
||||||
|
if left_normalized.casefold() == right_normalized.casefold():
|
||||||
|
return True
|
||||||
|
left_features = set(_near_duplicate_features(left_normalized, shingle_size))
|
||||||
|
right_features = set(_near_duplicate_features(right_normalized, shingle_size))
|
||||||
|
union = left_features | right_features
|
||||||
|
similarity = len(left_features & right_features) / len(union) if union else 1.0
|
||||||
|
if similarity >= similarity_threshold:
|
||||||
|
return True
|
||||||
|
return fingerprints_are_near_duplicate(
|
||||||
|
near_duplicate_fingerprint(left_normalized, shingle_size=shingle_size),
|
||||||
|
near_duplicate_fingerprint(right_normalized, shingle_size=shingle_size),
|
||||||
|
max_hamming_distance=max_hamming_distance,
|
||||||
|
)
|
||||||
|
|
||||||
|
def record_fingerprint(record: Mapping[str, Any]) -> str:
|
||||||
|
"""计算与字典键顺序无关的稳定记录指纹。"""
|
||||||
|
|
||||||
|
canonical = {
|
||||||
|
"instruction": normalize_text(str(record.get("instruction") or "")),
|
||||||
|
"input": normalize_text(str(record.get("input") or "")),
|
||||||
|
"output": normalize_text(str(record.get("output") or "")),
|
||||||
|
}
|
||||||
|
raw = json.dumps(canonical, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||||
|
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def _readability_score(text: str) -> float:
|
||||||
|
if not text:
|
||||||
|
return 0.0
|
||||||
|
nonspace = [char for char in text if not char.isspace()]
|
||||||
|
if not nonspace:
|
||||||
|
return 0.0
|
||||||
|
printable_ratio = sum(char.isprintable() for char in nonspace) / len(nonspace)
|
||||||
|
useful_ratio = sum(
|
||||||
|
char.isalnum() or "\u3400" <= char <= "\u9fff" or unicodedata.category(char).startswith("P")
|
||||||
|
for char in nonspace
|
||||||
|
) / len(nonspace)
|
||||||
|
return round(100 * (0.65 * printable_ratio + 0.35 * useful_ratio), 2)
|
||||||
|
|
||||||
|
|
||||||
|
def _internal_duplicate_score(text: str) -> float:
|
||||||
|
units = [unit.strip().lower() for unit in re.split(r"[\n。!?!?;;]+", text) if unit.strip()]
|
||||||
|
if len(units) <= 1:
|
||||||
|
return 100.0
|
||||||
|
return round(100 * len(set(units)) / len(units), 2)
|
||||||
|
|
||||||
|
|
||||||
|
def _source_relevance_score(record: Mapping[str, Any], source_content: str) -> float:
|
||||||
|
"""估算结果与来源文本的词元覆盖率。
|
||||||
|
|
||||||
|
这是无外部模型依赖、可重复的首版评分。没有来源文本(例如人工新增结果)
|
||||||
|
时不扣分;存在来源时,以结果中的有效词元被来源覆盖的比例计分。
|
||||||
|
"""
|
||||||
|
|
||||||
|
source = normalize_text(source_content)
|
||||||
|
if not source:
|
||||||
|
return 100.0
|
||||||
|
candidate = normalize_text(
|
||||||
|
"\n".join(
|
||||||
|
str(record.get(field) or "") for field in ("instruction", "input", "output")
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def semantic_tokens(text: str) -> set[str]:
|
||||||
|
return {
|
||||||
|
token.lower()
|
||||||
|
for token in _TOKEN_PATTERN.findall(text)
|
||||||
|
if token.isalnum() or "\u3400" <= token <= "\u9fff"
|
||||||
|
}
|
||||||
|
|
||||||
|
source_tokens = semantic_tokens(source)
|
||||||
|
candidate_tokens = semantic_tokens(candidate)
|
||||||
|
if not candidate_tokens:
|
||||||
|
return 0.0
|
||||||
|
if not source_tokens:
|
||||||
|
return 0.0
|
||||||
|
return round(100 * len(candidate_tokens & source_tokens) / len(candidate_tokens), 2)
|
||||||
|
|
||||||
|
|
||||||
|
def score_quality(
|
||||||
|
record: Mapping[str, Any],
|
||||||
|
*,
|
||||||
|
min_output_length: int = 20,
|
||||||
|
source_content: str = "",
|
||||||
|
known_fingerprints: Iterable[str] = (),
|
||||||
|
threshold: float = 60.0,
|
||||||
|
) -> QualityScore:
|
||||||
|
"""按完整性、长度、可读性、来源相关性和重复度计算质量分。"""
|
||||||
|
|
||||||
|
if min_output_length <= 0:
|
||||||
|
raise ValueError("min_output_length must be greater than 0")
|
||||||
|
if not 0 <= threshold <= 100:
|
||||||
|
raise ValueError("threshold must be in [0, 100]")
|
||||||
|
|
||||||
|
instruction = normalize_text(str(record.get("instruction") or ""))
|
||||||
|
input_text = normalize_text(str(record.get("input") or ""))
|
||||||
|
output = normalize_text(str(record.get("output") or ""))
|
||||||
|
flags: list[str] = []
|
||||||
|
|
||||||
|
completeness = 100.0
|
||||||
|
if not instruction:
|
||||||
|
completeness -= 50
|
||||||
|
flags.append("missing_instruction")
|
||||||
|
if not output:
|
||||||
|
completeness -= 50
|
||||||
|
flags.append("missing_output")
|
||||||
|
|
||||||
|
output_length = len(output)
|
||||||
|
length_score = round(min(100.0, output_length / min_output_length * 100), 2)
|
||||||
|
if output_length < min_output_length:
|
||||||
|
flags.append("output_too_short")
|
||||||
|
|
||||||
|
readability = _readability_score("\n".join((instruction, input_text, output)))
|
||||||
|
if readability < 70:
|
||||||
|
flags.append("low_readability")
|
||||||
|
|
||||||
|
relevance = _source_relevance_score(record, source_content)
|
||||||
|
if source_content and relevance < 30:
|
||||||
|
flags.append("low_source_relevance")
|
||||||
|
|
||||||
|
fingerprint = record_fingerprint(record)
|
||||||
|
known = set(known_fingerprints)
|
||||||
|
duplicate = 0.0 if fingerprint in known else _internal_duplicate_score(output)
|
||||||
|
if duplicate == 0:
|
||||||
|
flags.append("duplicate_record")
|
||||||
|
elif duplicate < 70:
|
||||||
|
flags.append("repetitive_output")
|
||||||
|
|
||||||
|
overall = round(
|
||||||
|
completeness * 0.35
|
||||||
|
+ length_score * 0.20
|
||||||
|
+ readability * 0.20
|
||||||
|
+ relevance * 0.15
|
||||||
|
+ duplicate * 0.10,
|
||||||
|
2,
|
||||||
|
)
|
||||||
|
hard_valid = bool(instruction and output)
|
||||||
|
return QualityScore(
|
||||||
|
overall=overall,
|
||||||
|
completeness=completeness,
|
||||||
|
length=length_score,
|
||||||
|
readability=readability,
|
||||||
|
relevance=relevance,
|
||||||
|
duplicate=duplicate,
|
||||||
|
is_valid=hard_valid and overall >= threshold,
|
||||||
|
flags=tuple(flags),
|
||||||
|
fingerprint=fingerprint,
|
||||||
|
)
|
||||||
@@ -0,0 +1,809 @@
|
|||||||
|
"""数据处理算法 - 结构化数据处理。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
from collections.abc import Iterable, Mapping, Sequence
|
||||||
|
from copy import deepcopy
|
||||||
|
from datetime import date, datetime, time
|
||||||
|
from decimal import Decimal
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from .parsers.json_parser import _extract_structured_records_with_locators
|
||||||
|
from .quality import (
|
||||||
|
_deduplicate_structured_entries,
|
||||||
|
content_quality_flags,
|
||||||
|
estimate_token_count,
|
||||||
|
)
|
||||||
|
from .text_utils import _normalize_field_name, normalize_text, structured_json_dumps
|
||||||
|
from .transforms import stable_split_assignments
|
||||||
|
from .types import (
|
||||||
|
_CHINESE_NAME_CONTEXT_PATTERN,
|
||||||
|
_EMAIL_PATTERN,
|
||||||
|
_ENGLISH_NAME_CONTEXT_PATTERN,
|
||||||
|
_ID_CARD_PATTERN,
|
||||||
|
_IDENTITY_FIELD_PATTERN,
|
||||||
|
_MAX_STRUCTURED_DEPTH,
|
||||||
|
_MAX_STRUCTURED_FIELDS,
|
||||||
|
_NAME_FIELD_NAMES,
|
||||||
|
_PHONE_PATTERN,
|
||||||
|
_STRUCTURED_OPTIONS,
|
||||||
|
MAX_QA_PAIRS_PER_ITEM,
|
||||||
|
ProcessedStructuredRecord,
|
||||||
|
StructuredPreprocessOption,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _canonical_value(value: Any) -> Any:
|
||||||
|
if value is None or isinstance(value, (bool, int)):
|
||||||
|
return value
|
||||||
|
if isinstance(value, Decimal):
|
||||||
|
if not value.is_finite():
|
||||||
|
raise ValueError("non-finite JSON number is not allowed")
|
||||||
|
return value
|
||||||
|
if isinstance(value, float):
|
||||||
|
if not math.isfinite(value):
|
||||||
|
raise ValueError("non-finite JSON number is not allowed")
|
||||||
|
return value
|
||||||
|
if isinstance(value, str):
|
||||||
|
return normalize_text(value)
|
||||||
|
if isinstance(value, (datetime, date, time)):
|
||||||
|
return value.isoformat()
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
normalized: dict[str, Any] = {}
|
||||||
|
for key, item in sorted(value.items(), key=lambda pair: str(pair[0])):
|
||||||
|
normalized_key = normalize_text(str(key))
|
||||||
|
if not normalized_key:
|
||||||
|
raise ValueError("structured record contains an empty field name")
|
||||||
|
if normalized_key in normalized:
|
||||||
|
raise ValueError(
|
||||||
|
f"structured record fields collide after normalization: {normalized_key}"
|
||||||
|
)
|
||||||
|
normalized[normalized_key] = _canonical_value(item)
|
||||||
|
return normalized
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return [_canonical_value(item) for item in value]
|
||||||
|
if isinstance(value, (set, frozenset)):
|
||||||
|
items = [_canonical_value(item) for item in value]
|
||||||
|
return sorted(
|
||||||
|
items,
|
||||||
|
key=lambda item: structured_json_dumps(item, sort_keys=True),
|
||||||
|
)
|
||||||
|
if isinstance(value, (bytes, bytearray, memoryview)):
|
||||||
|
return bytes(value).hex()
|
||||||
|
return normalize_text(str(value))
|
||||||
|
|
||||||
|
|
||||||
|
def _is_empty_value(value: Any) -> bool:
|
||||||
|
if value is None:
|
||||||
|
return True
|
||||||
|
if isinstance(value, str):
|
||||||
|
return not normalize_text(value)
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return not value or all(_is_empty_value(item) for item in value.values())
|
||||||
|
if isinstance(value, (list, tuple, set, frozenset)):
|
||||||
|
return not value or all(_is_empty_value(item) for item in value)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _is_name_field(field: Any) -> bool:
|
||||||
|
raw_field = normalize_text(str(field))
|
||||||
|
if not raw_field:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 只匹配明确表示自然人姓名的字段,避免将 table_name、product_name、
|
||||||
|
# chinese_name 等业务名称或元数据字段误判为个人敏感信息。
|
||||||
|
if _normalize_field_name(raw_field, "snake_case") in _NAME_FIELD_NAMES:
|
||||||
|
return True
|
||||||
|
|
||||||
|
# detect_structure 会使用点号生成扁平化路径(例如 profile.name);此时仅
|
||||||
|
# 判断最后一个路径段,不能退回到宽泛的 ``*_name`` 后缀匹配。
|
||||||
|
if "." not in raw_field:
|
||||||
|
return False
|
||||||
|
leaf_field = raw_field.rsplit(".", 1)[-1]
|
||||||
|
return _normalize_field_name(leaf_field, "snake_case") in _NAME_FIELD_NAMES
|
||||||
|
|
||||||
|
|
||||||
|
def _embedded_structure(value: Any) -> Any:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
return value
|
||||||
|
candidate = value.strip()
|
||||||
|
if not candidate or candidate[0] not in "[{":
|
||||||
|
return value
|
||||||
|
try:
|
||||||
|
parsed = json.loads(candidate)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return value
|
||||||
|
return parsed if isinstance(parsed, (Mapping, list)) else value
|
||||||
|
|
||||||
|
|
||||||
|
def _structured_options(options: Iterable[str] | Mapping[str, Any]) -> set[str]:
|
||||||
|
if isinstance(options, str):
|
||||||
|
raise TypeError("options must be an iterable or mapping of option names")
|
||||||
|
if isinstance(options, Mapping):
|
||||||
|
enabled = {str(key) for key, value in options.items() if bool(value)}
|
||||||
|
else:
|
||||||
|
enabled = {str(option) for option in options}
|
||||||
|
unknown = enabled - _STRUCTURED_OPTIONS
|
||||||
|
if unknown:
|
||||||
|
raise ValueError(f"unsupported structured preprocess options: {', '.join(sorted(unknown))}")
|
||||||
|
return enabled
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_invalid_structured_entries(
|
||||||
|
entries: Sequence[ProcessedStructuredRecord],
|
||||||
|
) -> list[ProcessedStructuredRecord]:
|
||||||
|
if not entries:
|
||||||
|
return []
|
||||||
|
fields: list[str] = []
|
||||||
|
for entry in entries:
|
||||||
|
record = entry.record
|
||||||
|
for field in record:
|
||||||
|
if field not in fields:
|
||||||
|
fields.append(field)
|
||||||
|
active_fields = [
|
||||||
|
field
|
||||||
|
for field in fields
|
||||||
|
if any(not _is_empty_value(entry.record.get(field)) for entry in entries)
|
||||||
|
]
|
||||||
|
if not active_fields:
|
||||||
|
return []
|
||||||
|
cleaned: list[ProcessedStructuredRecord] = []
|
||||||
|
for entry in entries:
|
||||||
|
record = entry.record
|
||||||
|
values = {field: deepcopy(record.get(field)) for field in active_fields}
|
||||||
|
# 清洗只依据整行是否为空。外键、父级 ID 等字段天然允许为空,不能
|
||||||
|
# 因为字段名以 *_id 结尾就把它们全部提升为联合必填项。
|
||||||
|
if all(_is_empty_value(value) for value in values.values()):
|
||||||
|
continue
|
||||||
|
cleaned.append(ProcessedStructuredRecord(entry.source_index, values))
|
||||||
|
return cleaned
|
||||||
|
|
||||||
|
|
||||||
|
def _percentile(values: Sequence[float], fraction: float) -> float:
|
||||||
|
if not values:
|
||||||
|
raise ValueError("cannot calculate a percentile of an empty sequence")
|
||||||
|
ordered = sorted(values)
|
||||||
|
position = (len(ordered) - 1) * fraction
|
||||||
|
lower = math.floor(position)
|
||||||
|
upper = math.ceil(position)
|
||||||
|
if lower == upper:
|
||||||
|
return ordered[lower]
|
||||||
|
weight = position - lower
|
||||||
|
return ordered[lower] * (1 - weight) + ordered[upper] * weight
|
||||||
|
|
||||||
|
|
||||||
|
def _filter_anomalous_structured_entries(
|
||||||
|
entries: Sequence[ProcessedStructuredRecord],
|
||||||
|
*,
|
||||||
|
iqr_multiplier: float = 1.5,
|
||||||
|
) -> list[ProcessedStructuredRecord]:
|
||||||
|
"""按字段级数值 IQR、乱码和极端文本长度过滤异常记录。"""
|
||||||
|
|
||||||
|
if iqr_multiplier <= 0:
|
||||||
|
raise ValueError("iqr_multiplier must be greater than 0")
|
||||||
|
numeric_values: dict[str, list[float]] = {}
|
||||||
|
text_lengths: dict[str, list[float]] = {}
|
||||||
|
for entry in entries:
|
||||||
|
record = entry.record
|
||||||
|
for field, value in record.items():
|
||||||
|
if (
|
||||||
|
isinstance(value, (int, float))
|
||||||
|
and not isinstance(value, bool)
|
||||||
|
and not _IDENTITY_FIELD_PATTERN.search(
|
||||||
|
_normalize_field_name(field, "snake_case")
|
||||||
|
)
|
||||||
|
):
|
||||||
|
number = float(value)
|
||||||
|
if math.isfinite(number):
|
||||||
|
numeric_values.setdefault(field, []).append(number)
|
||||||
|
elif isinstance(value, str) and value:
|
||||||
|
text_lengths.setdefault(field, []).append(float(len(value)))
|
||||||
|
|
||||||
|
numeric_bounds: dict[str, tuple[float, float]] = {}
|
||||||
|
for field, values in numeric_values.items():
|
||||||
|
# 小样本不做统计异常判断,避免把合法长尾值误删。
|
||||||
|
if len(values) < 8:
|
||||||
|
continue
|
||||||
|
first_quartile = _percentile(values, 0.25)
|
||||||
|
third_quartile = _percentile(values, 0.75)
|
||||||
|
spread = third_quartile - first_quartile
|
||||||
|
numeric_bounds[field] = (
|
||||||
|
first_quartile - iqr_multiplier * spread,
|
||||||
|
third_quartile + iqr_multiplier * spread,
|
||||||
|
)
|
||||||
|
|
||||||
|
text_upper_bounds: dict[str, float] = {}
|
||||||
|
for field, lengths in text_lengths.items():
|
||||||
|
if len(lengths) < 8:
|
||||||
|
continue
|
||||||
|
first_quartile = _percentile(lengths, 0.25)
|
||||||
|
third_quartile = _percentile(lengths, 0.75)
|
||||||
|
spread = third_quartile - first_quartile
|
||||||
|
text_upper_bounds[field] = max(512.0, third_quartile + 3 * spread)
|
||||||
|
|
||||||
|
accepted: list[ProcessedStructuredRecord] = []
|
||||||
|
for entry in entries:
|
||||||
|
record = entry.record
|
||||||
|
anomalous = False
|
||||||
|
for field, value in record.items():
|
||||||
|
if (
|
||||||
|
isinstance(value, (int, float))
|
||||||
|
and not isinstance(value, bool)
|
||||||
|
and not _IDENTITY_FIELD_PATTERN.search(
|
||||||
|
_normalize_field_name(field, "snake_case")
|
||||||
|
)
|
||||||
|
):
|
||||||
|
number = float(value)
|
||||||
|
if not math.isfinite(number):
|
||||||
|
anomalous = True
|
||||||
|
break
|
||||||
|
bounds = numeric_bounds.get(field)
|
||||||
|
if bounds and not bounds[0] <= number <= bounds[1]:
|
||||||
|
anomalous = True
|
||||||
|
break
|
||||||
|
if isinstance(value, str):
|
||||||
|
flags = content_quality_flags(value, min_chars=0, min_tokens=0)
|
||||||
|
if {"content_too_long", "mojibake", "low_printable_ratio"} & set(flags):
|
||||||
|
anomalous = True
|
||||||
|
break
|
||||||
|
upper_bound = text_upper_bounds.get(field)
|
||||||
|
if upper_bound is not None and len(value) > upper_bound:
|
||||||
|
anomalous = True
|
||||||
|
break
|
||||||
|
if not anomalous:
|
||||||
|
accepted.append(
|
||||||
|
ProcessedStructuredRecord(
|
||||||
|
entry.source_index,
|
||||||
|
deepcopy(dict(record)),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return accepted
|
||||||
|
|
||||||
|
|
||||||
|
def _preview_content(item: Mapping[str, Any]) -> str:
|
||||||
|
for field in ("edited_content", "editedContent", "original_content", "originalContent", "content"):
|
||||||
|
value = item.get(field)
|
||||||
|
if value is not None:
|
||||||
|
return normalize_text(str(value))
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _standard_fields(content: str) -> tuple[str, str, str]:
|
||||||
|
if not content:
|
||||||
|
return "", "", ""
|
||||||
|
|
||||||
|
try:
|
||||||
|
payload = json.loads(content)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
payload = None
|
||||||
|
if isinstance(payload, Mapping):
|
||||||
|
instruction = next(
|
||||||
|
(
|
||||||
|
str(payload[key])
|
||||||
|
for key in ("instruction", "question", "prompt")
|
||||||
|
if payload.get(key) is not None
|
||||||
|
),
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
input_text = next(
|
||||||
|
(str(payload[key]) for key in ("input", "context") if payload.get(key) is not None),
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
output = next(
|
||||||
|
(str(payload[key]) for key in ("output", "answer", "response") if payload.get(key) is not None),
|
||||||
|
"",
|
||||||
|
)
|
||||||
|
if instruction or output:
|
||||||
|
return normalize_text(instruction), normalize_text(input_text), normalize_text(output)
|
||||||
|
|
||||||
|
question_answer = re.match(
|
||||||
|
r"^\s*(?:问|question)\s*[::]\s*(.+?)(?:\n|\r\n?)\s*(?:答|answer)\s*[::]\s*(.+)\s*$",
|
||||||
|
content,
|
||||||
|
flags=re.IGNORECASE | re.DOTALL,
|
||||||
|
)
|
||||||
|
if question_answer:
|
||||||
|
return normalize_text(question_answer.group(1)), "", normalize_text(question_answer.group(2))
|
||||||
|
|
||||||
|
lines = [line.strip() for line in content.splitlines() if line.strip()]
|
||||||
|
first_line = re.sub(r"^(?:问|question)\s*[::]\s*", "", lines[0], flags=re.IGNORECASE)
|
||||||
|
output = normalize_text("\n".join(lines[1:])) if len(lines) > 1 else normalize_text(content)
|
||||||
|
return normalize_text(first_line), "", output
|
||||||
|
|
||||||
|
|
||||||
|
def _protected_markdown_ranges(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
preserve_code_blocks: bool,
|
||||||
|
preserve_tables: bool,
|
||||||
|
preserve_lists: bool,
|
||||||
|
) -> list[tuple[int, int]]:
|
||||||
|
"""找出不应从中间切开的 Markdown 代码块、表格和连续列表。"""
|
||||||
|
|
||||||
|
lines: list[tuple[int, int, str]] = []
|
||||||
|
cursor = 0
|
||||||
|
for raw_line in text.splitlines(keepends=True):
|
||||||
|
end = cursor + len(raw_line)
|
||||||
|
lines.append((cursor, end, raw_line.rstrip("\r\n")))
|
||||||
|
cursor = end
|
||||||
|
if cursor < len(text) or not lines:
|
||||||
|
lines.append((cursor, len(text), text[cursor:]))
|
||||||
|
|
||||||
|
ranges: list[tuple[int, int]] = []
|
||||||
|
code_line_indexes: set[int] = set()
|
||||||
|
if preserve_code_blocks:
|
||||||
|
open_block: tuple[int, str, int] | None = None
|
||||||
|
for index, (start, end, content) in enumerate(lines):
|
||||||
|
fence = re.match(r"^\s*(`{3,}|~{3,})", content)
|
||||||
|
if not fence:
|
||||||
|
continue
|
||||||
|
marker = fence.group(1)[0]
|
||||||
|
length = len(fence.group(1))
|
||||||
|
if open_block is None:
|
||||||
|
open_block = (index, marker, length)
|
||||||
|
continue
|
||||||
|
first_index, open_marker, open_length = open_block
|
||||||
|
if marker == open_marker and length >= open_length:
|
||||||
|
ranges.append((lines[first_index][0], end))
|
||||||
|
code_line_indexes.update(range(first_index, index + 1))
|
||||||
|
open_block = None
|
||||||
|
if open_block is not None:
|
||||||
|
first_index = open_block[0]
|
||||||
|
ranges.append((lines[first_index][0], len(text)))
|
||||||
|
code_line_indexes.update(range(first_index, len(lines)))
|
||||||
|
|
||||||
|
if preserve_tables:
|
||||||
|
index = 0
|
||||||
|
while index + 1 < len(lines):
|
||||||
|
if index in code_line_indexes:
|
||||||
|
index += 1
|
||||||
|
continue
|
||||||
|
header = lines[index][2].strip()
|
||||||
|
separator = lines[index + 1][2].strip().strip("|")
|
||||||
|
cells = [cell.strip() for cell in separator.split("|")]
|
||||||
|
if (
|
||||||
|
"|" not in header
|
||||||
|
or len(cells) < 2
|
||||||
|
or not all(re.fullmatch(r":?-{3,}:?", cell) for cell in cells)
|
||||||
|
):
|
||||||
|
index += 1
|
||||||
|
continue
|
||||||
|
end_index = index + 1
|
||||||
|
while (
|
||||||
|
end_index + 1 < len(lines)
|
||||||
|
and end_index + 1 not in code_line_indexes
|
||||||
|
and lines[end_index + 1][2].strip()
|
||||||
|
and "|" in lines[end_index + 1][2]
|
||||||
|
):
|
||||||
|
end_index += 1
|
||||||
|
ranges.append((lines[index][0], lines[end_index][1]))
|
||||||
|
index = end_index + 1
|
||||||
|
|
||||||
|
if preserve_lists:
|
||||||
|
list_pattern = re.compile(r"^\s*(?:[-+*]|\d+[.)])\s+\S")
|
||||||
|
continuation_pattern = re.compile(r"^\s{2,}\S")
|
||||||
|
index = 0
|
||||||
|
while index < len(lines):
|
||||||
|
if index in code_line_indexes or not list_pattern.match(lines[index][2]):
|
||||||
|
index += 1
|
||||||
|
continue
|
||||||
|
end_index = index
|
||||||
|
item_count = 1
|
||||||
|
while end_index + 1 < len(lines) and end_index + 1 not in code_line_indexes:
|
||||||
|
next_line = lines[end_index + 1][2]
|
||||||
|
if list_pattern.match(next_line):
|
||||||
|
item_count += 1
|
||||||
|
end_index += 1
|
||||||
|
elif continuation_pattern.match(next_line):
|
||||||
|
end_index += 1
|
||||||
|
else:
|
||||||
|
break
|
||||||
|
if item_count >= 2:
|
||||||
|
ranges.append((lines[index][0], lines[end_index][1]))
|
||||||
|
index = end_index + 1
|
||||||
|
|
||||||
|
merged: list[tuple[int, int]] = []
|
||||||
|
for start, end in sorted(ranges):
|
||||||
|
if merged and start < merged[-1][1]:
|
||||||
|
merged[-1] = (merged[-1][0], max(merged[-1][1], end))
|
||||||
|
else:
|
||||||
|
merged.append((start, end))
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
|
def extract_structured_records(text: str, file_format: str) -> list[dict[str, Any]]:
|
||||||
|
"""从 JSON、JSONL 或 CSV 中提取规范化记录。"""
|
||||||
|
|
||||||
|
records, _ = _extract_structured_records_with_locators(text, file_format)
|
||||||
|
return records
|
||||||
|
|
||||||
|
def canonical_record_json(record: Mapping[str, Any]) -> str:
|
||||||
|
"""生成与字段顺序无关、可用于比较和落库的 canonical JSON。"""
|
||||||
|
|
||||||
|
if not isinstance(record, Mapping):
|
||||||
|
raise TypeError("record must be a mapping")
|
||||||
|
return structured_json_dumps(_canonical_value(record), sort_keys=True)
|
||||||
|
|
||||||
|
def normalize_structured_record(
|
||||||
|
record: Mapping[str, Any],
|
||||||
|
*,
|
||||||
|
field_name_style: str = "snake_case",
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""规范字段名、Unicode/空白、容器类型和不可 JSON 化的标量。"""
|
||||||
|
|
||||||
|
if not isinstance(record, Mapping):
|
||||||
|
raise TypeError("record must be a mapping")
|
||||||
|
normalized: dict[str, Any] = {}
|
||||||
|
for key, value in record.items():
|
||||||
|
normalized_key = _normalize_field_name(key, field_name_style)
|
||||||
|
if not normalized_key:
|
||||||
|
raise ValueError("structured record contains an empty field name")
|
||||||
|
if normalized_key in normalized:
|
||||||
|
raise ValueError(
|
||||||
|
f"structured record fields collide after normalization: {normalized_key}"
|
||||||
|
)
|
||||||
|
normalized[normalized_key] = _canonical_value(value)
|
||||||
|
return dict(sorted(normalized.items()))
|
||||||
|
|
||||||
|
def flatten_structured_record(
|
||||||
|
record: Mapping[str, Any],
|
||||||
|
*,
|
||||||
|
separator: str = ".",
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""把嵌套对象展平;数组保留为 canonical JSON 兼容值。"""
|
||||||
|
|
||||||
|
if not isinstance(record, Mapping):
|
||||||
|
raise TypeError("record must be a mapping")
|
||||||
|
if not separator:
|
||||||
|
raise ValueError("separator cannot be empty")
|
||||||
|
flattened: dict[str, Any] = {}
|
||||||
|
|
||||||
|
def visit(value: Any, path: tuple[str, ...], depth: int) -> None:
|
||||||
|
if depth > _MAX_STRUCTURED_DEPTH:
|
||||||
|
raise ValueError(
|
||||||
|
f"structured record nesting exceeds {_MAX_STRUCTURED_DEPTH} levels"
|
||||||
|
)
|
||||||
|
value = _embedded_structure(value)
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
if not value and path:
|
||||||
|
key = separator.join(path)
|
||||||
|
flattened[key] = {}
|
||||||
|
return
|
||||||
|
for child_key, child_value in value.items():
|
||||||
|
normalized_key = normalize_text(str(child_key))
|
||||||
|
if not normalized_key:
|
||||||
|
raise ValueError("structured record contains an empty field name")
|
||||||
|
visit(child_value, (*path, normalized_key), depth + 1)
|
||||||
|
return
|
||||||
|
key = separator.join(path)
|
||||||
|
if key in flattened:
|
||||||
|
raise ValueError(f"structured record fields collide while flattening: {key}")
|
||||||
|
flattened[key] = _canonical_value(value)
|
||||||
|
if len(flattened) > _MAX_STRUCTURED_FIELDS:
|
||||||
|
raise ValueError(
|
||||||
|
f"structured record exceeds {_MAX_STRUCTURED_FIELDS} flattened fields"
|
||||||
|
)
|
||||||
|
|
||||||
|
for field, value in record.items():
|
||||||
|
field_name = normalize_text(str(field))
|
||||||
|
if not field_name:
|
||||||
|
raise ValueError("structured record contains an empty field name")
|
||||||
|
visit(value, (field_name,), 1)
|
||||||
|
return flattened
|
||||||
|
|
||||||
|
def filter_anomalous_structured_records(
|
||||||
|
records: Sequence[Mapping[str, Any]],
|
||||||
|
*,
|
||||||
|
iqr_multiplier: float = 1.5,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""按字段级数值 IQR、乱码和极端文本长度过滤异常记录。"""
|
||||||
|
|
||||||
|
entries = [
|
||||||
|
ProcessedStructuredRecord(index, deepcopy(dict(record)))
|
||||||
|
for index, record in enumerate(records)
|
||||||
|
]
|
||||||
|
return [
|
||||||
|
entry.record
|
||||||
|
for entry in _filter_anomalous_structured_entries(
|
||||||
|
entries,
|
||||||
|
iqr_multiplier=iqr_multiplier,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
def desensitize_pii(text: str) -> tuple[str, dict[str, int]]:
|
||||||
|
"""掩码邮箱、手机号、身份证号及有明确上下文的姓名。"""
|
||||||
|
|
||||||
|
if not isinstance(text, str):
|
||||||
|
raise TypeError("text must be str")
|
||||||
|
counts: dict[str, int] = {"email": 0, "phone": 0, "id_card": 0}
|
||||||
|
|
||||||
|
def replace(pattern: re.Pattern[str], replacement: str, kind: str, value: str) -> str:
|
||||||
|
def replacer(_: re.Match[str]) -> str:
|
||||||
|
counts[kind] += 1
|
||||||
|
return replacement
|
||||||
|
|
||||||
|
return pattern.sub(replacer, value)
|
||||||
|
|
||||||
|
masked = replace(_EMAIL_PATTERN, "[EMAIL]", "email", text)
|
||||||
|
masked = replace(_ID_CARD_PATTERN, "[ID_CARD]", "id_card", masked)
|
||||||
|
masked = replace(_PHONE_PATTERN, "[PHONE]", "phone", masked)
|
||||||
|
|
||||||
|
def replace_context_name(match: re.Match[str]) -> str:
|
||||||
|
counts["name"] = counts.get("name", 0) + 1
|
||||||
|
return f"{match.group('label')}{match.group('separator')}[NAME]"
|
||||||
|
|
||||||
|
masked = _CHINESE_NAME_CONTEXT_PATTERN.sub(replace_context_name, masked)
|
||||||
|
masked = _ENGLISH_NAME_CONTEXT_PATTERN.sub(replace_context_name, masked)
|
||||||
|
counts["total"] = sum(counts.values())
|
||||||
|
return masked, counts
|
||||||
|
|
||||||
|
def desensitize_structured_record(
|
||||||
|
record: Mapping[str, Any],
|
||||||
|
) -> tuple[dict[str, Any], dict[str, int]]:
|
||||||
|
"""递归脱敏结构化姓名字段及任意文本中的手机号、邮箱、身份证号。"""
|
||||||
|
|
||||||
|
if not isinstance(record, Mapping):
|
||||||
|
raise TypeError("record must be a mapping")
|
||||||
|
counts: dict[str, int] = {"email": 0, "phone": 0, "id_card": 0}
|
||||||
|
|
||||||
|
def add_counts(values: Mapping[str, int]) -> None:
|
||||||
|
for kind, count in values.items():
|
||||||
|
if kind != "total" and count:
|
||||||
|
counts[kind] = counts.get(kind, 0) + count
|
||||||
|
|
||||||
|
def visit(value: Any, field: Any = "") -> Any:
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return {key: visit(item, key) for key, item in value.items()}
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [visit(item, field) for item in value]
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return [visit(item, field) for item in value]
|
||||||
|
if _is_name_field(field) and not _is_empty_value(value):
|
||||||
|
counts["name"] = counts.get("name", 0) + 1
|
||||||
|
return "[NAME]"
|
||||||
|
if isinstance(value, str):
|
||||||
|
masked, found = desensitize_pii(value)
|
||||||
|
add_counts(found)
|
||||||
|
return masked
|
||||||
|
return deepcopy(value)
|
||||||
|
|
||||||
|
masked = {key: visit(value, key) for key, value in record.items()}
|
||||||
|
counts["total"] = sum(counts.values())
|
||||||
|
return masked, counts
|
||||||
|
|
||||||
|
def preprocess_structured_records_with_lineage(
|
||||||
|
records: Iterable[Mapping[str, Any]],
|
||||||
|
options: Iterable[str] | Mapping[str, Any],
|
||||||
|
) -> list[ProcessedStructuredRecord]:
|
||||||
|
"""执行结构化预处理,并保留每条结果在原始输入中的稳定索引。"""
|
||||||
|
|
||||||
|
enabled = _structured_options(options)
|
||||||
|
current: list[ProcessedStructuredRecord] = []
|
||||||
|
for source_index, record in enumerate(records):
|
||||||
|
if not isinstance(record, Mapping):
|
||||||
|
if "clean_invalid" in enabled:
|
||||||
|
continue
|
||||||
|
raise TypeError("structured records must contain mappings")
|
||||||
|
value = deepcopy(dict(record))
|
||||||
|
if "detect_structure" in enabled:
|
||||||
|
value = flatten_structured_record(value)
|
||||||
|
if "normalize_format" in enabled:
|
||||||
|
value = normalize_structured_record(value)
|
||||||
|
current.append(ProcessedStructuredRecord(source_index, value))
|
||||||
|
|
||||||
|
if "clean_invalid" in enabled:
|
||||||
|
current = _clean_invalid_structured_entries(current)
|
||||||
|
if "filter_anomaly" in enabled:
|
||||||
|
current = _filter_anomalous_structured_entries(current)
|
||||||
|
if "deduplicate" in enabled:
|
||||||
|
current = _deduplicate_structured_entries(current)
|
||||||
|
if "desensitize" in enabled:
|
||||||
|
current = [
|
||||||
|
ProcessedStructuredRecord(
|
||||||
|
entry.source_index,
|
||||||
|
desensitize_structured_record(entry.record)[0],
|
||||||
|
)
|
||||||
|
for entry in current
|
||||||
|
]
|
||||||
|
return current
|
||||||
|
|
||||||
|
def preprocess_structured_records(
|
||||||
|
records: Iterable[Mapping[str, Any]],
|
||||||
|
options: Iterable[str] | Mapping[str, Any],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""按界面选项执行确定性、无副作用的结构化数据预处理。"""
|
||||||
|
|
||||||
|
return [
|
||||||
|
entry.record
|
||||||
|
for entry in preprocess_structured_records_with_lineage(records, options)
|
||||||
|
]
|
||||||
|
|
||||||
|
def protected_context_ranges(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
preserve_code_blocks: bool = True,
|
||||||
|
preserve_tables: bool = True,
|
||||||
|
preserve_lists: bool = True,
|
||||||
|
) -> tuple[tuple[int, int], ...]:
|
||||||
|
"""返回代码块、表格和列表的不可拆分区间。
|
||||||
|
|
||||||
|
返回的偏移量基于规范化后的文本;调用方应在同一份 ``normalize_text``
|
||||||
|
结果上使用这些区间。
|
||||||
|
"""
|
||||||
|
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
return tuple(
|
||||||
|
_protected_markdown_ranges(
|
||||||
|
normalized,
|
||||||
|
preserve_code_blocks=preserve_code_blocks,
|
||||||
|
preserve_tables=preserve_tables,
|
||||||
|
preserve_lists=preserve_lists,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def expand_to_context_boundaries(
|
||||||
|
text: str,
|
||||||
|
start: int,
|
||||||
|
end: int,
|
||||||
|
*,
|
||||||
|
preserve_paragraph: bool = True,
|
||||||
|
preserve_code_blocks: bool = True,
|
||||||
|
preserve_tables: bool = True,
|
||||||
|
preserve_lists: bool = True,
|
||||||
|
) -> tuple[int, int]:
|
||||||
|
"""将一个文本区间扩展到段落及受保护 Markdown 结构边界。"""
|
||||||
|
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
if isinstance(start, bool) or isinstance(end, bool):
|
||||||
|
raise TypeError("start and end must be integers")
|
||||||
|
if not isinstance(start, int) or not isinstance(end, int):
|
||||||
|
raise TypeError("start and end must be integers")
|
||||||
|
if not 0 <= start <= end <= len(normalized):
|
||||||
|
raise ValueError("start and end must define a valid normalized text range")
|
||||||
|
|
||||||
|
expanded_start = start
|
||||||
|
expanded_end = end
|
||||||
|
if preserve_paragraph and normalized:
|
||||||
|
paragraph_start = normalized.rfind("\n\n", 0, start)
|
||||||
|
expanded_start = 0 if paragraph_start < 0 else paragraph_start + 2
|
||||||
|
paragraph_end = normalized.find("\n\n", end)
|
||||||
|
expanded_end = len(normalized) if paragraph_end < 0 else paragraph_end
|
||||||
|
|
||||||
|
ranges = protected_context_ranges(
|
||||||
|
normalized,
|
||||||
|
preserve_code_blocks=preserve_code_blocks,
|
||||||
|
preserve_tables=preserve_tables,
|
||||||
|
preserve_lists=preserve_lists,
|
||||||
|
)
|
||||||
|
changed = True
|
||||||
|
while changed:
|
||||||
|
changed = False
|
||||||
|
for range_start, range_end in ranges:
|
||||||
|
overlaps = range_start < expanded_end and range_end > expanded_start
|
||||||
|
contains_boundary = (
|
||||||
|
range_start <= expanded_start < range_end
|
||||||
|
or range_start < expanded_end <= range_end
|
||||||
|
)
|
||||||
|
if not overlaps and not contains_boundary:
|
||||||
|
continue
|
||||||
|
next_start = min(expanded_start, range_start)
|
||||||
|
next_end = max(expanded_end, range_end)
|
||||||
|
if (next_start, next_end) != (expanded_start, expanded_end):
|
||||||
|
expanded_start, expanded_end = next_start, next_end
|
||||||
|
changed = True
|
||||||
|
return expanded_start, expanded_end
|
||||||
|
|
||||||
|
def merge_short_blocks(
|
||||||
|
blocks: Iterable[str],
|
||||||
|
*,
|
||||||
|
min_token_count: int = 100,
|
||||||
|
separator: str = "\n\n",
|
||||||
|
) -> list[str]:
|
||||||
|
"""按原顺序合并短内容块,并把末尾残块归入前一块。"""
|
||||||
|
|
||||||
|
if isinstance(min_token_count, bool) or not isinstance(min_token_count, int):
|
||||||
|
raise TypeError("min_token_count must be an integer")
|
||||||
|
if min_token_count <= 0:
|
||||||
|
raise ValueError("min_token_count must be greater than 0")
|
||||||
|
if not isinstance(separator, str):
|
||||||
|
raise TypeError("separator must be str")
|
||||||
|
|
||||||
|
merged: list[str] = []
|
||||||
|
pending: list[str] = []
|
||||||
|
pending_tokens = 0
|
||||||
|
for block in blocks:
|
||||||
|
if not isinstance(block, str):
|
||||||
|
raise TypeError("blocks must contain strings")
|
||||||
|
normalized = normalize_text(block)
|
||||||
|
if not normalized:
|
||||||
|
continue
|
||||||
|
token_count = estimate_token_count(normalized)
|
||||||
|
if not pending and token_count >= min_token_count:
|
||||||
|
merged.append(normalized)
|
||||||
|
continue
|
||||||
|
pending.append(normalized)
|
||||||
|
pending_tokens += token_count
|
||||||
|
if pending_tokens >= min_token_count:
|
||||||
|
merged.append(separator.join(pending))
|
||||||
|
pending = []
|
||||||
|
pending_tokens = 0
|
||||||
|
|
||||||
|
if pending:
|
||||||
|
tail = separator.join(pending)
|
||||||
|
if merged:
|
||||||
|
merged[-1] = separator.join((merged[-1], tail))
|
||||||
|
else:
|
||||||
|
merged.append(tail)
|
||||||
|
return merged
|
||||||
|
|
||||||
|
def generate_standard_records(
|
||||||
|
preview_items: Iterable[Mapping[str, Any]],
|
||||||
|
*,
|
||||||
|
qa_pairs_per_item: int = 1,
|
||||||
|
semantic_enrichment: bool = False,
|
||||||
|
split: Mapping[str, int] | None = None,
|
||||||
|
split_seed: str = "",
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""把预览内容确定性转换为标准 instruction/input/output 记录。
|
||||||
|
|
||||||
|
该函数只负责本地标准化,不冒充 LLM;服务层可将其作为无模型模式或
|
||||||
|
LLM 响应解析后的统一落库步骤。
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not 1 <= qa_pairs_per_item <= MAX_QA_PAIRS_PER_ITEM:
|
||||||
|
raise ValueError(
|
||||||
|
f"qa_pairs_per_item must be in [1, {MAX_QA_PAIRS_PER_ITEM}]"
|
||||||
|
)
|
||||||
|
prefixes = (
|
||||||
|
"请结合实际情况说明:",
|
||||||
|
"请用通俗易懂的方式说明:",
|
||||||
|
"请从实际应用角度说明:",
|
||||||
|
"请简洁自然地说明:",
|
||||||
|
"请详细解答:",
|
||||||
|
)
|
||||||
|
results: list[dict[str, Any]] = []
|
||||||
|
for item_index, item in enumerate(preview_items):
|
||||||
|
content = _preview_content(item)
|
||||||
|
instruction, input_text, output = _standard_fields(content)
|
||||||
|
preview_id = str(item.get("id") or f"preview-{item_index + 1}")
|
||||||
|
for variant_index in range(qa_pairs_per_item):
|
||||||
|
variant_instruction = instruction
|
||||||
|
if variant_index:
|
||||||
|
if semantic_enrichment:
|
||||||
|
prefix = prefixes[variant_index % len(prefixes)]
|
||||||
|
if variant_index >= len(prefixes):
|
||||||
|
prefix = (
|
||||||
|
f"{prefix.removesuffix(':')}"
|
||||||
|
f"(问法 {variant_index + 1}):"
|
||||||
|
)
|
||||||
|
variant_instruction = f"{prefix}{instruction}"
|
||||||
|
else:
|
||||||
|
variant_instruction = f"{instruction}(问法 {variant_index + 1})"
|
||||||
|
raw_id = f"{preview_id}:{variant_index + 1}"
|
||||||
|
result_id = f"result_{hashlib.sha256(raw_id.encode('utf-8')).hexdigest()[:16]}"
|
||||||
|
status = "valid" if variant_instruction and output else "invalid"
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
"id": result_id,
|
||||||
|
"preview_item_id": preview_id,
|
||||||
|
"instruction": variant_instruction,
|
||||||
|
"input": input_text,
|
||||||
|
"output": output,
|
||||||
|
"original_instruction": variant_instruction,
|
||||||
|
"original_input": input_text,
|
||||||
|
"original_output": output,
|
||||||
|
"status": status,
|
||||||
|
"split": "train",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assignments = stable_split_assignments(
|
||||||
|
[str(result["id"]) for result in results],
|
||||||
|
split,
|
||||||
|
seed=split_seed,
|
||||||
|
)
|
||||||
|
for result, assignment in zip(results, assignments, strict=True):
|
||||||
|
result["split"] = assignment
|
||||||
|
return results
|
||||||
328
backend/app/modules/data_process/algorithms/text_utils.py
Normal file
328
backend/app/modules/data_process/algorithms/text_utils.py
Normal file
@@ -0,0 +1,328 @@
|
|||||||
|
"""数据处理算法 - 文本处理工具。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import csv
|
||||||
|
import io
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
from datetime import date, datetime, time
|
||||||
|
from decimal import Decimal
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Literal
|
||||||
|
from collections.abc import Mapping
|
||||||
|
|
||||||
|
from .types import ParsedText, TextFormat
|
||||||
|
|
||||||
|
# 格式别名映射
|
||||||
|
_FORMAT_ALIASES: dict[str, TextFormat] = {
|
||||||
|
"json": "json",
|
||||||
|
"jsonl": "jsonl",
|
||||||
|
"ndjson": "jsonl",
|
||||||
|
"csv": "csv",
|
||||||
|
"tsv": "csv",
|
||||||
|
"md": "markdown",
|
||||||
|
"markdown": "markdown",
|
||||||
|
"txt": "txt",
|
||||||
|
"text": "txt",
|
||||||
|
"pdf": "pdf",
|
||||||
|
"docx": "docx",
|
||||||
|
"xlsx": "xlsx",
|
||||||
|
"pptx": "pptx",
|
||||||
|
}
|
||||||
|
|
||||||
|
# 旧版 Office 格式映射
|
||||||
|
_LEGACY_OFFICE_FORMATS: dict[str, str] = {
|
||||||
|
"doc": "docx",
|
||||||
|
"xls": "xlsx",
|
||||||
|
"ppt": "pptx",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Office Open XML 格式集合
|
||||||
|
_OFFICE_OPEN_XML_FORMATS = {"docx", "xlsx", "pptx"}
|
||||||
|
|
||||||
|
# 文本提取限制
|
||||||
|
_MAX_EXTRACTED_TEXT_CHARS = 20_000_000
|
||||||
|
|
||||||
|
def decode_utf8(raw: bytes | bytearray | memoryview | str) -> str:
|
||||||
|
"""严格解码 UTF-8 文本,并移除可选 BOM。
|
||||||
|
|
||||||
|
不使用 ``errors='replace'``,避免上传内容损坏后仍被静默接收。
|
||||||
|
"""
|
||||||
|
|
||||||
|
if isinstance(raw, str):
|
||||||
|
return raw.removeprefix("\ufeff")
|
||||||
|
if not isinstance(raw, (bytes, bytearray, memoryview)):
|
||||||
|
raise TypeError("raw must be bytes-like or str")
|
||||||
|
try:
|
||||||
|
return bytes(raw).decode("utf-8-sig")
|
||||||
|
except UnicodeDecodeError as exc:
|
||||||
|
raise ValueError(f"content is not valid UTF-8 at byte {exc.start}") from exc
|
||||||
|
|
||||||
|
def parse_utf8_text(raw: bytes | bytearray | memoryview | str) -> str:
|
||||||
|
"""``decode_utf8`` 的语义化别名,供上传服务直接调用。"""
|
||||||
|
|
||||||
|
return decode_utf8(raw)
|
||||||
|
|
||||||
|
def normalize_text(text: str) -> str:
|
||||||
|
"""规范 Unicode、换行和行尾空白,同时保留段落结构。"""
|
||||||
|
|
||||||
|
if not isinstance(text, str):
|
||||||
|
raise TypeError("text must be str")
|
||||||
|
normalized = unicodedata.normalize("NFKC", text.removeprefix("\ufeff"))
|
||||||
|
normalized = normalized.replace("\r\n", "\n").replace("\r", "\n")
|
||||||
|
normalized = "".join(
|
||||||
|
char
|
||||||
|
for char in normalized
|
||||||
|
if char in {"\n", "\t"} or not unicodedata.category(char).startswith("C")
|
||||||
|
)
|
||||||
|
lines = [re.sub(r"[\t \f\v]+$", "", line) for line in normalized.split("\n")]
|
||||||
|
return "\n".join(lines).strip()
|
||||||
|
|
||||||
|
def _normalize_format(value: str | None) -> TextFormat | None:
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
normalized = value.strip().lower().removeprefix(".")
|
||||||
|
if normalized in _LEGACY_OFFICE_FORMATS:
|
||||||
|
replacement = _LEGACY_OFFICE_FORMATS[normalized]
|
||||||
|
raise ValueError(
|
||||||
|
f"legacy .{normalized} format is not supported; "
|
||||||
|
f"convert the file to .{replacement} and upload again"
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
return _FORMAT_ALIASES[normalized]
|
||||||
|
except KeyError as exc:
|
||||||
|
raise ValueError(f"unsupported text format: {value}") from exc
|
||||||
|
|
||||||
|
def detect_text_format(
|
||||||
|
*,
|
||||||
|
filename: str | None = None,
|
||||||
|
text: str = "",
|
||||||
|
file_format: str | None = None,
|
||||||
|
) -> TextFormat:
|
||||||
|
"""按显式格式、扩展名和内容特征依次识别文本格式。"""
|
||||||
|
|
||||||
|
explicit = _normalize_format(file_format)
|
||||||
|
if explicit:
|
||||||
|
return explicit
|
||||||
|
|
||||||
|
if filename:
|
||||||
|
suffix = Path(filename).suffix.lower().removeprefix(".")
|
||||||
|
if suffix in _LEGACY_OFFICE_FORMATS:
|
||||||
|
replacement = _LEGACY_OFFICE_FORMATS[suffix]
|
||||||
|
raise ValueError(
|
||||||
|
f"legacy .{suffix} format is not supported; "
|
||||||
|
f"convert the file to .{replacement} and upload again"
|
||||||
|
)
|
||||||
|
detected = _FORMAT_ALIASES.get(suffix)
|
||||||
|
if detected:
|
||||||
|
return detected
|
||||||
|
|
||||||
|
stripped = text.strip()
|
||||||
|
if stripped:
|
||||||
|
if stripped[0] in "[{":
|
||||||
|
try:
|
||||||
|
json.loads(stripped)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
return "json"
|
||||||
|
|
||||||
|
nonempty_lines = [line for line in stripped.splitlines() if line.strip()]
|
||||||
|
if len(nonempty_lines) > 1:
|
||||||
|
try:
|
||||||
|
for line in nonempty_lines:
|
||||||
|
json.loads(line)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
return "jsonl"
|
||||||
|
|
||||||
|
if re.search(r"(?m)^(?:#{1,6}\s+|```|~~~)", stripped) or re.search(
|
||||||
|
r"(?m)^\s*\|.+\|\s*$", stripped
|
||||||
|
):
|
||||||
|
return "markdown"
|
||||||
|
|
||||||
|
sample = stripped[:8192]
|
||||||
|
try:
|
||||||
|
dialect = csv.Sniffer().sniff(sample, delimiters=",\t;")
|
||||||
|
rows = list(csv.reader(io.StringIO(sample), dialect))
|
||||||
|
if len(rows) >= 2 and len(rows[0]) >= 2:
|
||||||
|
return "csv"
|
||||||
|
except csv.Error:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return "txt"
|
||||||
|
|
||||||
|
def _append_bounded_text(parts: list[str], value: Any, total: int) -> int:
|
||||||
|
text = normalize_text(str(value or ""))
|
||||||
|
if not text:
|
||||||
|
return total
|
||||||
|
total += len(text)
|
||||||
|
if total > _MAX_EXTRACTED_TEXT_CHARS:
|
||||||
|
raise ValueError(
|
||||||
|
f"extracted document text exceeds {_MAX_EXTRACTED_TEXT_CHARS} characters"
|
||||||
|
)
|
||||||
|
parts.append(text)
|
||||||
|
return total
|
||||||
|
|
||||||
|
def _normalize_spreadsheet_value(value: Any) -> Any:
|
||||||
|
if value is None:
|
||||||
|
return ""
|
||||||
|
if isinstance(value, str):
|
||||||
|
return normalize_text(value)
|
||||||
|
if isinstance(value, (datetime, date, time)):
|
||||||
|
return value.isoformat()
|
||||||
|
if isinstance(value, (bool, int, float)):
|
||||||
|
return value
|
||||||
|
return normalize_text(str(value))
|
||||||
|
|
||||||
|
def _normalize_value(value: Any) -> Any:
|
||||||
|
if isinstance(value, str):
|
||||||
|
return normalize_text(value)
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return {normalize_text(str(key)): _normalize_value(item) for key, item in value.items()}
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [_normalize_value(item) for item in value]
|
||||||
|
return value
|
||||||
|
|
||||||
|
def parse_text_content(
|
||||||
|
raw: bytes | bytearray | memoryview | str,
|
||||||
|
*,
|
||||||
|
filename: str | None = None,
|
||||||
|
file_format: str | None = None,
|
||||||
|
) -> ParsedText:
|
||||||
|
"""安全解析 UTF-8 文本、文本型 PDF 和现代 Office 文件。"""
|
||||||
|
|
||||||
|
# 解析器依赖 text_utils(normalize_text 等),这里延迟导入以断开循环依赖。
|
||||||
|
from .parsers.json_parser import _extract_structured_records_with_locators
|
||||||
|
from .parsers.office import _extract_docx_text, _extract_pptx_text, _extract_xlsx_records
|
||||||
|
from .parsers.pdf import _extract_pdf_text
|
||||||
|
|
||||||
|
detected_format = detect_text_format(
|
||||||
|
filename=filename,
|
||||||
|
text="",
|
||||||
|
file_format=file_format,
|
||||||
|
)
|
||||||
|
if detected_format in _OFFICE_OPEN_XML_FORMATS or detected_format == "pdf":
|
||||||
|
binary = _binary_bytes(raw, detected_format)
|
||||||
|
if detected_format == "pdf":
|
||||||
|
text = _extract_pdf_text(binary)
|
||||||
|
return ParsedText(format=detected_format, text=text, records=())
|
||||||
|
if detected_format == "docx":
|
||||||
|
text = _extract_docx_text(binary)
|
||||||
|
return ParsedText(format=detected_format, text=text, records=())
|
||||||
|
if detected_format == "pptx":
|
||||||
|
text = _extract_pptx_text(binary)
|
||||||
|
return ParsedText(format=detected_format, text=text, records=())
|
||||||
|
|
||||||
|
records, record_locators = _extract_xlsx_records(binary)
|
||||||
|
text = "\n".join(
|
||||||
|
json.dumps(record, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
for record in records
|
||||||
|
)
|
||||||
|
return ParsedText(
|
||||||
|
format=detected_format,
|
||||||
|
text=normalize_text(text),
|
||||||
|
records=tuple(records),
|
||||||
|
record_locators=tuple(record_locators),
|
||||||
|
)
|
||||||
|
|
||||||
|
decoded_text = decode_utf8(raw)
|
||||||
|
detected_format = detect_text_format(
|
||||||
|
filename=filename,
|
||||||
|
text=decoded_text,
|
||||||
|
file_format=file_format,
|
||||||
|
)
|
||||||
|
# JSON/JSONL 是有损规范化的禁区:NFKC、控制字符删除或 trim 都可能改变字段值、
|
||||||
|
# 掩盖非法输入,甚至把原本合法的字符串变成语法错误。其他格式保持历史行为。
|
||||||
|
text = (
|
||||||
|
decoded_text
|
||||||
|
if detected_format in {"json", "jsonl"}
|
||||||
|
else normalize_text(decoded_text)
|
||||||
|
)
|
||||||
|
records: list[dict[str, Any]] = []
|
||||||
|
record_locators: list[dict[str, Any]] = []
|
||||||
|
if detected_format in {"json", "jsonl", "csv"}:
|
||||||
|
records, record_locators = _extract_structured_records_with_locators(
|
||||||
|
text,
|
||||||
|
detected_format,
|
||||||
|
)
|
||||||
|
return ParsedText(
|
||||||
|
format=detected_format,
|
||||||
|
text=text,
|
||||||
|
records=tuple(records),
|
||||||
|
record_locators=tuple(record_locators),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _normalize_field_name(value: Any, style: str) -> str:
|
||||||
|
name = normalize_text(str(value))
|
||||||
|
if style == "preserve":
|
||||||
|
return name
|
||||||
|
if style == "lower":
|
||||||
|
return name.lower()
|
||||||
|
if style != "snake_case":
|
||||||
|
raise ValueError("field_name_style must be snake_case, lower or preserve")
|
||||||
|
name = re.sub(r"(?<=[a-z0-9])(?=[A-Z])", "_", name)
|
||||||
|
name = re.sub(r"[^\w\u3400-\u4dbf\u4e00-\u9fff]+", "_", name, flags=re.UNICODE)
|
||||||
|
return re.sub(r"_+", "_", name).strip("_").lower()
|
||||||
|
|
||||||
|
def structured_json_dumps(value: Any, *, sort_keys: bool = False) -> str:
|
||||||
|
"""序列化紧凑 JSON,并把 ``Decimal`` 保持为原值对应的 JSON 数字。
|
||||||
|
|
||||||
|
标准库会要求先把 ``Decimal`` 转成 float 或字符串;前者可能静默舍入,
|
||||||
|
后者会改变 JSON 类型。这里直接输出有限 Decimal 的十进制表示。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def serialize(item: Any) -> str:
|
||||||
|
if item is None:
|
||||||
|
return "null"
|
||||||
|
if item is True:
|
||||||
|
return "true"
|
||||||
|
if item is False:
|
||||||
|
return "false"
|
||||||
|
if isinstance(item, int):
|
||||||
|
return str(item)
|
||||||
|
if isinstance(item, Decimal):
|
||||||
|
if not item.is_finite():
|
||||||
|
raise ValueError("non-finite JSON number is not allowed")
|
||||||
|
return str(item)
|
||||||
|
if isinstance(item, float):
|
||||||
|
if not math.isfinite(item):
|
||||||
|
raise ValueError("non-finite JSON number is not allowed")
|
||||||
|
return json.dumps(item, allow_nan=False)
|
||||||
|
if isinstance(item, str):
|
||||||
|
return json.dumps(item, ensure_ascii=False)
|
||||||
|
if isinstance(item, Mapping):
|
||||||
|
pairs: list[tuple[str, Any]] = []
|
||||||
|
seen_keys: set[str] = set()
|
||||||
|
for key, child in item.items():
|
||||||
|
if not isinstance(key, str):
|
||||||
|
raise TypeError("JSON object keys must be strings")
|
||||||
|
if key in seen_keys:
|
||||||
|
raise ValueError(f"duplicate JSON object key: {key!r}")
|
||||||
|
seen_keys.add(key)
|
||||||
|
pairs.append((key, child))
|
||||||
|
if sort_keys:
|
||||||
|
pairs.sort(key=lambda pair: pair[0])
|
||||||
|
return "{" + ",".join(
|
||||||
|
f"{json.dumps(key, ensure_ascii=False)}:{serialize(child)}"
|
||||||
|
for key, child in pairs
|
||||||
|
) + "}"
|
||||||
|
if isinstance(item, (list, tuple)):
|
||||||
|
return "[" + ",".join(serialize(child) for child in item) + "]"
|
||||||
|
raise TypeError(f"value of type {type(item).__name__} is not JSON serializable")
|
||||||
|
|
||||||
|
return serialize(value)
|
||||||
|
|
||||||
|
def _binary_bytes(
|
||||||
|
raw: bytes | bytearray | memoryview | str,
|
||||||
|
file_format: TextFormat,
|
||||||
|
) -> bytes:
|
||||||
|
if isinstance(raw, str):
|
||||||
|
raise ValueError(f"{file_format.upper()} content must be uploaded as binary data")
|
||||||
|
if not isinstance(raw, (bytes, bytearray, memoryview)):
|
||||||
|
raise TypeError("raw must be bytes-like or str")
|
||||||
|
return bytes(raw)
|
||||||
81
backend/app/modules/data_process/algorithms/transforms.py
Normal file
81
backend/app/modules/data_process/algorithms/transforms.py
Normal file
@@ -0,0 +1,81 @@
|
|||||||
|
"""数据处理算法 - 数据集转换和分割。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import math
|
||||||
|
from collections.abc import Mapping, Sequence
|
||||||
|
|
||||||
|
from .types import DatasetSplit
|
||||||
|
|
||||||
|
|
||||||
|
def stable_split(
|
||||||
|
value: str | int,
|
||||||
|
split: Mapping[str, int] | None = None,
|
||||||
|
*,
|
||||||
|
seed: str = "",
|
||||||
|
) -> DatasetSplit:
|
||||||
|
"""按稳定哈希将记录划分到 train/validation/test。"""
|
||||||
|
|
||||||
|
ratios = dict(split or {"train": 80, "validation": 10, "test": 10})
|
||||||
|
required = {"train", "validation", "test"}
|
||||||
|
if set(ratios) != required:
|
||||||
|
raise ValueError("split must contain exactly train, validation and test")
|
||||||
|
if any(isinstance(value, bool) or not isinstance(value, int) or value < 0 for value in ratios.values()):
|
||||||
|
raise ValueError("split ratios must be non-negative integers")
|
||||||
|
if sum(ratios.values()) != 100:
|
||||||
|
raise ValueError("split ratios must sum to 100")
|
||||||
|
|
||||||
|
digest = hashlib.sha256(f"{seed}:{value}".encode("utf-8")).digest()
|
||||||
|
bucket = int.from_bytes(digest[:8], "big") % 10_000
|
||||||
|
train_boundary = ratios["train"] * 100
|
||||||
|
validation_boundary = train_boundary + ratios["validation"] * 100
|
||||||
|
if bucket < train_boundary:
|
||||||
|
return "train"
|
||||||
|
if bucket < validation_boundary:
|
||||||
|
return "validation"
|
||||||
|
return "test"
|
||||||
|
|
||||||
|
def stable_split_assignments(
|
||||||
|
values: Sequence[str | int],
|
||||||
|
split: Mapping[str, int] | None = None,
|
||||||
|
*,
|
||||||
|
seed: str = "",
|
||||||
|
) -> list[DatasetSplit]:
|
||||||
|
"""按稳定顺序和精确配额批量划分数据集。
|
||||||
|
|
||||||
|
单条哈希分桶只能在大样本下近似比例。这里先按哈希稳定排序,再用
|
||||||
|
最大余数法计算各切分配额,确保小数据集也严格遵循配置比例。
|
||||||
|
"""
|
||||||
|
|
||||||
|
ratios = dict(split or {"train": 80, "validation": 10, "test": 10})
|
||||||
|
# 复用单条划分的参数校验,避免两套规则逐渐漂移。
|
||||||
|
stable_split("validation", ratios, seed=seed)
|
||||||
|
if not values:
|
||||||
|
return []
|
||||||
|
|
||||||
|
split_order: tuple[DatasetSplit, ...] = ("train", "validation", "test")
|
||||||
|
exact = {name: len(values) * ratios[name] / 100 for name in split_order}
|
||||||
|
quotas = {name: math.floor(exact[name]) for name in split_order}
|
||||||
|
remaining = len(values) - sum(quotas.values())
|
||||||
|
remainder_order = sorted(
|
||||||
|
split_order,
|
||||||
|
key=lambda name: (-(exact[name] - quotas[name]), split_order.index(name)),
|
||||||
|
)
|
||||||
|
for name in remainder_order[:remaining]:
|
||||||
|
quotas[name] += 1
|
||||||
|
|
||||||
|
ranked_indices = sorted(
|
||||||
|
range(len(values)),
|
||||||
|
key=lambda index: (
|
||||||
|
hashlib.sha256(f"{seed}:{values[index]}".encode("utf-8")).digest(),
|
||||||
|
index,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
assignments: list[DatasetSplit] = ["train"] * len(values)
|
||||||
|
cursor = 0
|
||||||
|
for name in split_order:
|
||||||
|
for index in ranked_indices[cursor : cursor + quotas[name]]:
|
||||||
|
assignments[index] = name
|
||||||
|
cursor += quotas[name]
|
||||||
|
return assignments
|
||||||
266
backend/app/modules/data_process/algorithms/types.py
Normal file
266
backend/app/modules/data_process/algorithms/types.py
Normal file
@@ -0,0 +1,266 @@
|
|||||||
|
"""数据处理模块使用的无副作用算法。
|
||||||
|
|
||||||
|
本模块不访问数据库、文件系统或网络,便于 API、后台任务和测试共同复用。
|
||||||
|
所有偏移量均为 Python 字符串偏移量。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import csv
|
||||||
|
import hashlib
|
||||||
|
import io
|
||||||
|
import json
|
||||||
|
import math
|
||||||
|
import re
|
||||||
|
import unicodedata
|
||||||
|
import xml.etree.ElementTree as ET
|
||||||
|
import zipfile
|
||||||
|
from collections import Counter
|
||||||
|
from collections.abc import Iterable, Mapping, Sequence
|
||||||
|
from copy import deepcopy
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from datetime import date, datetime, time
|
||||||
|
from decimal import Decimal
|
||||||
|
from pathlib import Path, PurePosixPath
|
||||||
|
from typing import Any, Literal
|
||||||
|
from urllib.parse import unquote, urlsplit
|
||||||
|
|
||||||
|
from docx import Document
|
||||||
|
from docx.oxml.table import CT_Tbl
|
||||||
|
from docx.oxml.text.paragraph import CT_P
|
||||||
|
from docx.table import Table
|
||||||
|
from docx.text.paragraph import Paragraph
|
||||||
|
from openpyxl import load_workbook
|
||||||
|
from openpyxl.utils.cell import range_boundaries
|
||||||
|
from pptx import Presentation
|
||||||
|
from pypdf import PdfReader
|
||||||
|
|
||||||
|
# 避免循环导入:直接定义常量而不是从 constants 导入
|
||||||
|
MAX_QA_PAIRS_PER_ITEM = 50
|
||||||
|
|
||||||
|
TextFormat = Literal[
|
||||||
|
"json",
|
||||||
|
"jsonl",
|
||||||
|
"csv",
|
||||||
|
"markdown",
|
||||||
|
"txt",
|
||||||
|
"pdf",
|
||||||
|
"docx",
|
||||||
|
"xlsx",
|
||||||
|
"pptx",
|
||||||
|
]
|
||||||
|
DatasetSplit = Literal["train", "validation", "test"]
|
||||||
|
StructuredPreprocessOption = Literal[
|
||||||
|
"clean_invalid",
|
||||||
|
"detect_structure",
|
||||||
|
"deduplicate",
|
||||||
|
"normalize_format",
|
||||||
|
"filter_anomaly",
|
||||||
|
"desensitize",
|
||||||
|
]
|
||||||
|
|
||||||
|
SUPPORTED_TEXT_FORMATS: tuple[TextFormat, ...] = (
|
||||||
|
"json",
|
||||||
|
"jsonl",
|
||||||
|
"csv",
|
||||||
|
"markdown",
|
||||||
|
"txt",
|
||||||
|
"pdf",
|
||||||
|
"docx",
|
||||||
|
"xlsx",
|
||||||
|
"pptx",
|
||||||
|
)
|
||||||
|
|
||||||
|
_FORMAT_ALIASES: dict[str, TextFormat] = {
|
||||||
|
"json": "json",
|
||||||
|
"jsonl": "jsonl",
|
||||||
|
"ndjson": "jsonl",
|
||||||
|
"csv": "csv",
|
||||||
|
"tsv": "csv",
|
||||||
|
"md": "markdown",
|
||||||
|
"markdown": "markdown",
|
||||||
|
"txt": "txt",
|
||||||
|
"text": "txt",
|
||||||
|
"pdf": "pdf",
|
||||||
|
"docx": "docx",
|
||||||
|
"xlsx": "xlsx",
|
||||||
|
"pptx": "pptx",
|
||||||
|
}
|
||||||
|
_LEGACY_OFFICE_FORMATS: dict[str, str] = {
|
||||||
|
"doc": "docx",
|
||||||
|
"xls": "xlsx",
|
||||||
|
"ppt": "pptx",
|
||||||
|
}
|
||||||
|
_OFFICE_OPEN_XML_FORMATS = {"docx", "xlsx", "pptx"}
|
||||||
|
_MAX_ARCHIVE_ENTRIES = 10_000
|
||||||
|
_MAX_ARCHIVE_UNCOMPRESSED_BYTES = 512 * 1024 * 1024
|
||||||
|
_MAX_ARCHIVE_ENTRY_BYTES = 128 * 1024 * 1024
|
||||||
|
_MAX_ARCHIVE_COMPRESSION_RATIO = 200
|
||||||
|
_MAX_EXTRACTED_TEXT_CHARS = 20_000_000
|
||||||
|
_MAX_PDF_PAGES = 2_000
|
||||||
|
_MAX_PRESENTATION_SLIDES = 2_000
|
||||||
|
_MAX_WORKBOOK_SHEETS = 100
|
||||||
|
_MAX_WORKBOOK_ROWS = 100_000
|
||||||
|
_MAX_WORKBOOK_SCANNED_ROWS = 200_000
|
||||||
|
_MAX_WORKBOOK_COLUMNS = 256
|
||||||
|
_MAX_WORKBOOK_CELLS = 2_000_000
|
||||||
|
_MAX_WORKBOOK_HEADER_ROWS = 8
|
||||||
|
_MAX_WORKBOOK_HEADER_SCAN_ROWS = 64
|
||||||
|
_MAX_WORKBOOK_MERGED_RANGES = 100_000
|
||||||
|
_MAX_STRUCTURED_FIELDS = 1_024
|
||||||
|
_MAX_STRUCTURED_DEPTH = 16
|
||||||
|
_MAX_JSON_DEPTH = 64
|
||||||
|
_MAX_ANOMALY_TEXT_CHARS = 1_000_000
|
||||||
|
_STRUCTURED_OPTIONS = {
|
||||||
|
"clean_invalid",
|
||||||
|
"detect_structure",
|
||||||
|
"deduplicate",
|
||||||
|
"normalize_format",
|
||||||
|
"filter_anomaly",
|
||||||
|
"desensitize",
|
||||||
|
}
|
||||||
|
_IDENTITY_FIELD_PATTERN = re.compile(r"(?:^|[._])(?:id|uuid|key|code)$|(?:^|[._]).+_id$")
|
||||||
|
_MOJIBAKE_MARKERS = ("\ufffd", "锟斤拷", "烫烫烫", "屯屯屯", "Ã", "Â", "â€")
|
||||||
|
_JSON_RECORD_ARRAY_KEYS = ("records", "data", "items", "rows")
|
||||||
|
_JSON_ENVELOPE_KEYS = ("response", "payload")
|
||||||
|
_JSON_WRAPPER_METADATA_KEYS = frozenset(
|
||||||
|
{
|
||||||
|
"page",
|
||||||
|
"page_size",
|
||||||
|
"pageSize",
|
||||||
|
"per_page",
|
||||||
|
"perPage",
|
||||||
|
"total",
|
||||||
|
"total_count",
|
||||||
|
"totalCount",
|
||||||
|
"count",
|
||||||
|
"offset",
|
||||||
|
"limit",
|
||||||
|
"cursor",
|
||||||
|
"next_cursor",
|
||||||
|
"nextCursor",
|
||||||
|
"has_more",
|
||||||
|
"hasMore",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
_JSON_RESPONSE_METADATA_KEYS = _JSON_WRAPPER_METADATA_KEYS | {
|
||||||
|
"success",
|
||||||
|
"status",
|
||||||
|
"code",
|
||||||
|
"message",
|
||||||
|
"error",
|
||||||
|
}
|
||||||
|
_NAME_FIELD_NAMES = {
|
||||||
|
"name",
|
||||||
|
"full_name",
|
||||||
|
"fullname",
|
||||||
|
"real_name",
|
||||||
|
"contact_name",
|
||||||
|
"customer_name",
|
||||||
|
"recipient_name",
|
||||||
|
"姓名",
|
||||||
|
"中文姓名",
|
||||||
|
"真实姓名",
|
||||||
|
"联系人",
|
||||||
|
"联系人姓名",
|
||||||
|
"客户姓名",
|
||||||
|
"收件人",
|
||||||
|
"收件人姓名",
|
||||||
|
}
|
||||||
|
|
||||||
|
_EMAIL_PATTERN = re.compile(
|
||||||
|
r"(?<![\w.+-])[A-Za-z0-9.!#$%&'*+/=?^_`{|}~-]+"
|
||||||
|
r"@[A-Za-z0-9](?:[A-Za-z0-9-]{0,61}[A-Za-z0-9])?"
|
||||||
|
r"(?:\.[A-Za-z0-9](?:[A-Za-z0-9-]{0,61}[A-Za-z0-9])?)+(?![\w.-])"
|
||||||
|
)
|
||||||
|
_PHONE_PATTERN = re.compile(r"(?<!\d)(?:(?:\+|00)?86[-\s]?)?1[3-9]\d{9}(?!\d)")
|
||||||
|
_ID_CARD_PATTERN = re.compile(r"(?<!\d)(?:\d{17}[\dXx]|\d{15})(?!\d)")
|
||||||
|
_CHINESE_NAME_CONTEXT_PATTERN = re.compile(
|
||||||
|
r"(?P<label>姓名|真实姓名|联系人(?:姓名)?|收件人)"
|
||||||
|
r"(?P<separator>\s*(?:[::=]|为)\s*|\s+)"
|
||||||
|
r"(?P<name>[\u3400-\u4dbf\u4e00-\u9fff·]{2,8})"
|
||||||
|
)
|
||||||
|
_ENGLISH_NAME_CONTEXT_PATTERN = re.compile(
|
||||||
|
r"(?im)(?P<label>full\s+name|contact\s+name|name)"
|
||||||
|
r"(?P<separator>\s*[:=]\s*)"
|
||||||
|
r"(?P<name>[A-Za-z][A-Za-z'’-]*(?:[ \t]+[A-Za-z][A-Za-z'’-]*){0,3})"
|
||||||
|
)
|
||||||
|
_TOKEN_PATTERN = re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]|[A-Za-z0-9_]+|[^\s]")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ParsedText:
|
||||||
|
"""文本、文档或工作簿的统一解析结果。"""
|
||||||
|
|
||||||
|
format: TextFormat
|
||||||
|
text: str
|
||||||
|
records: tuple[dict[str, Any], ...]
|
||||||
|
record_locators: tuple[dict[str, Any], ...] = ()
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class ProcessedStructuredRecord:
|
||||||
|
"""保留原始记录索引的结构化预处理结果。"""
|
||||||
|
|
||||||
|
source_index: int
|
||||||
|
record: dict[str, Any]
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class PdfPageText:
|
||||||
|
"""PDF 物理页在统一提取文本中的字符范围。"""
|
||||||
|
|
||||||
|
page_number: int
|
||||||
|
text: str
|
||||||
|
source_start: int
|
||||||
|
source_end: int
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class DocumentNoiseSpan:
|
||||||
|
"""PDF 中可安全从展示内容移除的文本范围。"""
|
||||||
|
|
||||||
|
start: int
|
||||||
|
end: int
|
||||||
|
kind: Literal["page_number", "repeated_margin", "table_of_contents"]
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class QualityScore:
|
||||||
|
"""标准 instruction/input/output 记录的可解释质量分。"""
|
||||||
|
|
||||||
|
overall: float
|
||||||
|
completeness: float
|
||||||
|
length: float
|
||||||
|
readability: float
|
||||||
|
relevance: float
|
||||||
|
duplicate: float
|
||||||
|
is_valid: bool
|
||||||
|
flags: tuple[str, ...]
|
||||||
|
fingerprint: str
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class DocumentHeading:
|
||||||
|
"""文档标题的位置和层级。"""
|
||||||
|
|
||||||
|
level: int
|
||||||
|
title: str
|
||||||
|
line_number: int
|
||||||
|
start: int
|
||||||
|
end: int
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class DocumentStructure:
|
||||||
|
"""无需外部模型即可复现的文档结构摘要。"""
|
||||||
|
|
||||||
|
line_count: int
|
||||||
|
paragraph_count: int
|
||||||
|
headings: tuple[DocumentHeading, ...]
|
||||||
|
code_block_count: int
|
||||||
|
table_block_count: int
|
||||||
|
list_block_count: int
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class _PdfLine:
|
||||||
|
text: str
|
||||||
|
start: int
|
||||||
|
end: int
|
||||||
|
|
||||||
|
class _DuplicateJsonKeyError(ValueError):
|
||||||
|
"""严格 JSON 解析时发现同一对象内的重复键。"""
|
||||||
7
backend/app/modules/data_process/constants.py
Normal file
7
backend/app/modules/data_process/constants.py
Normal file
@@ -0,0 +1,7 @@
|
|||||||
|
"""数据处理模块的共享限制。"""
|
||||||
|
|
||||||
|
MAX_QA_PAIRS_PER_ITEM = 50
|
||||||
|
MODEL_GENERATION_BATCH_SIZE = 10
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["MAX_QA_PAIRS_PER_ITEM", "MODEL_GENERATION_BATCH_SIZE"]
|
||||||
158
backend/app/modules/data_process/dataset_format.py
Normal file
158
backend/app/modules/data_process/dataset_format.py
Normal file
@@ -0,0 +1,158 @@
|
|||||||
|
"""Dataset format validation for Alpaca, ShareGPT, DPO, CPT formats.
|
||||||
|
|
||||||
|
Used by the training preflight flow to validate that uploaded dataset files
|
||||||
|
conform to the declared format before submitting to the compute node.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
|
def _load_sample(path: str | None, content: str | None = None, max_samples: int = 20) -> list[dict[str, Any]]:
|
||||||
|
"""Load up to max_samples records from JSONL file path or raw content string."""
|
||||||
|
try:
|
||||||
|
if content is not None:
|
||||||
|
text = content.strip()
|
||||||
|
elif path:
|
||||||
|
with open(path, "r", encoding="utf-8") as fh:
|
||||||
|
text = fh.read().strip()
|
||||||
|
else:
|
||||||
|
return []
|
||||||
|
except Exception:
|
||||||
|
return []
|
||||||
|
|
||||||
|
if not text:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# 先按整文件 JSON(数组/单对象)解析,兼容 .json;失败再按 jsonl 逐行解析
|
||||||
|
try:
|
||||||
|
value = json.loads(text)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
value = None
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [item for item in value[:max_samples] if isinstance(item, dict)]
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return [value]
|
||||||
|
|
||||||
|
lines = text.splitlines()[:max_samples]
|
||||||
|
records: list[dict[str, Any]] = []
|
||||||
|
for line in lines:
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
record = json.loads(line)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
continue
|
||||||
|
if isinstance(record, dict):
|
||||||
|
records.append(record)
|
||||||
|
return records
|
||||||
|
|
||||||
|
|
||||||
|
def _check_alpaca(records: list[dict[str, Any]]) -> list[str]:
|
||||||
|
"""Validate Alpaca format: requires 'instruction' field."""
|
||||||
|
errors: list[str] = []
|
||||||
|
if not records:
|
||||||
|
errors.append("Alpaca 格式数据集无有效记录")
|
||||||
|
return errors
|
||||||
|
missing_instruction = sum(1 for r in records if not r.get("instruction"))
|
||||||
|
if missing_instruction:
|
||||||
|
errors.append(
|
||||||
|
f"Alpaca 格式要求每条记录包含 instruction 字段,"
|
||||||
|
f"前{len(records)}条中有{missing_instruction}条缺失"
|
||||||
|
)
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
def _check_sharegpt(records: list[dict[str, Any]]) -> list[str]:
|
||||||
|
"""Validate ShareGPT format: requires 'messages' (list of dicts with role/content)."""
|
||||||
|
errors: list[str] = []
|
||||||
|
if not records:
|
||||||
|
errors.append("ShareGPT 格式数据集无有效记录")
|
||||||
|
return errors
|
||||||
|
bad = 0
|
||||||
|
for r in records:
|
||||||
|
messages = r.get("messages")
|
||||||
|
if not isinstance(messages, list) or not messages:
|
||||||
|
bad += 1
|
||||||
|
continue
|
||||||
|
for msg in messages:
|
||||||
|
if not isinstance(msg, dict) or "role" not in msg or "content" not in msg:
|
||||||
|
bad += 1
|
||||||
|
break
|
||||||
|
if bad:
|
||||||
|
errors.append(
|
||||||
|
f"ShareGPT 格式要求每条记录包含 messages 列表,"
|
||||||
|
f"每条消息需有 role 和 content 字段,前{len(records)}条中有{bad}条不符合"
|
||||||
|
)
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
def _check_dpo(records: list[dict[str, Any]]) -> list[str]:
|
||||||
|
"""Validate DPO format: requires 'chosen' and 'rejected' fields."""
|
||||||
|
errors: list[str] = []
|
||||||
|
if not records:
|
||||||
|
errors.append("DPO 格式数据集无有效记录")
|
||||||
|
return errors
|
||||||
|
missing_chosen = sum(1 for r in records if not r.get("chosen"))
|
||||||
|
missing_rejected = sum(1 for r in records if not r.get("rejected"))
|
||||||
|
if missing_chosen:
|
||||||
|
errors.append(f"DPO 格式要求 chosen 字段,前{len(records)}条中有{missing_chosen}条缺失")
|
||||||
|
if missing_rejected:
|
||||||
|
errors.append(f"DPO 格式要求 rejected 字段,前{len(records)}条中有{missing_rejected}条缺失")
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
def _check_cpt(records: list[dict[str, Any]]) -> list[str]:
|
||||||
|
"""Validate CPT format: requires 'text' field, should NOT have instruction/output."""
|
||||||
|
errors: list[str] = []
|
||||||
|
if not records:
|
||||||
|
errors.append("CPT 格式数据集无有效记录")
|
||||||
|
return errors
|
||||||
|
missing_text = sum(1 for r in records if not r.get("text"))
|
||||||
|
has_instruction = sum(1 for r in records if r.get("instruction") or r.get("output"))
|
||||||
|
if missing_text:
|
||||||
|
errors.append(f"CPT 格式要求 text 字段,前{len(records)}条中有{missing_text}条缺失")
|
||||||
|
if has_instruction:
|
||||||
|
errors.append(
|
||||||
|
f"CPT 格式不应包含 instruction/output 字段(疑似 Alpaca 格式),"
|
||||||
|
f"前{len(records)}条中有{has_instruction}条包含此类字段"
|
||||||
|
)
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
FORMAT_VALIDATORS = {
|
||||||
|
"alpaca": _check_alpaca,
|
||||||
|
"alpaca_jsonl": _check_alpaca,
|
||||||
|
"sharegpt": _check_sharegpt,
|
||||||
|
"dpo": _check_dpo,
|
||||||
|
"cpt": _check_cpt,
|
||||||
|
"pt": _check_cpt,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def validate_dataset_format(
|
||||||
|
dataset_format: str,
|
||||||
|
content: str | None = None,
|
||||||
|
path: str | None = None,
|
||||||
|
max_samples: int = 20,
|
||||||
|
) -> list[str]:
|
||||||
|
"""Validate dataset content against expected format.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
dataset_format: One of 'alpaca', 'sharegpt', 'dpo', 'cpt'.
|
||||||
|
content: Raw file content (JSONL text). Mutually exclusive with path.
|
||||||
|
path: File path to read content from.
|
||||||
|
max_samples: Maximum records to sample for validation.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of error messages (empty if valid).
|
||||||
|
"""
|
||||||
|
fmt = str(dataset_format).lower().strip()
|
||||||
|
validator = FORMAT_VALIDATORS.get(fmt)
|
||||||
|
if not validator:
|
||||||
|
return [f"不支持的数据集格式: {dataset_format},支持的格式: {', '.join(sorted(FORMAT_VALIDATORS))}"]
|
||||||
|
records = _load_sample(path=path, content=content, max_samples=max_samples)
|
||||||
|
return validator(records)
|
||||||
509
backend/app/modules/data_process/document_chunking.py
Normal file
509
backend/app/modules/data_process/document_chunking.py
Normal file
@@ -0,0 +1,509 @@
|
|||||||
|
"""基于 Docling 与 LlamaIndex 的文档切分实现。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import threading
|
||||||
|
import unicodedata
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from functools import lru_cache
|
||||||
|
from io import BytesIO
|
||||||
|
from typing import Any, Literal
|
||||||
|
|
||||||
|
import tiktoken
|
||||||
|
from docling_core.transforms.chunker.hierarchical_chunker import ChunkingSerializerProvider
|
||||||
|
from llama_index.core import Document
|
||||||
|
from llama_index.core.base.embeddings.base import BaseEmbedding
|
||||||
|
from llama_index.core.node_parser import SemanticSplitterNodeParser, SentenceSplitter
|
||||||
|
|
||||||
|
from app.modules.data_process.algorithms import normalize_text
|
||||||
|
|
||||||
|
ChunkMethod = Literal["layout_hybrid", "semantic", "fixed"]
|
||||||
|
|
||||||
|
_PAGE_FURNITURE = re.compile(
|
||||||
|
r"(?m)^\s*(?:第\s*\d+\s*页\s*共\s*\d+\s*页|[-—–]?\s*\d+\s*[//]\s*\d+\s*[-—–]?)\s*$"
|
||||||
|
)
|
||||||
|
_COMPACT_CHARACTER = re.compile(r"[\w\u3400-\u4dbf\u4e00-\u9fff]", re.UNICODE)
|
||||||
|
_CONVERTER_LOCK = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class DocumentChunk:
|
||||||
|
"""切片正文及其在原文件中的可追溯信息。"""
|
||||||
|
|
||||||
|
original_content: str
|
||||||
|
contextualized_content: str
|
||||||
|
source_start: int | None
|
||||||
|
source_end: int | None
|
||||||
|
source_start_line: int | None
|
||||||
|
source_end_line: int | None
|
||||||
|
token_count: int
|
||||||
|
heading_path: tuple[str, ...] = ()
|
||||||
|
source_pages: tuple[int, ...] = ()
|
||||||
|
doc_item_refs: tuple[str, ...] = ()
|
||||||
|
source_bboxes: tuple[dict[str, Any], ...] = ()
|
||||||
|
|
||||||
|
|
||||||
|
def _sentence_chunks(text: str) -> list[str]:
|
||||||
|
"""提供稳定的中英文句界,避免 LlamaIndex 默认分词器下载额外资源。"""
|
||||||
|
|
||||||
|
boundary = re.compile(
|
||||||
|
r".*?(?:\n\s*\n|[。!?!?;;](?:[\"'”’)】》]*)|\.(?:\s+|$)|$)",
|
||||||
|
re.DOTALL,
|
||||||
|
)
|
||||||
|
return [part for part in boundary.findall(text) if part]
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def _tokenizer() -> tiktoken.Encoding:
|
||||||
|
"""加载 cl100k_base 编码器,优先在线下载,失败时使用本地缓存以支持离线环境。"""
|
||||||
|
import os
|
||||||
|
import base64
|
||||||
|
|
||||||
|
# 先设置缓存目录环境变量
|
||||||
|
offline_cache = os.path.expanduser("~/.cache/tiktoken")
|
||||||
|
os.environ.setdefault("TIKTOKEN_CACHE_DIR", offline_cache)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 尝试标准方式加载
|
||||||
|
return tiktoken.get_encoding("cl100k_base")
|
||||||
|
except Exception:
|
||||||
|
# 如果失败,尝试手动从本地文件构造
|
||||||
|
try:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
local_file = Path(offline_cache) / "9b5ad71b2ce5302211f9c61530b329a4922fc6a4"
|
||||||
|
if not local_file.exists():
|
||||||
|
# 尝试另一个可能的文件名
|
||||||
|
local_file = Path(offline_cache) / "cl100k_base.tiktoken"
|
||||||
|
|
||||||
|
if local_file.exists():
|
||||||
|
# 读取 BPE 文件内容
|
||||||
|
with open(local_file, "rb") as f:
|
||||||
|
contents = f.read()
|
||||||
|
|
||||||
|
# 解析 BPE 文件
|
||||||
|
mergeable_ranks = {}
|
||||||
|
for line in contents.splitlines():
|
||||||
|
if line:
|
||||||
|
token, rank = line.split()
|
||||||
|
mergeable_ranks[base64.b64decode(token)] = int(rank)
|
||||||
|
|
||||||
|
# 构造 Encoding 对象(模块顶部已 import tiktoken,
|
||||||
|
# 此处不能再 import tiktoken.core,否则会把 tiktoken
|
||||||
|
# 变成局部变量,使函数开头的 tiktoken.get_encoding 抛
|
||||||
|
# UnboundLocalError)
|
||||||
|
return tiktoken.core.Encoding(
|
||||||
|
name="cl100k_base",
|
||||||
|
pat_str=r"""'(?i:[sdmt]|ll|ve|re)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]++[\r\n]*|\s*[\r\n]|\s+(?!\S)|\s+""",
|
||||||
|
mergeable_ranks=mergeable_ranks,
|
||||||
|
special_tokens={
|
||||||
|
"<|endoftext|>": 100257,
|
||||||
|
"<|fim_prefix|>": 100258,
|
||||||
|
"<|fim_middle|>": 100259,
|
||||||
|
"<|fim_suffix|>": 100260,
|
||||||
|
"<|endofprompt|>": 100276,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
raise RuntimeError(
|
||||||
|
f"无法加载 cl100k_base 编码器\n"
|
||||||
|
f"请确保以下任一条件满足:\n"
|
||||||
|
f"1. 服务器可以访问网络\n"
|
||||||
|
f"2. 本地存在缓存文件: {offline_cache}/9b5ad71b2ce5302211f9c61530b329a4922fc6a4"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _text_chunks(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
chunk_size: int,
|
||||||
|
chunk_overlap: int,
|
||||||
|
) -> list[DocumentChunk]:
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
if not normalized:
|
||||||
|
return []
|
||||||
|
splitter = SentenceSplitter(
|
||||||
|
chunk_size=chunk_size,
|
||||||
|
chunk_overlap=chunk_overlap,
|
||||||
|
tokenizer=_tokenizer().encode,
|
||||||
|
chunking_tokenizer_fn=_sentence_chunks,
|
||||||
|
include_metadata=False,
|
||||||
|
include_prev_next_rel=False,
|
||||||
|
)
|
||||||
|
nodes = splitter.get_nodes_from_documents([Document(text=normalized)])
|
||||||
|
return _nodes_to_chunks(nodes, normalized)
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_fixed_text(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
chunk_size: int,
|
||||||
|
chunk_overlap: int,
|
||||||
|
) -> list[DocumentChunk]:
|
||||||
|
"""使用 LlamaIndex SentenceSplitter 按句界控制固定 Token 长度。"""
|
||||||
|
|
||||||
|
return _text_chunks(text, chunk_size=chunk_size, chunk_overlap=chunk_overlap)
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def _semantic_embedding_model() -> BaseEmbedding:
|
||||||
|
# 模型可在部署环境覆盖;默认模型体积较小且适合中英文语义边界判断。
|
||||||
|
from llama_index.embeddings.huggingface import HuggingFaceEmbedding
|
||||||
|
|
||||||
|
return HuggingFaceEmbedding(
|
||||||
|
model_name=os.getenv("DATA_PROCESS_EMBEDDING_MODEL", "BAAI/bge-small-zh-v1.5"),
|
||||||
|
device=os.getenv("DATA_PROCESS_EMBEDDING_DEVICE", "cpu"),
|
||||||
|
trust_remote_code=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_semantic_text(
|
||||||
|
text: str,
|
||||||
|
*,
|
||||||
|
chunk_size: int,
|
||||||
|
chunk_overlap: int,
|
||||||
|
breakpoint_percentile_threshold: int,
|
||||||
|
embed_model: BaseEmbedding | None = None,
|
||||||
|
) -> list[DocumentChunk]:
|
||||||
|
"""使用 LlamaIndex SemanticSplitter 识别主题跳变,再限制最大长度。"""
|
||||||
|
|
||||||
|
normalized = normalize_text(text)
|
||||||
|
if not normalized:
|
||||||
|
return []
|
||||||
|
splitter = SemanticSplitterNodeParser.from_defaults(
|
||||||
|
embed_model=embed_model or _semantic_embedding_model(),
|
||||||
|
breakpoint_percentile_threshold=breakpoint_percentile_threshold,
|
||||||
|
buffer_size=1,
|
||||||
|
sentence_splitter=_sentence_chunks,
|
||||||
|
include_metadata=False,
|
||||||
|
include_prev_next_rel=False,
|
||||||
|
)
|
||||||
|
semantic_nodes = splitter.get_nodes_from_documents([Document(text=normalized)])
|
||||||
|
result: list[DocumentChunk] = []
|
||||||
|
search_from = 0
|
||||||
|
for node in semantic_nodes:
|
||||||
|
content = node.get_content().strip()
|
||||||
|
if not content:
|
||||||
|
continue
|
||||||
|
start = _locate_text(normalized, content, search_from)
|
||||||
|
if start is None:
|
||||||
|
start = _locate_text(normalized, content, 0)
|
||||||
|
if start is None:
|
||||||
|
continue
|
||||||
|
if len(_tokenizer().encode(content)) <= chunk_size:
|
||||||
|
result.append(_make_text_chunk(normalized, start, start + len(content)))
|
||||||
|
else:
|
||||||
|
for child in _text_chunks(
|
||||||
|
content,
|
||||||
|
chunk_size=chunk_size,
|
||||||
|
chunk_overlap=chunk_overlap,
|
||||||
|
):
|
||||||
|
if child.source_start is None or child.source_end is None:
|
||||||
|
continue
|
||||||
|
result.append(
|
||||||
|
_make_text_chunk(
|
||||||
|
normalized,
|
||||||
|
start + child.source_start,
|
||||||
|
start + child.source_end,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
search_from = start + len(content)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _nodes_to_chunks(nodes: list[Any], source_text: str) -> list[DocumentChunk]:
|
||||||
|
chunks: list[DocumentChunk] = []
|
||||||
|
search_from = 0
|
||||||
|
for node in nodes:
|
||||||
|
content = node.get_content().strip()
|
||||||
|
if not content:
|
||||||
|
continue
|
||||||
|
raw_start = getattr(node, "start_char_idx", None)
|
||||||
|
raw_end = getattr(node, "end_char_idx", None)
|
||||||
|
if (
|
||||||
|
isinstance(raw_start, int)
|
||||||
|
and isinstance(raw_end, int)
|
||||||
|
and source_text[raw_start:raw_end].strip() == content
|
||||||
|
):
|
||||||
|
start = raw_start + len(source_text[raw_start:raw_end]) - len(source_text[raw_start:raw_end].lstrip())
|
||||||
|
else:
|
||||||
|
start = _locate_text(source_text, content, search_from)
|
||||||
|
if start is None:
|
||||||
|
start = _locate_text(source_text, content, 0)
|
||||||
|
if start is None:
|
||||||
|
continue
|
||||||
|
end = start + len(content)
|
||||||
|
chunks.append(_make_text_chunk(source_text, start, end))
|
||||||
|
search_from = max(search_from, end)
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
|
||||||
|
def _locate_text(source: str, content: str, start: int) -> int | None:
|
||||||
|
position = source.find(content, start)
|
||||||
|
return position if position >= 0 else None
|
||||||
|
|
||||||
|
|
||||||
|
def _make_text_chunk(source: str, start: int, end: int) -> DocumentChunk:
|
||||||
|
content = source[start:end]
|
||||||
|
return DocumentChunk(
|
||||||
|
original_content=content,
|
||||||
|
contextualized_content=content,
|
||||||
|
source_start=start,
|
||||||
|
source_end=end,
|
||||||
|
source_start_line=source.count("\n", 0, start) + 1,
|
||||||
|
source_end_line=source.count("\n", 0, max(start, end - 1)) + 1,
|
||||||
|
token_count=len(_tokenizer().encode(content)),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=1)
|
||||||
|
def _document_converter():
|
||||||
|
from docling.datamodel.base_models import InputFormat
|
||||||
|
from docling.datamodel.pipeline_options import PdfPipelineOptions
|
||||||
|
from docling.document_converter import DocumentConverter, PdfFormatOption
|
||||||
|
|
||||||
|
pipeline_options = PdfPipelineOptions()
|
||||||
|
pipeline_options.do_ocr = False
|
||||||
|
return DocumentConverter(
|
||||||
|
format_options={
|
||||||
|
InputFormat.PDF: PdfFormatOption(pipeline_options=pipeline_options),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _MarkdownSerializerProvider(ChunkingSerializerProvider):
|
||||||
|
def get_serializer(self, doc: Any):
|
||||||
|
from docling_core.transforms.chunker.hierarchical_chunker import ChunkingDocSerializer
|
||||||
|
from docling_core.transforms.serializer.markdown import (
|
||||||
|
MarkdownParams,
|
||||||
|
MarkdownTableSerializer,
|
||||||
|
)
|
||||||
|
from docling_core.types.doc import DocItemLabel
|
||||||
|
|
||||||
|
excluded = {
|
||||||
|
DocItemLabel.DOCUMENT_INDEX,
|
||||||
|
DocItemLabel.PAGE_HEADER,
|
||||||
|
DocItemLabel.PAGE_FOOTER,
|
||||||
|
}
|
||||||
|
return ChunkingDocSerializer(
|
||||||
|
doc=doc,
|
||||||
|
table_serializer=MarkdownTableSerializer(),
|
||||||
|
params=MarkdownParams(
|
||||||
|
labels=set(DocItemLabel) - excluded,
|
||||||
|
compact_tables=True,
|
||||||
|
image_placeholder="",
|
||||||
|
escape_html=False,
|
||||||
|
escape_underscores=False,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_layout_text(value: str) -> str:
|
||||||
|
return normalize_text(_PAGE_FURNITURE.sub("", value)).strip()
|
||||||
|
|
||||||
|
|
||||||
|
def _compact_with_offsets(value: str) -> tuple[str, list[int]]:
|
||||||
|
compact: list[str] = []
|
||||||
|
offsets: list[int] = []
|
||||||
|
for index, character in enumerate(unicodedata.normalize("NFKC", value)):
|
||||||
|
if _COMPACT_CHARACTER.fullmatch(character):
|
||||||
|
compact.append(character.casefold())
|
||||||
|
offsets.append(index)
|
||||||
|
return "".join(compact), offsets
|
||||||
|
|
||||||
|
|
||||||
|
def _project_layout_span(
|
||||||
|
source_text: str,
|
||||||
|
content: str,
|
||||||
|
*,
|
||||||
|
compact_source: str,
|
||||||
|
source_offsets: list[int],
|
||||||
|
compact_start: int,
|
||||||
|
) -> tuple[int | None, int | None, int]:
|
||||||
|
compact_content, _ = _compact_with_offsets(content)
|
||||||
|
if len(compact_content) < 4:
|
||||||
|
return None, None, compact_start
|
||||||
|
position = compact_source.find(compact_content, compact_start)
|
||||||
|
if position < 0:
|
||||||
|
position = compact_source.find(compact_content)
|
||||||
|
if position < 0:
|
||||||
|
return None, None, compact_start
|
||||||
|
start = source_offsets[position]
|
||||||
|
end = source_offsets[position + len(compact_content) - 1] + 1
|
||||||
|
while start > 0 and source_text[start - 1] not in "\r\n":
|
||||||
|
start -= 1
|
||||||
|
while end < len(source_text) and source_text[end] not in "\r\n":
|
||||||
|
end += 1
|
||||||
|
return start, end, position + len(compact_content)
|
||||||
|
|
||||||
|
|
||||||
|
def chunk_layout_document(
|
||||||
|
raw: bytes,
|
||||||
|
*,
|
||||||
|
filename: str,
|
||||||
|
source_text: str,
|
||||||
|
chunk_size: int,
|
||||||
|
) -> list[DocumentChunk]:
|
||||||
|
"""使用 Docling HybridChunker 按版面层级、列表与表格边界切分。"""
|
||||||
|
|
||||||
|
from docling.chunking import HybridChunker
|
||||||
|
from docling.datamodel.base_models import DocumentStream
|
||||||
|
from docling.exceptions import BaseError as DoclingError
|
||||||
|
from docling_core.transforms.chunker.tokenizer.openai import OpenAITokenizer
|
||||||
|
from docling_core.types.doc import DocItemLabel
|
||||||
|
|
||||||
|
try:
|
||||||
|
with _CONVERTER_LOCK:
|
||||||
|
conversion = _document_converter().convert(
|
||||||
|
DocumentStream(name=filename, stream=BytesIO(raw))
|
||||||
|
)
|
||||||
|
except DoclingError as exc:
|
||||||
|
raise ValueError(f"文档版面解析失败: {exc}") from exc
|
||||||
|
chunker = HybridChunker(
|
||||||
|
tokenizer=OpenAITokenizer(tokenizer=_tokenizer(), max_tokens=chunk_size),
|
||||||
|
serializer_provider=_MarkdownSerializerProvider(),
|
||||||
|
merge_peers=True,
|
||||||
|
repeat_table_header=True,
|
||||||
|
)
|
||||||
|
compact_source, source_offsets = _compact_with_offsets(source_text)
|
||||||
|
compact_start = 0
|
||||||
|
result: list[DocumentChunk] = []
|
||||||
|
excluded = {
|
||||||
|
DocItemLabel.DOCUMENT_INDEX,
|
||||||
|
DocItemLabel.PAGE_HEADER,
|
||||||
|
DocItemLabel.PAGE_FOOTER,
|
||||||
|
}
|
||||||
|
for raw_chunk in chunker.chunk(conversion.document):
|
||||||
|
doc_items = tuple(raw_chunk.meta.doc_items or ())
|
||||||
|
if doc_items and all(item.label in excluded for item in doc_items):
|
||||||
|
continue
|
||||||
|
content = _clean_layout_text(raw_chunk.text)
|
||||||
|
if not content:
|
||||||
|
continue
|
||||||
|
contextualized = _clean_layout_text(chunker.contextualize(raw_chunk)) or content
|
||||||
|
start, end, compact_start = _project_layout_span(
|
||||||
|
source_text,
|
||||||
|
content,
|
||||||
|
compact_source=compact_source,
|
||||||
|
source_offsets=source_offsets,
|
||||||
|
compact_start=compact_start,
|
||||||
|
)
|
||||||
|
original = source_text[start:end] if start is not None and end is not None else content
|
||||||
|
pages: set[int] = set()
|
||||||
|
refs: list[str] = []
|
||||||
|
bboxes: list[dict[str, Any]] = []
|
||||||
|
for item in doc_items:
|
||||||
|
refs.append(str(item.self_ref))
|
||||||
|
for provenance in item.prov or ():
|
||||||
|
pages.add(int(provenance.page_no))
|
||||||
|
bbox = provenance.bbox
|
||||||
|
bboxes.append(
|
||||||
|
{
|
||||||
|
"page": int(provenance.page_no),
|
||||||
|
"left": float(bbox.l),
|
||||||
|
"top": float(bbox.t),
|
||||||
|
"right": float(bbox.r),
|
||||||
|
"bottom": float(bbox.b),
|
||||||
|
"origin": str(bbox.coord_origin.value),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
result.append(
|
||||||
|
DocumentChunk(
|
||||||
|
original_content=original,
|
||||||
|
contextualized_content=contextualized,
|
||||||
|
source_start=start,
|
||||||
|
source_end=end,
|
||||||
|
source_start_line=(source_text.count("\n", 0, start) + 1 if start is not None else None),
|
||||||
|
source_end_line=(
|
||||||
|
source_text.count("\n", 0, max(start or 0, (end or 1) - 1)) + 1
|
||||||
|
if end is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
token_count=len(_tokenizer().encode(contextualized)),
|
||||||
|
heading_path=tuple(str(item) for item in (raw_chunk.meta.headings or ())),
|
||||||
|
source_pages=tuple(sorted(pages)),
|
||||||
|
doc_item_refs=tuple(refs),
|
||||||
|
source_bboxes=tuple(bboxes),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def merge_short_chunks(
|
||||||
|
chunks: list[DocumentChunk],
|
||||||
|
*,
|
||||||
|
source_text: str,
|
||||||
|
min_token_count: int,
|
||||||
|
max_token_count: int,
|
||||||
|
) -> list[DocumentChunk]:
|
||||||
|
"""在不突破长度上限的前提下,把过短块并入相邻内容。"""
|
||||||
|
|
||||||
|
result: list[DocumentChunk] = []
|
||||||
|
index = 0
|
||||||
|
while index < len(chunks):
|
||||||
|
current = chunks[index]
|
||||||
|
if current.token_count >= min_token_count:
|
||||||
|
result.append(current)
|
||||||
|
index += 1
|
||||||
|
continue
|
||||||
|
if index + 1 < len(chunks):
|
||||||
|
combined = _combine_chunks(current, chunks[index + 1], source_text)
|
||||||
|
if combined.token_count <= max_token_count:
|
||||||
|
result.append(combined)
|
||||||
|
index += 2
|
||||||
|
continue
|
||||||
|
if result:
|
||||||
|
combined = _combine_chunks(result[-1], current, source_text)
|
||||||
|
if combined.token_count <= max_token_count:
|
||||||
|
result[-1] = combined
|
||||||
|
index += 1
|
||||||
|
continue
|
||||||
|
result.append(current)
|
||||||
|
index += 1
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _combine_chunks(
|
||||||
|
left: DocumentChunk,
|
||||||
|
right: DocumentChunk,
|
||||||
|
source_text: str,
|
||||||
|
) -> DocumentChunk:
|
||||||
|
contextualized = "\n\n".join(
|
||||||
|
part for part in (left.contextualized_content, right.contextualized_content) if part
|
||||||
|
)
|
||||||
|
start = left.source_start
|
||||||
|
end = right.source_end
|
||||||
|
has_contiguous_source = (
|
||||||
|
start is not None
|
||||||
|
and left.source_end is not None
|
||||||
|
and right.source_start is not None
|
||||||
|
and end is not None
|
||||||
|
and left.source_end <= right.source_start
|
||||||
|
)
|
||||||
|
original = (
|
||||||
|
source_text[start:end]
|
||||||
|
if has_contiguous_source and start is not None and end is not None
|
||||||
|
else "\n\n".join(
|
||||||
|
part for part in (left.original_content, right.original_content) if part
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if not has_contiguous_source:
|
||||||
|
start = None
|
||||||
|
end = None
|
||||||
|
return DocumentChunk(
|
||||||
|
original_content=original,
|
||||||
|
contextualized_content=contextualized,
|
||||||
|
source_start=start,
|
||||||
|
source_end=end,
|
||||||
|
source_start_line=left.source_start_line if start is not None else None,
|
||||||
|
source_end_line=right.source_end_line if end is not None else None,
|
||||||
|
token_count=len(_tokenizer().encode(contextualized)),
|
||||||
|
heading_path=left.heading_path or right.heading_path,
|
||||||
|
source_pages=tuple(sorted(set(left.source_pages) | set(right.source_pages))),
|
||||||
|
doc_item_refs=left.doc_item_refs + right.doc_item_refs,
|
||||||
|
source_bboxes=left.source_bboxes + right.source_bboxes,
|
||||||
|
)
|
||||||
641
backend/app/modules/data_process/generation.py
Normal file
641
backend/app/modules/data_process/generation.py
Normal file
@@ -0,0 +1,641 @@
|
|||||||
|
"""数据处理任务的大模型生成适配器。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from collections.abc import Callable, Iterable, Mapping
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import urlsplit, urlunsplit
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from app.modules.data_process.algorithms import normalize_text, stable_split_assignments
|
||||||
|
from app.modules.data_process.constants import (
|
||||||
|
MAX_QA_PAIRS_PER_ITEM,
|
||||||
|
MODEL_GENERATION_BATCH_SIZE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelGenerationError(ValueError):
|
||||||
|
"""模型配置、响应或调用失败。"""
|
||||||
|
|
||||||
|
|
||||||
|
class _TerminalModelGenerationError(ModelGenerationError):
|
||||||
|
"""使用相同参数重试也无法恢复的模型响应错误。"""
|
||||||
|
|
||||||
|
|
||||||
|
OUTPUT_TYPE_STANDARD = "standard"
|
||||||
|
OUTPUT_TYPE_REASONING = "reasoning"
|
||||||
|
OUTPUT_TYPE_DPO = "dpo"
|
||||||
|
SUPPORTED_OUTPUT_TYPES = {
|
||||||
|
OUTPUT_TYPE_STANDARD,
|
||||||
|
OUTPUT_TYPE_REASONING,
|
||||||
|
OUTPUT_TYPE_DPO,
|
||||||
|
}
|
||||||
|
REASONING_DETAIL_NORMAL = "normal"
|
||||||
|
REASONING_DETAIL_DETAILED = "detailed"
|
||||||
|
SUPPORTED_REASONING_DETAILS = {
|
||||||
|
REASONING_DETAIL_NORMAL,
|
||||||
|
REASONING_DETAIL_DETAILED,
|
||||||
|
}
|
||||||
|
MINIMAX_M3_API_HOSTS = {"api.minimax.io", "api.minimaxi.com"}
|
||||||
|
MINIMAX_M3_MIN_COMPLETION_TOKENS = 4096
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# 问题表述风格规则:防止模型产出“请描述/请说明”式模板化问句。
|
||||||
|
_QUESTION_STYLE_RULE = (
|
||||||
|
"各条问题必须覆盖不同的信息点并使用不同的句式,只替换关键词套用同一句式视为重复。"
|
||||||
|
"问题表述要像真实用户自然提出的问题:具体、口语化、直奔信息点,"
|
||||||
|
"避免“请描述”“请说明”“根据文档”等模板化开头,"
|
||||||
|
"也不要把原文句子直接改成问句;多条问题时交替使用直接疑问、场景式提问、追问式等句式。"
|
||||||
|
"表述示例(仅示意风格,不要照搬内容):"
|
||||||
|
"避免——“请描述系统的权限控制机制”;"
|
||||||
|
"推荐——“不同角色能看到的菜单不一样,平台是怎么控制的?”"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 任务配置未提供提示语时的兜底,与前端内置默认提示语保持同等信息量。
|
||||||
|
_DEFAULT_GENERATION_PROMPT = (
|
||||||
|
"你是一名专业的数据生成专家。请基于来源内容生成高质量、"
|
||||||
|
"可直接用于监督微调的问答数据:问题聚焦核心信息点、"
|
||||||
|
"表述像真实用户自然提出的问题,具体、口语化,多条问题使用不同句式;"
|
||||||
|
"答案严格依据来源内容,准确、完整、语言自然,不引入来源之外的信息。"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_retryable_generation_error(exc: Exception) -> bool:
|
||||||
|
if isinstance(exc, _TerminalModelGenerationError):
|
||||||
|
return False
|
||||||
|
if isinstance(exc, httpx.HTTPStatusError):
|
||||||
|
status_code = exc.response.status_code
|
||||||
|
return status_code in {408, 425, 429} or status_code >= 500
|
||||||
|
if isinstance(exc, httpx.RequestError):
|
||||||
|
return True
|
||||||
|
return isinstance(exc, (json.JSONDecodeError, ModelGenerationError))
|
||||||
|
|
||||||
|
|
||||||
|
def _is_official_minimax_m3(endpoint: str, model_name: str) -> bool:
|
||||||
|
host = (urlsplit(endpoint).hostname or "").casefold()
|
||||||
|
return host in MINIMAX_M3_API_HOSTS and model_name.casefold() == "minimax-m3"
|
||||||
|
|
||||||
|
|
||||||
|
def chat_completions_url(value: str) -> str:
|
||||||
|
"""把域名、基础 URL 或完整地址统一为 chat completions 地址。"""
|
||||||
|
|
||||||
|
raw = (value or "").strip()
|
||||||
|
if not raw:
|
||||||
|
raise ModelGenerationError("generation model api_url is required")
|
||||||
|
if "://" not in raw:
|
||||||
|
raw = f"https://{raw}"
|
||||||
|
parsed = urlsplit(raw)
|
||||||
|
if parsed.scheme not in {"http", "https"} or not parsed.hostname:
|
||||||
|
raise ModelGenerationError("generation model api_url must be an HTTP(S) host or URL")
|
||||||
|
if parsed.username or parsed.password:
|
||||||
|
raise ModelGenerationError("generation model api_url must not contain credentials")
|
||||||
|
|
||||||
|
path = parsed.path.rstrip("/")
|
||||||
|
if path.endswith("/chat/completions"):
|
||||||
|
target_path = path
|
||||||
|
elif path.endswith("/v1"):
|
||||||
|
target_path = f"{path}/chat/completions"
|
||||||
|
elif not path:
|
||||||
|
target_path = "/v1/chat/completions"
|
||||||
|
else:
|
||||||
|
target_path = f"{path}/v1/chat/completions"
|
||||||
|
return urlunsplit((parsed.scheme, parsed.netloc, target_path, "", ""))
|
||||||
|
|
||||||
|
|
||||||
|
def _response_choice(payload: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||||
|
try:
|
||||||
|
choice = payload["choices"][0]
|
||||||
|
except (KeyError, IndexError, TypeError) as exc:
|
||||||
|
raise ModelGenerationError("模型响应缺少 choices[0]") from exc
|
||||||
|
if not isinstance(choice, Mapping):
|
||||||
|
raise ModelGenerationError("模型响应 choices[0] 不是对象")
|
||||||
|
return choice
|
||||||
|
|
||||||
|
|
||||||
|
def _response_finish_reason(payload: Mapping[str, Any]) -> str:
|
||||||
|
try:
|
||||||
|
return str(_response_choice(payload).get("finish_reason") or "").strip().lower()
|
||||||
|
except ModelGenerationError:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _response_content_length(payload: Mapping[str, Any]) -> int:
|
||||||
|
try:
|
||||||
|
message = _response_choice(payload).get("message")
|
||||||
|
if not isinstance(message, Mapping):
|
||||||
|
return 0
|
||||||
|
content = message.get("content")
|
||||||
|
if isinstance(content, str):
|
||||||
|
return len(content)
|
||||||
|
if isinstance(content, list):
|
||||||
|
return sum(
|
||||||
|
len(str(item.get("text") or ""))
|
||||||
|
for item in content
|
||||||
|
if isinstance(item, Mapping)
|
||||||
|
)
|
||||||
|
except ModelGenerationError:
|
||||||
|
pass
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def _raise_for_terminal_response(payload: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||||
|
choice = _response_choice(payload)
|
||||||
|
base_response = payload.get("base_resp")
|
||||||
|
status_code: Any = None
|
||||||
|
status_message = ""
|
||||||
|
if isinstance(base_response, Mapping):
|
||||||
|
status_code = base_response.get("status_code")
|
||||||
|
status_message = re.sub(
|
||||||
|
r"\s+", " ", str(base_response.get("status_msg") or "")
|
||||||
|
).strip()[:200]
|
||||||
|
|
||||||
|
if bool(payload.get("input_sensitive")) or status_code in {1026, "1026"}:
|
||||||
|
raise _TerminalModelGenerationError(
|
||||||
|
f"模型输入触发内容安全拦截(code={status_code or 1026})"
|
||||||
|
)
|
||||||
|
if bool(payload.get("output_sensitive")) or status_code in {1027, "1027"}:
|
||||||
|
raise _TerminalModelGenerationError(
|
||||||
|
f"模型输出触发内容安全拦截(code={status_code or 1027})"
|
||||||
|
)
|
||||||
|
|
||||||
|
finish_reason = str(choice.get("finish_reason") or "").strip().lower()
|
||||||
|
if finish_reason == "length":
|
||||||
|
raise _TerminalModelGenerationError(
|
||||||
|
"模型输出因达到 Token 上限被截断(finish_reason=length),"
|
||||||
|
"请提高最大输出长度后重试"
|
||||||
|
)
|
||||||
|
if finish_reason == "content_filter":
|
||||||
|
raise _TerminalModelGenerationError(
|
||||||
|
"模型输出被内容安全策略拦截(finish_reason=content_filter)"
|
||||||
|
)
|
||||||
|
if finish_reason in {"tool_calls", "function_call"}:
|
||||||
|
raise _TerminalModelGenerationError(
|
||||||
|
f"模型返回了当前生成任务不支持的工具调用(finish_reason={finish_reason})"
|
||||||
|
)
|
||||||
|
if status_code not in {None, "", 0, "0"}:
|
||||||
|
detail = f":{status_message}" if status_message else ""
|
||||||
|
raise _TerminalModelGenerationError(
|
||||||
|
f"模型服务返回业务错误(code={status_code}){detail}"
|
||||||
|
)
|
||||||
|
return choice
|
||||||
|
|
||||||
|
|
||||||
|
def _message_content(payload: Mapping[str, Any]) -> str:
|
||||||
|
choice = _raise_for_terminal_response(payload)
|
||||||
|
message = choice.get("message")
|
||||||
|
if not isinstance(message, Mapping):
|
||||||
|
raise ModelGenerationError("模型响应缺少 choices[0].message")
|
||||||
|
content = message.get("content")
|
||||||
|
if isinstance(content, str):
|
||||||
|
result = content
|
||||||
|
elif isinstance(content, list):
|
||||||
|
parts = [
|
||||||
|
str(item.get("text") or "")
|
||||||
|
for item in content
|
||||||
|
if isinstance(item, Mapping) and item.get("type") in {None, "text", "output_text"}
|
||||||
|
]
|
||||||
|
result = "".join(parts)
|
||||||
|
elif content is None:
|
||||||
|
result = ""
|
||||||
|
else:
|
||||||
|
raise ModelGenerationError("模型响应 content 必须是文本")
|
||||||
|
if not result.strip():
|
||||||
|
raise ModelGenerationError("模型返回的最终内容为空,未生成可解析的 JSON")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _json_documents(content: str) -> list[Any]:
|
||||||
|
decoder = json.JSONDecoder()
|
||||||
|
documents: list[Any] = []
|
||||||
|
cursor = 0
|
||||||
|
while cursor < len(content):
|
||||||
|
match = re.search(r"[\[{]", content[cursor:])
|
||||||
|
if not match:
|
||||||
|
break
|
||||||
|
start = cursor + match.start()
|
||||||
|
try:
|
||||||
|
value, end = decoder.raw_decode(content[start:])
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
cursor = start + 1
|
||||||
|
continue
|
||||||
|
if isinstance(value, (Mapping, list)):
|
||||||
|
documents.append(value)
|
||||||
|
cursor = start + max(end, 1)
|
||||||
|
return documents
|
||||||
|
|
||||||
|
|
||||||
|
def _json_payload(content: str) -> Any:
|
||||||
|
# 只移除模型在 JSON 之前自行输出的思考过程,不能破坏 JSON 字段中的训练内容。
|
||||||
|
cleaned = content.strip()
|
||||||
|
if re.match(r"^\s*<think>", cleaned, flags=re.IGNORECASE) and not re.match(
|
||||||
|
r"^\s*<think>[\s\S]*?</think>", cleaned, flags=re.IGNORECASE
|
||||||
|
):
|
||||||
|
raise ModelGenerationError("模型思考内容未闭合,响应可能已被截断")
|
||||||
|
cleaned = re.sub(
|
||||||
|
r"^\s*(?:<think>[\s\S]*?</think>\s*)+",
|
||||||
|
"",
|
||||||
|
cleaned,
|
||||||
|
count=1,
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
).strip()
|
||||||
|
fenced = re.fullmatch(r"```(?:json)?\s*([\s\S]*?)\s*```", cleaned, flags=re.IGNORECASE)
|
||||||
|
if fenced:
|
||||||
|
cleaned = fenced.group(1).strip()
|
||||||
|
try:
|
||||||
|
return json.loads(cleaned)
|
||||||
|
except json.JSONDecodeError as direct_error:
|
||||||
|
documents = _json_documents(cleaned)
|
||||||
|
if len(documents) == 1:
|
||||||
|
return documents[0]
|
||||||
|
if len(documents) > 1:
|
||||||
|
raise ModelGenerationError("模型响应包含多个 JSON 对象,无法确定应使用哪一个")
|
||||||
|
raise ModelGenerationError(
|
||||||
|
"模型响应中没有找到唯一且完整的 JSON 对象"
|
||||||
|
f"(第 {direct_error.lineno} 行,第 {direct_error.colno} 列)"
|
||||||
|
) from direct_error
|
||||||
|
|
||||||
|
|
||||||
|
def _result_items(payload: Any) -> list[Mapping[str, Any]]:
|
||||||
|
if isinstance(payload, list):
|
||||||
|
values = payload
|
||||||
|
elif isinstance(payload, Mapping):
|
||||||
|
nested = next(
|
||||||
|
(
|
||||||
|
payload[key]
|
||||||
|
for key in ("items", "results", "data", "records")
|
||||||
|
if isinstance(payload.get(key), list)
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
values = nested if isinstance(nested, list) else [payload]
|
||||||
|
else:
|
||||||
|
raise ModelGenerationError("model JSON must be an object or array")
|
||||||
|
items = [item for item in values if isinstance(item, Mapping)]
|
||||||
|
if not items:
|
||||||
|
raise ModelGenerationError("model JSON does not contain result objects")
|
||||||
|
return items
|
||||||
|
|
||||||
|
|
||||||
|
def _prompt_messages(
|
||||||
|
prompt: str,
|
||||||
|
content: str,
|
||||||
|
count: int,
|
||||||
|
*,
|
||||||
|
start_index: int,
|
||||||
|
total_count: int,
|
||||||
|
output_type: str,
|
||||||
|
reasoning_detail: str,
|
||||||
|
) -> list[dict[str, str]]:
|
||||||
|
end_index = start_index + count - 1
|
||||||
|
if output_type == OUTPUT_TYPE_REASONING:
|
||||||
|
schema = '{"items":[{"instruction":"...","input":"...","reasoning":"...","answer":"..."}]}'
|
||||||
|
detail_rule = (
|
||||||
|
"推理详细程度为“详细”:完整展开问题条件、来源依据、中间计算或推导,"
|
||||||
|
"并在得出答案前核对结论;每一步都必须能从来源内容中验证。"
|
||||||
|
if reasoning_detail == REASONING_DETAIL_DETAILED
|
||||||
|
else
|
||||||
|
"推理详细程度为“普通”:只保留得出答案所需的关键依据和必要步骤,"
|
||||||
|
"避免冗长复述、套话和无依据扩展。"
|
||||||
|
)
|
||||||
|
output_rule = (
|
||||||
|
"你正在生成用于训练推理模型的思维链数据,而不是普通问答数据。"
|
||||||
|
"instruction、reasoning 和 answer 均不得为空;reasoning 必须是基于来源内容、"
|
||||||
|
f"可核对的推理过程,answer 只写最终答案。{detail_rule}"
|
||||||
|
"这是思维链输出模式,即使其他提示语要求省略分析,也不得省略 reasoning。"
|
||||||
|
"不要自行添加 <think> 标签,系统会在保存时统一组装。"
|
||||||
|
)
|
||||||
|
elif output_type == OUTPUT_TYPE_DPO:
|
||||||
|
schema = (
|
||||||
|
'{"items":[{"instruction":"...","input":"...",'
|
||||||
|
'"chosen":"...","rejected":"..."}]}'
|
||||||
|
)
|
||||||
|
output_rule = (
|
||||||
|
"你正在生成用于直接偏好优化(DPO)的成对偏好数据。"
|
||||||
|
"instruction、chosen 和 rejected 均不得为空;chosen 必须是忠于来源、"
|
||||||
|
"准确完整的优选回答,rejected 必须是表面合理但存在明确质量缺陷的拒选回答。"
|
||||||
|
"两者不得相同;rejected 不得包含违法危险内容,也不得用空白、乱码或无关文本凑数。"
|
||||||
|
"不要输出分析过程或 <think> 标签。"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
schema = '{"items":[{"instruction":"...","input":"...","output":"..."}]}'
|
||||||
|
output_rule = (
|
||||||
|
"你正在生成标准监督微调问答数据。instruction 和 output 不得为空;"
|
||||||
|
"output 只写最终答案,禁止输出分析、推理过程或 <think> 标签。"
|
||||||
|
)
|
||||||
|
schema_instruction = (
|
||||||
|
f"必须只返回 JSON 对象,格式为 {schema};items 必须包含 {count} 条。"
|
||||||
|
f"这是总计 {total_count} 条中的第 {start_index}-{end_index} 条。"
|
||||||
|
f"{_QUESTION_STYLE_RULE}{output_rule}"
|
||||||
|
"不要输出 Markdown 代码围栏或 JSON 之外的说明。"
|
||||||
|
)
|
||||||
|
base_prompt = normalize_text(prompt) or _DEFAULT_GENERATION_PROMPT
|
||||||
|
if "{{ content }}" in base_prompt:
|
||||||
|
user_prompt = base_prompt.replace("{{ content }}", content)
|
||||||
|
return [
|
||||||
|
{"role": "system", "content": schema_instruction},
|
||||||
|
{"role": "user", "content": user_prompt},
|
||||||
|
]
|
||||||
|
return [
|
||||||
|
{"role": "system", "content": f"{base_prompt}\n{schema_instruction}"},
|
||||||
|
{"role": "user", "content": f"来源内容:\n{content}"},
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def generate_model_records(
|
||||||
|
preview_items: Iterable[Mapping[str, Any]],
|
||||||
|
*,
|
||||||
|
model: Mapping[str, Any],
|
||||||
|
config: Mapping[str, Any],
|
||||||
|
task_id: str,
|
||||||
|
split: Mapping[str, int],
|
||||||
|
qa_pairs_per_item: int,
|
||||||
|
client: httpx.Client | None = None,
|
||||||
|
on_progress: Callable[[int, int], None] | None = None,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""调用 OpenAI 兼容接口,将预览切片生成标准训练记录。
|
||||||
|
|
||||||
|
每个切片按安全批次调用模型;失败批次会产生一条可人工修复的
|
||||||
|
invalid 结果,已经成功的批次不会丢失。
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not 1 <= qa_pairs_per_item <= MAX_QA_PAIRS_PER_ITEM:
|
||||||
|
raise ModelGenerationError(f"qa_pairs_per_item must be in [1, {MAX_QA_PAIRS_PER_ITEM}]")
|
||||||
|
output_type = str(config.get("output_type") or OUTPUT_TYPE_STANDARD).strip().lower()
|
||||||
|
if output_type not in SUPPORTED_OUTPUT_TYPES:
|
||||||
|
raise ModelGenerationError(f"output_type must be one of {sorted(SUPPORTED_OUTPUT_TYPES)}")
|
||||||
|
reasoning_detail = str(
|
||||||
|
config.get("reasoning_detail") or REASONING_DETAIL_NORMAL
|
||||||
|
).strip().lower()
|
||||||
|
if reasoning_detail not in SUPPORTED_REASONING_DETAILS:
|
||||||
|
raise ModelGenerationError(
|
||||||
|
f"reasoning_detail must be one of {sorted(SUPPORTED_REASONING_DETAILS)}"
|
||||||
|
)
|
||||||
|
endpoint = chat_completions_url(str(model.get("api_url") or ""))
|
||||||
|
model_name = str(model.get("online_model_name") or model.get("name") or "").strip()
|
||||||
|
if not model_name:
|
||||||
|
raise ModelGenerationError("generation model name is required")
|
||||||
|
is_minimax_m3 = _is_official_minimax_m3(endpoint, model_name)
|
||||||
|
|
||||||
|
temperature = float(config.get("temperature", 0.7))
|
||||||
|
max_tokens = int(config.get("max_tokens", 1024))
|
||||||
|
timeout = max(1.0, min(120.0, float(config.get("request_timeout_seconds", 60))))
|
||||||
|
retries = max(0, min(5, int(config.get("generation_retries", 2))))
|
||||||
|
headers = {"Content-Type": "application/json"}
|
||||||
|
api_key = str(model.get("api_key") or "").strip()
|
||||||
|
if api_key:
|
||||||
|
headers["Authorization"] = f"Bearer {api_key}"
|
||||||
|
|
||||||
|
owns_client = client is None
|
||||||
|
http_client = client or httpx.Client(timeout=timeout)
|
||||||
|
results: list[dict[str, Any]] = []
|
||||||
|
try:
|
||||||
|
preview_list = list(preview_items)
|
||||||
|
total_items = len(preview_list)
|
||||||
|
for item_index, item in enumerate(preview_list):
|
||||||
|
preview_id = str(item.get("id") or f"preview-{item_index + 1}")
|
||||||
|
content = normalize_text(
|
||||||
|
str(item.get("edited_content") or item.get("original_content") or "")
|
||||||
|
)
|
||||||
|
for batch_offset in range(0, qa_pairs_per_item, MODEL_GENERATION_BATCH_SIZE):
|
||||||
|
batch_count = min(
|
||||||
|
MODEL_GENERATION_BATCH_SIZE,
|
||||||
|
qa_pairs_per_item - batch_offset,
|
||||||
|
)
|
||||||
|
batch_start = batch_offset + 1
|
||||||
|
batch_end = batch_offset + batch_count
|
||||||
|
request_payload: dict[str, Any] = {
|
||||||
|
"model": model_name,
|
||||||
|
"messages": _prompt_messages(
|
||||||
|
str(config.get("generation_prompt") or ""),
|
||||||
|
content,
|
||||||
|
batch_count,
|
||||||
|
start_index=batch_start,
|
||||||
|
total_count=qa_pairs_per_item,
|
||||||
|
output_type=output_type,
|
||||||
|
reasoning_detail=reasoning_detail,
|
||||||
|
),
|
||||||
|
"temperature": temperature,
|
||||||
|
}
|
||||||
|
if is_minimax_m3:
|
||||||
|
request_payload.update(
|
||||||
|
reasoning_split=True,
|
||||||
|
max_completion_tokens=max(
|
||||||
|
max_tokens,
|
||||||
|
MINIMAX_M3_MIN_COMPLETION_TOKENS,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
request_payload["max_tokens"] = max_tokens
|
||||||
|
if bool(config.get("json_mode", False)) and not is_minimax_m3:
|
||||||
|
request_payload["response_format"] = {"type": "json_object"}
|
||||||
|
|
||||||
|
last_error: Exception | None = None
|
||||||
|
generated_items: list[Mapping[str, Any]] | None = None
|
||||||
|
for _ in range(retries + 1):
|
||||||
|
try:
|
||||||
|
response = http_client.post(
|
||||||
|
endpoint,
|
||||||
|
headers=headers,
|
||||||
|
json=request_payload,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
body = response.json()
|
||||||
|
if not isinstance(body, Mapping):
|
||||||
|
raise ModelGenerationError("model response body must be a JSON object")
|
||||||
|
try:
|
||||||
|
candidate_items = _result_items(
|
||||||
|
_json_payload(_message_content(body))
|
||||||
|
)
|
||||||
|
except ModelGenerationError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"data process model response rejected task_id=%s model=%s "
|
||||||
|
"finish_reason=%s response_chars=%s input_sensitive=%s "
|
||||||
|
"output_sensitive=%s reason=%s",
|
||||||
|
task_id,
|
||||||
|
model_name,
|
||||||
|
_response_finish_reason(body) or "missing",
|
||||||
|
_response_content_length(body),
|
||||||
|
bool(body.get("input_sensitive")),
|
||||||
|
bool(body.get("output_sensitive")),
|
||||||
|
str(exc),
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
if len(candidate_items) < batch_count:
|
||||||
|
raise ModelGenerationError(
|
||||||
|
"model response contains fewer result objects than requested: "
|
||||||
|
f"expected {batch_count}, got {len(candidate_items)}"
|
||||||
|
)
|
||||||
|
generated_items = candidate_items
|
||||||
|
break
|
||||||
|
except (
|
||||||
|
httpx.HTTPError,
|
||||||
|
json.JSONDecodeError,
|
||||||
|
ModelGenerationError,
|
||||||
|
) as exc:
|
||||||
|
last_error = exc
|
||||||
|
if not _is_retryable_generation_error(exc):
|
||||||
|
break
|
||||||
|
|
||||||
|
if generated_items is None:
|
||||||
|
error_message = str(last_error or "model generation failed")[:2000]
|
||||||
|
failure_instruction = (
|
||||||
|
f"模型生成失败,请人工补充(第 {batch_start}-{batch_end} 条)"
|
||||||
|
)
|
||||||
|
result_id = (
|
||||||
|
"result_"
|
||||||
|
f"{hashlib.sha256(f'{preview_id}:error:{batch_start}'.encode()).hexdigest()[:16]}"
|
||||||
|
)
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
"id": result_id,
|
||||||
|
"preview_item_id": preview_id,
|
||||||
|
"instruction": failure_instruction,
|
||||||
|
"input": content,
|
||||||
|
"output": "",
|
||||||
|
"chosen": "",
|
||||||
|
"rejected": "",
|
||||||
|
"original_instruction": failure_instruction,
|
||||||
|
"original_input": content,
|
||||||
|
"original_output": "",
|
||||||
|
"original_chosen": "",
|
||||||
|
"original_rejected": "",
|
||||||
|
"status": "invalid",
|
||||||
|
"error": error_message,
|
||||||
|
"split": "train",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
for batch_index, value in enumerate(generated_items[:batch_count]):
|
||||||
|
variant_index = batch_offset + batch_index
|
||||||
|
instruction = normalize_text(
|
||||||
|
str(value.get("instruction") or value.get("question") or "")
|
||||||
|
)
|
||||||
|
input_text = normalize_text(
|
||||||
|
str(value.get("input") or value.get("context") or "")
|
||||||
|
)
|
||||||
|
chosen = ""
|
||||||
|
rejected = ""
|
||||||
|
if output_type == OUTPUT_TYPE_REASONING:
|
||||||
|
reasoning = normalize_text(
|
||||||
|
re.sub(
|
||||||
|
r"</?think>",
|
||||||
|
"",
|
||||||
|
str(value.get("reasoning") or value.get("analysis") or ""),
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
answer = normalize_text(
|
||||||
|
re.sub(
|
||||||
|
r"</?think>",
|
||||||
|
"",
|
||||||
|
str(
|
||||||
|
value.get("answer")
|
||||||
|
or value.get("final_answer")
|
||||||
|
or value.get("output")
|
||||||
|
or ""
|
||||||
|
),
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
output = (
|
||||||
|
f"<think>\n{reasoning}\n</think>\n{answer}"
|
||||||
|
if reasoning and answer
|
||||||
|
else answer or (f"<think>\n{reasoning}\n</think>" if reasoning else "")
|
||||||
|
)
|
||||||
|
valid = bool(instruction and reasoning and answer)
|
||||||
|
missing_error = "model result is missing instruction, reasoning or answer"
|
||||||
|
elif output_type == OUTPUT_TYPE_DPO:
|
||||||
|
chosen = normalize_text(str(value.get("chosen") or ""))
|
||||||
|
rejected = normalize_text(str(value.get("rejected") or ""))
|
||||||
|
chosen = normalize_text(
|
||||||
|
re.sub(
|
||||||
|
r"<think>[\s\S]*?(?:</think>|$)",
|
||||||
|
"",
|
||||||
|
chosen,
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
rejected = normalize_text(
|
||||||
|
re.sub(
|
||||||
|
r"<think>[\s\S]*?(?:</think>|$)",
|
||||||
|
"",
|
||||||
|
rejected,
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
output = chosen
|
||||||
|
valid = bool(
|
||||||
|
instruction
|
||||||
|
and chosen
|
||||||
|
and rejected
|
||||||
|
and chosen.strip() != rejected.strip()
|
||||||
|
)
|
||||||
|
missing_error = (
|
||||||
|
"model result is missing instruction, chosen or rejected, "
|
||||||
|
"or chosen equals rejected"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
output = normalize_text(
|
||||||
|
str(
|
||||||
|
value.get("output")
|
||||||
|
or value.get("answer")
|
||||||
|
or value.get("response")
|
||||||
|
or ""
|
||||||
|
)
|
||||||
|
)
|
||||||
|
output = normalize_text(
|
||||||
|
re.sub(
|
||||||
|
r"<think>[\s\S]*?(?:</think>|$)",
|
||||||
|
"",
|
||||||
|
output,
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
valid = bool(instruction and output)
|
||||||
|
missing_error = "model result is missing instruction or output"
|
||||||
|
raw_id = (
|
||||||
|
f"{preview_id}:{variant_index + 1}:{instruction}:"
|
||||||
|
f"{output}:{rejected}"
|
||||||
|
)
|
||||||
|
result_id = f"result_{hashlib.sha256(raw_id.encode()).hexdigest()[:16]}"
|
||||||
|
results.append(
|
||||||
|
{
|
||||||
|
"id": result_id,
|
||||||
|
"preview_item_id": preview_id,
|
||||||
|
"instruction": instruction,
|
||||||
|
"input": input_text,
|
||||||
|
"output": output,
|
||||||
|
"chosen": chosen,
|
||||||
|
"rejected": rejected,
|
||||||
|
"original_instruction": instruction,
|
||||||
|
"original_input": input_text,
|
||||||
|
"original_output": output,
|
||||||
|
"original_chosen": chosen,
|
||||||
|
"original_rejected": rejected,
|
||||||
|
"status": "valid" if valid else "invalid",
|
||||||
|
"error": (None if valid else missing_error),
|
||||||
|
"split": "train",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
if on_progress:
|
||||||
|
on_progress(item_index + 1, total_items)
|
||||||
|
finally:
|
||||||
|
if owns_client:
|
||||||
|
http_client.close()
|
||||||
|
assignments = stable_split_assignments(
|
||||||
|
[str(result["id"]) for result in results],
|
||||||
|
split,
|
||||||
|
seed=task_id,
|
||||||
|
)
|
||||||
|
for result, assignment in zip(results, assignments, strict=True):
|
||||||
|
result["split"] = assignment
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["ModelGenerationError", "chat_completions_url", "generate_model_records"]
|
||||||
308
backend/app/modules/data_process/office_preview.py
Normal file
308
backend/app/modules/data_process/office_preview.py
Normal file
@@ -0,0 +1,308 @@
|
|||||||
|
"""Word 与 Excel 原文件的安全、受限预览模型。
|
||||||
|
|
||||||
|
预览只返回浏览器绘制所需的结构化数据,不返回或执行 Office 包中的活动内容。
|
||||||
|
DOCX 的字符偏移与上传时的正文抽取规则保持一致,供前端定位当前切片。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import io
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from docx import Document
|
||||||
|
from docx.oxml.table import CT_Tbl
|
||||||
|
from docx.oxml.text.paragraph import CT_P
|
||||||
|
from docx.table import Table
|
||||||
|
from docx.text.paragraph import Paragraph
|
||||||
|
from openpyxl import load_workbook
|
||||||
|
|
||||||
|
from app.modules.data_process.algorithms import (
|
||||||
|
_MAX_WORKBOOK_COLUMNS,
|
||||||
|
_MAX_WORKBOOK_HEADER_SCAN_ROWS,
|
||||||
|
_infer_xlsx_header_region,
|
||||||
|
_normalize_spreadsheet_value,
|
||||||
|
_rewrite_xlsx_workbook_relationships,
|
||||||
|
_validate_office_archive,
|
||||||
|
_xlsx_sheet_merge_ranges,
|
||||||
|
normalize_text,
|
||||||
|
)
|
||||||
|
|
||||||
|
MAX_DOCX_PREVIEW_BLOCKS = 2_000
|
||||||
|
MAX_XLSX_PREVIEW_ROWS = 200
|
||||||
|
|
||||||
|
|
||||||
|
def _docx_alignment(paragraph: Paragraph) -> str:
|
||||||
|
value = paragraph.alignment
|
||||||
|
return {
|
||||||
|
0: "left",
|
||||||
|
1: "center",
|
||||||
|
2: "right",
|
||||||
|
3: "justify",
|
||||||
|
4: "distribute",
|
||||||
|
5: "justify",
|
||||||
|
7: "justify",
|
||||||
|
8: "distribute",
|
||||||
|
9: "distribute",
|
||||||
|
}.get(int(value) if value is not None else -1, "left")
|
||||||
|
|
||||||
|
|
||||||
|
def _docx_heading_level(paragraph: Paragraph) -> int | None:
|
||||||
|
style = paragraph.style
|
||||||
|
if style is None:
|
||||||
|
return None
|
||||||
|
style_name = str(style.name or "")
|
||||||
|
style_id = str(style.style_id or "")
|
||||||
|
match = re.search(r"(?:heading|标题)\s*([1-6])", f"{style_name} {style_id}", re.IGNORECASE)
|
||||||
|
return int(match.group(1)) if match else None
|
||||||
|
|
||||||
|
|
||||||
|
def build_docx_preview(raw: bytes) -> dict[str, Any]:
|
||||||
|
"""把 DOCX 转为保留标题、段落和表格顺序的浏览器预览模型。"""
|
||||||
|
|
||||||
|
_validate_office_archive(raw, "docx")
|
||||||
|
try:
|
||||||
|
document = Document(io.BytesIO(raw))
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"invalid DOCX file: {exc}") from exc
|
||||||
|
|
||||||
|
blocks: list[dict[str, Any]] = []
|
||||||
|
source_cursor = 0
|
||||||
|
has_source_content = False
|
||||||
|
rendered_blocks = 0
|
||||||
|
truncated = False
|
||||||
|
|
||||||
|
def source_range(value: str) -> tuple[str, int, int] | None:
|
||||||
|
nonlocal source_cursor, has_source_content
|
||||||
|
text = normalize_text(value)
|
||||||
|
if not text:
|
||||||
|
return None
|
||||||
|
if has_source_content:
|
||||||
|
source_cursor += 2
|
||||||
|
start = source_cursor
|
||||||
|
source_cursor += len(text)
|
||||||
|
has_source_content = True
|
||||||
|
return text, start, source_cursor
|
||||||
|
|
||||||
|
for child in document.element.body.iterchildren():
|
||||||
|
if rendered_blocks >= MAX_DOCX_PREVIEW_BLOCKS:
|
||||||
|
truncated = True
|
||||||
|
break
|
||||||
|
|
||||||
|
if isinstance(child, CT_P):
|
||||||
|
paragraph = Paragraph(child, document)
|
||||||
|
located = source_range(paragraph.text)
|
||||||
|
if located is None:
|
||||||
|
continue
|
||||||
|
text, start, end = located
|
||||||
|
style_name = str(paragraph.style.name or "") if paragraph.style else ""
|
||||||
|
blocks.append(
|
||||||
|
{
|
||||||
|
"type": "paragraph",
|
||||||
|
"text": text,
|
||||||
|
"style": style_name,
|
||||||
|
"heading_level": _docx_heading_level(paragraph),
|
||||||
|
"alignment": _docx_alignment(paragraph),
|
||||||
|
"is_list": "list" in style_name.casefold() or "列表" in style_name,
|
||||||
|
"source_start": start,
|
||||||
|
"source_end": end,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
rendered_blocks += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not isinstance(child, CT_Tbl):
|
||||||
|
continue
|
||||||
|
table = Table(child, document)
|
||||||
|
preview_rows: list[dict[str, Any]] = []
|
||||||
|
for row in table.rows:
|
||||||
|
if rendered_blocks >= MAX_DOCX_PREVIEW_BLOCKS:
|
||||||
|
truncated = True
|
||||||
|
break
|
||||||
|
cell_values = [normalize_text(cell.text) for cell in row.cells]
|
||||||
|
located = source_range("\t".join(cell_values))
|
||||||
|
if located is None:
|
||||||
|
continue
|
||||||
|
_, start, end = located
|
||||||
|
preview_rows.append(
|
||||||
|
{
|
||||||
|
"cells": cell_values,
|
||||||
|
"source_start": start,
|
||||||
|
"source_end": end,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
rendered_blocks += 1
|
||||||
|
if preview_rows:
|
||||||
|
blocks.append({"type": "table", "rows": preview_rows})
|
||||||
|
if truncated:
|
||||||
|
break
|
||||||
|
|
||||||
|
return {
|
||||||
|
"format": "docx",
|
||||||
|
"blocks": blocks,
|
||||||
|
"truncated": truncated,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def build_xlsx_preview(
|
||||||
|
raw: bytes,
|
||||||
|
*,
|
||||||
|
sheet_index: int = 0,
|
||||||
|
offset: int = 0,
|
||||||
|
limit: int = 100,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""按工作表分页返回 XLSX 的表头和记录网格。"""
|
||||||
|
|
||||||
|
if sheet_index < 0 or offset < 0:
|
||||||
|
raise ValueError("sheet_index and offset must be non-negative")
|
||||||
|
if limit < 1 or limit > MAX_XLSX_PREVIEW_ROWS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX preview limit must be in [1, {MAX_XLSX_PREVIEW_ROWS}]"
|
||||||
|
)
|
||||||
|
|
||||||
|
_validate_office_archive(raw, "xlsx")
|
||||||
|
merged_by_sheet, normalized_targets = _xlsx_sheet_merge_ranges(raw)
|
||||||
|
workbook_raw = (
|
||||||
|
_rewrite_xlsx_workbook_relationships(raw, normalized_targets)
|
||||||
|
if normalized_targets
|
||||||
|
else raw
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
workbook = load_workbook(
|
||||||
|
io.BytesIO(workbook_raw),
|
||||||
|
read_only=True,
|
||||||
|
data_only=True,
|
||||||
|
keep_links=False,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
raise ValueError(f"invalid XLSX file: {exc}") from exc
|
||||||
|
|
||||||
|
try:
|
||||||
|
sheets = [
|
||||||
|
{
|
||||||
|
"index": index,
|
||||||
|
"name": worksheet.title,
|
||||||
|
"state": worksheet.sheet_state,
|
||||||
|
}
|
||||||
|
for index, worksheet in enumerate(workbook.worksheets)
|
||||||
|
]
|
||||||
|
if not sheets:
|
||||||
|
raise ValueError("XLSX workbook contains no worksheets")
|
||||||
|
if sheet_index >= len(sheets):
|
||||||
|
raise ValueError("XLSX worksheet index is out of range")
|
||||||
|
|
||||||
|
worksheet = workbook.worksheets[sheet_index]
|
||||||
|
reset_dimensions = getattr(worksheet, "reset_dimensions", None)
|
||||||
|
if callable(reset_dimensions):
|
||||||
|
reset_dimensions()
|
||||||
|
row_iterator = enumerate(worksheet.iter_rows(values_only=True), start=1)
|
||||||
|
buffered_rows: dict[int, tuple[Any, ...]] = {}
|
||||||
|
|
||||||
|
def normalized_values(row: tuple[Any, ...]) -> list[Any]:
|
||||||
|
values = list(row)
|
||||||
|
while values and values[-1] in {None, ""}:
|
||||||
|
values.pop()
|
||||||
|
if len(values) > _MAX_WORKBOOK_COLUMNS:
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {worksheet.title!r} exceeds "
|
||||||
|
f"{_MAX_WORKBOOK_COLUMNS} columns"
|
||||||
|
)
|
||||||
|
return values
|
||||||
|
|
||||||
|
for row_number, row in row_iterator:
|
||||||
|
values = normalized_values(row)
|
||||||
|
if not values or all(value in {None, ""} for value in values):
|
||||||
|
continue
|
||||||
|
buffered_rows[row_number] = tuple(values)
|
||||||
|
if len(buffered_rows) >= _MAX_WORKBOOK_HEADER_SCAN_ROWS:
|
||||||
|
break
|
||||||
|
|
||||||
|
if not buffered_rows:
|
||||||
|
return {
|
||||||
|
"format": "xlsx",
|
||||||
|
"sheets": sheets,
|
||||||
|
"active_sheet": {
|
||||||
|
"index": sheet_index,
|
||||||
|
"name": worksheet.title,
|
||||||
|
"columns": [],
|
||||||
|
"rows": [],
|
||||||
|
"offset": offset,
|
||||||
|
"limit": limit,
|
||||||
|
"has_more": False,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, header_end_row, headers = _infer_xlsx_header_region(
|
||||||
|
worksheet.title,
|
||||||
|
buffered_rows,
|
||||||
|
merged_by_sheet.get(worksheet.title, ()),
|
||||||
|
)
|
||||||
|
preview_rows: list[dict[str, Any]] = []
|
||||||
|
record_index = 0
|
||||||
|
has_more = False
|
||||||
|
|
||||||
|
def append_row(row_number: int, values: tuple[Any, ...] | list[Any]) -> bool:
|
||||||
|
nonlocal record_index, has_more
|
||||||
|
row_values = list(values)
|
||||||
|
if len(row_values) > len(headers):
|
||||||
|
raise ValueError(
|
||||||
|
f"XLSX worksheet {worksheet.title!r} has a row wider than its header"
|
||||||
|
)
|
||||||
|
row_values.extend([None] * (len(headers) - len(row_values)))
|
||||||
|
record = {
|
||||||
|
header: _normalize_spreadsheet_value(value)
|
||||||
|
for header, value in zip(headers, row_values, strict=True)
|
||||||
|
}
|
||||||
|
if not any(value not in {"", None} for value in record.values()):
|
||||||
|
return False
|
||||||
|
current_index = record_index
|
||||||
|
record_index += 1
|
||||||
|
if current_index < offset:
|
||||||
|
return False
|
||||||
|
if len(preview_rows) >= limit:
|
||||||
|
has_more = True
|
||||||
|
return True
|
||||||
|
preview_rows.append(
|
||||||
|
{
|
||||||
|
"row_number": row_number,
|
||||||
|
"record_index": current_index,
|
||||||
|
"values": [record[header] for header in headers],
|
||||||
|
"record": record,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
for row_number, values in buffered_rows.items():
|
||||||
|
if row_number > header_end_row and append_row(row_number, values):
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
for row_number, row in row_iterator:
|
||||||
|
values = normalized_values(row)
|
||||||
|
if not values or all(value in {None, ""} for value in values):
|
||||||
|
continue
|
||||||
|
if append_row(row_number, values):
|
||||||
|
break
|
||||||
|
|
||||||
|
return {
|
||||||
|
"format": "xlsx",
|
||||||
|
"sheets": sheets,
|
||||||
|
"active_sheet": {
|
||||||
|
"index": sheet_index,
|
||||||
|
"name": worksheet.title,
|
||||||
|
"columns": headers,
|
||||||
|
"rows": preview_rows,
|
||||||
|
"offset": offset,
|
||||||
|
"limit": limit,
|
||||||
|
"has_more": has_more,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
finally:
|
||||||
|
workbook.close()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"MAX_DOCX_PREVIEW_BLOCKS",
|
||||||
|
"MAX_XLSX_PREVIEW_ROWS",
|
||||||
|
"build_docx_preview",
|
||||||
|
"build_xlsx_preview",
|
||||||
|
]
|
||||||
76
backend/app/modules/data_process/schema_cli.py
Normal file
76
backend/app/modules/data_process/schema_cli.py
Normal file
@@ -0,0 +1,76 @@
|
|||||||
|
"""数据处理运行表的显式检查与安装命令。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
from urllib.parse import urlsplit
|
||||||
|
|
||||||
|
from app.modules.data_process.store import DataProcessStore
|
||||||
|
|
||||||
|
REQUIRED_TASK_COLUMNS = (
|
||||||
|
"generation_run_id",
|
||||||
|
"results_confirmed",
|
||||||
|
"workflow_step",
|
||||||
|
"preview_status",
|
||||||
|
"preview_progress",
|
||||||
|
"preview_run_id",
|
||||||
|
"preview_failure_reason",
|
||||||
|
"preview_total_files",
|
||||||
|
"preview_completed_files",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _target_label(database_url: str) -> str:
|
||||||
|
parsed = urlsplit(database_url)
|
||||||
|
database = parsed.path.strip("/") or "(unknown)"
|
||||||
|
return f"{parsed.hostname or '(unknown)'}:{parsed.port or 5432}/{database}"
|
||||||
|
|
||||||
|
|
||||||
|
def _schema_ready(store: DataProcessStore) -> bool:
|
||||||
|
with store.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT COUNT(*) = %s AS ready
|
||||||
|
FROM information_schema.columns
|
||||||
|
WHERE table_schema=current_schema()
|
||||||
|
AND table_name='data_process_tasks'
|
||||||
|
AND column_name = ANY(%s)
|
||||||
|
""",
|
||||||
|
(len(REQUIRED_TASK_COLUMNS), list(REQUIRED_TASK_COLUMNS)),
|
||||||
|
).fetchone()
|
||||||
|
return bool(row and row["ready"])
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> int:
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="检查或显式安装数据处理运行表(不会由应用启动自动执行)"
|
||||||
|
)
|
||||||
|
action = parser.add_mutually_exclusive_group(required=True)
|
||||||
|
action.add_argument("--check", action="store_true", help="只读检查迁移是否已安装")
|
||||||
|
action.add_argument("--apply", action="store_true", help="执行 002 数据处理迁移")
|
||||||
|
parser.add_argument(
|
||||||
|
"--yes",
|
||||||
|
action="store_true",
|
||||||
|
help="确认允许修改 DATABASE_URL 指向的数据库;与 --apply 同时使用",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
store = DataProcessStore()
|
||||||
|
target = _target_label(store.database_url)
|
||||||
|
if args.check:
|
||||||
|
ready = _schema_ready(store)
|
||||||
|
print(f"数据处理 schema:{'已安装' if ready else '未安装'};目标:{target}")
|
||||||
|
return 0 if ready else 1
|
||||||
|
if not args.yes:
|
||||||
|
parser.error("--apply 必须同时提供 --yes,确认修改目标数据库")
|
||||||
|
|
||||||
|
print(f"正在安装数据处理 schema;目标:{target}")
|
||||||
|
store.ensure_schema()
|
||||||
|
if not _schema_ready(store):
|
||||||
|
raise RuntimeError("迁移执行后仍未检测到 generation_run_id")
|
||||||
|
print("数据处理 schema 安装完成")
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
567
backend/app/modules/data_process/storage.py
Normal file
567
backend/app/modules/data_process/storage.py
Normal file
@@ -0,0 +1,567 @@
|
|||||||
|
"""数据处理原始源文件的受控本地对象存储。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
import stat
|
||||||
|
import unicodedata
|
||||||
|
import uuid
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from functools import lru_cache
|
||||||
|
from pathlib import Path, PurePosixPath
|
||||||
|
from typing import Iterable, Iterator
|
||||||
|
from urllib.parse import quote, unquote, urlsplit
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessStorageError(ValueError):
|
||||||
|
"""本地对象引用或文件系统状态不安全。"""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class StagedSourceObject:
|
||||||
|
"""尚未发布的原始文件;绝对路径仅在存储模块内部流转。"""
|
||||||
|
|
||||||
|
reference: str
|
||||||
|
_temporary_path: Path
|
||||||
|
_relative_path: PurePosixPath
|
||||||
|
|
||||||
|
|
||||||
|
def _default_storage_root() -> Path:
|
||||||
|
data_root = os.getenv("YG_FT_DATA_ROOT", "").strip()
|
||||||
|
if data_root:
|
||||||
|
return Path(data_root).expanduser() / "data-process"
|
||||||
|
return Path(__file__).resolve().parents[3] / "storage" / "data-process"
|
||||||
|
|
||||||
|
|
||||||
|
def _configured_storage_root() -> Path:
|
||||||
|
configured = os.getenv("DATA_PROCESS_STORAGE_DIR", "").strip()
|
||||||
|
if not configured:
|
||||||
|
return _default_storage_root()
|
||||||
|
path = Path(configured).expanduser()
|
||||||
|
# 相对配置固定以 backend 目录为基准,
|
||||||
|
# 避免从不同 cwd 启动时写入不同位置。
|
||||||
|
return path if path.is_absolute() else Path(__file__).resolve().parents[3] / path
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_component(value: str, label: str) -> str:
|
||||||
|
if not value or value in {".", ".."} or len(value) > 128:
|
||||||
|
raise DataProcessStorageError(f"invalid {label}")
|
||||||
|
if not value[0].isalnum() or any(
|
||||||
|
not (character.isalnum() or character in {"-", "_", "."})
|
||||||
|
for character in value
|
||||||
|
):
|
||||||
|
raise DataProcessStorageError(f"invalid {label}")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_basename(value: str) -> str:
|
||||||
|
if not value or len(value.encode("utf-8")) > 255:
|
||||||
|
raise DataProcessStorageError("invalid source file name")
|
||||||
|
if value != Path(value).name or "/" in value or "\\" in value or "\x00" in value:
|
||||||
|
raise DataProcessStorageError("invalid source file name")
|
||||||
|
if value in {".", ".."} or any(
|
||||||
|
unicodedata.category(character).startswith("C") for character in value
|
||||||
|
):
|
||||||
|
raise DataProcessStorageError("invalid source file name")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
class LocalDataProcessStorage:
|
||||||
|
"""只允许访问配置根目录下的版本化原始文件。"""
|
||||||
|
|
||||||
|
def __init__(self, root: str | os.PathLike[str] | Path | None = None) -> None:
|
||||||
|
configured = Path(root) if root is not None else _configured_storage_root()
|
||||||
|
configured = configured.expanduser()
|
||||||
|
if configured.exists() and configured.is_symlink():
|
||||||
|
raise DataProcessStorageError("data process storage root must not be a symlink")
|
||||||
|
configured.mkdir(parents=True, exist_ok=True, mode=0o700)
|
||||||
|
self._root = configured.resolve(strict=True)
|
||||||
|
# StagedSourceObject 本身是普通 dataclass,不能只依赖其中的路径字段判断
|
||||||
|
# 来源;只接受由当前存储实例实际签发的对象,
|
||||||
|
# 避免调用方伪造暂存路径。
|
||||||
|
self._issued_staged_objects: dict[Path, StagedSourceObject] = {}
|
||||||
|
self._ensure_directory(self._root / ".staging")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def root(self) -> Path:
|
||||||
|
"""仅供运维和测试检查;API 响应不得序列化该属性。"""
|
||||||
|
|
||||||
|
return self._root
|
||||||
|
|
||||||
|
def new_batch_id(self) -> str:
|
||||||
|
return f"batch-{uuid.uuid4().hex}"
|
||||||
|
|
||||||
|
def stage_bytes(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
batch_id: str,
|
||||||
|
task_id: str,
|
||||||
|
source_file_id: str,
|
||||||
|
version: int,
|
||||||
|
name: str,
|
||||||
|
content: bytes,
|
||||||
|
) -> StagedSourceObject:
|
||||||
|
batch_id = _safe_component(batch_id, "batch id")
|
||||||
|
task_id = _safe_component(task_id, "task id")
|
||||||
|
source_file_id = _safe_component(source_file_id, "source file id")
|
||||||
|
if isinstance(version, bool) or not isinstance(version, int) or version < 1:
|
||||||
|
raise DataProcessStorageError("invalid source file version")
|
||||||
|
basename = _safe_basename(name)
|
||||||
|
if not isinstance(content, bytes):
|
||||||
|
raise TypeError("content must be bytes")
|
||||||
|
|
||||||
|
batch_directory = self._ensure_directory(self._root / ".staging" / batch_id)
|
||||||
|
temporary_path = batch_directory / f"{source_file_id}-{uuid.uuid4().hex}.tmp"
|
||||||
|
flags = os.O_CREAT | os.O_EXCL | os.O_WRONLY
|
||||||
|
if hasattr(os, "O_NOFOLLOW"):
|
||||||
|
flags |= os.O_NOFOLLOW
|
||||||
|
descriptor = os.open(temporary_path, flags, 0o600)
|
||||||
|
try:
|
||||||
|
with os.fdopen(descriptor, "wb", closefd=True) as stream:
|
||||||
|
stream.write(content)
|
||||||
|
stream.flush()
|
||||||
|
os.fsync(stream.fileno())
|
||||||
|
except Exception:
|
||||||
|
temporary_path.unlink(missing_ok=True)
|
||||||
|
raise
|
||||||
|
|
||||||
|
relative_path = PurePosixPath(
|
||||||
|
task_id,
|
||||||
|
source_file_id,
|
||||||
|
f"v{version}",
|
||||||
|
basename,
|
||||||
|
)
|
||||||
|
reference = (
|
||||||
|
"local://data-process/"
|
||||||
|
f"{task_id}/{source_file_id}/v{version}/{quote(basename, safe='')}"
|
||||||
|
)
|
||||||
|
staged = StagedSourceObject(reference, temporary_path, relative_path)
|
||||||
|
self._issued_staged_objects[temporary_path] = staged
|
||||||
|
return staged
|
||||||
|
|
||||||
|
def stage_copy(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
batch_id: str,
|
||||||
|
source_reference: str,
|
||||||
|
expected_source_task_id: str,
|
||||||
|
expected_source_file_id: str,
|
||||||
|
task_id: str,
|
||||||
|
source_file_id: str,
|
||||||
|
version: int,
|
||||||
|
name: str,
|
||||||
|
) -> StagedSourceObject:
|
||||||
|
"""为不可变源对象创建独立目录项,不把大文件重新读入内存。"""
|
||||||
|
|
||||||
|
batch_id = _safe_component(batch_id, "batch id")
|
||||||
|
task_id = _safe_component(task_id, "task id")
|
||||||
|
source_file_id = _safe_component(source_file_id, "source file id")
|
||||||
|
if isinstance(version, bool) or not isinstance(version, int) or version < 1:
|
||||||
|
raise DataProcessStorageError("invalid source file version")
|
||||||
|
basename = _safe_basename(name)
|
||||||
|
source_relative = self._relative_from_reference(source_reference)
|
||||||
|
if source_relative is None:
|
||||||
|
raise DataProcessStorageError("original source object is not available")
|
||||||
|
self._assert_expected_owner(
|
||||||
|
source_relative,
|
||||||
|
expected_task_id=expected_source_task_id,
|
||||||
|
expected_source_file_id=expected_source_file_id,
|
||||||
|
)
|
||||||
|
descriptor, source_info = self._open_read_descriptor(source_relative)
|
||||||
|
os.close(descriptor)
|
||||||
|
|
||||||
|
batch_directory = self._ensure_directory(self._root / ".staging" / batch_id)
|
||||||
|
temporary_path = batch_directory / f"{source_file_id}-{uuid.uuid4().hex}.tmp"
|
||||||
|
source_path = self._path_for_relative(source_relative)
|
||||||
|
try:
|
||||||
|
os.link(source_path, temporary_path, follow_symlinks=False)
|
||||||
|
copy_info = temporary_path.lstat()
|
||||||
|
if (
|
||||||
|
not stat.S_ISREG(copy_info.st_mode)
|
||||||
|
or source_info.st_dev != copy_info.st_dev
|
||||||
|
or source_info.st_ino != copy_info.st_ino
|
||||||
|
):
|
||||||
|
raise DataProcessStorageError("source storage object changed while copying")
|
||||||
|
except Exception:
|
||||||
|
temporary_path.unlink(missing_ok=True)
|
||||||
|
raise
|
||||||
|
|
||||||
|
relative_path = PurePosixPath(
|
||||||
|
task_id,
|
||||||
|
source_file_id,
|
||||||
|
f"v{version}",
|
||||||
|
basename,
|
||||||
|
)
|
||||||
|
reference = (
|
||||||
|
"local://data-process/"
|
||||||
|
f"{task_id}/{source_file_id}/v{version}/{quote(basename, safe='')}"
|
||||||
|
)
|
||||||
|
staged = StagedSourceObject(reference, temporary_path, relative_path)
|
||||||
|
self._issued_staged_objects[temporary_path] = staged
|
||||||
|
return staged
|
||||||
|
|
||||||
|
def publish(self, objects: Iterable[StagedSourceObject]) -> None:
|
||||||
|
staged = list(objects)
|
||||||
|
published: list[StagedSourceObject] = []
|
||||||
|
try:
|
||||||
|
seen_temporary_paths: set[Path] = set()
|
||||||
|
for item in staged:
|
||||||
|
self._validate_staged_object(item, require_file=True)
|
||||||
|
if item._temporary_path in seen_temporary_paths:
|
||||||
|
raise DataProcessStorageError("duplicate staged source object")
|
||||||
|
seen_temporary_paths.add(item._temporary_path)
|
||||||
|
for item in staged:
|
||||||
|
final_path = self._path_for_relative(item._relative_path)
|
||||||
|
self._ensure_directory(final_path.parent)
|
||||||
|
if final_path.exists() or final_path.is_symlink():
|
||||||
|
raise DataProcessStorageError("source storage object already exists")
|
||||||
|
os.link(item._temporary_path, final_path, follow_symlinks=False)
|
||||||
|
published.append(item)
|
||||||
|
item._temporary_path.unlink()
|
||||||
|
self._fsync_directory(final_path.parent)
|
||||||
|
except Exception:
|
||||||
|
for item in reversed(published):
|
||||||
|
try:
|
||||||
|
self.delete(item.reference)
|
||||||
|
except Exception:
|
||||||
|
# 回滚必须尽量处理其余对象,并保留真正的发布异常。
|
||||||
|
pass
|
||||||
|
for item in staged:
|
||||||
|
try:
|
||||||
|
self.discard([item])
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
self.discard(staged)
|
||||||
|
|
||||||
|
def discard(self, objects: Iterable[StagedSourceObject]) -> None:
|
||||||
|
staged = list(objects)
|
||||||
|
for item in staged:
|
||||||
|
self._validate_staged_object(item, require_file=False)
|
||||||
|
|
||||||
|
batch_directories: set[Path] = set()
|
||||||
|
first_error: Exception | None = None
|
||||||
|
for item in staged:
|
||||||
|
temporary_path = item._temporary_path
|
||||||
|
try:
|
||||||
|
temporary_path.unlink(missing_ok=True)
|
||||||
|
except Exception as exc:
|
||||||
|
if first_error is None:
|
||||||
|
first_error = exc
|
||||||
|
else:
|
||||||
|
self._issued_staged_objects.pop(temporary_path, None)
|
||||||
|
batch_directories.add(temporary_path.parent)
|
||||||
|
for directory in batch_directories:
|
||||||
|
self._remove_empty_directory(directory)
|
||||||
|
if first_error is not None:
|
||||||
|
raise first_error
|
||||||
|
|
||||||
|
def read(self, reference: str) -> bytes | None:
|
||||||
|
"""读取 local 引用;旧 ``db://`` 对象返回 ``None`` 由数据库正文兜底。"""
|
||||||
|
|
||||||
|
relative_path = self._relative_from_reference(reference)
|
||||||
|
if relative_path is None:
|
||||||
|
return None
|
||||||
|
descriptor, _ = self._open_read_descriptor(relative_path)
|
||||||
|
with os.fdopen(descriptor, "rb", closefd=True) as stream:
|
||||||
|
return stream.read()
|
||||||
|
|
||||||
|
def file_size(
|
||||||
|
self,
|
||||||
|
reference: str,
|
||||||
|
*,
|
||||||
|
expected_task_id: str,
|
||||||
|
expected_source_file_id: str,
|
||||||
|
) -> int | None:
|
||||||
|
"""返回受控 local 对象大小;旧 ``db://`` 对象没有原始文件。"""
|
||||||
|
|
||||||
|
relative_path = self._relative_from_reference(reference)
|
||||||
|
if relative_path is None:
|
||||||
|
return None
|
||||||
|
self._assert_expected_owner(
|
||||||
|
relative_path,
|
||||||
|
expected_task_id=expected_task_id,
|
||||||
|
expected_source_file_id=expected_source_file_id,
|
||||||
|
)
|
||||||
|
descriptor, info = self._open_read_descriptor(relative_path)
|
||||||
|
os.close(descriptor)
|
||||||
|
return info.st_size
|
||||||
|
|
||||||
|
def iter_bytes(
|
||||||
|
self,
|
||||||
|
reference: str,
|
||||||
|
*,
|
||||||
|
expected_task_id: str,
|
||||||
|
expected_source_file_id: str,
|
||||||
|
expected_size: int,
|
||||||
|
start: int = 0,
|
||||||
|
length: int | None = None,
|
||||||
|
chunk_size: int = 256 * 1024,
|
||||||
|
) -> Iterator[bytes]:
|
||||||
|
"""按范围流式读取原始文件,避免 PDF 预览把大文件整体载入内存。"""
|
||||||
|
|
||||||
|
relative_path = self._relative_from_reference(reference)
|
||||||
|
if relative_path is None:
|
||||||
|
raise DataProcessStorageError("original source object is not available")
|
||||||
|
self._assert_expected_owner(
|
||||||
|
relative_path,
|
||||||
|
expected_task_id=expected_task_id,
|
||||||
|
expected_source_file_id=expected_source_file_id,
|
||||||
|
)
|
||||||
|
if start < 0 or expected_size < 0 or chunk_size < 1:
|
||||||
|
raise DataProcessStorageError("invalid source byte range")
|
||||||
|
descriptor, info = self._open_read_descriptor(relative_path)
|
||||||
|
if info.st_size != expected_size:
|
||||||
|
os.close(descriptor)
|
||||||
|
raise DataProcessStorageError("source object size does not match metadata")
|
||||||
|
remaining = expected_size - start if length is None else length
|
||||||
|
if remaining < 0 or start + remaining > expected_size:
|
||||||
|
os.close(descriptor)
|
||||||
|
raise DataProcessStorageError("invalid source byte range")
|
||||||
|
with os.fdopen(descriptor, "rb", closefd=True) as stream:
|
||||||
|
stream.seek(start)
|
||||||
|
while remaining:
|
||||||
|
chunk = stream.read(min(chunk_size, remaining))
|
||||||
|
if not chunk:
|
||||||
|
raise DataProcessStorageError("source object ended unexpectedly")
|
||||||
|
remaining -= len(chunk)
|
||||||
|
yield chunk
|
||||||
|
|
||||||
|
def validate_owner(
|
||||||
|
self,
|
||||||
|
reference: str,
|
||||||
|
*,
|
||||||
|
expected_task_id: str,
|
||||||
|
expected_source_file_id: str,
|
||||||
|
) -> bool:
|
||||||
|
"""校验 local 引用归属;旧 ``db://`` 引用无需文件系统处理。"""
|
||||||
|
|
||||||
|
relative_path = self._relative_from_reference(reference)
|
||||||
|
if relative_path is None:
|
||||||
|
return False
|
||||||
|
self._assert_expected_owner(
|
||||||
|
relative_path,
|
||||||
|
expected_task_id=expected_task_id,
|
||||||
|
expected_source_file_id=expected_source_file_id,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _open_read_descriptor(
|
||||||
|
self,
|
||||||
|
relative_path: PurePosixPath,
|
||||||
|
) -> tuple[int, os.stat_result]:
|
||||||
|
path = self._path_for_relative(relative_path)
|
||||||
|
self._assert_controlled_parent(path)
|
||||||
|
try:
|
||||||
|
before_open = path.lstat()
|
||||||
|
except FileNotFoundError as exc:
|
||||||
|
raise DataProcessStorageError("source storage object does not exist") from exc
|
||||||
|
if stat.S_ISLNK(before_open.st_mode) or not stat.S_ISREG(before_open.st_mode):
|
||||||
|
raise DataProcessStorageError("source storage object is not a regular file")
|
||||||
|
flags = os.O_RDONLY
|
||||||
|
if hasattr(os, "O_NOFOLLOW"):
|
||||||
|
flags |= os.O_NOFOLLOW
|
||||||
|
descriptor = os.open(path, flags)
|
||||||
|
after_open = os.fstat(descriptor)
|
||||||
|
if (
|
||||||
|
not stat.S_ISREG(after_open.st_mode)
|
||||||
|
or before_open.st_dev != after_open.st_dev
|
||||||
|
or before_open.st_ino != after_open.st_ino
|
||||||
|
):
|
||||||
|
os.close(descriptor)
|
||||||
|
raise DataProcessStorageError("source storage object changed while opening")
|
||||||
|
return descriptor, after_open
|
||||||
|
|
||||||
|
def delete(
|
||||||
|
self,
|
||||||
|
reference: str,
|
||||||
|
*,
|
||||||
|
expected_task_id: str | None = None,
|
||||||
|
expected_source_file_id: str | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""删除受控 local 对象;旧 ``db://`` 引用保持不变。"""
|
||||||
|
|
||||||
|
relative_path = self._relative_from_reference(reference)
|
||||||
|
if relative_path is None:
|
||||||
|
return False
|
||||||
|
if (expected_task_id is None) != (expected_source_file_id is None):
|
||||||
|
raise DataProcessStorageError("both expected storage owner fields are required")
|
||||||
|
if expected_task_id is not None and expected_source_file_id is not None:
|
||||||
|
self._assert_expected_owner(
|
||||||
|
relative_path,
|
||||||
|
expected_task_id=expected_task_id,
|
||||||
|
expected_source_file_id=expected_source_file_id,
|
||||||
|
)
|
||||||
|
path = self._path_for_relative(relative_path)
|
||||||
|
self._assert_controlled_parent(path)
|
||||||
|
try:
|
||||||
|
info = path.lstat()
|
||||||
|
except FileNotFoundError:
|
||||||
|
return False
|
||||||
|
if stat.S_ISLNK(info.st_mode) or not stat.S_ISREG(info.st_mode):
|
||||||
|
raise DataProcessStorageError("refusing to delete a non-regular storage object")
|
||||||
|
path.unlink()
|
||||||
|
self._fsync_directory(path.parent)
|
||||||
|
for directory in (path.parent, path.parent.parent, path.parent.parent.parent):
|
||||||
|
self._remove_empty_directory(directory)
|
||||||
|
return True
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _assert_expected_owner(
|
||||||
|
relative_path: PurePosixPath,
|
||||||
|
*,
|
||||||
|
expected_task_id: str,
|
||||||
|
expected_source_file_id: str,
|
||||||
|
) -> None:
|
||||||
|
task_id = _safe_component(expected_task_id, "expected task id")
|
||||||
|
source_file_id = _safe_component(
|
||||||
|
expected_source_file_id,
|
||||||
|
"expected source file id",
|
||||||
|
)
|
||||||
|
if relative_path.parts[:2] != (task_id, source_file_id):
|
||||||
|
raise DataProcessStorageError("source storage object owner mismatch")
|
||||||
|
|
||||||
|
def _relative_from_reference(self, reference: str) -> PurePosixPath | None:
|
||||||
|
if reference.startswith("db://"):
|
||||||
|
return None
|
||||||
|
parsed = urlsplit(reference)
|
||||||
|
if parsed.scheme != "local" or parsed.netloc != "data-process":
|
||||||
|
raise DataProcessStorageError("unsupported source storage reference")
|
||||||
|
if parsed.query or parsed.fragment or "\\" in parsed.path:
|
||||||
|
raise DataProcessStorageError("unsafe source storage reference")
|
||||||
|
raw_parts = parsed.path.lstrip("/").split("/")
|
||||||
|
if len(raw_parts) != 4:
|
||||||
|
raise DataProcessStorageError("unsafe source storage reference")
|
||||||
|
if any(re.search(r"%(?![0-9A-Fa-f]{2})", part) for part in raw_parts):
|
||||||
|
raise DataProcessStorageError("unsafe source storage reference")
|
||||||
|
try:
|
||||||
|
decoded = [unquote(part, encoding="utf-8", errors="strict") for part in raw_parts]
|
||||||
|
except UnicodeDecodeError as exc:
|
||||||
|
raise DataProcessStorageError("unsafe source storage reference") from exc
|
||||||
|
if any("/" in part or "\\" in part for part in decoded):
|
||||||
|
raise DataProcessStorageError("unsafe source storage reference")
|
||||||
|
canonical_parts = [
|
||||||
|
quote(decoded[0], safe="-_."),
|
||||||
|
quote(decoded[1], safe="-_."),
|
||||||
|
quote(decoded[2], safe="-_."),
|
||||||
|
quote(decoded[3], safe=""),
|
||||||
|
]
|
||||||
|
if canonical_parts != raw_parts:
|
||||||
|
raise DataProcessStorageError("source storage reference is not canonical")
|
||||||
|
task_id = _safe_component(decoded[0], "task id")
|
||||||
|
source_file_id = _safe_component(decoded[1], "source file id")
|
||||||
|
version_component = decoded[2]
|
||||||
|
if not version_component.startswith("v") or not version_component[1:].isdigit():
|
||||||
|
raise DataProcessStorageError("invalid source file version")
|
||||||
|
version = int(version_component[1:])
|
||||||
|
if version < 1:
|
||||||
|
raise DataProcessStorageError("invalid source file version")
|
||||||
|
basename = _safe_basename(decoded[3])
|
||||||
|
return PurePosixPath(task_id, source_file_id, f"v{version}", basename)
|
||||||
|
|
||||||
|
def _path_for_relative(self, relative_path: PurePosixPath) -> Path:
|
||||||
|
if relative_path.is_absolute() or any(
|
||||||
|
part in {"", ".", ".."} for part in relative_path.parts
|
||||||
|
):
|
||||||
|
raise DataProcessStorageError("storage path escapes the configured root")
|
||||||
|
path = self._root.joinpath(*relative_path.parts)
|
||||||
|
self._assert_controlled_parent(path)
|
||||||
|
return path
|
||||||
|
|
||||||
|
def _validate_staged_object(
|
||||||
|
self,
|
||||||
|
item: StagedSourceObject,
|
||||||
|
*,
|
||||||
|
require_file: bool,
|
||||||
|
) -> None:
|
||||||
|
if not isinstance(item, StagedSourceObject):
|
||||||
|
raise DataProcessStorageError("invalid staged source object")
|
||||||
|
if self._issued_staged_objects.get(item._temporary_path) is not item:
|
||||||
|
raise DataProcessStorageError("staged source object was not issued by this storage")
|
||||||
|
expected_relative = self._relative_from_reference(item.reference)
|
||||||
|
if expected_relative is None or expected_relative != item._relative_path:
|
||||||
|
raise DataProcessStorageError("staged source object reference mismatch")
|
||||||
|
staging_root = self._root / ".staging"
|
||||||
|
try:
|
||||||
|
relative_temporary = item._temporary_path.relative_to(staging_root)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise DataProcessStorageError("staged source object escapes staging") from exc
|
||||||
|
if len(relative_temporary.parts) != 2:
|
||||||
|
raise DataProcessStorageError("invalid staged source object path")
|
||||||
|
_safe_component(relative_temporary.parts[0], "batch id")
|
||||||
|
_safe_basename(relative_temporary.parts[1])
|
||||||
|
self._assert_controlled_parent(item._temporary_path)
|
||||||
|
try:
|
||||||
|
info = item._temporary_path.lstat()
|
||||||
|
except FileNotFoundError:
|
||||||
|
if require_file:
|
||||||
|
raise DataProcessStorageError("staged source object does not exist") from None
|
||||||
|
return
|
||||||
|
if stat.S_ISLNK(info.st_mode) or not stat.S_ISREG(info.st_mode):
|
||||||
|
raise DataProcessStorageError("staged source object is not a regular file")
|
||||||
|
|
||||||
|
def _ensure_directory(self, directory: Path) -> Path:
|
||||||
|
try:
|
||||||
|
relative = directory.relative_to(self._root)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise DataProcessStorageError("storage path escapes the configured root") from exc
|
||||||
|
current = self._root
|
||||||
|
for component in relative.parts:
|
||||||
|
current = current / component
|
||||||
|
try:
|
||||||
|
current.mkdir(mode=0o700)
|
||||||
|
except FileExistsError:
|
||||||
|
pass
|
||||||
|
info = current.lstat()
|
||||||
|
if stat.S_ISLNK(info.st_mode) or not stat.S_ISDIR(info.st_mode):
|
||||||
|
raise DataProcessStorageError("storage path contains a symlink or non-directory")
|
||||||
|
return directory
|
||||||
|
|
||||||
|
def _assert_controlled_parent(self, path: Path) -> None:
|
||||||
|
try:
|
||||||
|
relative_parent = path.parent.relative_to(self._root)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise DataProcessStorageError("storage path escapes the configured root") from exc
|
||||||
|
current = self._root
|
||||||
|
for component in relative_parent.parts:
|
||||||
|
current = current / component
|
||||||
|
if not current.exists():
|
||||||
|
continue
|
||||||
|
info = current.lstat()
|
||||||
|
if stat.S_ISLNK(info.st_mode) or not stat.S_ISDIR(info.st_mode):
|
||||||
|
raise DataProcessStorageError("storage path contains a symlink or non-directory")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _fsync_directory(directory: Path) -> None:
|
||||||
|
# Windows 不支持以 O_RDONLY 打开目录做 fsync,跳过即可。
|
||||||
|
# 数据完整性在 Linux 生产环境保障,Windows 开发环境忽略。
|
||||||
|
if os.name == "nt":
|
||||||
|
return
|
||||||
|
descriptor = os.open(directory, os.O_RDONLY)
|
||||||
|
try:
|
||||||
|
os.fsync(descriptor)
|
||||||
|
finally:
|
||||||
|
os.close(descriptor)
|
||||||
|
|
||||||
|
def _remove_empty_directory(self, directory: Path) -> None:
|
||||||
|
if directory in {self._root, self._root / ".staging"}:
|
||||||
|
return
|
||||||
|
self._assert_controlled_parent(directory / "placeholder")
|
||||||
|
try:
|
||||||
|
directory.rmdir()
|
||||||
|
except (FileNotFoundError, OSError):
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache
|
||||||
|
def get_data_process_storage() -> LocalDataProcessStorage:
|
||||||
|
return LocalDataProcessStorage()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"DataProcessStorageError",
|
||||||
|
"LocalDataProcessStorage",
|
||||||
|
"StagedSourceObject",
|
||||||
|
"get_data_process_storage",
|
||||||
|
]
|
||||||
69
backend/app/modules/data_process/store/__init__.py
Normal file
69
backend/app/modules/data_process/store/__init__.py
Normal file
@@ -0,0 +1,69 @@
|
|||||||
|
"""数据处理存储层。"""
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
DataProcessStoreError,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
TASK_STATUSES,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_decode_row,
|
||||||
|
_preview_config_changed,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
)
|
||||||
|
from .tasks import TasksMixin
|
||||||
|
from .source_files import SourceFilesMixin
|
||||||
|
from .preview import PreviewMixin
|
||||||
|
from .generation import GenerationMixin
|
||||||
|
from .results import ResultsMixin
|
||||||
|
from .datasets import DatasetsMixin
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessStore(
|
||||||
|
StoreBase,
|
||||||
|
TasksMixin,
|
||||||
|
SourceFilesMixin,
|
||||||
|
PreviewMixin,
|
||||||
|
GenerationMixin,
|
||||||
|
ResultsMixin,
|
||||||
|
DatasetsMixin,
|
||||||
|
):
|
||||||
|
"""数据处理持久层。
|
||||||
|
|
||||||
|
构造函数不会连接数据库或执行迁移。部署方必须显式执行 002 SQL,
|
||||||
|
或在受控的管理命令中调用 :meth:`ensure_schema`,避免应用启动时
|
||||||
|
修改远程数据库。
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def get_data_process_store() -> DataProcessStore:
|
||||||
|
"""获取数据处理存储实例。"""
|
||||||
|
return DataProcessStore()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"DataProcessStore",
|
||||||
|
"DataProcessStoreError",
|
||||||
|
"NotFoundError",
|
||||||
|
"ConflictError",
|
||||||
|
"InvalidStateError",
|
||||||
|
"get_data_process_store",
|
||||||
|
"utcnow",
|
||||||
|
"new_id",
|
||||||
|
"repeat_task_id",
|
||||||
|
"TASK_STATUSES",
|
||||||
|
"EDITABLE_STATUSES",
|
||||||
|
"_decode_row",
|
||||||
|
"_preview_config_changed",
|
||||||
|
"_reasoning_output_is_valid",
|
||||||
|
"_source_storage_descriptor",
|
||||||
|
]
|
||||||
305
backend/app/modules/data_process/store/base.py
Normal file
305
backend/app/modules/data_process/store/base.py
Normal file
@@ -0,0 +1,305 @@
|
|||||||
|
"""数据处理存储层 - 基础设施。
|
||||||
|
|
||||||
|
包含:异常类、工具函数、常量定义、基类。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import re
|
||||||
|
import uuid
|
||||||
|
from collections.abc import Iterator
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from datetime import UTC, date, datetime
|
||||||
|
from decimal import Decimal
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
from psycopg.rows import dict_row
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
|
||||||
|
# 常量定义
|
||||||
|
TASK_STATUSES = {"pending", "running", "completed", "failed", "stopped"}
|
||||||
|
EDITABLE_STATUSES = {"pending", "failed", "stopped", "completed"}
|
||||||
|
ACTIVE_PREVIEW_STATUSES = {"queued", "running"}
|
||||||
|
WORKFLOW_STEPS = {"create", "model", "upload", "preview", "generate", "results"}
|
||||||
|
|
||||||
|
_PREVIEW_CONFIG_ALIASES = {
|
||||||
|
"preprocess_options": "preprocessOptions",
|
||||||
|
"chunk_method": "chunkMethod",
|
||||||
|
"chunk_size": "chunkSize",
|
||||||
|
"chunk_overlap": "chunkOverlap",
|
||||||
|
"min_chunk_size": "minChunkSize",
|
||||||
|
"semantic_breakpoint_percentile": "semanticBreakpointPercentile",
|
||||||
|
"preserve_tables": "preserveTables",
|
||||||
|
"preserve_code_blocks": "preserveCodeBlocks",
|
||||||
|
"preserve_lists": "preserveLists",
|
||||||
|
}
|
||||||
|
|
||||||
|
_UNSTRUCTURED_PREVIEW_DEFAULTS: dict[str, Any] = {
|
||||||
|
"chunk_method": "layout_hybrid",
|
||||||
|
"chunk_size": 800,
|
||||||
|
"chunk_overlap": 100,
|
||||||
|
"min_chunk_size": 100,
|
||||||
|
"semantic_breakpoint_percentile": 95,
|
||||||
|
"preserve_tables": True,
|
||||||
|
"preserve_code_blocks": True,
|
||||||
|
"preserve_lists": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
_REGENERATION_MARKER_KEY = "_regeneration_prepared"
|
||||||
|
_REPEAT_SOURCE_TASK_KEY = "_repeat_source_task_id"
|
||||||
|
_REPEAT_REQUEST_KEY = "_repeat_request_id"
|
||||||
|
_INTERNAL_CONFIG_KEYS = {
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# 异常类
|
||||||
|
class DataProcessStoreError(RuntimeError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class NotFoundError(DataProcessStoreError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class ConflictError(DataProcessStoreError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class InvalidStateError(DataProcessStoreError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
# 工具函数
|
||||||
|
def utcnow() -> str:
|
||||||
|
return datetime.now(UTC).replace(microsecond=0).isoformat().replace("+00:00", "Z")
|
||||||
|
|
||||||
|
|
||||||
|
def new_id(prefix: str) -> str:
|
||||||
|
return f"{prefix}_{uuid.uuid4().hex[:20]}"
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_task_id(source_task_id: str, request_id: str) -> str:
|
||||||
|
"""按源任务和请求幂等键生成稳定的新任务 ID。"""
|
||||||
|
digest = hashlib.sha256(f"{source_task_id}:{request_id}".encode()).hexdigest()
|
||||||
|
return f"dpt_{digest[:20]}"
|
||||||
|
|
||||||
|
|
||||||
|
def json_dumps(value: Any) -> str:
|
||||||
|
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
||||||
|
|
||||||
|
|
||||||
|
def _database_url(value: str) -> str:
|
||||||
|
return value.replace("postgresql+psycopg://", "postgresql://")
|
||||||
|
|
||||||
|
|
||||||
|
def _json_value(value: Any, default: Any) -> Any:
|
||||||
|
if value is None or value == "":
|
||||||
|
return default
|
||||||
|
if isinstance(value, (dict, list)):
|
||||||
|
return value
|
||||||
|
try:
|
||||||
|
return json.loads(value)
|
||||||
|
except (TypeError, json.JSONDecodeError):
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
|
def _task_output_type(task: dict[str, Any]) -> str:
|
||||||
|
config = _json_value(task.get("config"), {})
|
||||||
|
if not isinstance(config, dict):
|
||||||
|
return "standard"
|
||||||
|
return str(config.get("output_type") or config.get("outputType") or "standard")
|
||||||
|
|
||||||
|
|
||||||
|
def _task_reasoning_detail(task: dict[str, Any]) -> str:
|
||||||
|
config = _json_value(task.get("config"), {})
|
||||||
|
if not isinstance(config, dict):
|
||||||
|
return "normal"
|
||||||
|
return str(
|
||||||
|
config.get("reasoning_detail")
|
||||||
|
or config.get("reasoningDetail")
|
||||||
|
or "normal"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _reasoning_output_is_valid(value: Any) -> bool:
|
||||||
|
match = re.fullmatch(
|
||||||
|
r"\s*<think>\s*(?P<reasoning>[\s\S]*?)\s*</think>\s*(?P<answer>[\s\S]+?)\s*",
|
||||||
|
str(value or ""),
|
||||||
|
flags=re.IGNORECASE,
|
||||||
|
)
|
||||||
|
return bool(
|
||||||
|
match
|
||||||
|
and match.group("reasoning").strip()
|
||||||
|
and match.group("answer").strip()
|
||||||
|
and all(
|
||||||
|
tag not in part.lower()
|
||||||
|
for tag in ("<think", "</think")
|
||||||
|
for part in (match.group("reasoning"), match.group("answer"))
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _dpo_fields_are_valid(row: dict[str, Any]) -> bool:
|
||||||
|
chosen = str(row.get("chosen") or "").strip()
|
||||||
|
rejected = str(row.get("rejected") or "").strip()
|
||||||
|
return bool(chosen and rejected and chosen != rejected)
|
||||||
|
|
||||||
|
|
||||||
|
def _preview_config_value(config: dict[str, Any], key: str, default: Any) -> Any:
|
||||||
|
if key in config:
|
||||||
|
return config[key]
|
||||||
|
return config.get(_PREVIEW_CONFIG_ALIASES[key], default)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalized_preprocess_options(config: dict[str, Any]) -> Any:
|
||||||
|
value = _preview_config_value(config, "preprocess_options", [])
|
||||||
|
if isinstance(value, (list, tuple, set)):
|
||||||
|
return tuple(sorted({str(item) for item in value}))
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _preview_config_projection(
|
||||||
|
process_type: str,
|
||||||
|
config: dict[str, Any],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""只投影会改变预览切片的配置。
|
||||||
|
|
||||||
|
生成模型、提示词、温度等参数不影响源文切片,因此不应该
|
||||||
|
破坏用户已经校对过的预览内容。
|
||||||
|
"""
|
||||||
|
projection: dict[str, Any] = {
|
||||||
|
"preprocess_options": _normalized_preprocess_options(config),
|
||||||
|
}
|
||||||
|
if process_type != "unstructured":
|
||||||
|
return projection
|
||||||
|
for key, default in _UNSTRUCTURED_PREVIEW_DEFAULTS.items():
|
||||||
|
projection[key] = _preview_config_value(config, key, default)
|
||||||
|
return projection
|
||||||
|
|
||||||
|
|
||||||
|
def _preview_config_changed(
|
||||||
|
process_type: str,
|
||||||
|
current_config: dict[str, Any],
|
||||||
|
next_config: dict[str, Any],
|
||||||
|
) -> bool:
|
||||||
|
return _preview_config_projection(process_type, current_config) != _preview_config_projection(
|
||||||
|
process_type, next_config
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _regeneration_marker(task: dict[str, Any]) -> dict[str, Any] | None:
|
||||||
|
config = task.get("config")
|
||||||
|
if not isinstance(config, dict):
|
||||||
|
return None
|
||||||
|
marker = config.get(_REGENERATION_MARKER_KEY)
|
||||||
|
if not isinstance(marker, dict) or marker.get("prepared") is not True:
|
||||||
|
return None
|
||||||
|
return marker
|
||||||
|
|
||||||
|
|
||||||
|
def _is_regeneration_prepared(task: dict[str, Any]) -> bool:
|
||||||
|
return _regeneration_marker(task) is not None
|
||||||
|
|
||||||
|
|
||||||
|
def _business_config(config: dict[str, Any] | None) -> dict[str, Any]:
|
||||||
|
"""过滤只供服务端维护的工作流标记。"""
|
||||||
|
return {
|
||||||
|
key: value
|
||||||
|
for key, value in (config or {}).items()
|
||||||
|
if key not in _INTERNAL_CONFIG_KEYS
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _public_task(item: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||||
|
"""从 API 任务快照中移除服务端内部工作流标记。"""
|
||||||
|
if item is None:
|
||||||
|
return None
|
||||||
|
public = dict(item)
|
||||||
|
config = public.get("config")
|
||||||
|
if isinstance(config, dict):
|
||||||
|
public["config"] = _business_config(config)
|
||||||
|
return public
|
||||||
|
|
||||||
|
|
||||||
|
def _serialize_value(value: Any) -> Any:
|
||||||
|
if isinstance(value, (datetime, date)):
|
||||||
|
return value.isoformat().replace("+00:00", "Z")
|
||||||
|
if isinstance(value, Decimal):
|
||||||
|
return float(value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _source_storage_descriptor(
|
||||||
|
payload: dict[str, Any],
|
||||||
|
task_id: str,
|
||||||
|
file_id: str,
|
||||||
|
) -> tuple[str, dict[str, Any]]:
|
||||||
|
storage_object_id = str(
|
||||||
|
payload.get("storage_object_id")
|
||||||
|
or f"db://data-process/{task_id}/{file_id}/v1"
|
||||||
|
)
|
||||||
|
expected_local_prefix = f"local://data-process/{task_id}/{file_id}/v1/"
|
||||||
|
expected_database_reference = f"db://data-process/{task_id}/{file_id}/v1"
|
||||||
|
if storage_object_id.startswith(expected_local_prefix) and len(storage_object_id) > len(
|
||||||
|
expected_local_prefix
|
||||||
|
):
|
||||||
|
storage_backend = "local"
|
||||||
|
elif storage_object_id == expected_database_reference:
|
||||||
|
storage_backend = "database"
|
||||||
|
elif storage_object_id.startswith(("local://data-process/", "db://data-process/")):
|
||||||
|
raise DataProcessStoreError("source storage object owner mismatch")
|
||||||
|
else:
|
||||||
|
raise DataProcessStoreError("unsupported source storage object reference")
|
||||||
|
metadata = {
|
||||||
|
**(payload.get("metadata") or {}),
|
||||||
|
"storage_backend": storage_backend,
|
||||||
|
}
|
||||||
|
return storage_object_id, metadata
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_row(row: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
item = {key: _serialize_value(value) for key, value in row.items()}
|
||||||
|
for key, default in {
|
||||||
|
"config": {},
|
||||||
|
"metadata": {},
|
||||||
|
"quality_score": {},
|
||||||
|
"versions": [],
|
||||||
|
"output_datasets": [],
|
||||||
|
}.items():
|
||||||
|
if key in item:
|
||||||
|
item[key] = _json_value(item[key], default)
|
||||||
|
return item
|
||||||
|
|
||||||
|
|
||||||
|
class StoreBase:
|
||||||
|
"""数据处理存储基类。"""
|
||||||
|
|
||||||
|
def __init__(self, database_url: str | None = None) -> None:
|
||||||
|
self.database_url = _database_url(database_url or get_settings().database_url)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def connect(self) -> Iterator[psycopg.Connection[dict[str, Any]]]:
|
||||||
|
with psycopg.connect(self.database_url, row_factory=dict_row) as conn:
|
||||||
|
try:
|
||||||
|
yield conn
|
||||||
|
conn.commit()
|
||||||
|
except Exception:
|
||||||
|
conn.rollback()
|
||||||
|
raise
|
||||||
|
|
||||||
|
def ensure_schema(self) -> None:
|
||||||
|
"""显式安装数据处理表;API 路由和应用启动流程不会调用此方法。"""
|
||||||
|
schema_path = Path(__file__).resolve().parents[3] / "db" / "sql" / "002_data_process.sql"
|
||||||
|
sql = schema_path.read_text(encoding="utf-8")
|
||||||
|
with self.connect() as conn, conn.cursor() as cursor:
|
||||||
|
cursor.execute(sql)
|
||||||
555
backend/app/modules/data_process/store/datasets.py
Normal file
555
backend/app/modules/data_process/store/datasets.py
Normal file
@@ -0,0 +1,555 @@
|
|||||||
|
"""数据处理存储层 - 数据集发布。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
from app.modules.storage.minio_store import get_object_storage
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
_json_value,
|
||||||
|
_decode_row,
|
||||||
|
_public_task,
|
||||||
|
_business_config,
|
||||||
|
_preview_config_value,
|
||||||
|
_preview_config_changed,
|
||||||
|
_preview_config_projection,
|
||||||
|
_normalized_preprocess_options,
|
||||||
|
_regeneration_marker,
|
||||||
|
_is_regeneration_prepared,
|
||||||
|
_task_output_type,
|
||||||
|
_task_reasoning_detail,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_dpo_fields_are_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
)
|
||||||
|
|
||||||
|
from ..algorithms import stable_split_assignments
|
||||||
|
|
||||||
|
class DatasetsMixin:
|
||||||
|
"""数据集发布 Mixin。"""
|
||||||
|
|
||||||
|
def get_generation_model(self, model_id: str) -> dict[str, Any]:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id, name, type, purpose, model_source, description, path,
|
||||||
|
api_url, api_key, online_model_name, create_time
|
||||||
|
FROM models WHERE id=%s
|
||||||
|
""",
|
||||||
|
(model_id,),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("generation model not found")
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def save_generation_model_snapshot(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
model_snapshot: dict[str, Any],
|
||||||
|
*,
|
||||||
|
generation_run_id: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
# API 密钥仅用于本次调用,绝不能进入任务配置、详情响应或审计快照。
|
||||||
|
safe_snapshot = {
|
||||||
|
key: value for key, value in model_snapshot.items() if key != "api_key"
|
||||||
|
}
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if (
|
||||||
|
task["status"] != "running"
|
||||||
|
or task.get("generation_run_id") != generation_run_id
|
||||||
|
):
|
||||||
|
raise InvalidStateError("generation run is no longer active")
|
||||||
|
config = dict(task.get("config") or {})
|
||||||
|
config["generation_model_snapshot"] = safe_snapshot
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks SET config=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND generation_run_id=%s RETURNING *
|
||||||
|
""",
|
||||||
|
(json_dumps(config), utcnow(), task_id, generation_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def publish(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
"""按精确配额发布训练、验证、测试三个独立数据集。"""
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if _is_regeneration_prepared(task):
|
||||||
|
raise InvalidStateError(
|
||||||
|
"regeneration must start and complete before publishing"
|
||||||
|
)
|
||||||
|
if task["status"] != "completed":
|
||||||
|
raise InvalidStateError("only a completed task can be published")
|
||||||
|
if not task.get("results_confirmed"):
|
||||||
|
raise InvalidStateError("results must be confirmed before publishing")
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT * FROM data_process_results
|
||||||
|
WHERE task_id=%s ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchall()
|
||||||
|
if not rows:
|
||||||
|
raise InvalidStateError("task has no results to publish")
|
||||||
|
invalid_count = sum(
|
||||||
|
1
|
||||||
|
for row in rows
|
||||||
|
if row["status"] == "invalid"
|
||||||
|
or not str(row.get("instruction") or "").strip()
|
||||||
|
or not str(row.get("output") or "").strip()
|
||||||
|
or (
|
||||||
|
_task_output_type(task) == "reasoning"
|
||||||
|
and not _reasoning_output_is_valid(row.get("output"))
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
_task_output_type(task) == "dpo"
|
||||||
|
and not _dpo_fields_are_valid(row)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if invalid_count:
|
||||||
|
raise InvalidStateError(f"task contains {invalid_count} invalid results")
|
||||||
|
|
||||||
|
now = utcnow()
|
||||||
|
requested_split = payload.get("split") or {
|
||||||
|
"train": 80,
|
||||||
|
"validation": 10,
|
||||||
|
"test": 10,
|
||||||
|
}
|
||||||
|
assignments = stable_split_assignments(
|
||||||
|
[str(row["id"]) for row in rows],
|
||||||
|
requested_split,
|
||||||
|
seed=task_id,
|
||||||
|
)
|
||||||
|
if _task_output_type(task) == "dpo":
|
||||||
|
records = [
|
||||||
|
{
|
||||||
|
"instruction": row["instruction"],
|
||||||
|
"input": row["input"],
|
||||||
|
"chosen": row["chosen"],
|
||||||
|
"rejected": row["rejected"],
|
||||||
|
"split": assignment,
|
||||||
|
}
|
||||||
|
for row, assignment in zip(rows, assignments, strict=True)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
records = [
|
||||||
|
{
|
||||||
|
"instruction": row["instruction"],
|
||||||
|
"input": row["input"],
|
||||||
|
"output": row["output"],
|
||||||
|
"split": assignment,
|
||||||
|
}
|
||||||
|
for row, assignment in zip(rows, assignments, strict=True)
|
||||||
|
]
|
||||||
|
split_order = ("train", "validation", "test")
|
||||||
|
split_counts = {
|
||||||
|
split_name: assignments.count(split_name) for split_name in split_order
|
||||||
|
}
|
||||||
|
split_specs: list[dict[str, Any]] = []
|
||||||
|
use_minio = bool(get_settings().minio_enabled)
|
||||||
|
for split_name in split_order:
|
||||||
|
split_records = [
|
||||||
|
(source_row, record)
|
||||||
|
for source_row, record in zip(rows, records, strict=True)
|
||||||
|
if record["split"] == split_name
|
||||||
|
]
|
||||||
|
file_id = new_id("dfile")
|
||||||
|
version_id = new_id("dfv")
|
||||||
|
content = "".join(
|
||||||
|
json_dumps(record) + "\n" for _, record in split_records
|
||||||
|
)
|
||||||
|
raw = content.encode("utf-8")
|
||||||
|
split_specs.append(
|
||||||
|
{
|
||||||
|
"split": split_name,
|
||||||
|
"records": split_records,
|
||||||
|
"file_id": file_id,
|
||||||
|
"version_id": version_id,
|
||||||
|
"content": content,
|
||||||
|
"raw": raw,
|
||||||
|
"checksum": hashlib.sha256(raw).hexdigest(),
|
||||||
|
"storage_object_id": (
|
||||||
|
f"db://data-process/{task_id}/{file_id}/v1"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
source_result_ids = [row["id"] for row in rows]
|
||||||
|
common_metadata = {
|
||||||
|
"source": "data_process",
|
||||||
|
"storage_backend": "minio" if use_minio else "database",
|
||||||
|
"source_task_id": task_id,
|
||||||
|
"output_type": _task_output_type(task),
|
||||||
|
"reasoning_detail": _task_reasoning_detail(task),
|
||||||
|
"source_file_ids": [item["id"] for item in self._source_ids(conn, task_id)],
|
||||||
|
"source_result_ids": source_result_ids,
|
||||||
|
"format": (
|
||||||
|
"dpo"
|
||||||
|
if _task_output_type(task) == "dpo"
|
||||||
|
else payload.get("format") or "alpaca_jsonl"
|
||||||
|
),
|
||||||
|
"split": requested_split,
|
||||||
|
}
|
||||||
|
|
||||||
|
existing_datasets = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT * FROM datasets
|
||||||
|
WHERE source='task' AND deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
source_task_id=%s
|
||||||
|
OR (source_task_id IS NULL AND task_id=%s)
|
||||||
|
)
|
||||||
|
ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(task_id, task_id),
|
||||||
|
).fetchall()
|
||||||
|
existing_by_split: dict[str, dict[str, Any]] = {}
|
||||||
|
primary_existing = None
|
||||||
|
for existing in existing_datasets:
|
||||||
|
existing_metadata = _json_value(existing.get("metadata"), {})
|
||||||
|
existing_split = str(existing_metadata.get("dataset_split") or "")
|
||||||
|
if existing_split in split_order:
|
||||||
|
existing_by_split[existing_split] = existing
|
||||||
|
if str(existing["id"]) == str(task.get("output_dataset_id") or ""):
|
||||||
|
primary_existing = existing
|
||||||
|
if primary_existing and "train" not in existing_by_split:
|
||||||
|
# 兼容旧版“一个数据集包含三个文件”的发布物,原数据集复用为训练集。
|
||||||
|
existing_by_split["train"] = primary_existing
|
||||||
|
|
||||||
|
existing_group_metadata = _json_value(
|
||||||
|
(primary_existing or {}).get("metadata"), {}
|
||||||
|
)
|
||||||
|
base_dataset_name = str(
|
||||||
|
existing_group_metadata.get("base_dataset_name")
|
||||||
|
or payload["dataset_name"]
|
||||||
|
).strip()
|
||||||
|
for suffix in ("-训练集", "-验证集", "-测试集"):
|
||||||
|
if base_dataset_name.endswith(suffix):
|
||||||
|
base_dataset_name = base_dataset_name[: -len(suffix)].rstrip()
|
||||||
|
break
|
||||||
|
split_group_id = str(
|
||||||
|
existing_group_metadata.get("split_group_id")
|
||||||
|
or f"dsg_{hashlib.sha256(task_id.encode()).hexdigest()[:20]}"
|
||||||
|
)
|
||||||
|
dataset_ids = {
|
||||||
|
spec["split"]: str(existing_by_split[spec["split"]]["id"])
|
||||||
|
if spec["split"] in existing_by_split
|
||||||
|
else new_id("dataset")
|
||||||
|
for spec in split_specs
|
||||||
|
}
|
||||||
|
created_any = any(
|
||||||
|
spec["split"] not in existing_by_split for spec in split_specs
|
||||||
|
)
|
||||||
|
split_labels = {
|
||||||
|
"train": "训练集",
|
||||||
|
"validation": "验证集",
|
||||||
|
"test": "测试集",
|
||||||
|
}
|
||||||
|
dataset_types = {"train": "train", "validation": "val", "test": "test"}
|
||||||
|
published_datasets: list[dict[str, Any]] = []
|
||||||
|
try:
|
||||||
|
for spec in split_specs:
|
||||||
|
split_name = str(spec["split"])
|
||||||
|
dataset_id = dataset_ids[split_name]
|
||||||
|
storage_object_id = str(spec["storage_object_id"])
|
||||||
|
if use_minio:
|
||||||
|
file_name = f"{base_dataset_name}.{split_name}.jsonl"
|
||||||
|
object_key = f"datasets/{dataset_id}/versions/{spec['version_id']}/{file_name}"
|
||||||
|
uploaded = get_object_storage().put_bytes(
|
||||||
|
object_key,
|
||||||
|
spec["raw"],
|
||||||
|
"application/jsonl",
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO storage_objects
|
||||||
|
(id, resource_type, resource_id, version_id, bucket, object_key,
|
||||||
|
file_name, content_type, checksum_sha256, byte_size, status,
|
||||||
|
created_by, create_time)
|
||||||
|
VALUES (%s, 'dataset', %s, %s, %s, %s, %s, %s, %s, %s,
|
||||||
|
'available', %s, %s)
|
||||||
|
ON CONFLICT (resource_type, resource_id, version_id, object_key)
|
||||||
|
DO UPDATE SET file_name=EXCLUDED.file_name,
|
||||||
|
content_type=EXCLUDED.content_type,
|
||||||
|
checksum_sha256=EXCLUDED.checksum_sha256,
|
||||||
|
byte_size=EXCLUDED.byte_size,
|
||||||
|
status='available',
|
||||||
|
created_by=EXCLUDED.created_by
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
new_id("object"),
|
||||||
|
dataset_id,
|
||||||
|
spec["version_id"],
|
||||||
|
uploaded["bucket"],
|
||||||
|
object_key,
|
||||||
|
file_name,
|
||||||
|
"application/jsonl",
|
||||||
|
spec["checksum"],
|
||||||
|
len(spec["raw"]),
|
||||||
|
payload.get("created_by") or task.get("created_by"),
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
storage_object_id = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id FROM storage_objects
|
||||||
|
WHERE resource_type='dataset' AND resource_id=%s
|
||||||
|
AND version_id=%s AND object_key=%s
|
||||||
|
""",
|
||||||
|
(dataset_id, spec["version_id"], object_key),
|
||||||
|
).fetchone()["id"]
|
||||||
|
existing_dataset = existing_by_split.get(split_name)
|
||||||
|
dataset_metadata = {
|
||||||
|
**common_metadata,
|
||||||
|
"base_dataset_name": base_dataset_name,
|
||||||
|
"dataset_split": split_name,
|
||||||
|
"split_group_id": split_group_id,
|
||||||
|
"split_dataset_ids": dataset_ids,
|
||||||
|
"split_counts": {
|
||||||
|
name: split_counts[name] if name == split_name else 0
|
||||||
|
for name in split_order
|
||||||
|
},
|
||||||
|
}
|
||||||
|
dataset_name = f"{base_dataset_name}-{split_labels[split_name]}"
|
||||||
|
if existing_dataset:
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM dataset_records WHERE dataset_id=%s", (dataset_id,)
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""DELETE FROM dataset_file_versions
|
||||||
|
WHERE dataset_file_id IN
|
||||||
|
(SELECT id FROM dataset_files WHERE dataset_id=%s)""",
|
||||||
|
(dataset_id,),
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM dataset_files WHERE dataset_id=%s", (dataset_id,)
|
||||||
|
)
|
||||||
|
dataset = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE datasets
|
||||||
|
SET name=%s, type=%s, storage_type=%s, size=%s, size_bytes=%s,
|
||||||
|
count=%s, record_count=%s, description=%s, metadata=%s,
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
dataset_name,
|
||||||
|
dataset_types[split_name],
|
||||||
|
"minio" if use_minio else (payload.get("storage_type") or "local"),
|
||||||
|
f"{len(spec['raw'])} B",
|
||||||
|
len(spec["raw"]),
|
||||||
|
len(spec["records"]),
|
||||||
|
len(spec["records"]),
|
||||||
|
payload.get("description") or task.get("description") or "",
|
||||||
|
json_dumps(dataset_metadata),
|
||||||
|
now,
|
||||||
|
dataset_id,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
else:
|
||||||
|
dataset = conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO datasets
|
||||||
|
(id, name, type, storage_type, source, task_id, source_task_id,
|
||||||
|
size, size_bytes, count, record_count, description, metadata,
|
||||||
|
tenant_id, project_id, owner_id, created_by, create_time,
|
||||||
|
created_at, updated_at)
|
||||||
|
VALUES (
|
||||||
|
%s, %s, %s, %s, 'task', %s, %s,
|
||||||
|
%s, %s, %s, %s, %s, %s,
|
||||||
|
%s, %s, %s, %s, %s, %s, %s
|
||||||
|
)
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
dataset_id,
|
||||||
|
dataset_name,
|
||||||
|
dataset_types[split_name],
|
||||||
|
"minio" if use_minio else (payload.get("storage_type") or "local"),
|
||||||
|
task_id,
|
||||||
|
task_id,
|
||||||
|
f"{len(spec['raw'])} B",
|
||||||
|
len(spec["raw"]),
|
||||||
|
len(spec["records"]),
|
||||||
|
len(spec["records"]),
|
||||||
|
payload.get("description") or task.get("description") or "",
|
||||||
|
json_dumps(dataset_metadata),
|
||||||
|
task.get("tenant_id"),
|
||||||
|
task.get("project_id"),
|
||||||
|
task.get("owner_id"),
|
||||||
|
payload.get("created_by") or task.get("created_by"),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
file_metadata = {**dataset_metadata, "file_split": split_name}
|
||||||
|
version = {
|
||||||
|
"id": spec["version_id"],
|
||||||
|
"version_no": 1,
|
||||||
|
"version": 1,
|
||||||
|
"description": f"data process {split_name} publish",
|
||||||
|
"checksum_sha256": spec["checksum"],
|
||||||
|
"size_bytes": len(spec["raw"]),
|
||||||
|
"record_count": len(spec["records"]),
|
||||||
|
"created_at": now,
|
||||||
|
"create_time": now,
|
||||||
|
"source_task_id": task_id,
|
||||||
|
"storage_object_id": storage_object_id,
|
||||||
|
}
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO dataset_files
|
||||||
|
(id, dataset_id, name, storage_object_id, size, content,
|
||||||
|
active_version_id, versions, create_time, current_version_id,
|
||||||
|
size_bytes, record_count, file_format, checksum_sha256, version_no,
|
||||||
|
source_task_id, tenant_id, project_id, created_by, metadata,
|
||||||
|
created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
||||||
|
%s, %s, 1, %s, %s, %s, %s, %s, %s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
spec["file_id"],
|
||||||
|
dataset_id,
|
||||||
|
f"{base_dataset_name}.{split_name}.jsonl",
|
||||||
|
storage_object_id,
|
||||||
|
f"{len(spec['raw'])} B",
|
||||||
|
spec["content"],
|
||||||
|
spec["version_id"],
|
||||||
|
json_dumps([version]),
|
||||||
|
now,
|
||||||
|
spec["version_id"],
|
||||||
|
len(spec["raw"]),
|
||||||
|
len(spec["records"]),
|
||||||
|
"jsonl",
|
||||||
|
spec["checksum"],
|
||||||
|
task_id,
|
||||||
|
task.get("tenant_id"),
|
||||||
|
task.get("project_id"),
|
||||||
|
payload.get("created_by") or task.get("created_by"),
|
||||||
|
json_dumps(file_metadata),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO dataset_file_versions
|
||||||
|
(id, dataset_file_id, version_no, storage_object_id, content_preview,
|
||||||
|
description, size_bytes, record_count, checksum_sha256,
|
||||||
|
source_task_id, metadata, created_by, created_at)
|
||||||
|
VALUES (%s, %s, 1, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
spec["version_id"],
|
||||||
|
spec["file_id"],
|
||||||
|
storage_object_id,
|
||||||
|
spec["content"][:2000],
|
||||||
|
f"data process {split_name} publish",
|
||||||
|
len(spec["raw"]),
|
||||||
|
len(spec["records"]),
|
||||||
|
spec["checksum"],
|
||||||
|
task_id,
|
||||||
|
json_dumps(file_metadata),
|
||||||
|
payload.get("created_by") or task.get("created_by"),
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for line_number, (source_row, record) in enumerate(
|
||||||
|
spec["records"], start=1
|
||||||
|
):
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO dataset_records
|
||||||
|
(id, dataset_id, dataset_file_id, version_id, line_no, split,
|
||||||
|
instruction, input, output, raw, status, source_task_id,
|
||||||
|
source_result_id, preview_item_id, created_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
||||||
|
%s, %s, %s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
new_id("drec"),
|
||||||
|
dataset_id,
|
||||||
|
spec["file_id"],
|
||||||
|
spec["version_id"],
|
||||||
|
line_number,
|
||||||
|
record["split"],
|
||||||
|
record["instruction"],
|
||||||
|
record["input"],
|
||||||
|
record.get("output") or record.get("chosen") or "",
|
||||||
|
json_dumps(
|
||||||
|
{
|
||||||
|
**record,
|
||||||
|
"source_task_id": task_id,
|
||||||
|
"source_result_id": source_row["id"],
|
||||||
|
"preview_item_id": source_row.get("preview_item_id"),
|
||||||
|
}
|
||||||
|
),
|
||||||
|
source_row["status"],
|
||||||
|
task_id,
|
||||||
|
source_row["id"],
|
||||||
|
source_row.get("preview_item_id"),
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
published_datasets.append(_decode_row(dataset) or {})
|
||||||
|
|
||||||
|
except psycopg.errors.UniqueViolation as exc:
|
||||||
|
raise ConflictError("dataset name already exists") from exc
|
||||||
|
train_dataset_id = dataset_ids.get("train")
|
||||||
|
if not train_dataset_id:
|
||||||
|
raise InvalidStateError("published split does not contain training data")
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET output_dataset_id=%s, updated_at=%s, updated_by=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(train_dataset_id, now, payload.get("created_by"), task_id),
|
||||||
|
)
|
||||||
|
train_dataset = next(
|
||||||
|
item
|
||||||
|
for item in published_datasets
|
||||||
|
if _json_value(item.get("metadata"), {}).get("dataset_split") == "train"
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"dataset": train_dataset,
|
||||||
|
"datasets": published_datasets,
|
||||||
|
"output_datasets": published_datasets,
|
||||||
|
"created": created_any,
|
||||||
|
"split_counts": split_counts,
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _source_ids(
|
||||||
|
conn: psycopg.Connection[dict[str, Any]], task_id: str
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
return conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchall()
|
||||||
275
backend/app/modules/data_process/store/generation.py
Normal file
275
backend/app/modules/data_process/store/generation.py
Normal file
@@ -0,0 +1,275 @@
|
|||||||
|
"""数据处理存储层 - 生成管理。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
_json_value,
|
||||||
|
_decode_row,
|
||||||
|
_public_task,
|
||||||
|
_business_config,
|
||||||
|
_preview_config_value,
|
||||||
|
_preview_config_changed,
|
||||||
|
_preview_config_projection,
|
||||||
|
_normalized_preprocess_options,
|
||||||
|
_regeneration_marker,
|
||||||
|
_is_regeneration_prepared,
|
||||||
|
_task_output_type,
|
||||||
|
_task_reasoning_detail,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_dpo_fields_are_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
DataProcessStoreError,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class GenerationMixin:
|
||||||
|
"""生成管理 Mixin。"""
|
||||||
|
|
||||||
|
def _invalidate_results(
|
||||||
|
self,
|
||||||
|
conn: psycopg.Connection[dict[str, Any]],
|
||||||
|
task: dict[str, Any],
|
||||||
|
task_id: str,
|
||||||
|
now: str,
|
||||||
|
) -> None:
|
||||||
|
if _is_regeneration_prepared(task):
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_process_tasks SET updated_at=%s WHERE id=%s",
|
||||||
|
(now, task_id),
|
||||||
|
)
|
||||||
|
return
|
||||||
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='pending', progress=20, output_count=0, filtered_count=0,
|
||||||
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
||||||
|
generation_run_id=NULL, results_confirmed=FALSE, updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(now, task_id),
|
||||||
|
)
|
||||||
|
|
||||||
|
def start_generation(self, task_id: str, *, replace_existing: bool = True) -> dict[str, Any]:
|
||||||
|
if not replace_existing:
|
||||||
|
raise DataProcessStoreError("incremental generation is not supported")
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
regeneration_prepared = _is_regeneration_prepared(task)
|
||||||
|
if task.get("output_dataset_id") and not regeneration_prepared:
|
||||||
|
raise InvalidStateError("published task cannot be regenerated")
|
||||||
|
if task["status"] == "running":
|
||||||
|
raise ConflictError("data process task is already running")
|
||||||
|
if task.get("preview_status") in ACTIVE_PREVIEW_STATUSES:
|
||||||
|
raise ConflictError("preview is still running")
|
||||||
|
preview_count = conn.execute(
|
||||||
|
"SELECT COUNT(*) AS count FROM data_process_preview_items WHERE task_id=%s",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()["count"]
|
||||||
|
if not preview_count:
|
||||||
|
raise InvalidStateError("preview must be built before generation")
|
||||||
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
||||||
|
now = utcnow()
|
||||||
|
generation_run_id = new_id("dprun")
|
||||||
|
next_config = dict(task.get("config") or {})
|
||||||
|
next_config.pop(_REGENERATION_MARKER_KEY, None)
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET config=%s, status='running', progress=30, failure_reason=NULL,
|
||||||
|
started_at=%s, completed_at=NULL, output_dataset_id=NULL,
|
||||||
|
output_count=0, filtered_count=0, duplicate_count=0, error_count=0,
|
||||||
|
generation_run_id=%s, results_confirmed=FALSE,
|
||||||
|
workflow_step='generate', updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(json_dumps(next_config), now, generation_run_id, now, task_id),
|
||||||
|
).fetchone()
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def stop_task(self, task_id: str) -> dict[str, Any]:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if task["status"] != "running":
|
||||||
|
raise InvalidStateError("only a running task can be stopped")
|
||||||
|
now = utcnow()
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='stopped', failure_reason=NULL, generation_run_id=NULL,
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s RETURNING *
|
||||||
|
""",
|
||||||
|
(now, task_id),
|
||||||
|
).fetchone()
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def generation_is_running(self, task_id: str, generation_run_id: str) -> bool:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT status, generation_run_id
|
||||||
|
FROM data_process_tasks
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
return bool(
|
||||||
|
row
|
||||||
|
and row.get("status") == "running"
|
||||||
|
and row.get("generation_run_id") == generation_run_id
|
||||||
|
)
|
||||||
|
|
||||||
|
def update_generation_progress(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
generation_run_id: str,
|
||||||
|
processed_count: int,
|
||||||
|
total_count: int,
|
||||||
|
) -> bool:
|
||||||
|
ratio = processed_count / max(1, total_count)
|
||||||
|
progress = min(95.0, 30.0 + ratio * 65.0)
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET progress=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND status='running' AND generation_run_id=%s
|
||||||
|
RETURNING id
|
||||||
|
""",
|
||||||
|
(progress, utcnow(), task_id, generation_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return row is not None
|
||||||
|
|
||||||
|
def complete_generation(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
results: Sequence[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
generation_run_id: str,
|
||||||
|
filtered_count: int = 0,
|
||||||
|
duplicate_count: int = 0,
|
||||||
|
error_count: int = 0,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
now = utcnow()
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if (
|
||||||
|
task["status"] != "running"
|
||||||
|
or task.get("generation_run_id") != generation_run_id
|
||||||
|
):
|
||||||
|
return task
|
||||||
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
||||||
|
for result in results:
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_results
|
||||||
|
(id, task_id, preview_item_id, instruction, input, output,
|
||||||
|
chosen, rejected, original_instruction, original_input,
|
||||||
|
original_output, original_chosen, original_rejected, status, error,
|
||||||
|
split, quality_score, created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
||||||
|
%s, %s, %s, %s, %s, %s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
result.get("id") or new_id("dpr"),
|
||||||
|
task_id,
|
||||||
|
result.get("preview_item_id"),
|
||||||
|
result.get("instruction") or "",
|
||||||
|
result.get("input") or "",
|
||||||
|
result.get("output") or "",
|
||||||
|
result.get("chosen") or "",
|
||||||
|
result.get("rejected") or "",
|
||||||
|
result.get("original_instruction", result.get("instruction") or ""),
|
||||||
|
result.get("original_input", result.get("input") or ""),
|
||||||
|
result.get("original_output", result.get("output") or ""),
|
||||||
|
result.get("original_chosen", result.get("chosen") or ""),
|
||||||
|
result.get("original_rejected", result.get("rejected") or ""),
|
||||||
|
result.get("status") or "valid",
|
||||||
|
result.get("error"),
|
||||||
|
result.get("split"),
|
||||||
|
json_dumps(result.get("quality_score") or {}),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='completed', progress=100, output_count=%s, filtered_count=%s,
|
||||||
|
duplicate_count=%s, error_count=%s, failure_reason=NULL,
|
||||||
|
completed_at=%s, generation_run_id=NULL, results_confirmed=FALSE,
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s AND generation_run_id=%s RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
len(results),
|
||||||
|
filtered_count,
|
||||||
|
duplicate_count,
|
||||||
|
error_count,
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
task_id,
|
||||||
|
generation_run_id,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def mark_failed(
|
||||||
|
self, task_id: str, reason: str, *, generation_run_id: str
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if (
|
||||||
|
task["status"] != "running"
|
||||||
|
or task.get("generation_run_id") != generation_run_id
|
||||||
|
):
|
||||||
|
return task
|
||||||
|
now = utcnow()
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='failed', failure_reason=%s, completed_at=%s,
|
||||||
|
generation_run_id=NULL, updated_at=%s
|
||||||
|
WHERE id=%s AND generation_run_id=%s RETURNING *
|
||||||
|
""",
|
||||||
|
(reason[:4000], now, now, task_id, generation_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def progress(self, task_id: str) -> dict[str, Any]:
|
||||||
|
task = self.get_task(task_id)
|
||||||
|
return {
|
||||||
|
"task_id": task["id"],
|
||||||
|
"status": task["status"],
|
||||||
|
"progress": float(task.get("progress") or 0),
|
||||||
|
"input_count": int(task.get("input_count") or 0),
|
||||||
|
"output_count": int(task.get("output_count") or 0),
|
||||||
|
"filtered_count": int(task.get("filtered_count") or 0),
|
||||||
|
"duplicate_count": int(task.get("duplicate_count") or 0),
|
||||||
|
"error_count": int(task.get("error_count") or 0),
|
||||||
|
"failure_reason": task.get("failure_reason"),
|
||||||
|
"results_confirmed": bool(task.get("results_confirmed")),
|
||||||
|
"started_at": task.get("started_at"),
|
||||||
|
"completed_at": task.get("completed_at"),
|
||||||
|
}
|
||||||
539
backend/app/modules/data_process/store/preview.py
Normal file
539
backend/app/modules/data_process/store/preview.py
Normal file
@@ -0,0 +1,539 @@
|
|||||||
|
"""数据处理存储层 - 预览管理。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
_json_value,
|
||||||
|
_decode_row,
|
||||||
|
_public_task,
|
||||||
|
_business_config,
|
||||||
|
_preview_config_value,
|
||||||
|
_preview_config_changed,
|
||||||
|
_preview_config_projection,
|
||||||
|
_normalized_preprocess_options,
|
||||||
|
_regeneration_marker,
|
||||||
|
_is_regeneration_prepared,
|
||||||
|
_task_output_type,
|
||||||
|
_task_reasoning_detail,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_dpo_fields_are_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
_serialize_value,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
)
|
||||||
|
from ..algorithms import estimate_token_count # noqa: E402
|
||||||
|
|
||||||
|
|
||||||
|
class PreviewMixin:
|
||||||
|
"""预览管理 Mixin。"""
|
||||||
|
|
||||||
|
def replace_preview_items(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
items: Sequence[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
source_file_ids: Sequence[str] | None = None,
|
||||||
|
preview_run_id: str | None = None,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
selected_ids = (
|
||||||
|
list(dict.fromkeys(str(file_id) for file_id in source_file_ids))
|
||||||
|
if source_file_ids is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if selected_ids is not None:
|
||||||
|
if not selected_ids or any(not file_id for file_id in selected_ids):
|
||||||
|
raise ValueError("source_file_ids must contain non-empty ids")
|
||||||
|
selected_set = set(selected_ids)
|
||||||
|
unexpected = {
|
||||||
|
str(item.get("source_file_id") or "")
|
||||||
|
for item in items
|
||||||
|
if str(item.get("source_file_id") or "") not in selected_set
|
||||||
|
}
|
||||||
|
if unexpected:
|
||||||
|
raise ValueError("preview items contain an unselected source file")
|
||||||
|
preview_file_count = len(selected_ids) if selected_ids is not None else len(
|
||||||
|
{str(item.get("source_file_id") or "") for item in items}
|
||||||
|
)
|
||||||
|
is_direct_build = preview_run_id is None
|
||||||
|
|
||||||
|
now = utcnow()
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if is_direct_build:
|
||||||
|
self._ensure_editable(task)
|
||||||
|
elif (
|
||||||
|
task.get("preview_run_id") != preview_run_id
|
||||||
|
or task.get("preview_status") != "running"
|
||||||
|
):
|
||||||
|
raise InvalidStateError("preview run is no longer active")
|
||||||
|
regeneration_prepared = _is_regeneration_prepared(task)
|
||||||
|
if not regeneration_prepared:
|
||||||
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
||||||
|
if selected_ids is None:
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM data_process_preview_items WHERE task_id=%s", (task_id,)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL AND id=ANY(%s)
|
||||||
|
""",
|
||||||
|
(task_id, selected_ids),
|
||||||
|
).fetchall()
|
||||||
|
found = {str(row["id"]) for row in rows}
|
||||||
|
missing = set(selected_ids) - found
|
||||||
|
if missing:
|
||||||
|
raise NotFoundError(
|
||||||
|
f"source files not found: {', '.join(sorted(missing))}"
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
DELETE FROM data_process_preview_items
|
||||||
|
WHERE task_id=%s AND source_file_id=ANY(%s)
|
||||||
|
""",
|
||||||
|
(task_id, selected_ids),
|
||||||
|
)
|
||||||
|
created: list[dict[str, Any]] = []
|
||||||
|
for item in items:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_preview_items
|
||||||
|
(id, task_id, source_file_id, original_content, edited_content,
|
||||||
|
source_start, source_end, source_start_line, source_end_line,
|
||||||
|
token_count, status, quality_score, created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
item.get("id") or new_id("dpp"),
|
||||||
|
task_id,
|
||||||
|
item.get("source_file_id"),
|
||||||
|
item.get("original_content") or "",
|
||||||
|
item.get("edited_content", item.get("original_content") or ""),
|
||||||
|
item.get("source_start"),
|
||||||
|
item.get("source_end"),
|
||||||
|
item.get("source_start_line"),
|
||||||
|
item.get("source_end_line"),
|
||||||
|
max(0, int(item.get("token_count") or 0)),
|
||||||
|
item.get("status") or "original",
|
||||||
|
json_dumps(item.get("quality_score") or {}),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
created.append(_decode_row(row) or {})
|
||||||
|
if regeneration_prepared:
|
||||||
|
if is_direct_build:
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET workflow_step='preview', preview_status='completed',
|
||||||
|
preview_progress=100, preview_run_id=NULL,
|
||||||
|
preview_failure_reason=NULL, preview_total_files=%s,
|
||||||
|
preview_completed_files=%s, updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(preview_file_count, preview_file_count, now, task_id),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
conn.execute(
|
||||||
|
"UPDATE data_process_tasks SET updated_at=%s WHERE id=%s",
|
||||||
|
(now, task_id),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='pending', progress=20, output_count=0, filtered_count=0,
|
||||||
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
||||||
|
results_confirmed=FALSE,
|
||||||
|
workflow_step=CASE WHEN %s THEN 'preview' ELSE workflow_step END,
|
||||||
|
preview_status=CASE WHEN %s THEN 'completed' ELSE preview_status END,
|
||||||
|
preview_progress=CASE WHEN %s THEN 100 ELSE preview_progress END,
|
||||||
|
preview_run_id=CASE WHEN %s THEN NULL ELSE preview_run_id END,
|
||||||
|
preview_failure_reason=CASE WHEN %s THEN NULL ELSE preview_failure_reason END,
|
||||||
|
preview_total_files=CASE WHEN %s THEN %s ELSE preview_total_files END,
|
||||||
|
preview_completed_files=CASE WHEN %s THEN %s ELSE preview_completed_files END,
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
is_direct_build,
|
||||||
|
is_direct_build,
|
||||||
|
is_direct_build,
|
||||||
|
is_direct_build,
|
||||||
|
is_direct_build,
|
||||||
|
is_direct_build,
|
||||||
|
preview_file_count,
|
||||||
|
is_direct_build,
|
||||||
|
preview_file_count,
|
||||||
|
now,
|
||||||
|
task_id,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return created
|
||||||
|
|
||||||
|
def start_preview(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
*,
|
||||||
|
source_file_ids: Sequence[str] | None = None,
|
||||||
|
) -> tuple[dict[str, Any], list[str]]:
|
||||||
|
"""创建一轮持久化切分任务,并返回本轮固定的源文件集合。"""
|
||||||
|
|
||||||
|
requested_ids = (
|
||||||
|
list(dict.fromkeys(str(file_id) for file_id in source_file_ids))
|
||||||
|
if source_file_ids is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if requested_ids is not None and (
|
||||||
|
not requested_ids or any(not file_id for file_id in requested_ids)
|
||||||
|
):
|
||||||
|
raise ValueError("source_file_ids must contain non-empty ids")
|
||||||
|
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
if requested_ids is None:
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL
|
||||||
|
ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchall()
|
||||||
|
else:
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL AND id=ANY(%s)
|
||||||
|
ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(task_id, requested_ids),
|
||||||
|
).fetchall()
|
||||||
|
selected_ids = [str(row["id"]) for row in rows]
|
||||||
|
if not selected_ids:
|
||||||
|
raise InvalidStateError("at least one source file is required")
|
||||||
|
if requested_ids is not None:
|
||||||
|
missing = set(requested_ids) - set(selected_ids)
|
||||||
|
if missing:
|
||||||
|
raise NotFoundError(
|
||||||
|
f"source files not found: {', '.join(sorted(missing))}"
|
||||||
|
)
|
||||||
|
|
||||||
|
regeneration_prepared = _is_regeneration_prepared(task)
|
||||||
|
if not regeneration_prepared:
|
||||||
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
||||||
|
preview_run_id = new_id("dpprun")
|
||||||
|
now = utcnow()
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status=CASE WHEN %s THEN status ELSE 'pending' END,
|
||||||
|
progress=CASE WHEN %s THEN progress ELSE 0 END,
|
||||||
|
output_count=CASE WHEN %s THEN output_count ELSE 0 END,
|
||||||
|
filtered_count=CASE WHEN %s THEN filtered_count ELSE 0 END,
|
||||||
|
duplicate_count=CASE WHEN %s THEN duplicate_count ELSE 0 END,
|
||||||
|
error_count=CASE WHEN %s THEN error_count ELSE 0 END,
|
||||||
|
failure_reason=CASE WHEN %s THEN failure_reason ELSE NULL END,
|
||||||
|
results_confirmed=CASE WHEN %s THEN results_confirmed ELSE FALSE END,
|
||||||
|
workflow_step='upload', preview_status='queued', preview_progress=0,
|
||||||
|
preview_run_id=%s, preview_failure_reason=NULL,
|
||||||
|
preview_total_files=%s, preview_completed_files=0,
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
regeneration_prepared,
|
||||||
|
preview_run_id,
|
||||||
|
len(selected_ids),
|
||||||
|
now,
|
||||||
|
task_id,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
return _public_task(_decode_row(row)) or {}, selected_ids
|
||||||
|
|
||||||
|
def mark_preview_running(self, task_id: str, preview_run_id: str) -> bool:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET preview_status='running', updated_at=%s
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
AND preview_status='queued' AND preview_run_id=%s
|
||||||
|
RETURNING id
|
||||||
|
""",
|
||||||
|
(utcnow(), task_id, preview_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return row is not None
|
||||||
|
|
||||||
|
def preview_is_running(self, task_id: str, preview_run_id: str) -> bool:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT preview_status, preview_run_id
|
||||||
|
FROM data_process_tasks
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
return bool(
|
||||||
|
row
|
||||||
|
and row.get("preview_status") in ACTIVE_PREVIEW_STATUSES
|
||||||
|
and row.get("preview_run_id") == preview_run_id
|
||||||
|
)
|
||||||
|
|
||||||
|
def update_preview_progress(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
preview_run_id: str,
|
||||||
|
completed_files: int,
|
||||||
|
total_files: int,
|
||||||
|
) -> bool:
|
||||||
|
total = max(1, total_files)
|
||||||
|
completed = min(max(0, completed_files), total)
|
||||||
|
progress = completed / total * 100
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET preview_progress=%s, preview_completed_files=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
AND preview_status='running' AND preview_run_id=%s
|
||||||
|
RETURNING id
|
||||||
|
""",
|
||||||
|
(progress, completed, utcnow(), task_id, preview_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return row is not None
|
||||||
|
|
||||||
|
def complete_preview(self, task_id: str, preview_run_id: str) -> bool:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET workflow_step='preview', preview_status='completed',
|
||||||
|
preview_progress=100, preview_run_id=NULL,
|
||||||
|
preview_failure_reason=NULL,
|
||||||
|
preview_completed_files=preview_total_files, updated_at=%s
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
AND preview_status='running' AND preview_run_id=%s
|
||||||
|
RETURNING id
|
||||||
|
""",
|
||||||
|
(utcnow(), task_id, preview_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return row is not None
|
||||||
|
|
||||||
|
def mark_preview_failed(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
reason: str,
|
||||||
|
*,
|
||||||
|
preview_run_id: str,
|
||||||
|
) -> bool:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET preview_status='failed', preview_run_id=NULL,
|
||||||
|
preview_failure_reason=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
AND preview_status IN ('queued', 'running') AND preview_run_id=%s
|
||||||
|
RETURNING id
|
||||||
|
""",
|
||||||
|
(reason[:4000], utcnow(), task_id, preview_run_id),
|
||||||
|
).fetchone()
|
||||||
|
return row is not None
|
||||||
|
|
||||||
|
def preview_progress(self, task_id: str) -> dict[str, Any]:
|
||||||
|
task = self.get_task(task_id)
|
||||||
|
return {
|
||||||
|
"task_id": task["id"],
|
||||||
|
"workflow_step": task.get("workflow_step") or "create",
|
||||||
|
"preview_status": task.get("preview_status") or "idle",
|
||||||
|
"preview_progress": float(task.get("preview_progress") or 0),
|
||||||
|
"preview_run_id": task.get("preview_run_id"),
|
||||||
|
"preview_failure_reason": task.get("preview_failure_reason"),
|
||||||
|
"preview_total_files": int(task.get("preview_total_files") or 0),
|
||||||
|
"preview_completed_files": int(task.get("preview_completed_files") or 0),
|
||||||
|
}
|
||||||
|
|
||||||
|
def list_preview_items(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
*,
|
||||||
|
source_file_id: str | None = None,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 200,
|
||||||
|
keyword: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
self.get_task(task_id)
|
||||||
|
clauses = ["task_id=%s"]
|
||||||
|
params: list[Any] = [task_id]
|
||||||
|
if source_file_id:
|
||||||
|
clauses.append("source_file_id=%s")
|
||||||
|
params.append(source_file_id)
|
||||||
|
if keyword:
|
||||||
|
clauses.append("(original_content ILIKE %s OR edited_content ILIKE %s)")
|
||||||
|
pattern = f"%{keyword.strip()}%"
|
||||||
|
params.extend([pattern, pattern])
|
||||||
|
where = " AND ".join(clauses)
|
||||||
|
with self.connect() as conn:
|
||||||
|
total = conn.execute(
|
||||||
|
f"SELECT COUNT(*) AS count FROM data_process_preview_items WHERE {where}", params
|
||||||
|
).fetchone()["count"]
|
||||||
|
rows = conn.execute(
|
||||||
|
f"""
|
||||||
|
SELECT * FROM data_process_preview_items
|
||||||
|
WHERE {where}
|
||||||
|
ORDER BY source_file_id NULLS LAST, source_start NULLS LAST, created_at, id
|
||||||
|
LIMIT %s OFFSET %s
|
||||||
|
""",
|
||||||
|
[*params, page_size, (page - 1) * page_size],
|
||||||
|
).fetchall()
|
||||||
|
return {
|
||||||
|
"items": [_decode_row(row) for row in rows],
|
||||||
|
"total": int(total),
|
||||||
|
"page": page,
|
||||||
|
"page_size": page_size,
|
||||||
|
}
|
||||||
|
|
||||||
|
def get_preview_item(self, task_id: str, preview_id: str) -> dict[str, Any]:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT * FROM data_process_preview_items WHERE id=%s AND task_id=%s",
|
||||||
|
(preview_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("preview item not found")
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def create_preview_item(self, task_id: str, item: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
now = utcnow()
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
if item.get("source_file_id"):
|
||||||
|
source = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id FROM data_process_source_files
|
||||||
|
WHERE id=%s AND task_id=%s AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(item["source_file_id"], task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not source:
|
||||||
|
raise NotFoundError("source file not found")
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_preview_items
|
||||||
|
(id, task_id, source_file_id, original_content, edited_content,
|
||||||
|
source_start, source_end, source_start_line, source_end_line,
|
||||||
|
token_count, status, quality_score, created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
new_id("dpp"),
|
||||||
|
task_id,
|
||||||
|
item.get("source_file_id"),
|
||||||
|
item.get("original_content") or "",
|
||||||
|
item.get("edited_content") or "",
|
||||||
|
item.get("source_start"),
|
||||||
|
item.get("source_end"),
|
||||||
|
item.get("source_start_line"),
|
||||||
|
item.get("source_end_line"),
|
||||||
|
max(0, int(item.get("token_count") or 0)),
|
||||||
|
item.get("status") or "manual",
|
||||||
|
json_dumps(item.get("quality_score") or {}),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
self._invalidate_results(conn, task, task_id, now)
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def update_preview_item(
|
||||||
|
self, task_id: str, preview_id: str, payload: dict[str, Any]
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
existing = conn.execute(
|
||||||
|
"SELECT * FROM data_process_preview_items WHERE id=%s AND task_id=%s",
|
||||||
|
(preview_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not existing:
|
||||||
|
raise NotFoundError("preview item not found")
|
||||||
|
expected_updated_at = payload.get("expected_updated_at")
|
||||||
|
current_updated_at = _serialize_value(existing.get("updated_at"))
|
||||||
|
if expected_updated_at and expected_updated_at != current_updated_at:
|
||||||
|
raise ConflictError("preview item was modified by another request")
|
||||||
|
edited = payload["edited_content"]
|
||||||
|
status = payload.get("status")
|
||||||
|
if not status:
|
||||||
|
if not edited.strip():
|
||||||
|
status = "invalid"
|
||||||
|
elif edited == existing["original_content"]:
|
||||||
|
status = "original"
|
||||||
|
else:
|
||||||
|
status = "modified"
|
||||||
|
now = utcnow()
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_preview_items
|
||||||
|
SET edited_content=%s, token_count=%s, status=%s, quality_score=%s,
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s AND task_id=%s
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
edited,
|
||||||
|
estimate_token_count(edited),
|
||||||
|
status,
|
||||||
|
json_dumps(payload.get("quality_score") or {}),
|
||||||
|
now,
|
||||||
|
preview_id,
|
||||||
|
task_id,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
self._invalidate_results(conn, task, task_id, now)
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def delete_preview_item(self, task_id: str, preview_id: str) -> None:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
row = conn.execute(
|
||||||
|
"DELETE FROM data_process_preview_items WHERE id=%s AND task_id=%s RETURNING id",
|
||||||
|
(preview_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("preview item not found")
|
||||||
|
self._invalidate_results(conn, task, task_id, utcnow())
|
||||||
324
backend/app/modules/data_process/store/results.py
Normal file
324
backend/app/modules/data_process/store/results.py
Normal file
@@ -0,0 +1,324 @@
|
|||||||
|
"""数据处理存储层 - 结果管理。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
_json_value,
|
||||||
|
_decode_row,
|
||||||
|
_public_task,
|
||||||
|
_business_config,
|
||||||
|
_preview_config_value,
|
||||||
|
_preview_config_changed,
|
||||||
|
_preview_config_projection,
|
||||||
|
_normalized_preprocess_options,
|
||||||
|
_regeneration_marker,
|
||||||
|
_is_regeneration_prepared,
|
||||||
|
_task_output_type,
|
||||||
|
_task_reasoning_detail,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_dpo_fields_are_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
_serialize_value,
|
||||||
|
DataProcessStoreError,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ResultsMixin:
|
||||||
|
"""结果管理 Mixin。"""
|
||||||
|
|
||||||
|
def confirm_results(self, task_id: str) -> dict[str, Any]:
|
||||||
|
"""确认第六步结果,确认前再次校验所有生成记录。"""
|
||||||
|
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if task["status"] != "completed":
|
||||||
|
raise InvalidStateError("only a completed task can confirm results")
|
||||||
|
if task.get("workflow_step") != "results":
|
||||||
|
raise InvalidStateError("workflow must be on results before confirmation")
|
||||||
|
if task.get("results_confirmed"):
|
||||||
|
return task
|
||||||
|
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT status, instruction, output, chosen, rejected
|
||||||
|
FROM data_process_results
|
||||||
|
WHERE task_id=%s
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchall()
|
||||||
|
if not rows:
|
||||||
|
raise InvalidStateError("task has no results to confirm")
|
||||||
|
invalid_count = sum(
|
||||||
|
1
|
||||||
|
for row in rows
|
||||||
|
if row["status"] == "invalid"
|
||||||
|
or not str(row.get("instruction") or "").strip()
|
||||||
|
or not str(row.get("output") or "").strip()
|
||||||
|
or (
|
||||||
|
_task_output_type(task) == "reasoning"
|
||||||
|
and not _reasoning_output_is_valid(row.get("output"))
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
_task_output_type(task) == "dpo"
|
||||||
|
and not _dpo_fields_are_valid(row)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if invalid_count:
|
||||||
|
raise InvalidStateError(
|
||||||
|
f"task contains {invalid_count} invalid results"
|
||||||
|
)
|
||||||
|
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET results_confirmed=TRUE, updated_at=%s
|
||||||
|
WHERE id=%s RETURNING *
|
||||||
|
""",
|
||||||
|
(utcnow(), task_id),
|
||||||
|
).fetchone()
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def list_results(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
*,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 100,
|
||||||
|
status: str | None = None,
|
||||||
|
split: str | None = None,
|
||||||
|
keyword: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
self.get_task(task_id)
|
||||||
|
clauses = ["task_id=%s"]
|
||||||
|
params: list[Any] = [task_id]
|
||||||
|
if status:
|
||||||
|
clauses.append("status=%s")
|
||||||
|
params.append(status)
|
||||||
|
if split:
|
||||||
|
clauses.append("split=%s")
|
||||||
|
params.append(split)
|
||||||
|
if keyword:
|
||||||
|
clauses.append("(instruction ILIKE %s OR input ILIKE %s OR output ILIKE %s)")
|
||||||
|
pattern = f"%{keyword.strip()}%"
|
||||||
|
params.extend([pattern, pattern, pattern])
|
||||||
|
where = " AND ".join(clauses)
|
||||||
|
with self.connect() as conn:
|
||||||
|
total = conn.execute(
|
||||||
|
f"SELECT COUNT(*) AS count FROM data_process_results WHERE {where}", params
|
||||||
|
).fetchone()["count"]
|
||||||
|
rows = conn.execute(
|
||||||
|
f"""
|
||||||
|
SELECT * FROM data_process_results WHERE {where}
|
||||||
|
ORDER BY created_at, id LIMIT %s OFFSET %s
|
||||||
|
""",
|
||||||
|
[*params, page_size, (page - 1) * page_size],
|
||||||
|
).fetchall()
|
||||||
|
return {
|
||||||
|
"items": [_decode_row(row) for row in rows],
|
||||||
|
"total": int(total),
|
||||||
|
"page": page,
|
||||||
|
"page_size": page_size,
|
||||||
|
}
|
||||||
|
|
||||||
|
def get_result(self, task_id: str, result_id: str) -> dict[str, Any]:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT * FROM data_process_results WHERE id=%s AND task_id=%s",
|
||||||
|
(result_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("data process result not found")
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def update_result(
|
||||||
|
self, task_id: str, result_id: str, payload: dict[str, Any]
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
allowed = {
|
||||||
|
"instruction", "input", "output", "chosen", "rejected", "quality_score"
|
||||||
|
}
|
||||||
|
values = {key: value for key, value in payload.items() if key in allowed}
|
||||||
|
if "quality_score" in values:
|
||||||
|
values["quality_score"] = json_dumps(values["quality_score"])
|
||||||
|
if not values:
|
||||||
|
raise DataProcessStoreError("no result fields supplied")
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if task["status"] == "running":
|
||||||
|
raise InvalidStateError("results cannot be edited while generation is running")
|
||||||
|
if task.get("output_dataset_id"):
|
||||||
|
raise InvalidStateError("published results cannot be edited")
|
||||||
|
current = conn.execute(
|
||||||
|
"SELECT * FROM data_process_results WHERE id=%s AND task_id=%s",
|
||||||
|
(result_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not current:
|
||||||
|
raise NotFoundError("data process result not found")
|
||||||
|
expected_updated_at = payload.get("expected_updated_at")
|
||||||
|
current_updated_at = _serialize_value(current.get("updated_at"))
|
||||||
|
if expected_updated_at and expected_updated_at != current_updated_at:
|
||||||
|
raise ConflictError("data process result was modified by another request")
|
||||||
|
output_type = _task_output_type(task)
|
||||||
|
if output_type == "dpo" and "chosen" in values:
|
||||||
|
values["output"] = values["chosen"]
|
||||||
|
merged = {**current, **values}
|
||||||
|
quality = payload.get("quality_score") or {}
|
||||||
|
instruction_valid = bool(str(merged.get("instruction") or "").strip())
|
||||||
|
output_valid = bool(str(merged.get("output") or "").strip())
|
||||||
|
reasoning_valid = (
|
||||||
|
output_type != "reasoning"
|
||||||
|
or _reasoning_output_is_valid(merged.get("output"))
|
||||||
|
)
|
||||||
|
dpo_valid = output_type != "dpo" or _dpo_fields_are_valid(merged)
|
||||||
|
hard_valid = instruction_valid and output_valid and reasoning_valid and dpo_valid
|
||||||
|
quality_valid = bool(quality.get("is_valid", hard_valid))
|
||||||
|
changed = any(
|
||||||
|
str(merged.get(field) or "")
|
||||||
|
!= str(merged.get(f"original_{field}") or "")
|
||||||
|
for field in (
|
||||||
|
("instruction", "input", "chosen", "rejected")
|
||||||
|
if output_type == "dpo"
|
||||||
|
else ("instruction", "input", "output")
|
||||||
|
)
|
||||||
|
)
|
||||||
|
status = "invalid" if not hard_valid or not quality_valid else (
|
||||||
|
"modified" if changed else "valid"
|
||||||
|
)
|
||||||
|
values["status"] = status
|
||||||
|
flags = quality.get("flags") if isinstance(quality, dict) else None
|
||||||
|
format_error = (
|
||||||
|
"思维链输出必须包含非空的 <think>...</think> 推理过程和最终答案"
|
||||||
|
if instruction_valid and output_valid and not reasoning_valid
|
||||||
|
else "DPO 输出必须包含不同的非空 Chosen 和 Rejected 回答"
|
||||||
|
if instruction_valid and not dpo_valid
|
||||||
|
else "Instruction 和 Output 不能为空"
|
||||||
|
if not instruction_valid or not output_valid
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
values["error"] = ", ".join(str(flag) for flag in flags or []) or (
|
||||||
|
format_error
|
||||||
|
or ("quality validation failed" if status == "invalid" else None)
|
||||||
|
)
|
||||||
|
values["updated_at"] = utcnow()
|
||||||
|
assignments = ", ".join(f"{key}=%s" for key in values)
|
||||||
|
row = conn.execute(
|
||||||
|
f"""UPDATE data_process_results SET {assignments}
|
||||||
|
WHERE id=%s AND task_id=%s RETURNING *""",
|
||||||
|
[*values.values(), result_id, task_id],
|
||||||
|
).fetchone()
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET error_count=(
|
||||||
|
SELECT COUNT(*) FROM data_process_results
|
||||||
|
WHERE task_id=%s AND status='invalid'
|
||||||
|
), updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(task_id, utcnow(), task_id),
|
||||||
|
)
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def replace_generated_result(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
result_id: str,
|
||||||
|
replacement: dict[str, Any],
|
||||||
|
*,
|
||||||
|
expected_updated_at: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""用新模型结果原位替换失败项,并将新内容设为恢复基线。"""
|
||||||
|
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if task["status"] != "completed" or task.get("workflow_step") != "results":
|
||||||
|
raise InvalidStateError("task is not editing generation results")
|
||||||
|
if task.get("results_confirmed"):
|
||||||
|
raise InvalidStateError("confirmed results cannot be regenerated")
|
||||||
|
if task.get("output_dataset_id"):
|
||||||
|
raise InvalidStateError("published results cannot be regenerated")
|
||||||
|
|
||||||
|
current = conn.execute(
|
||||||
|
"""SELECT * FROM data_process_results
|
||||||
|
WHERE id=%s AND task_id=%s FOR UPDATE""",
|
||||||
|
(result_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not current:
|
||||||
|
raise NotFoundError("data process result not found")
|
||||||
|
if current.get("status") != "invalid":
|
||||||
|
raise InvalidStateError("only an invalid result can be regenerated")
|
||||||
|
current_updated_at = _serialize_value(current.get("updated_at"))
|
||||||
|
if expected_updated_at != current_updated_at:
|
||||||
|
raise ConflictError("data process result was modified by another request")
|
||||||
|
|
||||||
|
instruction = str(replacement.get("instruction") or "").strip()
|
||||||
|
input_text = str(replacement.get("input") or "").strip()
|
||||||
|
output = str(replacement.get("output") or "").strip()
|
||||||
|
chosen = str(replacement.get("chosen") or "").strip()
|
||||||
|
rejected = str(replacement.get("rejected") or "").strip()
|
||||||
|
quality_score = replacement.get("quality_score") or {}
|
||||||
|
if not instruction or not output or not bool(quality_score.get("is_valid")):
|
||||||
|
raise InvalidStateError("regenerated result did not pass quality validation")
|
||||||
|
if _task_output_type(task) == "reasoning" and not _reasoning_output_is_valid(output):
|
||||||
|
raise InvalidStateError("regenerated reasoning result has an invalid output format")
|
||||||
|
if _task_output_type(task) == "dpo" and not _dpo_fields_are_valid(replacement):
|
||||||
|
raise InvalidStateError("regenerated DPO result has invalid preference fields")
|
||||||
|
|
||||||
|
now = utcnow()
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_results
|
||||||
|
SET instruction=%s, input=%s, output=%s, chosen=%s, rejected=%s,
|
||||||
|
original_instruction=%s, original_input=%s, original_output=%s,
|
||||||
|
original_chosen=%s, original_rejected=%s,
|
||||||
|
status='valid', error=NULL, quality_score=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND task_id=%s
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
instruction,
|
||||||
|
input_text,
|
||||||
|
output,
|
||||||
|
chosen,
|
||||||
|
rejected,
|
||||||
|
instruction,
|
||||||
|
input_text,
|
||||||
|
output,
|
||||||
|
chosen,
|
||||||
|
rejected,
|
||||||
|
json_dumps(quality_score),
|
||||||
|
now,
|
||||||
|
result_id,
|
||||||
|
task_id,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET error_count=(
|
||||||
|
SELECT COUNT(*) FROM data_process_results
|
||||||
|
WHERE task_id=%s AND status='invalid'
|
||||||
|
), updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(task_id, now, task_id),
|
||||||
|
)
|
||||||
|
return _decode_row(row) or {}
|
||||||
321
backend/app/modules/data_process/store/source_files.py
Normal file
321
backend/app/modules/data_process/store/source_files.py
Normal file
@@ -0,0 +1,321 @@
|
|||||||
|
"""数据处理存储层 - 源文件管理。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
_json_value,
|
||||||
|
_decode_row,
|
||||||
|
_public_task,
|
||||||
|
_business_config,
|
||||||
|
_preview_config_value,
|
||||||
|
_preview_config_changed,
|
||||||
|
_preview_config_projection,
|
||||||
|
_normalized_preprocess_options,
|
||||||
|
_regeneration_marker,
|
||||||
|
_is_regeneration_prepared,
|
||||||
|
_task_output_type,
|
||||||
|
_task_reasoning_detail,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_dpo_fields_are_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
DataProcessStoreError,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SourceFilesMixin:
|
||||||
|
"""源文件管理 Mixin。"""
|
||||||
|
|
||||||
|
def list_source_files(self, task_id: str) -> list[dict[str, Any]]:
|
||||||
|
self.get_task(task_id)
|
||||||
|
with self.connect() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id, task_id, storage_object_id, name, size_bytes, record_count,
|
||||||
|
file_format, checksum_sha256, version_no, content_preview, metadata,
|
||||||
|
tenant_id, project_id, created_by, created_at, updated_at
|
||||||
|
FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL
|
||||||
|
ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchall()
|
||||||
|
return [_decode_row(row) or {} for row in rows]
|
||||||
|
|
||||||
|
def add_source_file(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
*,
|
||||||
|
name: str,
|
||||||
|
content: str,
|
||||||
|
raw_size: int,
|
||||||
|
checksum_sha256: str,
|
||||||
|
file_format: str,
|
||||||
|
record_count: int,
|
||||||
|
metadata: dict[str, Any] | None = None,
|
||||||
|
created_by: str | None = None,
|
||||||
|
source_file_id: str | None = None,
|
||||||
|
storage_object_id: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
return self.add_source_files(
|
||||||
|
task_id,
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"name": name,
|
||||||
|
"content": content,
|
||||||
|
"raw_size": raw_size,
|
||||||
|
"checksum_sha256": checksum_sha256,
|
||||||
|
"file_format": file_format,
|
||||||
|
"record_count": record_count,
|
||||||
|
"metadata": metadata or {},
|
||||||
|
"created_by": created_by,
|
||||||
|
"id": source_file_id,
|
||||||
|
"storage_object_id": storage_object_id,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
def add_source_files(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
files: Sequence[dict[str, Any]],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""在同一事务中登记一个上传批次,任一文件失败则全部回滚。"""
|
||||||
|
|
||||||
|
if not files:
|
||||||
|
raise DataProcessStoreError("at least one source file is required")
|
||||||
|
now = utcnow()
|
||||||
|
created: list[dict[str, Any]] = []
|
||||||
|
try:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
for payload in files:
|
||||||
|
file_id = str(payload.get("id") or new_id("dpsf"))
|
||||||
|
storage_object_id, metadata_payload = _source_storage_descriptor(
|
||||||
|
payload,
|
||||||
|
task_id,
|
||||||
|
file_id,
|
||||||
|
)
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_source_files
|
||||||
|
(id, task_id, storage_object_id, name, size_bytes, record_count,
|
||||||
|
file_format, checksum_sha256, version_no, content, content_preview,
|
||||||
|
metadata, tenant_id, project_id,
|
||||||
|
created_by, created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, 1, %s, %s, %s, %s,
|
||||||
|
%s, %s, %s, %s)
|
||||||
|
RETURNING id, task_id, storage_object_id, name, size_bytes, record_count,
|
||||||
|
file_format, checksum_sha256, version_no, content_preview, metadata,
|
||||||
|
tenant_id, project_id, created_by, created_at, updated_at
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
file_id,
|
||||||
|
task_id,
|
||||||
|
storage_object_id,
|
||||||
|
payload["name"],
|
||||||
|
payload["raw_size"],
|
||||||
|
payload["record_count"],
|
||||||
|
payload["file_format"],
|
||||||
|
payload["checksum_sha256"],
|
||||||
|
payload["content"],
|
||||||
|
str(payload["content"])[:2000],
|
||||||
|
json_dumps(metadata_payload),
|
||||||
|
task.get("tenant_id"),
|
||||||
|
task.get("project_id"),
|
||||||
|
payload.get("created_by") or task.get("created_by"),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
created.append(_decode_row(row) or {})
|
||||||
|
preview_row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT COUNT(*) AS count FROM data_process_preview_items
|
||||||
|
WHERE task_id=%s
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
preview_count = int((preview_row or {}).get("count") or 0)
|
||||||
|
if _is_regeneration_prepared(task):
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET input_count=(
|
||||||
|
SELECT COALESCE(SUM(record_count), 0)
|
||||||
|
FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL
|
||||||
|
), workflow_step='upload', preview_status='idle',
|
||||||
|
preview_progress=0, preview_run_id=NULL,
|
||||||
|
preview_failure_reason=NULL, preview_total_files=0,
|
||||||
|
preview_completed_files=0, updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(task_id, now, task_id),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# 前端会只对本次新增的源文件构建预览,因此保留旧文件切片,
|
||||||
|
# 但普通未发布任务的旧生成结果已经不再有效。
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM data_process_results WHERE task_id=%s", (task_id,)
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='pending', progress=%s, output_count=0, filtered_count=0,
|
||||||
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
||||||
|
generation_run_id=NULL, results_confirmed=FALSE,
|
||||||
|
workflow_step='upload', preview_status='idle',
|
||||||
|
preview_progress=0, preview_run_id=NULL,
|
||||||
|
preview_failure_reason=NULL, preview_total_files=0,
|
||||||
|
preview_completed_files=0,
|
||||||
|
started_at=NULL, completed_at=NULL,
|
||||||
|
input_count=(
|
||||||
|
SELECT COALESCE(SUM(record_count), 0)
|
||||||
|
FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL
|
||||||
|
), updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(20 if preview_count else 0, task_id, now, task_id),
|
||||||
|
)
|
||||||
|
except psycopg.errors.UniqueViolation as exc:
|
||||||
|
raise ConflictError(
|
||||||
|
"the same source file content is already attached to this task"
|
||||||
|
) from exc
|
||||||
|
return created
|
||||||
|
|
||||||
|
def get_source_file(
|
||||||
|
self, task_id: str, file_id: str, *, include_content: bool = True
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
# 先验证父任务仍然可见,避免软删除任务后通过已知文件 ID 读取正文。
|
||||||
|
self.get_task(task_id)
|
||||||
|
content_column = ", content" if include_content else ""
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
f"""
|
||||||
|
SELECT id, task_id, storage_object_id, name, size_bytes, record_count,
|
||||||
|
file_format, checksum_sha256, version_no, content_preview, metadata,
|
||||||
|
tenant_id, project_id, created_by, created_at, updated_at{content_column}
|
||||||
|
FROM data_process_source_files
|
||||||
|
WHERE id=%s AND task_id=%s AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(file_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("source file not found")
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
def source_content_window(
|
||||||
|
self, task_id: str, file_id: str, offset: int, limit: int
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
source_file = self.get_source_file(task_id, file_id, include_content=True)
|
||||||
|
content = str(source_file.pop("content", ""))
|
||||||
|
window = content[offset : offset + limit]
|
||||||
|
return {
|
||||||
|
"file": source_file,
|
||||||
|
"content": window,
|
||||||
|
"offset": offset,
|
||||||
|
"limit": limit,
|
||||||
|
"total_chars": len(content),
|
||||||
|
"has_more": offset + len(window) < len(content),
|
||||||
|
}
|
||||||
|
|
||||||
|
def source_content_lines(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
file_id: str,
|
||||||
|
start_line: int,
|
||||||
|
line_count: int,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
source_file = self.get_source_file(task_id, file_id, include_content=True)
|
||||||
|
content = str(source_file.pop("content", ""))
|
||||||
|
lines = content.splitlines(keepends=True)
|
||||||
|
start_index = min(len(lines), start_line - 1)
|
||||||
|
selected = lines[start_index : start_index + line_count]
|
||||||
|
end_line = start_index + len(selected)
|
||||||
|
return {
|
||||||
|
"file": source_file,
|
||||||
|
"content": "".join(selected),
|
||||||
|
"start_line": start_line,
|
||||||
|
"end_line": end_line,
|
||||||
|
"line_count": len(selected),
|
||||||
|
"total_lines": len(lines),
|
||||||
|
"has_more": end_line < len(lines),
|
||||||
|
}
|
||||||
|
|
||||||
|
def delete_source_file(self, task_id: str, file_id: str) -> None:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_source_files
|
||||||
|
SET deleted_at=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND task_id=%s AND deleted_at IS NULL
|
||||||
|
RETURNING id
|
||||||
|
""",
|
||||||
|
(utcnow(), utcnow(), file_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("source file not found")
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM data_process_preview_items WHERE source_file_id=%s", (file_id,)
|
||||||
|
)
|
||||||
|
now = utcnow()
|
||||||
|
if _is_regeneration_prepared(task):
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET input_count=(SELECT COALESCE(SUM(record_count), 0)
|
||||||
|
FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL),
|
||||||
|
workflow_step='upload', preview_status='idle',
|
||||||
|
preview_progress=0, preview_run_id=NULL,
|
||||||
|
preview_failure_reason=NULL, preview_total_files=0,
|
||||||
|
preview_completed_files=0, updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(task_id, now, task_id),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM data_process_results WHERE task_id=%s", (task_id,)
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='pending', progress=0, output_count=0, filtered_count=0,
|
||||||
|
duplicate_count=0, error_count=0, failure_reason=NULL,
|
||||||
|
results_confirmed=FALSE,
|
||||||
|
workflow_step='upload', preview_status='idle',
|
||||||
|
preview_progress=0, preview_run_id=NULL,
|
||||||
|
preview_failure_reason=NULL, preview_total_files=0,
|
||||||
|
preview_completed_files=0,
|
||||||
|
input_count=(SELECT COALESCE(SUM(record_count), 0)
|
||||||
|
FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL),
|
||||||
|
updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(task_id, now, task_id),
|
||||||
|
)
|
||||||
871
backend/app/modules/data_process/store/tasks.py
Normal file
871
backend/app/modules/data_process/store/tasks.py
Normal file
@@ -0,0 +1,871 @@
|
|||||||
|
"""数据处理存储层 - 任务管理。"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import psycopg
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
StoreBase,
|
||||||
|
utcnow,
|
||||||
|
new_id,
|
||||||
|
repeat_task_id,
|
||||||
|
json_dumps,
|
||||||
|
_json_value,
|
||||||
|
_decode_row,
|
||||||
|
_public_task,
|
||||||
|
_business_config,
|
||||||
|
_preview_config_value,
|
||||||
|
_preview_config_changed,
|
||||||
|
_preview_config_projection,
|
||||||
|
_normalized_preprocess_options,
|
||||||
|
_regeneration_marker,
|
||||||
|
_is_regeneration_prepared,
|
||||||
|
_task_output_type,
|
||||||
|
_task_reasoning_detail,
|
||||||
|
_reasoning_output_is_valid,
|
||||||
|
_dpo_fields_are_valid,
|
||||||
|
_source_storage_descriptor,
|
||||||
|
_serialize_value,
|
||||||
|
NotFoundError,
|
||||||
|
ConflictError,
|
||||||
|
InvalidStateError,
|
||||||
|
EDITABLE_STATUSES,
|
||||||
|
ACTIVE_PREVIEW_STATUSES,
|
||||||
|
WORKFLOW_STEPS,
|
||||||
|
_REGENERATION_MARKER_KEY,
|
||||||
|
_REPEAT_SOURCE_TASK_KEY,
|
||||||
|
_REPEAT_REQUEST_KEY,
|
||||||
|
_INTERNAL_CONFIG_KEYS,
|
||||||
|
)
|
||||||
|
|
||||||
|
from ..algorithms import estimate_token_count
|
||||||
|
|
||||||
|
class TasksMixin:
|
||||||
|
"""任务管理 Mixin。"""
|
||||||
|
|
||||||
|
def list_tasks(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
page: int = 1,
|
||||||
|
page_size: int = 20,
|
||||||
|
keyword: str | None = None,
|
||||||
|
status: str | None = None,
|
||||||
|
process_type: str | None = None,
|
||||||
|
tenant_id: str | None = None,
|
||||||
|
project_id: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
clauses = ["task.deleted_at IS NULL"]
|
||||||
|
params: list[Any] = []
|
||||||
|
if keyword:
|
||||||
|
clauses.append("(task.name ILIKE %s OR COALESCE(task.description, '') ILIKE %s)")
|
||||||
|
pattern = f"%{keyword.strip()}%"
|
||||||
|
params.extend([pattern, pattern])
|
||||||
|
if status:
|
||||||
|
clauses.append("task.status = %s")
|
||||||
|
params.append(status)
|
||||||
|
if process_type:
|
||||||
|
clauses.append("task.process_type = %s")
|
||||||
|
params.append(process_type)
|
||||||
|
if tenant_id:
|
||||||
|
clauses.append("task.tenant_id = %s")
|
||||||
|
params.append(tenant_id)
|
||||||
|
if project_id:
|
||||||
|
clauses.append("task.project_id = %s")
|
||||||
|
params.append(project_id)
|
||||||
|
where = " AND ".join(clauses)
|
||||||
|
with self.connect() as conn:
|
||||||
|
total = conn.execute(
|
||||||
|
f"SELECT COUNT(*) AS count FROM data_process_tasks task WHERE {where}",
|
||||||
|
params,
|
||||||
|
).fetchone()["count"]
|
||||||
|
rows = conn.execute(
|
||||||
|
f"""
|
||||||
|
SELECT task.*,
|
||||||
|
(SELECT COUNT(*) FROM data_process_source_files source_file
|
||||||
|
WHERE source_file.task_id=task.id
|
||||||
|
AND source_file.deleted_at IS NULL) AS source_file_count
|
||||||
|
FROM data_process_tasks task
|
||||||
|
WHERE {where}
|
||||||
|
ORDER BY task.created_at DESC, task.id DESC
|
||||||
|
LIMIT %s OFFSET %s
|
||||||
|
""",
|
||||||
|
[*params, page_size, (page - 1) * page_size],
|
||||||
|
).fetchall()
|
||||||
|
return {
|
||||||
|
"items": [_public_task(_decode_row(row)) for row in rows],
|
||||||
|
"total": int(total),
|
||||||
|
"page": page,
|
||||||
|
"page_size": page_size,
|
||||||
|
}
|
||||||
|
|
||||||
|
def create_task(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
task_id = new_id("dpt")
|
||||||
|
now = utcnow()
|
||||||
|
try:
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_tasks
|
||||||
|
(id, name, description, status, process_type, source_dataset_id, config,
|
||||||
|
progress, results_confirmed, tenant_id, project_id, owner_id, created_by, updated_by,
|
||||||
|
created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, 'pending', %s, %s, %s, 0, FALSE,
|
||||||
|
%s, %s, %s, %s, %s, %s, %s)
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
task_id,
|
||||||
|
payload["name"],
|
||||||
|
payload.get("description") or "",
|
||||||
|
payload["process_type"],
|
||||||
|
payload.get("source_dataset_id"),
|
||||||
|
json_dumps(_business_config(payload.get("config"))),
|
||||||
|
payload.get("tenant_id"),
|
||||||
|
payload.get("project_id"),
|
||||||
|
payload.get("owner_id"),
|
||||||
|
payload.get("created_by"),
|
||||||
|
payload.get("created_by"),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
except psycopg.errors.UniqueViolation as exc:
|
||||||
|
raise ConflictError("data process task name already exists") from exc
|
||||||
|
return _public_task(_decode_row(row)) or {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _repeat_response(
|
||||||
|
conn: psycopg.Connection[dict[str, Any]],
|
||||||
|
row: dict[str, Any],
|
||||||
|
*,
|
||||||
|
source_task_id: str,
|
||||||
|
created: bool,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
task_id = str(row["id"])
|
||||||
|
counts = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT
|
||||||
|
(SELECT COUNT(*) FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL) AS source_file_count,
|
||||||
|
(SELECT COUNT(*) FROM data_process_preview_items
|
||||||
|
WHERE task_id=%s) AS preview_count
|
||||||
|
""",
|
||||||
|
(task_id, task_id),
|
||||||
|
).fetchone() or {}
|
||||||
|
task = _public_task(_decode_row(row)) or {}
|
||||||
|
task["source_file_count"] = int(counts.get("source_file_count") or 0)
|
||||||
|
task["preview_count"] = int(counts.get("preview_count") or 0)
|
||||||
|
return {
|
||||||
|
"task": task,
|
||||||
|
"source_task_id": source_task_id,
|
||||||
|
"created": created,
|
||||||
|
"copied_source_file_count": task["source_file_count"],
|
||||||
|
"copied_preview_count": task["preview_count"],
|
||||||
|
}
|
||||||
|
|
||||||
|
def find_repeated_task(
|
||||||
|
self,
|
||||||
|
source_task_id: str,
|
||||||
|
request_id: str,
|
||||||
|
) -> dict[str, Any] | None:
|
||||||
|
"""查找同一幂等请求已创建的新任务。"""
|
||||||
|
|
||||||
|
task_id = repeat_task_id(source_task_id, request_id)
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"SELECT * FROM data_process_tasks WHERE id=%s",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
if row is None:
|
||||||
|
return None
|
||||||
|
decoded = _decode_row(row) or {}
|
||||||
|
config = decoded.get("config") or {}
|
||||||
|
if (
|
||||||
|
config.get(_REPEAT_SOURCE_TASK_KEY) != source_task_id
|
||||||
|
or config.get(_REPEAT_REQUEST_KEY) != request_id
|
||||||
|
):
|
||||||
|
raise ConflictError("再次生成请求与现有任务冲突")
|
||||||
|
if decoded.get("deleted_at"):
|
||||||
|
raise ConflictError("此次再次生成创建的任务已被删除,请重新发起")
|
||||||
|
return self._repeat_response(
|
||||||
|
conn,
|
||||||
|
row,
|
||||||
|
source_task_id=source_task_id,
|
||||||
|
created=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def repeat_task(
|
||||||
|
self,
|
||||||
|
source_task_id: str,
|
||||||
|
*,
|
||||||
|
expected_updated_at: str,
|
||||||
|
request_id: str,
|
||||||
|
file_copies: dict[str, dict[str, str]],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""复制已确认任务的配置、源文件和预览,结果与发布数据保持独立。"""
|
||||||
|
|
||||||
|
task_id = repeat_task_id(source_task_id, request_id)
|
||||||
|
now = utcnow()
|
||||||
|
try:
|
||||||
|
with self.connect() as conn:
|
||||||
|
existing = conn.execute(
|
||||||
|
"SELECT * FROM data_process_tasks WHERE id=%s FOR UPDATE",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
if existing is not None:
|
||||||
|
decoded = _decode_row(existing) or {}
|
||||||
|
config = decoded.get("config") or {}
|
||||||
|
if (
|
||||||
|
config.get(_REPEAT_SOURCE_TASK_KEY) != source_task_id
|
||||||
|
or config.get(_REPEAT_REQUEST_KEY) != request_id
|
||||||
|
):
|
||||||
|
raise ConflictError("再次生成请求与现有任务冲突")
|
||||||
|
if decoded.get("deleted_at"):
|
||||||
|
raise ConflictError("此次再次生成创建的任务已被删除,请重新发起")
|
||||||
|
return self._repeat_response(
|
||||||
|
conn,
|
||||||
|
existing,
|
||||||
|
source_task_id=source_task_id,
|
||||||
|
created=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
source_task = self._task_in_connection(
|
||||||
|
conn,
|
||||||
|
source_task_id,
|
||||||
|
for_update=True,
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
source_task.get("status") != "completed"
|
||||||
|
or source_task.get("results_confirmed") is False
|
||||||
|
):
|
||||||
|
raise InvalidStateError("只有已完成并确认结果的任务可以再次生成")
|
||||||
|
if source_task.get("preview_status") in ACTIVE_PREVIEW_STATUSES:
|
||||||
|
raise ConflictError("源任务仍在处理切分,暂时不能再次生成")
|
||||||
|
if expected_updated_at != _serialize_value(source_task.get("updated_at")):
|
||||||
|
raise ConflictError("源任务已被其他操作修改,请刷新后重试")
|
||||||
|
|
||||||
|
source_files = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT * FROM data_process_source_files
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL
|
||||||
|
ORDER BY created_at, id
|
||||||
|
""",
|
||||||
|
(source_task_id,),
|
||||||
|
).fetchall()
|
||||||
|
source_file_ids = {str(row["id"]) for row in source_files}
|
||||||
|
if source_file_ids != set(file_copies):
|
||||||
|
raise ConflictError("源文件快照已变化,请刷新后重试")
|
||||||
|
previews = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT * FROM data_process_preview_items
|
||||||
|
WHERE task_id=%s
|
||||||
|
ORDER BY source_file_id NULLS LAST, source_start NULLS LAST,
|
||||||
|
created_at, id
|
||||||
|
""",
|
||||||
|
(source_task_id,),
|
||||||
|
).fetchall()
|
||||||
|
if not previews:
|
||||||
|
raise InvalidStateError("源任务没有可用于再次生成的切分结果")
|
||||||
|
|
||||||
|
suffix = f"(再次生成-{task_id[-6:]})"
|
||||||
|
base_name = str(source_task.get("name") or "数据处理任务")
|
||||||
|
repeated_name = f"{base_name[: max(1, 150 - len(suffix))]}{suffix}"
|
||||||
|
repeated_config = _business_config(source_task.get("config") or {})
|
||||||
|
repeated_config[_REPEAT_SOURCE_TASK_KEY] = source_task_id
|
||||||
|
repeated_config[_REPEAT_REQUEST_KEY] = request_id
|
||||||
|
input_count = sum(int(row.get("record_count") or 0) for row in source_files)
|
||||||
|
task_row = conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_tasks
|
||||||
|
(id, name, description, status, process_type, source_dataset_id,
|
||||||
|
output_dataset_id, config, progress, input_count, output_count,
|
||||||
|
filtered_count, duplicate_count, error_count, failure_reason,
|
||||||
|
generation_run_id, results_confirmed, workflow_step,
|
||||||
|
preview_status, preview_progress, preview_run_id,
|
||||||
|
preview_failure_reason, preview_total_files,
|
||||||
|
preview_completed_files, tenant_id, project_id, owner_id,
|
||||||
|
approval_status, created_by, updated_by, created_at, updated_at)
|
||||||
|
VALUES
|
||||||
|
(%s, %s, %s, 'pending', %s, %s, NULL, %s, 20, %s, 0,
|
||||||
|
0, 0, 0, NULL, NULL, FALSE, 'preview', 'completed', 100,
|
||||||
|
NULL, NULL, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
task_id,
|
||||||
|
repeated_name,
|
||||||
|
source_task.get("description") or "",
|
||||||
|
source_task["process_type"],
|
||||||
|
source_task.get("source_dataset_id"),
|
||||||
|
json_dumps(repeated_config),
|
||||||
|
input_count,
|
||||||
|
len(source_files),
|
||||||
|
len(source_files),
|
||||||
|
source_task.get("tenant_id"),
|
||||||
|
source_task.get("project_id"),
|
||||||
|
source_task.get("owner_id"),
|
||||||
|
source_task.get("approval_status") or "not_required",
|
||||||
|
source_task.get("created_by"),
|
||||||
|
source_task.get("created_by"),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
file_id_map: dict[str, str] = {}
|
||||||
|
for source in source_files:
|
||||||
|
old_file_id = str(source["id"])
|
||||||
|
copy = file_copies[old_file_id]
|
||||||
|
new_file_id = str(copy["id"])
|
||||||
|
storage_object_id, metadata = _source_storage_descriptor(
|
||||||
|
{
|
||||||
|
"storage_object_id": copy["storage_object_id"],
|
||||||
|
"metadata": _json_value(source.get("metadata"), {}),
|
||||||
|
},
|
||||||
|
task_id,
|
||||||
|
new_file_id,
|
||||||
|
)
|
||||||
|
file_id_map[old_file_id] = new_file_id
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_source_files
|
||||||
|
(id, task_id, storage_object_id, name, size_bytes, record_count,
|
||||||
|
file_format, checksum_sha256, version_no, content,
|
||||||
|
content_preview, metadata, tenant_id, project_id, created_by,
|
||||||
|
created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, 1, %s, %s, %s,
|
||||||
|
%s, %s, %s, %s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
new_file_id,
|
||||||
|
task_id,
|
||||||
|
storage_object_id,
|
||||||
|
source["name"],
|
||||||
|
source.get("size_bytes") or 0,
|
||||||
|
source.get("record_count") or 0,
|
||||||
|
source.get("file_format"),
|
||||||
|
source["checksum_sha256"],
|
||||||
|
source.get("content") or "",
|
||||||
|
source.get("content_preview"),
|
||||||
|
json_dumps(metadata),
|
||||||
|
source_task.get("tenant_id"),
|
||||||
|
source_task.get("project_id"),
|
||||||
|
source.get("created_by") or source_task.get("created_by"),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
for preview in previews:
|
||||||
|
old_source_file_id = preview.get("source_file_id")
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_preview_items
|
||||||
|
(id, task_id, source_file_id, original_content, edited_content,
|
||||||
|
source_start, source_end, source_start_line, source_end_line,
|
||||||
|
token_count, status, quality_score, created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
||||||
|
%s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
new_id("dpp"),
|
||||||
|
task_id,
|
||||||
|
file_id_map.get(str(old_source_file_id))
|
||||||
|
if old_source_file_id
|
||||||
|
else None,
|
||||||
|
preview.get("original_content") or "",
|
||||||
|
preview.get("edited_content") or "",
|
||||||
|
preview.get("source_start"),
|
||||||
|
preview.get("source_end"),
|
||||||
|
preview.get("source_start_line"),
|
||||||
|
preview.get("source_end_line"),
|
||||||
|
max(0, int(preview.get("token_count") or 0)),
|
||||||
|
preview.get("status") or "original",
|
||||||
|
json_dumps(_json_value(preview.get("quality_score"), {})),
|
||||||
|
now,
|
||||||
|
now,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return self._repeat_response(
|
||||||
|
conn,
|
||||||
|
task_row or {},
|
||||||
|
source_task_id=source_task_id,
|
||||||
|
created=True,
|
||||||
|
)
|
||||||
|
except psycopg.errors.UniqueViolation as exc:
|
||||||
|
raise ConflictError("再次生成任务名称或请求发生冲突,请重试") from exc
|
||||||
|
|
||||||
|
def get_task(self, task_id: str, *, for_update: bool = False) -> dict[str, Any]:
|
||||||
|
lock = " FOR UPDATE" if for_update else ""
|
||||||
|
with self.connect() as conn:
|
||||||
|
if for_update:
|
||||||
|
row = conn.execute(
|
||||||
|
f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
else:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT task.*,
|
||||||
|
(SELECT COUNT(*) FROM data_process_source_files source
|
||||||
|
WHERE source.task_id=task.id AND source.deleted_at IS NULL)
|
||||||
|
AS source_file_count,
|
||||||
|
(SELECT COUNT(*) FROM data_process_preview_items preview
|
||||||
|
WHERE preview.task_id=task.id) AS preview_count,
|
||||||
|
(SELECT COALESCE(json_agg(json_build_object(
|
||||||
|
'id', dataset.id,
|
||||||
|
'name', dataset.name,
|
||||||
|
'type', dataset.type,
|
||||||
|
'count', dataset.count,
|
||||||
|
'dataset_split', CASE dataset.type
|
||||||
|
WHEN 'train' THEN 'train'
|
||||||
|
WHEN 'val' THEN 'validation'
|
||||||
|
WHEN 'test' THEN 'test'
|
||||||
|
ELSE NULL
|
||||||
|
END
|
||||||
|
) ORDER BY CASE dataset.type
|
||||||
|
WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4 END), '[]'::json)
|
||||||
|
FROM datasets dataset
|
||||||
|
WHERE dataset.source='task'
|
||||||
|
AND dataset.deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
dataset.source_task_id=task.id
|
||||||
|
OR (dataset.source_task_id IS NULL AND dataset.task_id=task.id)
|
||||||
|
))
|
||||||
|
AS output_datasets,
|
||||||
|
CASE
|
||||||
|
WHEN task.started_at IS NOT NULL AND task.completed_at IS NOT NULL
|
||||||
|
THEN EXTRACT(EPOCH FROM (task.completed_at - task.started_at))
|
||||||
|
ELSE NULL
|
||||||
|
END AS duration_seconds
|
||||||
|
FROM data_process_tasks task
|
||||||
|
WHERE task.id=%s AND task.deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("data process task not found")
|
||||||
|
return _public_task(_decode_row(row)) or {}
|
||||||
|
|
||||||
|
def _task_in_connection(
|
||||||
|
self,
|
||||||
|
conn: psycopg.Connection[dict[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
*,
|
||||||
|
for_update: bool = False,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
lock = " FOR UPDATE" if for_update else ""
|
||||||
|
row = conn.execute(
|
||||||
|
f"SELECT * FROM data_process_tasks WHERE id=%s AND deleted_at IS NULL{lock}",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("data process task not found")
|
||||||
|
return _decode_row(row) or {}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _ensure_editable(task: dict[str, Any]) -> None:
|
||||||
|
if task["status"] not in EDITABLE_STATUSES:
|
||||||
|
raise InvalidStateError(f"task cannot be edited while status is {task['status']}")
|
||||||
|
if task.get("preview_status") in ACTIVE_PREVIEW_STATUSES:
|
||||||
|
raise InvalidStateError("task cannot be edited while preview is running")
|
||||||
|
if task.get("output_dataset_id") and not _is_regeneration_prepared(task):
|
||||||
|
raise InvalidStateError("published task cannot be edited")
|
||||||
|
|
||||||
|
def update_workflow_step(self, task_id: str, workflow_step: str) -> dict[str, Any]:
|
||||||
|
"""独立保存向导位置,不触发配置或结果失效逻辑。"""
|
||||||
|
|
||||||
|
if workflow_step not in WORKFLOW_STEPS:
|
||||||
|
raise ValueError("invalid data process workflow step")
|
||||||
|
with self.connect() as conn:
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET workflow_step=%s, updated_at=%s
|
||||||
|
WHERE id=%s AND deleted_at IS NULL
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(workflow_step, utcnow(), task_id),
|
||||||
|
).fetchone()
|
||||||
|
if not row:
|
||||||
|
raise NotFoundError("data process task not found")
|
||||||
|
return _public_task(_decode_row(row)) or {}
|
||||||
|
|
||||||
|
def recover_legacy_aborted_regeneration(self, task_id: str) -> dict[str, Any]:
|
||||||
|
"""恢复旧版在真正开始生成前误删的上一轮结果。
|
||||||
|
|
||||||
|
旧实现会在 ``POST /regenerate`` 时立即把已发布任务置为 pending、
|
||||||
|
清空结果并解除输出指针。三个已发布数据集仍是独立完整产物,因此只在
|
||||||
|
这个特征完全匹配时,使用其记录恢复结果和任务状态。该操作幂等,不会
|
||||||
|
触碰正常的新建待生成任务或已经开始的新一轮生成。
|
||||||
|
"""
|
||||||
|
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if (
|
||||||
|
task.get("status") != "pending"
|
||||||
|
or task.get("generation_run_id")
|
||||||
|
or task.get("output_dataset_id")
|
||||||
|
or int(task.get("output_count") or 0) != 0
|
||||||
|
):
|
||||||
|
return {"recovered": False, "result_count": 0}
|
||||||
|
|
||||||
|
result_count = int(
|
||||||
|
(
|
||||||
|
conn.execute(
|
||||||
|
"SELECT COUNT(*) AS count FROM data_process_results WHERE task_id=%s",
|
||||||
|
(task_id,),
|
||||||
|
).fetchone()
|
||||||
|
or {}
|
||||||
|
).get("count")
|
||||||
|
or 0
|
||||||
|
)
|
||||||
|
if result_count:
|
||||||
|
return {"recovered": False, "result_count": result_count}
|
||||||
|
|
||||||
|
datasets = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id, type, count, created_at
|
||||||
|
FROM datasets
|
||||||
|
WHERE source='task' AND deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
source_task_id=%s
|
||||||
|
OR (source_task_id IS NULL AND task_id=%s)
|
||||||
|
)
|
||||||
|
ORDER BY CASE type
|
||||||
|
WHEN 'train' THEN 1 WHEN 'val' THEN 2 WHEN 'test' THEN 3 ELSE 4
|
||||||
|
END, created_at, id
|
||||||
|
""",
|
||||||
|
(task_id, task_id),
|
||||||
|
).fetchall()
|
||||||
|
train_dataset = next(
|
||||||
|
(dataset for dataset in datasets if dataset.get("type") == "train"),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not train_dataset:
|
||||||
|
return {"recovered": False, "result_count": 0}
|
||||||
|
|
||||||
|
dataset_ids = [str(dataset["id"]) for dataset in datasets]
|
||||||
|
records = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT id, dataset_id, line_no, split, instruction, input, output,
|
||||||
|
raw, status, source_result_id, preview_item_id, created_at
|
||||||
|
FROM dataset_records
|
||||||
|
WHERE dataset_id = ANY(%s)
|
||||||
|
ORDER BY created_at, dataset_id, line_no NULLS LAST, id
|
||||||
|
""",
|
||||||
|
(dataset_ids,),
|
||||||
|
).fetchall()
|
||||||
|
if not records:
|
||||||
|
return {"recovered": False, "result_count": 0}
|
||||||
|
|
||||||
|
preview_rows = conn.execute(
|
||||||
|
"SELECT id FROM data_process_preview_items WHERE task_id=%s",
|
||||||
|
(task_id,),
|
||||||
|
).fetchall()
|
||||||
|
preview_ids = {str(row["id"]) for row in preview_rows}
|
||||||
|
used_result_ids: set[str] = set()
|
||||||
|
recovered_count = 0
|
||||||
|
for record in records:
|
||||||
|
raw = _json_value(record.get("raw"), {})
|
||||||
|
raw = raw if isinstance(raw, dict) else {}
|
||||||
|
candidate_id = str(
|
||||||
|
record.get("source_result_id")
|
||||||
|
or raw.get("source_result_id")
|
||||||
|
or ""
|
||||||
|
)
|
||||||
|
result_id = (
|
||||||
|
candidate_id
|
||||||
|
if candidate_id and candidate_id not in used_result_ids
|
||||||
|
else new_id("dpr")
|
||||||
|
)
|
||||||
|
used_result_ids.add(result_id)
|
||||||
|
candidate_preview_id = str(
|
||||||
|
record.get("preview_item_id")
|
||||||
|
or raw.get("preview_item_id")
|
||||||
|
or ""
|
||||||
|
)
|
||||||
|
preview_item_id = (
|
||||||
|
candidate_preview_id if candidate_preview_id in preview_ids else None
|
||||||
|
)
|
||||||
|
instruction = str(record.get("instruction") or raw.get("instruction") or "")
|
||||||
|
input_text = str(record.get("input") or raw.get("input") or "")
|
||||||
|
chosen = str(raw.get("chosen") or "")
|
||||||
|
rejected = str(raw.get("rejected") or "")
|
||||||
|
output = str(
|
||||||
|
record.get("output") or raw.get("output") or chosen or ""
|
||||||
|
)
|
||||||
|
split = str(record.get("split") or raw.get("split") or "") or None
|
||||||
|
status = str(record.get("status") or "valid")
|
||||||
|
if status not in {"valid", "modified", "invalid"}:
|
||||||
|
status = "valid"
|
||||||
|
created_at = record.get("created_at") or utcnow()
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
INSERT INTO data_process_results
|
||||||
|
(id, task_id, preview_item_id, instruction, input, output,
|
||||||
|
chosen, rejected, original_instruction, original_input,
|
||||||
|
original_output, original_chosen, original_rejected, status,
|
||||||
|
error, split, quality_score, created_at, updated_at)
|
||||||
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
||||||
|
%s, %s, NULL, %s, '{}', %s, %s)
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
result_id,
|
||||||
|
task_id,
|
||||||
|
preview_item_id,
|
||||||
|
instruction,
|
||||||
|
input_text,
|
||||||
|
output,
|
||||||
|
chosen,
|
||||||
|
rejected,
|
||||||
|
instruction,
|
||||||
|
input_text,
|
||||||
|
output,
|
||||||
|
chosen,
|
||||||
|
rejected,
|
||||||
|
status,
|
||||||
|
split,
|
||||||
|
created_at,
|
||||||
|
created_at,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE dataset_records
|
||||||
|
SET source_result_id=%s, preview_item_id=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(result_id, preview_item_id, record["id"]),
|
||||||
|
)
|
||||||
|
recovered_count += 1
|
||||||
|
|
||||||
|
now = utcnow()
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status='completed', progress=100, output_dataset_id=%s,
|
||||||
|
output_count=%s, filtered_count=0, duplicate_count=0,
|
||||||
|
error_count=(SELECT COUNT(*) FROM data_process_results
|
||||||
|
WHERE task_id=%s AND status='invalid'),
|
||||||
|
failure_reason=NULL, results_confirmed=TRUE,
|
||||||
|
workflow_step='results', updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
train_dataset["id"],
|
||||||
|
recovered_count,
|
||||||
|
task_id,
|
||||||
|
now,
|
||||||
|
task_id,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return {"recovered": True, "result_count": recovered_count}
|
||||||
|
|
||||||
|
def update_task(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
allowed = {
|
||||||
|
"name",
|
||||||
|
"description",
|
||||||
|
"process_type",
|
||||||
|
"source_dataset_id",
|
||||||
|
}
|
||||||
|
values: dict[str, Any] = {key: value for key, value in payload.items() if key in allowed}
|
||||||
|
if not values and payload.get("config") is None:
|
||||||
|
return self.get_task(task_id)
|
||||||
|
try:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
self._ensure_editable(task)
|
||||||
|
regeneration_prepared = _is_regeneration_prepared(task)
|
||||||
|
if regeneration_prepared and any(
|
||||||
|
key in payload and payload.get(key) != task.get(key)
|
||||||
|
for key in ("process_type", "source_dataset_id")
|
||||||
|
):
|
||||||
|
raise InvalidStateError(
|
||||||
|
"process type and source dataset cannot change during regeneration"
|
||||||
|
)
|
||||||
|
if payload.get("config") is not None:
|
||||||
|
next_config = _business_config(payload["config"])
|
||||||
|
current_config = dict(task.get("config") or {})
|
||||||
|
for key in _INTERNAL_CONFIG_KEYS:
|
||||||
|
if key in current_config:
|
||||||
|
next_config[key] = current_config[key]
|
||||||
|
values["config"] = json_dumps(next_config)
|
||||||
|
invalidates_results = (
|
||||||
|
("config" in payload and payload.get("config") != task.get("config"))
|
||||||
|
or (
|
||||||
|
"process_type" in payload
|
||||||
|
and payload.get("process_type") != task.get("process_type")
|
||||||
|
)
|
||||||
|
or (
|
||||||
|
"source_dataset_id" in payload
|
||||||
|
and payload.get("source_dataset_id") != task.get("source_dataset_id")
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if invalidates_results and not regeneration_prepared:
|
||||||
|
values.update(
|
||||||
|
{
|
||||||
|
"status": "pending",
|
||||||
|
"progress": 0,
|
||||||
|
"output_count": 0,
|
||||||
|
"filtered_count": 0,
|
||||||
|
"duplicate_count": 0,
|
||||||
|
"error_count": 0,
|
||||||
|
"failure_reason": None,
|
||||||
|
"generation_run_id": None,
|
||||||
|
"results_confirmed": False,
|
||||||
|
"preview_status": "idle",
|
||||||
|
"preview_progress": 0,
|
||||||
|
"preview_run_id": None,
|
||||||
|
"preview_failure_reason": None,
|
||||||
|
"preview_total_files": 0,
|
||||||
|
"preview_completed_files": 0,
|
||||||
|
"started_at": None,
|
||||||
|
"completed_at": None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
conn.execute("DELETE FROM data_process_results WHERE task_id=%s", (task_id,))
|
||||||
|
conn.execute(
|
||||||
|
"DELETE FROM data_process_preview_items WHERE task_id=%s", (task_id,)
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
"process_type" in payload
|
||||||
|
and payload.get("process_type") != task.get("process_type")
|
||||||
|
):
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_source_files
|
||||||
|
SET deleted_at=%s, updated_at=%s
|
||||||
|
WHERE task_id=%s AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(utcnow(), utcnow(), task_id),
|
||||||
|
)
|
||||||
|
values["input_count"] = 0
|
||||||
|
values["updated_at"] = utcnow()
|
||||||
|
assignments = ", ".join(f"{key}=%s" for key in values)
|
||||||
|
row = conn.execute(
|
||||||
|
f"UPDATE data_process_tasks SET {assignments} WHERE id=%s RETURNING *",
|
||||||
|
[*values.values(), task_id],
|
||||||
|
).fetchone()
|
||||||
|
except psycopg.errors.UniqueViolation as exc:
|
||||||
|
raise ConflictError("data process task name already exists") from exc
|
||||||
|
return _public_task(_decode_row(row)) or {}
|
||||||
|
|
||||||
|
def prepare_regeneration(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
payload: dict[str, Any],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""非破坏性地保存重新生成配置。
|
||||||
|
|
||||||
|
准备阶段保留任务当前状态、结果、切片及已发布数据集。真正开始
|
||||||
|
生成时,才在同一事务内切换运行状态并清理上一轮结果。
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 先修复曾被旧版 prepare 提前清空的任务,再建立新的非破坏性草稿标记。
|
||||||
|
self.recover_legacy_aborted_regeneration(task_id)
|
||||||
|
try:
|
||||||
|
with self.connect() as conn:
|
||||||
|
task = self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
if task["status"] == "running":
|
||||||
|
raise ConflictError("running task cannot be prepared for regeneration")
|
||||||
|
if task.get("preview_status") in ACTIVE_PREVIEW_STATUSES:
|
||||||
|
raise ConflictError("running preview cannot be prepared for regeneration")
|
||||||
|
|
||||||
|
current_updated_at = _serialize_value(task.get("updated_at"))
|
||||||
|
if payload["expected_updated_at"] != current_updated_at:
|
||||||
|
raise ConflictError("data process task was modified by another request")
|
||||||
|
|
||||||
|
process_type = str(payload["process_type"])
|
||||||
|
if process_type != str(task["process_type"]):
|
||||||
|
raise InvalidStateError("process_type cannot be changed during regeneration")
|
||||||
|
|
||||||
|
current_config = dict(task.get("config") or {})
|
||||||
|
next_config = _business_config(payload.get("config"))
|
||||||
|
for key in (_REPEAT_SOURCE_TASK_KEY, _REPEAT_REQUEST_KEY):
|
||||||
|
if key in current_config:
|
||||||
|
next_config[key] = current_config[key]
|
||||||
|
preview_invalidated = _preview_config_changed(
|
||||||
|
process_type,
|
||||||
|
current_config,
|
||||||
|
next_config,
|
||||||
|
)
|
||||||
|
now = utcnow()
|
||||||
|
next_config[_REGENERATION_MARKER_KEY] = {
|
||||||
|
"prepared": True,
|
||||||
|
"preview_invalidated": preview_invalidated,
|
||||||
|
"prepared_at": now,
|
||||||
|
}
|
||||||
|
# 002 迁移前发布的数据集只有 task_id。先补齐新关联字段,保证
|
||||||
|
# 解除任务输出指针后,详情和后续重新发布仍能定位原来的三份数据集。
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE datasets
|
||||||
|
SET source_task_id=%s, updated_at=%s
|
||||||
|
WHERE source='task' AND source_task_id IS NULL AND task_id=%s
|
||||||
|
AND deleted_at IS NULL
|
||||||
|
""",
|
||||||
|
(task_id, now, task_id),
|
||||||
|
)
|
||||||
|
published_row = conn.execute(
|
||||||
|
"""
|
||||||
|
SELECT EXISTS(
|
||||||
|
SELECT 1 FROM datasets
|
||||||
|
WHERE source='task' AND deleted_at IS NULL
|
||||||
|
AND (
|
||||||
|
source_task_id=%s
|
||||||
|
OR (source_task_id IS NULL AND task_id=%s)
|
||||||
|
)
|
||||||
|
) AS exists
|
||||||
|
""",
|
||||||
|
(task_id, task_id),
|
||||||
|
).fetchone()
|
||||||
|
published_outputs_preserved = bool(task.get("output_dataset_id")) or bool(
|
||||||
|
published_row and published_row.get("exists")
|
||||||
|
)
|
||||||
|
|
||||||
|
row = conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET name=%s, description=%s, config=%s, updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
RETURNING *
|
||||||
|
""",
|
||||||
|
(
|
||||||
|
payload["name"],
|
||||||
|
payload.get("description") or "",
|
||||||
|
json_dumps(next_config),
|
||||||
|
now,
|
||||||
|
task_id,
|
||||||
|
),
|
||||||
|
).fetchone()
|
||||||
|
except psycopg.errors.UniqueViolation as exc:
|
||||||
|
raise ConflictError("data process task name already exists") from exc
|
||||||
|
return {
|
||||||
|
"task": _public_task(_decode_row(row)) or {},
|
||||||
|
"preview_invalidated": preview_invalidated,
|
||||||
|
"published_outputs_preserved": published_outputs_preserved,
|
||||||
|
}
|
||||||
|
|
||||||
|
def delete_task(self, task_id: str, *, deleted_by: str | None = None) -> None:
|
||||||
|
with self.connect() as conn:
|
||||||
|
self._task_in_connection(conn, task_id, for_update=True)
|
||||||
|
now = utcnow()
|
||||||
|
conn.execute(
|
||||||
|
"""
|
||||||
|
UPDATE data_process_tasks
|
||||||
|
SET status=CASE WHEN status='running' THEN 'stopped' ELSE status END,
|
||||||
|
generation_run_id=NULL,
|
||||||
|
preview_status=CASE
|
||||||
|
WHEN preview_status IN ('queued', 'running') THEN 'cancelled'
|
||||||
|
ELSE preview_status
|
||||||
|
END,
|
||||||
|
preview_run_id=NULL,
|
||||||
|
deleted_at=%s, deleted_by=%s, updated_at=%s
|
||||||
|
WHERE id=%s
|
||||||
|
""",
|
||||||
|
(now, deleted_by, now, task_id),
|
||||||
|
)
|
||||||
1
backend/app/modules/dataset/__init__.py
Normal file
1
backend/app/modules/dataset/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Dataset management module."""
|
||||||
1
backend/app/modules/engine_registry/__init__.py
Normal file
1
backend/app/modules/engine_registry/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Training engine registry module."""
|
||||||
1
backend/app/modules/eval/__init__.py
Normal file
1
backend/app/modules/eval/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Evaluation module."""
|
||||||
1
backend/app/modules/file_gateway/__init__.py
Normal file
1
backend/app/modules/file_gateway/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Application-side file gateway module."""
|
||||||
1
backend/app/modules/fine_tune/__init__.py
Normal file
1
backend/app/modules/fine_tune/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Fine-tuning task module."""
|
||||||
1
backend/app/modules/gpu/__init__.py
Normal file
1
backend/app/modules/gpu/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""GPU assignment management module."""
|
||||||
77
backend/app/modules/gpu/router.py
Normal file
77
backend/app/modules/gpu/router.py
Normal file
@@ -0,0 +1,77 @@
|
|||||||
|
"""GPU 算力分配管理路由。"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Depends, Request
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.core.auth import get_current_user, is_admin
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/compute", tags=["gpu-assignment"])
|
||||||
|
|
||||||
|
|
||||||
|
def _actor_id(request: Request) -> str | None:
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
if token.startswith("platform-token-"):
|
||||||
|
return token[len("platform-token-"):]
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/gpu-assignments")
|
||||||
|
def list_assignments(current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
"""查看全部分配关系(仅 admin)。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
return ok(get_platform_store().gpu_assignments())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/gpu-assignments")
|
||||||
|
def assign_gpus(
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""批量分配 GPU(仅 admin)。body: { assignments: [{ node_id, gpu_index, user_id }] }"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
assignments = payload.get("assignments") or []
|
||||||
|
if not assignments:
|
||||||
|
raise fail(400, "assignments 不能为空")
|
||||||
|
actor = _actor_id(request) if request else None
|
||||||
|
result = get_platform_store().assign_gpus(assignments, assigned_by=actor)
|
||||||
|
get_platform_store().record_audit(
|
||||||
|
action="gpu.assign",
|
||||||
|
actor_id=actor,
|
||||||
|
target_type="gpu",
|
||||||
|
detail=f"count={len(assignments)}",
|
||||||
|
)
|
||||||
|
return ok(result)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/gpu-assignments/{assignment_id}")
|
||||||
|
def unassign_gpu(
|
||||||
|
assignment_id: str,
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""撤销 GPU 分配(仅 admin)。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
get_platform_store().unassign_gpu(assignment_id)
|
||||||
|
actor = _actor_id(request) if request else None
|
||||||
|
get_platform_store().record_audit(
|
||||||
|
action="gpu.unassign",
|
||||||
|
actor_id=actor,
|
||||||
|
target_type="gpu",
|
||||||
|
target_id=assignment_id,
|
||||||
|
)
|
||||||
|
return ok({"deleted": assignment_id})
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/my-gpus")
|
||||||
|
def my_gpus(current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
"""查看当前用户可用的 GPU 列表。"""
|
||||||
|
return ok(get_platform_store().gpu_assignments_for_user(current_user["id"]))
|
||||||
1
backend/app/modules/inference/__init__.py
Normal file
1
backend/app/modules/inference/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Inference and compare module."""
|
||||||
1
backend/app/modules/model/__init__.py
Normal file
1
backend/app/modules/model/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Model registry module."""
|
||||||
1
backend/app/modules/project/__init__.py
Normal file
1
backend/app/modules/project/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Project workspace and member module."""
|
||||||
244
backend/app/modules/project/router.py
Normal file
244
backend/app/modules/project/router.py
Normal file
@@ -0,0 +1,244 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Depends, Request
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.core.auth import filter_accessible_resource_ids, get_current_user, has_resource_access, is_admin
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/projects", tags=["project"])
|
||||||
|
|
||||||
|
|
||||||
|
def _actor(request: Request) -> str | None:
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
return token or None
|
||||||
|
|
||||||
|
|
||||||
|
def _require_no_pending_approval(resource_type: str, resource_id: str) -> None:
|
||||||
|
"""第 4 周:写操作审批拦截——存在待审批实例时拒绝执行。"""
|
||||||
|
store = get_platform_store()
|
||||||
|
pending = [
|
||||||
|
i for i in store.approval_instances(status="pending")
|
||||||
|
if i["resource_type"] == resource_type and i["resource_id"] == resource_id
|
||||||
|
]
|
||||||
|
if pending:
|
||||||
|
raise fail(409, "存在待审批的变更,请先完成审批")
|
||||||
|
|
||||||
|
|
||||||
|
def _require_approval_or_admin(
|
||||||
|
resource_type: str,
|
||||||
|
resource_id: str,
|
||||||
|
current_user: dict[str, Any],
|
||||||
|
action_desc: str = "",
|
||||||
|
) -> dict[str, Any] | None:
|
||||||
|
"""高风险操作审批旁路:admin 直接放行,普通用户创建审批实例(code=202)。"""
|
||||||
|
if is_admin(current_user):
|
||||||
|
return None
|
||||||
|
store = get_platform_store()
|
||||||
|
instance = store.create_approval_instance({
|
||||||
|
"resource_type": resource_type,
|
||||||
|
"resource_id": resource_id,
|
||||||
|
"applicant_id": current_user.get("id"),
|
||||||
|
"template_id": None,
|
||||||
|
})
|
||||||
|
return {
|
||||||
|
"code": 202,
|
||||||
|
"message": f"操作已提交审批,等待管理员批准:{action_desc}",
|
||||||
|
"data": {"approval_required": True, "approval_id": instance["id"]},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("")
|
||||||
|
def list_projects(
|
||||||
|
tenant_id: str = "default",
|
||||||
|
status: str | None = None,
|
||||||
|
keyword: str | None = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
projects = store.projects(tenant_id=tenant_id, status=status, keyword=keyword)
|
||||||
|
# #1 ACL 过滤:admin 直接放行,普通用户只能看到自己被授权的项目
|
||||||
|
accessible_ids = set(
|
||||||
|
filter_accessible_resource_ids("project", [p["id"] for p in projects], current_user)
|
||||||
|
)
|
||||||
|
filtered = [p for p in projects if p["id"] in accessible_ids]
|
||||||
|
return ok(filtered)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("")
|
||||||
|
def create_project(payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
proj = store.create_project(payload)
|
||||||
|
store.record_audit(
|
||||||
|
action="project.create",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project",
|
||||||
|
target_id=proj["id"],
|
||||||
|
tenant_id=proj.get("tenant_id"),
|
||||||
|
detail=f"name={proj.get('name')}",
|
||||||
|
)
|
||||||
|
return ok(proj)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{project_id}")
|
||||||
|
def get_project(project_id: str, current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
# #2 访问控制:普通用户无 read 权限则拒绝
|
||||||
|
if not has_resource_access("project", project_id, current_user, "read"):
|
||||||
|
raise fail(403, "no permission to access this project")
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().project(project_id))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "project not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{project_id}")
|
||||||
|
def update_project(
|
||||||
|
project_id: str,
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if not has_resource_access("project", project_id, current_user, "write"):
|
||||||
|
raise fail(403, "no permission to update this project")
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
proj = store.update_project(project_id, payload)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "project not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="project.update",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project",
|
||||||
|
target_id=project_id,
|
||||||
|
tenant_id=proj.get("tenant_id"),
|
||||||
|
detail=f"fields={','.join(payload.keys())}",
|
||||||
|
)
|
||||||
|
return ok(proj)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{project_id}/archive")
|
||||||
|
def archive_project(
|
||||||
|
project_id: str,
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
_require_no_pending_approval("project", project_id)
|
||||||
|
pending = _require_approval_or_admin("project", project_id, current_user, f"归档项目 {project_id}")
|
||||||
|
if pending:
|
||||||
|
return pending
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
proj = store.archive_project(project_id)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "project not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="project.archive",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project",
|
||||||
|
target_id=project_id,
|
||||||
|
tenant_id=proj.get("tenant_id"),
|
||||||
|
)
|
||||||
|
return ok(proj)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{project_id}")
|
||||||
|
def delete_project(
|
||||||
|
project_id: str,
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
_require_no_pending_approval("project", project_id)
|
||||||
|
pending = _require_approval_or_admin("project", project_id, current_user, f"删除项目 {project_id}")
|
||||||
|
if pending:
|
||||||
|
return pending
|
||||||
|
store = get_platform_store()
|
||||||
|
store.delete_project(project_id)
|
||||||
|
store.record_audit(
|
||||||
|
action="project.delete",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project",
|
||||||
|
target_id=project_id,
|
||||||
|
)
|
||||||
|
return ok(None)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{project_id}/members")
|
||||||
|
def list_members(project_id: str, current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
if not has_resource_access("project", project_id, current_user, "read"):
|
||||||
|
raise fail(403, "no permission to access this project")
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().project_members(project_id))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "project not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/{project_id}/members")
|
||||||
|
def add_member(
|
||||||
|
project_id: str,
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if not has_resource_access("project", project_id, current_user, "write"):
|
||||||
|
raise fail(403, "no permission to manage members of this project")
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
member = store.add_project_member(project_id, payload)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "project not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="project.member.add",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project.member",
|
||||||
|
target_id=project_id,
|
||||||
|
detail=f"user_id={payload.get('user_id')},role={payload.get('role')}",
|
||||||
|
)
|
||||||
|
return ok(member)
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{project_id}/members/{user_id}")
|
||||||
|
def update_member(
|
||||||
|
project_id: str,
|
||||||
|
user_id: str,
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if not has_resource_access("project", project_id, current_user, "write"):
|
||||||
|
raise fail(403, "no permission to manage members of this project")
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
member = store.update_project_member_role(project_id, user_id, payload)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "project or member not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="project.member.update",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project.member",
|
||||||
|
target_id=project_id,
|
||||||
|
detail=f"user_id={user_id},role={payload.get('role')}",
|
||||||
|
)
|
||||||
|
return ok(member)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{project_id}/members/{user_id}")
|
||||||
|
def remove_member(
|
||||||
|
project_id: str,
|
||||||
|
user_id: str,
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if not has_resource_access("project", project_id, current_user, "write"):
|
||||||
|
raise fail(403, "no permission to manage members of this project")
|
||||||
|
store = get_platform_store()
|
||||||
|
store.remove_project_member(project_id, user_id)
|
||||||
|
store.record_audit(
|
||||||
|
action="project.member.remove",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="project.member",
|
||||||
|
target_id=project_id,
|
||||||
|
detail=f"user_id={user_id}",
|
||||||
|
)
|
||||||
|
return ok(None)
|
||||||
1
backend/app/modules/resource/__init__.py
Normal file
1
backend/app/modules/resource/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Resource access control list (ACL) module."""
|
||||||
59
backend/app/modules/resource/router.py
Normal file
59
backend/app/modules/resource/router.py
Normal file
@@ -0,0 +1,59 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Request, Depends
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
from app.core.auth import get_current_user, has_resource_access, is_admin
|
||||||
|
from app.core.audit import audit_log, AuditActions
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/resources", tags=["resource"])
|
||||||
|
|
||||||
|
|
||||||
|
def _actor(request: Request) -> str | None:
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
return token or None
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{resource_type}/{resource_id}/acl")
|
||||||
|
def get_acl(resource_type: str, resource_id: str, current_user: dict = Depends(get_current_user)) -> dict[str, Any]:
|
||||||
|
"""查询资源 ACL,返回按主体分组的权限列表。"""
|
||||||
|
if not has_resource_access(resource_type, resource_id, current_user, "read"):
|
||||||
|
raise fail(403, "no permission to access resource ACL")
|
||||||
|
return ok(get_platform_store().resource_acl(resource_type, resource_id))
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{resource_type}/{resource_id}/acl")
|
||||||
|
@audit_log(
|
||||||
|
action=AuditActions.GRANT_ACL,
|
||||||
|
target_type="",
|
||||||
|
detail_template="设置资源授权: {resource_type}/{resource_id}",
|
||||||
|
)
|
||||||
|
def set_acl(
|
||||||
|
resource_type: str,
|
||||||
|
resource_id: str,
|
||||||
|
payload: dict[str, Any] = Body(...),
|
||||||
|
request: Request = None,
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""设置资源 ACL,body: { entries: [{ subject_type, subject_id, permissions: [] }] }"""
|
||||||
|
if not is_admin(current_user) and not has_resource_access(resource_type, resource_id, current_user, "write"):
|
||||||
|
raise fail(403, "only resource owner or admin can update ACL")
|
||||||
|
entries = payload.get("entries") or []
|
||||||
|
allowed = {"read", "write", "execute", "download", "delete", "admin"}
|
||||||
|
for entry in entries:
|
||||||
|
if entry.get("principal_type") not in {"user", "role"} or not entry.get("principal_id"):
|
||||||
|
raise fail(400, "invalid ACL principal")
|
||||||
|
if any(permission not in allowed for permission in entry.get("permissions") or []):
|
||||||
|
raise fail(400, "invalid ACL permission")
|
||||||
|
result = get_platform_store().set_resource_acl(resource_type, resource_id, entries)
|
||||||
|
get_platform_store().record_audit(
|
||||||
|
action="resource.acl.set",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type=resource_type,
|
||||||
|
target_id=resource_id,
|
||||||
|
detail=f"entries={len(entries)}",
|
||||||
|
)
|
||||||
|
return ok(result)
|
||||||
1
backend/app/modules/retention/__init__.py
Normal file
1
backend/app/modules/retention/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Retention policy and cleanup module."""
|
||||||
75
backend/app/modules/retention/router.py
Normal file
75
backend/app/modules/retention/router.py
Normal file
@@ -0,0 +1,75 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Request
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/retention-policies", tags=["retention"])
|
||||||
|
|
||||||
|
|
||||||
|
def _actor(request: Request) -> str | None:
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
return token or None
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("")
|
||||||
|
def list_policies() -> dict[str, Any]:
|
||||||
|
return ok(get_platform_store().retention_policies())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("")
|
||||||
|
def create_policy(payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||||
|
if not payload.get("name"):
|
||||||
|
raise fail(400, "name 必填")
|
||||||
|
policy = get_platform_store().create_retention_policy(payload)
|
||||||
|
get_platform_store().record_audit(
|
||||||
|
action="retention.create",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="retention_policy",
|
||||||
|
target_id=policy["id"],
|
||||||
|
detail=f"name={policy.get('name')}",
|
||||||
|
)
|
||||||
|
return ok(policy)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{policy_id}")
|
||||||
|
def get_policy(policy_id: str) -> dict[str, Any]:
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().retention_policy(policy_id))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "retention policy not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{policy_id}")
|
||||||
|
def update_policy(
|
||||||
|
policy_id: str, payload: dict[str, Any] = Body(...), request: Request = None
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
policy = store.update_retention_policy(policy_id, payload)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "retention policy not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="retention.update",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="retention_policy",
|
||||||
|
target_id=policy_id,
|
||||||
|
detail=f"fields={','.join(payload.keys())}",
|
||||||
|
)
|
||||||
|
return ok(policy)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{policy_id}")
|
||||||
|
def delete_policy(policy_id: str, request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
store.delete_retention_policy(policy_id)
|
||||||
|
store.record_audit(
|
||||||
|
action="retention.delete",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="retention_policy",
|
||||||
|
target_id=policy_id,
|
||||||
|
)
|
||||||
|
return ok({"deleted": policy_id})
|
||||||
1
backend/app/modules/storage/__init__.py
Normal file
1
backend/app/modules/storage/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Central object storage integration."""
|
||||||
68
backend/app/modules/storage/minio_store.py
Normal file
68
backend/app/modules/storage/minio_store.py
Normal file
@@ -0,0 +1,68 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from datetime import timedelta
|
||||||
|
from functools import lru_cache
|
||||||
|
from io import BytesIO
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from minio import Minio
|
||||||
|
from minio.error import S3Error
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
|
||||||
|
|
||||||
|
class ObjectStorageError(RuntimeError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class MinioObjectStorage:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
settings = get_settings()
|
||||||
|
endpoint = settings.minio_endpoint.replace("http://", "").replace("https://", "").rstrip("/")
|
||||||
|
self.client = Minio(endpoint, access_key=settings.minio_access_key, secret_key=settings.minio_secret_key, secure=settings.minio_secure)
|
||||||
|
self.bucket = settings.minio_bucket
|
||||||
|
|
||||||
|
def _ensure_enabled(self) -> None:
|
||||||
|
if not get_settings().minio_enabled:
|
||||||
|
raise ObjectStorageError("MinIO object storage is disabled")
|
||||||
|
|
||||||
|
def ensure_bucket(self) -> None:
|
||||||
|
self._ensure_enabled()
|
||||||
|
try:
|
||||||
|
if not self.client.bucket_exists(self.bucket):
|
||||||
|
self.client.make_bucket(self.bucket)
|
||||||
|
except S3Error as exc:
|
||||||
|
raise ObjectStorageError(str(exc)) from exc
|
||||||
|
|
||||||
|
def presigned_put(self, object_key: str, expires_seconds: int = 3600) -> str:
|
||||||
|
self._ensure_enabled()
|
||||||
|
self.ensure_bucket()
|
||||||
|
return self.client.presigned_put_object(self.bucket, object_key, expires=timedelta(seconds=expires_seconds))
|
||||||
|
|
||||||
|
def presigned_get(self, object_key: str, expires_seconds: int = 3600) -> str:
|
||||||
|
self._ensure_enabled()
|
||||||
|
self.ensure_bucket()
|
||||||
|
return self.client.presigned_get_object(self.bucket, object_key, expires=timedelta(seconds=expires_seconds))
|
||||||
|
|
||||||
|
def stat(self, object_key: str) -> dict[str, Any]:
|
||||||
|
self._ensure_enabled()
|
||||||
|
self.ensure_bucket()
|
||||||
|
try:
|
||||||
|
result = self.client.stat_object(self.bucket, object_key)
|
||||||
|
return {"object_key": object_key, "byte_size": result.size, "etag": result.etag, "last_modified": result.last_modified.isoformat() if result.last_modified else None}
|
||||||
|
except S3Error as exc:
|
||||||
|
raise ObjectStorageError(str(exc)) from exc
|
||||||
|
|
||||||
|
def put_bytes(self, object_key: str, content: bytes, content_type: str = "application/octet-stream") -> dict[str, Any]:
|
||||||
|
self._ensure_enabled()
|
||||||
|
self.ensure_bucket()
|
||||||
|
try:
|
||||||
|
result = self.client.put_object(self.bucket, object_key, BytesIO(content), len(content), content_type=content_type)
|
||||||
|
return {"bucket": self.bucket, "object_key": object_key, "etag": result.etag, "byte_size": len(content)}
|
||||||
|
except S3Error as exc:
|
||||||
|
raise ObjectStorageError(str(exc)) from exc
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache
|
||||||
|
def get_object_storage() -> MinioObjectStorage:
|
||||||
|
return MinioObjectStorage()
|
||||||
1
backend/app/modules/system/__init__.py
Normal file
1
backend/app/modules/system/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""System health, metrics and logs module."""
|
||||||
264
backend/app/modules/system/router.py
Normal file
264
backend/app/modules/system/router.py
Normal file
@@ -0,0 +1,264 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Query, Request, Depends
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
|
||||||
|
from app.db.platform_store import ALL_PERMISSIONS, get_platform_store
|
||||||
|
from app.core.auth import get_current_user, is_admin
|
||||||
|
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/system", tags=["system"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/audit/visit")
|
||||||
|
def record_visit(payload: dict = Body(...), request: Request = None) -> dict:
|
||||||
|
"""记录用户访问业务模块的行为,用于看板用户操作分布统计。"""
|
||||||
|
action = str(payload.get("action") or payload.get("module") or "").strip()
|
||||||
|
if not action:
|
||||||
|
return {"code": 0, "message": "ok", "data": {"recorded": False}}
|
||||||
|
actor_id = ""
|
||||||
|
if request is not None:
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
if token.startswith("platform-token-"):
|
||||||
|
actor_id = token[len("platform-token-"):]
|
||||||
|
get_platform_store().record_audit(
|
||||||
|
action=action,
|
||||||
|
actor_id=actor_id or None,
|
||||||
|
target_type="module",
|
||||||
|
target_id=action,
|
||||||
|
detail=str(payload.get("detail") or ""),
|
||||||
|
)
|
||||||
|
return {"code": 0, "message": "ok", "data": {"recorded": True}}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/permissions/codes")
|
||||||
|
def permission_codes(current_user: dict = Depends(get_current_user)) -> dict:
|
||||||
|
"""返回平台权限码清单(权限码接口)。"""
|
||||||
|
return {"code": 0, "message": "ok", "data": {"codes": ALL_PERMISSIONS}}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/permissions")
|
||||||
|
def permissions_overview(current_user: dict = Depends(get_current_user)) -> dict:
|
||||||
|
"""返回权限码清单与角色定义。"""
|
||||||
|
store = get_platform_store()
|
||||||
|
return {
|
||||||
|
"code": 0,
|
||||||
|
"message": "ok",
|
||||||
|
"data": {"codes": ALL_PERMISSIONS, "roles": store.roles()},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/audit-logs")
|
||||||
|
def audit_logs(
|
||||||
|
tenant_id: str | None = Query(default=None, description="租户 ID"),
|
||||||
|
project_id: str | None = Query(default=None, description="项目 ID"),
|
||||||
|
actor_id: str | None = Query(default=None, description="操作人 ID"),
|
||||||
|
action: str | None = Query(default=None, description="动作类型"),
|
||||||
|
target_type: str | None = Query(default=None, description="目标类型"),
|
||||||
|
start_time: str | None = Query(default=None, description="ISO8601 起始时间"),
|
||||||
|
end_time: str | None = Query(default=None, description="ISO8601 结束时间"),
|
||||||
|
limit: int = Query(default=50, ge=1, le=200),
|
||||||
|
offset: int = Query(default=0, ge=0),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict:
|
||||||
|
"""审计日志查询:按租户/项目/操作人/动作/目标类型/时间范围分页过滤。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
from app.api.v1.endpoints.platform import fail
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
store = get_platform_store()
|
||||||
|
result = store.audit_logs(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
project_id=project_id,
|
||||||
|
actor_id=actor_id,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
start_time=start_time,
|
||||||
|
end_time=end_time,
|
||||||
|
limit=limit,
|
||||||
|
offset=offset,
|
||||||
|
)
|
||||||
|
return {"code": 0, "message": "ok", "data": result}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/audit-logs/export")
|
||||||
|
def audit_logs_export(
|
||||||
|
tenant_id: str | None = Query(default=None, description="租户 ID"),
|
||||||
|
project_id: str | None = Query(default=None, description="项目 ID"),
|
||||||
|
actor_id: str | None = Query(default=None, description="操作人 ID"),
|
||||||
|
action: str | None = Query(default=None, description="动作类型"),
|
||||||
|
target_type: str | None = Query(default=None, description="目标类型"),
|
||||||
|
start_time: str | None = Query(default=None, description="ISO8601 起始时间"),
|
||||||
|
end_time: str | None = Query(default=None, description="ISO8601 结束时间"),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> StreamingResponse:
|
||||||
|
"""审计日志导出:返回 CSV 流,与应用查询相同的过滤条件。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
from app.api.v1.endpoints.platform import fail
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
store = get_platform_store()
|
||||||
|
result = store.audit_logs(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
project_id=project_id,
|
||||||
|
actor_id=actor_id,
|
||||||
|
action=action,
|
||||||
|
target_type=target_type,
|
||||||
|
start_time=start_time,
|
||||||
|
end_time=end_time,
|
||||||
|
limit=10000,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
items = result["items"]
|
||||||
|
columns = ["time", "tenant_id", "project_id", "actor_id", "action", "target_type", "target_id", "detail", "client_ip"]
|
||||||
|
header = ",".join(columns) + "\n"
|
||||||
|
|
||||||
|
def iter_rows():
|
||||||
|
yield header
|
||||||
|
for row in items:
|
||||||
|
yield ",".join(f'"{str(row.get(c, "") or "")}"' for c in columns) + "\n"
|
||||||
|
|
||||||
|
return StreamingResponse(
|
||||||
|
iter_rows(),
|
||||||
|
media_type="text/csv",
|
||||||
|
headers={"Content-Disposition": "attachment; filename=audit_logs.csv"},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ===================== 操作日志 =====================
|
||||||
|
|
||||||
|
@router.get("/operation-logs")
|
||||||
|
def operation_logs(
|
||||||
|
user_id: str | None = Query(default=None, description="按用户 ID 筛选"),
|
||||||
|
module: str | None = Query(default=None, description="按模块筛选: fine-tune/model-eval/model-inference/dataset/data-convert/model-manage"),
|
||||||
|
action: str | None = Query(default=None, description="按动作筛选: create/start/stop/delete/upload/convert/merge"),
|
||||||
|
status: str | None = Query(default=None, description="按状态筛选: success/failure(不传则查全部)"),
|
||||||
|
keyword: str | None = Query(default=None, description="关键字搜索报错信息(error_message)"),
|
||||||
|
start_time: str | None = Query(default=None, description="ISO8601 起始时间"),
|
||||||
|
end_time: str | None = Query(default=None, description="ISO8601 结束时间"),
|
||||||
|
limit: int = Query(default=50, ge=1, le=200),
|
||||||
|
offset: int = Query(default=0, ge=0),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict:
|
||||||
|
"""操作日志查询:按用户/模块/动作/状态/关键字/时间范围分页过滤。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
from app.api.v1.endpoints.platform import fail
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
store = get_platform_store()
|
||||||
|
conditions = []
|
||||||
|
params: list = []
|
||||||
|
if user_id:
|
||||||
|
conditions.append("user_id = %s")
|
||||||
|
params.append(user_id)
|
||||||
|
if module:
|
||||||
|
conditions.append("module = %s")
|
||||||
|
params.append(module)
|
||||||
|
if action:
|
||||||
|
conditions.append("action = %s")
|
||||||
|
params.append(action)
|
||||||
|
if status:
|
||||||
|
conditions.append("status = %s")
|
||||||
|
params.append(status)
|
||||||
|
if keyword:
|
||||||
|
conditions.append("(error_message ILIKE %s OR error_type ILIKE %s)")
|
||||||
|
params.append(f"%{keyword}%")
|
||||||
|
params.append(f"%{keyword}%")
|
||||||
|
if start_time:
|
||||||
|
conditions.append("create_time >= %s")
|
||||||
|
params.append(start_time)
|
||||||
|
if end_time:
|
||||||
|
conditions.append("create_time <= %s")
|
||||||
|
params.append(end_time)
|
||||||
|
where = " WHERE " + " AND ".join(conditions) if conditions else ""
|
||||||
|
with store.connect() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
f"SELECT * FROM operation_logs{where} ORDER BY create_time DESC LIMIT %s OFFSET %s",
|
||||||
|
tuple(params + [limit, offset]),
|
||||||
|
).fetchall()
|
||||||
|
total = conn.execute(f"SELECT COUNT(*) FROM operation_logs{where}", tuple(params)).fetchone()[0]
|
||||||
|
return {"code": 0, "message": "ok", "data": {"items": [dict(r) for r in rows], "total": total}}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/operation-logs/stats")
|
||||||
|
def operation_logs_stats(
|
||||||
|
start_time: str | None = Query(default=None, description="ISO8601 起始时间"),
|
||||||
|
end_time: str | None = Query(default=None, description="ISO8601 结束时间"),
|
||||||
|
current_user: dict = Depends(get_current_user),
|
||||||
|
) -> dict:
|
||||||
|
"""操作日志统计:总操作数、成功数、失败数、失败率、各模块失败分布、最近错误列表。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
from app.api.v1.endpoints.platform import fail
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
store = get_platform_store()
|
||||||
|
conditions = []
|
||||||
|
params: list = []
|
||||||
|
if start_time:
|
||||||
|
conditions.append("create_time >= %s")
|
||||||
|
params.append(start_time)
|
||||||
|
if end_time:
|
||||||
|
conditions.append("create_time <= %s")
|
||||||
|
params.append(end_time)
|
||||||
|
where = " WHERE " + " AND ".join(conditions) if conditions else ""
|
||||||
|
failure_where = where + " AND status = 'failure'" if where else " WHERE status = 'failure'"
|
||||||
|
|
||||||
|
with store.connect() as conn:
|
||||||
|
# 总计
|
||||||
|
row = conn.execute(
|
||||||
|
f"SELECT status, COUNT(*) as cnt FROM operation_logs{where} GROUP BY status", tuple(params)
|
||||||
|
).fetchall()
|
||||||
|
total_count = 0
|
||||||
|
success_count = 0
|
||||||
|
failure_count = 0
|
||||||
|
for r in row:
|
||||||
|
total_count += r["cnt"]
|
||||||
|
if r["status"] == "success":
|
||||||
|
success_count = r["cnt"]
|
||||||
|
elif r["status"] == "failure":
|
||||||
|
failure_count = r["cnt"]
|
||||||
|
failure_rate = round(failure_count / total_count * 100, 2) if total_count > 0 else 0
|
||||||
|
|
||||||
|
# 各模块失败数
|
||||||
|
module_stats = conn.execute(
|
||||||
|
f"SELECT module, COUNT(*) as cnt FROM operation_logs{failure_where} GROUP BY module ORDER BY cnt DESC",
|
||||||
|
tuple(params),
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
# 各异常类型分布
|
||||||
|
error_type_stats = conn.execute(
|
||||||
|
f"SELECT error_type, COUNT(*) as cnt FROM operation_logs{failure_where} AND error_type IS NOT NULL GROUP BY error_type ORDER BY cnt DESC LIMIT 10",
|
||||||
|
tuple(params),
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
# 最近的 10 条错误
|
||||||
|
recent_errors = conn.execute(
|
||||||
|
f"SELECT * FROM operation_logs{failure_where} ORDER BY create_time DESC LIMIT 10",
|
||||||
|
tuple(params),
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"code": 0,
|
||||||
|
"message": "ok",
|
||||||
|
"data": {
|
||||||
|
"total": total_count,
|
||||||
|
"success": success_count,
|
||||||
|
"failure": failure_count,
|
||||||
|
"failure_rate": failure_rate,
|
||||||
|
"module_failures": [{"module": r["module"], "count": r["cnt"]} for r in module_stats],
|
||||||
|
"error_types": [{"type": r["error_type"], "count": r["cnt"]} for r in error_type_stats],
|
||||||
|
"recent_errors": [dict(r) for r in recent_errors],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/operation-logs/modules")
|
||||||
|
def operation_log_modules(current_user: dict = Depends(get_current_user)) -> dict:
|
||||||
|
"""返回操作日志中出现的模块列表(用于筛选下拉框)。"""
|
||||||
|
if not is_admin(current_user):
|
||||||
|
from app.api.v1.endpoints.platform import fail
|
||||||
|
raise fail(403, "admin permission required")
|
||||||
|
store = get_platform_store()
|
||||||
|
with store.connect() as conn:
|
||||||
|
rows = conn.execute(
|
||||||
|
"SELECT DISTINCT module FROM operation_logs WHERE module IS NOT NULL ORDER BY module"
|
||||||
|
).fetchall()
|
||||||
|
modules = [{"value": r["module"], "label": r["module"]} for r in rows]
|
||||||
|
return {"code": 0, "message": "ok", "data": modules}
|
||||||
1
backend/app/modules/tenant/__init__.py
Normal file
1
backend/app/modules/tenant/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Tenant management module."""
|
||||||
116
backend/app/modules/tenant/router.py
Normal file
116
backend/app/modules/tenant/router.py
Normal file
@@ -0,0 +1,116 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Body, Request
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import ok, fail
|
||||||
|
from app.db.platform_store import get_platform_store
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/tenants", tags=["tenant"])
|
||||||
|
|
||||||
|
|
||||||
|
def _actor(request: Request) -> str | None:
|
||||||
|
auth = request.headers.get("Authorization", "")
|
||||||
|
token = auth.replace("Bearer ", "").strip()
|
||||||
|
return token or None
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("")
|
||||||
|
def list_tenants() -> dict[str, Any]:
|
||||||
|
return ok(get_platform_store().tenants())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("")
|
||||||
|
def create_tenant(payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
tenant = store.create_tenant(payload)
|
||||||
|
except KeyError as e:
|
||||||
|
raise fail(400, f"missing field: {e}")
|
||||||
|
store.record_audit(
|
||||||
|
action="tenant.create",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="tenant",
|
||||||
|
target_id=tenant["id"],
|
||||||
|
tenant_id=tenant["id"],
|
||||||
|
detail=f"name={tenant.get('name')}",
|
||||||
|
)
|
||||||
|
return ok(tenant)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/{tenant_id}")
|
||||||
|
def get_tenant(tenant_id: str) -> dict[str, Any]:
|
||||||
|
try:
|
||||||
|
return ok(get_platform_store().tenant(tenant_id))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "tenant not found")
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{tenant_id}")
|
||||||
|
def update_tenant(tenant_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
tenant = store.update_tenant(tenant_id, payload)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "tenant not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="tenant.update",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="tenant",
|
||||||
|
target_id=tenant_id,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
detail=f"fields={','.join(payload.keys())}",
|
||||||
|
)
|
||||||
|
return ok(tenant)
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{tenant_id}/quota")
|
||||||
|
def set_quota(tenant_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
tenant = store.set_tenant_quota(tenant_id, payload)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "tenant not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="tenant.quota.set",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="tenant",
|
||||||
|
target_id=tenant_id,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
)
|
||||||
|
return ok(tenant)
|
||||||
|
|
||||||
|
|
||||||
|
@router.put("/{tenant_id}/retention-policy")
|
||||||
|
def set_retention(tenant_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
tenant = store.set_tenant_retention(tenant_id, payload.get("retention_policy_id"))
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "tenant not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="tenant.retention.set",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="tenant",
|
||||||
|
target_id=tenant_id,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
)
|
||||||
|
return ok(tenant)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{tenant_id}")
|
||||||
|
def delete_tenant(tenant_id: str, request: Request = None) -> dict[str, Any]:
|
||||||
|
store = get_platform_store()
|
||||||
|
try:
|
||||||
|
tenant = store.delete_tenant(tenant_id)
|
||||||
|
except KeyError:
|
||||||
|
raise fail(404, "tenant not found")
|
||||||
|
store.record_audit(
|
||||||
|
action="tenant.delete",
|
||||||
|
actor_id=_actor(request) if request else None,
|
||||||
|
target_type="tenant",
|
||||||
|
target_id=tenant_id,
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
detail=f"name={tenant.get('name')}",
|
||||||
|
)
|
||||||
|
return ok(tenant)
|
||||||
1
backend/app/schemas/__init__.py
Normal file
1
backend/app/schemas/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Shared schemas package."""
|
||||||
417
backend/app/schemas/data_process.py
Normal file
417
backend/app/schemas/data_process.py
Normal file
@@ -0,0 +1,417 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from enum import StrEnum
|
||||||
|
from typing import Any, Literal
|
||||||
|
from urllib.parse import parse_qs, urlsplit
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||||
|
|
||||||
|
from app.modules.data_process.constants import MAX_QA_PAIRS_PER_ITEM
|
||||||
|
|
||||||
|
|
||||||
|
def _config_value(config: dict[str, Any], snake_name: str, camel_name: str, default: Any) -> Any:
|
||||||
|
if snake_name in config:
|
||||||
|
return config[snake_name]
|
||||||
|
return config.get(camel_name, default)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_process_config(config: dict[str, Any]) -> None:
|
||||||
|
output_type = _config_value(config, "output_type", "outputType", "standard")
|
||||||
|
if output_type not in {"standard", "reasoning", "dpo"}:
|
||||||
|
raise ValueError("output_type must be one of: standard, reasoning, dpo")
|
||||||
|
|
||||||
|
source_mode = _config_value(config, "source_mode", "sourceMode", "local")
|
||||||
|
if source_mode not in {"local", "external"}:
|
||||||
|
raise ValueError("source_mode must be one of: local, external")
|
||||||
|
external_source = _config_value(config, "external_source", "externalSource", None)
|
||||||
|
if external_source is not None:
|
||||||
|
if not isinstance(external_source, dict):
|
||||||
|
raise ValueError("external_source must be an object")
|
||||||
|
if any(
|
||||||
|
key.lower() in {"password", "secret", "token", "api_key"}
|
||||||
|
for key in external_source
|
||||||
|
):
|
||||||
|
raise ValueError("external_source must not persist credentials")
|
||||||
|
external_url = str(external_source.get("url") or "").strip()
|
||||||
|
if external_url:
|
||||||
|
parsed_external_url = urlsplit(external_url)
|
||||||
|
sensitive_query_keys = {"password", "secret", "token", "api_key", "user", "username"}
|
||||||
|
if parsed_external_url.username or parsed_external_url.password or (
|
||||||
|
set(parse_qs(parsed_external_url.query)) & sensitive_query_keys
|
||||||
|
):
|
||||||
|
raise ValueError("external_source URL must not contain credentials")
|
||||||
|
|
||||||
|
chunk_method = _config_value(config, "chunk_method", "chunkMethod", "layout_hybrid")
|
||||||
|
if not isinstance(chunk_method, str) or chunk_method not in {
|
||||||
|
"layout_hybrid",
|
||||||
|
"semantic",
|
||||||
|
"fixed",
|
||||||
|
}:
|
||||||
|
raise ValueError("chunk_method must be one of: layout_hybrid, semantic, fixed")
|
||||||
|
|
||||||
|
semantic_percentile = _config_value(
|
||||||
|
config,
|
||||||
|
"semantic_breakpoint_percentile",
|
||||||
|
"semanticBreakpointPercentile",
|
||||||
|
95,
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
isinstance(semantic_percentile, bool)
|
||||||
|
or not isinstance(semantic_percentile, int)
|
||||||
|
or not 1 <= semantic_percentile <= 99
|
||||||
|
):
|
||||||
|
raise ValueError("semantic_breakpoint_percentile must be an integer in [1, 99]")
|
||||||
|
|
||||||
|
split = _config_value(config, "dataset_split", "datasetSplit", None)
|
||||||
|
if split is not None:
|
||||||
|
if not isinstance(split, dict) or set(split) != {"train", "validation", "test"}:
|
||||||
|
raise ValueError("dataset_split must contain train, validation and test")
|
||||||
|
values = list(split.values())
|
||||||
|
if any(isinstance(value, bool) or not isinstance(value, int) for value in values):
|
||||||
|
raise ValueError("dataset_split values must be integers")
|
||||||
|
if any(value < 0 or value > 100 for value in values) or sum(values) != 100:
|
||||||
|
raise ValueError("dataset_split values must be in [0, 100] and total 100")
|
||||||
|
|
||||||
|
chunk_fields = {
|
||||||
|
"chunk_size",
|
||||||
|
"chunkSize",
|
||||||
|
"chunk_overlap",
|
||||||
|
"chunkOverlap",
|
||||||
|
"min_chunk_size",
|
||||||
|
"minChunkSize",
|
||||||
|
}
|
||||||
|
if chunk_fields.intersection(config):
|
||||||
|
chunk_size = _config_value(config, "chunk_size", "chunkSize", 800)
|
||||||
|
overlap = _config_value(config, "chunk_overlap", "chunkOverlap", 100)
|
||||||
|
minimum = _config_value(config, "min_chunk_size", "minChunkSize", 100)
|
||||||
|
if any(
|
||||||
|
isinstance(value, bool) or not isinstance(value, int)
|
||||||
|
for value in (chunk_size, overlap, minimum)
|
||||||
|
):
|
||||||
|
raise ValueError("chunk_size, chunk_overlap and min_chunk_size must be integers")
|
||||||
|
if not 16 <= chunk_size <= 32_768:
|
||||||
|
raise ValueError("chunk_size must be in [16, 32768]")
|
||||||
|
if overlap < 0 or overlap >= chunk_size:
|
||||||
|
raise ValueError("chunk_overlap must be in [0, chunk_size)")
|
||||||
|
if minimum <= 0 or minimum > chunk_size or overlap + minimum > chunk_size:
|
||||||
|
raise ValueError("min_chunk_size and chunk_overlap exceed chunk_size")
|
||||||
|
|
||||||
|
temperature = _config_value(config, "temperature", "temperature", None)
|
||||||
|
if temperature is not None:
|
||||||
|
if isinstance(temperature, bool) or not isinstance(temperature, (int, float)):
|
||||||
|
raise ValueError("temperature must be a number")
|
||||||
|
if not 0 <= float(temperature) <= 2:
|
||||||
|
raise ValueError("temperature must be in [0, 2]")
|
||||||
|
|
||||||
|
max_tokens = _config_value(config, "max_tokens", "maxTokens", None)
|
||||||
|
if max_tokens is not None:
|
||||||
|
if isinstance(max_tokens, bool) or not isinstance(max_tokens, int):
|
||||||
|
raise ValueError("max_tokens must be an integer")
|
||||||
|
if not 1 <= max_tokens <= 32_768:
|
||||||
|
raise ValueError("max_tokens must be in [1, 32768]")
|
||||||
|
|
||||||
|
for snake_name, camel_name in (
|
||||||
|
("qa_pairs_per_row", "qaPairsPerRow"),
|
||||||
|
("qa_pairs_per_chunk", "qaPairsPerChunk"),
|
||||||
|
):
|
||||||
|
pairs = _config_value(config, snake_name, camel_name, None)
|
||||||
|
if pairs is None:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
isinstance(pairs, bool)
|
||||||
|
or not isinstance(pairs, int)
|
||||||
|
or not 1 <= pairs <= MAX_QA_PAIRS_PER_ITEM
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"{snake_name} must be an integer in [1, {MAX_QA_PAIRS_PER_ITEM}]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessStatus(StrEnum):
|
||||||
|
pending = "pending"
|
||||||
|
running = "running"
|
||||||
|
completed = "completed"
|
||||||
|
failed = "failed"
|
||||||
|
stopped = "stopped"
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessWorkflowStep(StrEnum):
|
||||||
|
create = "create"
|
||||||
|
model = "model"
|
||||||
|
upload = "upload"
|
||||||
|
preview = "preview"
|
||||||
|
generate = "generate"
|
||||||
|
results = "results"
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessPreviewStatus(StrEnum):
|
||||||
|
idle = "idle"
|
||||||
|
queued = "queued"
|
||||||
|
running = "running"
|
||||||
|
completed = "completed"
|
||||||
|
failed = "failed"
|
||||||
|
cancelled = "cancelled"
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessType(StrEnum):
|
||||||
|
structured = "structured"
|
||||||
|
unstructured = "unstructured"
|
||||||
|
external = "external"
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessTaskCreate(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
name: str = Field(min_length=1, max_length=150)
|
||||||
|
description: str = ""
|
||||||
|
process_type: ProcessType
|
||||||
|
source_dataset_id: str | None = None
|
||||||
|
config: dict[str, Any] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
@field_validator("name")
|
||||||
|
@classmethod
|
||||||
|
def normalize_name(cls, value: str) -> str:
|
||||||
|
value = value.strip()
|
||||||
|
if not value:
|
||||||
|
raise ValueError("task name cannot be empty")
|
||||||
|
return value
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_config(self) -> "DataProcessTaskCreate":
|
||||||
|
_validate_process_config(self.config)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessTaskUpdate(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
name: str | None = Field(default=None, min_length=1, max_length=150)
|
||||||
|
description: str | None = None
|
||||||
|
process_type: ProcessType | None = None
|
||||||
|
source_dataset_id: str | None = None
|
||||||
|
config: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
@field_validator("name")
|
||||||
|
@classmethod
|
||||||
|
def normalize_name(cls, value: str | None) -> str | None:
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
value = value.strip()
|
||||||
|
if not value:
|
||||||
|
raise ValueError("task name cannot be empty")
|
||||||
|
return value
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_config(self) -> "DataProcessTaskUpdate":
|
||||||
|
if self.config is not None:
|
||||||
|
_validate_process_config(self.config)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessWorkflowStepUpdate(BaseModel):
|
||||||
|
"""仅保存创建向导位置,不修改配置或使下游产物失效。"""
|
||||||
|
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
workflow_step: DataProcessWorkflowStep
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessRegenerateRequest(BaseModel):
|
||||||
|
"""以一份完整配置准备任务重新生成。
|
||||||
|
|
||||||
|
``expected_updated_at`` 用于防止详情页的旧快照覆盖其他人刚刚
|
||||||
|
保存的配置。重新生成不允许改变处理类型,避免旧源文件在新解析
|
||||||
|
规则下被静默误用。
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
name: str = Field(min_length=1, max_length=150)
|
||||||
|
description: str
|
||||||
|
process_type: ProcessType
|
||||||
|
config: dict[str, Any]
|
||||||
|
expected_updated_at: str = Field(min_length=1)
|
||||||
|
|
||||||
|
@field_validator("name")
|
||||||
|
@classmethod
|
||||||
|
def normalize_name(cls, value: str) -> str:
|
||||||
|
value = value.strip()
|
||||||
|
if not value:
|
||||||
|
raise ValueError("task name cannot be empty")
|
||||||
|
return value
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_config(self) -> "DataProcessRegenerateRequest":
|
||||||
|
_validate_process_config(self.config)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class DataProcessRepeatRequest(BaseModel):
|
||||||
|
"""按已确认任务的完整快照创建一批独立的新生成结果。"""
|
||||||
|
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
expected_updated_at: str = Field(min_length=1)
|
||||||
|
request_id: str = Field(
|
||||||
|
min_length=8,
|
||||||
|
max_length=80,
|
||||||
|
pattern=r"^[A-Za-z0-9_-]+$",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PreviewBuildRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
replace_existing: Literal[True] = True
|
||||||
|
source_file_ids: list[str] | None = None
|
||||||
|
source_file_id: str | None = None
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_source_file_selection(self) -> "PreviewBuildRequest":
|
||||||
|
if self.source_file_ids is not None and self.source_file_id is not None:
|
||||||
|
raise ValueError("source_file_id and source_file_ids cannot be used together")
|
||||||
|
values = self.source_file_ids
|
||||||
|
if values is None and self.source_file_id is not None:
|
||||||
|
values = [self.source_file_id]
|
||||||
|
if values is None:
|
||||||
|
return self
|
||||||
|
normalized = list(dict.fromkeys(str(value).strip() for value in values))
|
||||||
|
if not normalized or any(not value for value in normalized):
|
||||||
|
raise ValueError("at least one non-empty source file id is required")
|
||||||
|
self.source_file_ids = normalized
|
||||||
|
self.source_file_id = None
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class PreviewItemCreate(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
source_file_id: str | None = None
|
||||||
|
original_content: str = ""
|
||||||
|
edited_content: str = ""
|
||||||
|
source_start: int | None = Field(default=None, ge=0)
|
||||||
|
source_end: int | None = Field(default=None, ge=0)
|
||||||
|
source_start_line: int | None = Field(default=None, ge=1)
|
||||||
|
source_end_line: int | None = Field(default=None, ge=1)
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_ranges(self) -> "PreviewItemCreate":
|
||||||
|
if self.source_start is not None and self.source_end is not None:
|
||||||
|
if self.source_end < self.source_start:
|
||||||
|
raise ValueError("source_end must be greater than or equal to source_start")
|
||||||
|
if self.source_start_line is not None and self.source_end_line is not None:
|
||||||
|
if self.source_end_line < self.source_start_line:
|
||||||
|
raise ValueError(
|
||||||
|
"source_end_line must be greater than or equal to source_start_line"
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class PreviewItemUpdate(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
edited_content: str
|
||||||
|
expected_updated_at: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class GenerateRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
replace_existing: Literal[True] = True
|
||||||
|
|
||||||
|
|
||||||
|
class ExternalSourceRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
type: str = Field(min_length=1, max_length=30)
|
||||||
|
url: str = Field(min_length=1, max_length=2048)
|
||||||
|
auth_mode: Literal["none", "basic"] = "none"
|
||||||
|
username: str | None = Field(default=None, max_length=150)
|
||||||
|
password: str | None = Field(default=None, max_length=500)
|
||||||
|
limit: int = Field(default=1000, ge=1, le=100_000)
|
||||||
|
connect_timeout_seconds: int = Field(default=5, ge=1, le=30)
|
||||||
|
statement_timeout_seconds: int = Field(default=30, ge=1, le=300)
|
||||||
|
ssl_mode: Literal["disable", "prefer", "require", "verify-ca", "verify-full"] = "prefer"
|
||||||
|
|
||||||
|
|
||||||
|
class ExternalPullRequest(ExternalSourceRequest):
|
||||||
|
query: str | None = Field(default=None, max_length=20_000)
|
||||||
|
file_name: str = Field(default="external-data.jsonl", min_length=1, max_length=255)
|
||||||
|
|
||||||
|
@field_validator("file_name")
|
||||||
|
@classmethod
|
||||||
|
def validate_file_name(cls, value: str) -> str:
|
||||||
|
name = value.strip()
|
||||||
|
if not name.lower().endswith((".jsonl", ".ndjson")):
|
||||||
|
raise ValueError("external pull file_name must end with .jsonl or .ndjson")
|
||||||
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
class ResultUpdate(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
instruction: str | None = None
|
||||||
|
input: str | None = None
|
||||||
|
output: str | None = None
|
||||||
|
chosen: str | None = None
|
||||||
|
rejected: str | None = None
|
||||||
|
expected_updated_at: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class ResultRegenerateRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
expected_updated_at: str = Field(min_length=1, max_length=100)
|
||||||
|
|
||||||
|
|
||||||
|
class ResultBatchRegenerateItem(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
result_id: str = Field(min_length=1, max_length=100)
|
||||||
|
expected_updated_at: str = Field(min_length=1, max_length=100)
|
||||||
|
|
||||||
|
|
||||||
|
class ResultBatchRegenerateRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
items: list[ResultBatchRegenerateItem] = Field(min_length=1, max_length=100)
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_unique_results(self) -> "ResultBatchRegenerateRequest":
|
||||||
|
result_ids = [item.result_id for item in self.items]
|
||||||
|
if len(result_ids) != len(set(result_ids)):
|
||||||
|
raise ValueError("result_id values must be unique")
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class DatasetSplit(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
train: int = Field(default=80, ge=0, le=100)
|
||||||
|
validation: int = Field(default=10, ge=0, le=100)
|
||||||
|
test: int = Field(default=10, ge=0, le=100)
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_total(self) -> "DatasetSplit":
|
||||||
|
if self.train + self.validation + self.test != 100:
|
||||||
|
raise ValueError("dataset split must total 100")
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class PublishRequest(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="forbid")
|
||||||
|
|
||||||
|
dataset_name: str = Field(min_length=1, max_length=150)
|
||||||
|
dataset_type: Literal["train", "test", "eval", "val", "other"] = "train"
|
||||||
|
storage_type: Literal["local"] = "local"
|
||||||
|
split: DatasetSplit = Field(default_factory=DatasetSplit)
|
||||||
|
format: Literal["alpaca_jsonl", "jsonl", "dpo"] = "alpaca_jsonl"
|
||||||
|
description: str = ""
|
||||||
|
|
||||||
|
@field_validator("dataset_name")
|
||||||
|
@classmethod
|
||||||
|
def normalize_dataset_name(cls, value: str) -> str:
|
||||||
|
value = value.strip()
|
||||||
|
if not value:
|
||||||
|
raise ValueError("dataset name cannot be empty")
|
||||||
|
return value
|
||||||
1
backend/app/services/__init__.py
Normal file
1
backend/app/services/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Cross-module services package."""
|
||||||
1
backend/app/workers/__init__.py
Normal file
1
backend/app/workers/__init__.py
Normal file
@@ -0,0 +1 @@
|
|||||||
|
"""Background workers package."""
|
||||||
33
backend/app/workers/compute_poller.py
Normal file
33
backend/app/workers/compute_poller.py
Normal file
@@ -0,0 +1,33 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from app.core.config import get_settings
|
||||||
|
from app.core.logging import get_logger
|
||||||
|
from app.modules.compute_gateway.sync import poll_compute_jobs_once
|
||||||
|
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
async def run_compute_poller() -> None:
|
||||||
|
settings = get_settings()
|
||||||
|
if settings.compute_mode == "simulator" or settings.compute_status_sync_mode != "polling":
|
||||||
|
logger.info("compute poller disabled", extra={"compute_mode": settings.compute_mode})
|
||||||
|
return
|
||||||
|
|
||||||
|
interval = max(3, settings.compute_poll_interval_seconds)
|
||||||
|
logger.info("compute poller started", extra={"interval_seconds": interval})
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
result = await poll_compute_jobs_once()
|
||||||
|
if result["failed"]:
|
||||||
|
logger.warning("compute polling reported failures", extra={"result": result})
|
||||||
|
elif result["synced"]:
|
||||||
|
logger.debug("compute jobs synchronized", extra={"result": result})
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
logger.info("compute poller stopped")
|
||||||
|
raise
|
||||||
|
except Exception as exc: # noqa: BLE001 - keep background polling alive
|
||||||
|
logger.exception("compute poller failed", extra={"error": str(exc)})
|
||||||
|
await asyncio.sleep(interval)
|
||||||
41
backend/pyproject.toml
Normal file
41
backend/pyproject.toml
Normal file
@@ -0,0 +1,41 @@
|
|||||||
|
[project]
|
||||||
|
name = "yg-ft-backend"
|
||||||
|
version = "0.1.0"
|
||||||
|
description = "Backend service for the model fine-tuning platform"
|
||||||
|
requires-python = ">=3.12"
|
||||||
|
dependencies = [
|
||||||
|
"fastapi>=0.111.0",
|
||||||
|
"uvicorn[standard]>=0.30.0",
|
||||||
|
"python-multipart>=0.0.9",
|
||||||
|
"pydantic>=2.7.0",
|
||||||
|
"sqlalchemy>=2.0.30",
|
||||||
|
"psycopg[binary]>=3.2.1",
|
||||||
|
"psycopg-pool>=3.2.1",
|
||||||
|
"alembic>=1.13.1",
|
||||||
|
"redis>=5.0.4",
|
||||||
|
"httpx>=0.27.0",
|
||||||
|
"PyJWT>=2.8.0",
|
||||||
|
"passlib[bcrypt]>=1.7.4",
|
||||||
|
"python-dotenv>=1.0.1",
|
||||||
|
"pypdf[crypto]>=5.0.0",
|
||||||
|
"python-docx>=1.1.2",
|
||||||
|
"openpyxl>=3.1.5",
|
||||||
|
"python-pptx>=1.0.2",
|
||||||
|
"llama-index-core==0.14.23",
|
||||||
|
"llama-index-embeddings-huggingface==0.6.1",
|
||||||
|
"docling==2.115.0",
|
||||||
|
"tiktoken>=0.7.0",
|
||||||
|
]
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = [
|
||||||
|
"pytest>=8.2.0",
|
||||||
|
"ruff>=0.5.0",
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
line-length = 100
|
||||||
|
target-version = "py312"
|
||||||
|
|
||||||
|
[tool.pytest.ini_options]
|
||||||
|
testpaths = ["tests"]
|
||||||
26
backend/requirements.txt
Normal file
26
backend/requirements.txt
Normal file
@@ -0,0 +1,26 @@
|
|||||||
|
fastapi>=0.111.0
|
||||||
|
uvicorn[standard]>=0.30.0
|
||||||
|
python-multipart>=0.0.9
|
||||||
|
pydantic>=2.7.0
|
||||||
|
sqlalchemy>=2.0.30
|
||||||
|
psycopg[binary]>=3.2.1
|
||||||
|
psycopg-pool>=3.2.1
|
||||||
|
alembic>=1.13.1
|
||||||
|
redis>=5.0.4
|
||||||
|
httpx>=0.27.0
|
||||||
|
minio>=7.2.7
|
||||||
|
PyJWT>=2.8.0
|
||||||
|
passlib[bcrypt]>=1.7.4
|
||||||
|
python-dotenv>=1.0.1
|
||||||
|
pypdf[crypto]>=5.0.0
|
||||||
|
python-docx>=1.1.2
|
||||||
|
openpyxl>=3.1.5
|
||||||
|
python-pptx>=1.0.2
|
||||||
|
llama-index-core==0.14.23
|
||||||
|
llama-index-embeddings-huggingface==0.6.1
|
||||||
|
docling==2.115.0
|
||||||
|
tiktoken>=0.7.0
|
||||||
|
|
||||||
|
# 测试与代码检查
|
||||||
|
pytest>=8.2.0
|
||||||
|
ruff>=0.5.0
|
||||||
276
backend/tests/test_compare_inference_async.py
Normal file
276
backend/tests/test_compare_inference_async.py
Normal file
@@ -0,0 +1,276 @@
|
|||||||
|
"""
|
||||||
|
模型推理异步加载改造的单元测试。
|
||||||
|
|
||||||
|
覆盖:
|
||||||
|
- model_compare_load:异步派发,立即返回 starting + 节点信息(不等待加载完成)
|
||||||
|
- model_compare_delete:先删记录,卸载失败也不阻塞删除
|
||||||
|
- reconcile_inference_loads:starting -> ready/error/idle/不可达的状态迁移与封顶
|
||||||
|
- _unload_from_compute_node:任务感知,只命中记录中的节点
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import model_compare_delete, model_compare_load
|
||||||
|
import app.api.v1.endpoints.platform as platform
|
||||||
|
from app.modules.compute_gateway.client import ComputeNodeClient
|
||||||
|
from app.modules.compute_gateway.sync import MAX_STARTING_ATTEMPTS, reconcile_inference_loads
|
||||||
|
|
||||||
|
|
||||||
|
class FakeInferenceStore:
|
||||||
|
"""内存 store,仅实现推理加载/对账用到的接口。"""
|
||||||
|
|
||||||
|
def __init__(self, tasks: list[dict[str, Any]] | None = None, nodes: list[dict[str, Any]] | None = None) -> None:
|
||||||
|
self._tasks: dict[str, dict[str, Any]] = {t["id"]: dict(t) for t in (tasks or [])}
|
||||||
|
self._nodes = nodes or []
|
||||||
|
self._inference_nodes: set[str] = set()
|
||||||
|
|
||||||
|
def compare_task(self, task_id: str) -> dict[str, Any]:
|
||||||
|
if task_id not in self._tasks:
|
||||||
|
raise KeyError(task_id)
|
||||||
|
return dict(self._tasks[task_id])
|
||||||
|
|
||||||
|
def compare_tasks(self) -> list[dict[str, Any]]:
|
||||||
|
return [dict(t) for t in self._tasks.values()]
|
||||||
|
|
||||||
|
def update_compare_task(self, task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
current = self._tasks[task_id]
|
||||||
|
merged = {**current, **payload, "id": task_id}
|
||||||
|
self._tasks[task_id] = merged
|
||||||
|
return dict(merged)
|
||||||
|
|
||||||
|
def delete_compare_task(self, task_id: str) -> None:
|
||||||
|
self._tasks.pop(task_id, None)
|
||||||
|
|
||||||
|
def compute_nodes(self) -> list[dict[str, Any]]:
|
||||||
|
return [dict(n) for n in self._nodes]
|
||||||
|
|
||||||
|
def model(self, model_id: str) -> dict[str, Any]:
|
||||||
|
raise KeyError(model_id)
|
||||||
|
|
||||||
|
def trained_models(self) -> list[dict[str, Any]]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
def mark_inference_loaded(self, node_id: str) -> None:
|
||||||
|
self._inference_nodes.add(node_id)
|
||||||
|
|
||||||
|
def mark_inference_unloaded(self, node_id: str) -> None:
|
||||||
|
self._inference_nodes.discard(node_id)
|
||||||
|
|
||||||
|
def is_inference_loaded(self, node_id: str) -> bool:
|
||||||
|
return node_id in self._inference_nodes
|
||||||
|
|
||||||
|
|
||||||
|
def _node(node_id: str, code: str = "") -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": node_id,
|
||||||
|
"code": code or node_id,
|
||||||
|
"name": code or node_id,
|
||||||
|
"api_base_url": f"http://{code or node_id}:19100",
|
||||||
|
"enabled": True,
|
||||||
|
"scheduler_status": "online",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _task(task_id: str, *, node_id: str | None = None, load_status: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": task_id,
|
||||||
|
"name": f"task-{task_id}",
|
||||||
|
"status": "pending",
|
||||||
|
"models": [
|
||||||
|
{"model_id": "m_1", "model_name": "qwen", "model_path": "/models/qwen", "node_id": node_id}
|
||||||
|
],
|
||||||
|
"load_status": load_status or {"loaded_models": []},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _fake_inference_load(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
return {"loaded": False, "status": "loading", "request_id": "req-1"}
|
||||||
|
|
||||||
|
|
||||||
|
async def _fake_inference_unload(self) -> dict[str, Any]:
|
||||||
|
return {"unloaded": True, "status": "idle"}
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_store(monkeypatch, store: FakeInferenceStore) -> None:
|
||||||
|
monkeypatch.setattr(platform, "get_platform_store", lambda: store)
|
||||||
|
monkeypatch.setattr(platform, "get_settings", lambda: SimpleNamespace(compute_mode="real"))
|
||||||
|
|
||||||
|
|
||||||
|
def test_select_eval_node_prefers_model_node(monkeypatch) -> None:
|
||||||
|
from app.api.v1.endpoints.platform import _select_eval_node
|
||||||
|
|
||||||
|
store = FakeInferenceStore(nodes=[_node("n1"), _node("n2")])
|
||||||
|
# 指定模型所在节点时优先返回该节点
|
||||||
|
assert _select_eval_node(store, "n2")["id"] == "n2"
|
||||||
|
# 无指定节点时回退到第一个在线节点
|
||||||
|
assert _select_eval_node(store, None)["id"] == "n1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_select_eval_node_returns_none_when_model_node_offline(monkeypatch) -> None:
|
||||||
|
from app.api.v1.endpoints.platform import _select_eval_node
|
||||||
|
|
||||||
|
nodes = [_node("n1"), _node("n2")]
|
||||||
|
nodes[1]["enabled"] = False
|
||||||
|
store = FakeInferenceStore(nodes=nodes)
|
||||||
|
# 模型所在节点不可用 → 明确失败,不派发到其它节点
|
||||||
|
assert _select_eval_node(store, "n2") is None
|
||||||
|
# 无指定节点时仍回退第一个在线节点
|
||||||
|
assert _select_eval_node(store, None)["id"] == "n1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_model_compare_load_dispatches_and_returns_starting(monkeypatch) -> None:
|
||||||
|
store = FakeInferenceStore(tasks=[_task("t1", node_id="n1")], nodes=[_node("n1")])
|
||||||
|
_patch_store(monkeypatch, store)
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_load", _fake_inference_load)
|
||||||
|
|
||||||
|
result = asyncio.run(model_compare_load("t1"))
|
||||||
|
assert result["code"] == 0
|
||||||
|
updated = result["data"]
|
||||||
|
assert updated["status"] == "starting"
|
||||||
|
items = updated["load_status"]["loaded_models"]
|
||||||
|
assert items[0]["status"] == "starting"
|
||||||
|
assert items[0]["node_id"] == "n1"
|
||||||
|
assert "n1" in store._inference_nodes
|
||||||
|
|
||||||
|
|
||||||
|
def test_model_compare_load_marks_error_when_all_nodes_fail(monkeypatch) -> None:
|
||||||
|
store = FakeInferenceStore(tasks=[_task("t1", node_id="n1")], nodes=[_node("n1")])
|
||||||
|
_patch_store(monkeypatch, store)
|
||||||
|
|
||||||
|
async def _raise(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
raise RuntimeError("conn refused")
|
||||||
|
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_load", _raise)
|
||||||
|
|
||||||
|
result = asyncio.run(model_compare_load("t1"))
|
||||||
|
updated = result["data"]
|
||||||
|
assert updated["status"] == "failed"
|
||||||
|
assert updated["load_status"]["loaded_models"][0]["status"] == "error"
|
||||||
|
assert "conn refused" in updated["load_status"]["loaded_models"][0]["error"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_model_compare_delete_removes_record_even_if_unload_raises(monkeypatch) -> None:
|
||||||
|
task = _task(
|
||||||
|
"t1",
|
||||||
|
node_id="n1",
|
||||||
|
load_status={"loaded_models": [{"model_id": "m_1", "status": "ready", "node_id": "n1"}]},
|
||||||
|
)
|
||||||
|
store = FakeInferenceStore(tasks=[task], nodes=[_node("n1")])
|
||||||
|
_patch_store(monkeypatch, store)
|
||||||
|
|
||||||
|
async def _raise(self) -> dict[str, Any]:
|
||||||
|
raise RuntimeError("unload boom")
|
||||||
|
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_unload", _raise)
|
||||||
|
|
||||||
|
result = asyncio.run(model_compare_delete("t1"))
|
||||||
|
assert result["data"] == {"deleted": "t1"}
|
||||||
|
assert "t1" not in store._tasks
|
||||||
|
# finally 中仍清掉了节点标记
|
||||||
|
assert "n1" not in store._inference_nodes
|
||||||
|
|
||||||
|
|
||||||
|
def test_unload_from_compute_node_only_hits_recorded_node(monkeypatch) -> None:
|
||||||
|
task = _task(
|
||||||
|
"t1",
|
||||||
|
load_status={"loaded_models": [{"model_id": "m_1", "status": "ready", "node_id": "n1"}]},
|
||||||
|
)
|
||||||
|
store = FakeInferenceStore(tasks=[task], nodes=[_node("n1"), _node("n2")])
|
||||||
|
_patch_store(monkeypatch, store)
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_unload", _fake_inference_unload)
|
||||||
|
|
||||||
|
from app.api.v1.endpoints.platform import _unload_from_compute_node
|
||||||
|
|
||||||
|
result = asyncio.run(_unload_from_compute_node(store, task=task))
|
||||||
|
assert result["unloaded"] is True
|
||||||
|
# 只命中任务记录中的节点 n1,n2 未被卸载
|
||||||
|
assert [r["node_id"] for r in result["nodes"]] == ["n1"]
|
||||||
|
assert "n1" not in store._inference_nodes
|
||||||
|
|
||||||
|
|
||||||
|
async def _status_ready(self) -> dict[str, Any]:
|
||||||
|
return {"loaded": True, "status": "ready", "model_name": "qwen"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_reconcile_transitions_starting_to_ready(monkeypatch) -> None:
|
||||||
|
task = _task(
|
||||||
|
"t1",
|
||||||
|
node_id="n1",
|
||||||
|
load_status={"loaded_models": [{"model_id": "m_1", "status": "starting", "node_id": "n1"}]},
|
||||||
|
)
|
||||||
|
store = FakeInferenceStore(tasks=[task], nodes=[_node("n1")])
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_status", _status_ready)
|
||||||
|
|
||||||
|
reconciled = asyncio.run(reconcile_inference_loads(store))
|
||||||
|
assert reconciled == [{"task_id": "t1", "status": "loaded"}]
|
||||||
|
updated = store._tasks["t1"]
|
||||||
|
assert updated["status"] == "loaded"
|
||||||
|
assert updated["load_status"]["loaded_models"][0]["status"] == "ready"
|
||||||
|
assert "n1" in store._inference_nodes
|
||||||
|
|
||||||
|
|
||||||
|
def test_reconcile_transitions_to_error_and_failed(monkeypatch) -> None:
|
||||||
|
async def _status_error(self) -> dict[str, Any]:
|
||||||
|
return {"loaded": False, "status": "error", "error": "CUDA out of memory"}
|
||||||
|
|
||||||
|
task = _task(
|
||||||
|
"t1",
|
||||||
|
node_id="n1",
|
||||||
|
load_status={"loaded_models": [{"model_id": "m_1", "status": "starting", "node_id": "n1"}]},
|
||||||
|
)
|
||||||
|
store = FakeInferenceStore(tasks=[task], nodes=[_node("n1")])
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_status", _status_error)
|
||||||
|
|
||||||
|
reconciled = asyncio.run(reconcile_inference_loads(store))
|
||||||
|
assert reconciled == [{"task_id": "t1", "status": "failed"}]
|
||||||
|
item = store._tasks["t1"]["load_status"]["loaded_models"][0]
|
||||||
|
assert item["status"] == "error"
|
||||||
|
assert "CUDA out of memory" in item["error"]
|
||||||
|
assert "n1" not in store._inference_nodes
|
||||||
|
|
||||||
|
|
||||||
|
def test_reconcile_idle_marks_model_disappeared(monkeypatch) -> None:
|
||||||
|
async def _status_idle(self) -> dict[str, Any]:
|
||||||
|
return {"loaded": False, "status": "idle"}
|
||||||
|
|
||||||
|
task = _task(
|
||||||
|
"t1",
|
||||||
|
node_id="n1",
|
||||||
|
load_status={"loaded_models": [{"model_id": "m_1", "status": "starting", "node_id": "n1"}]},
|
||||||
|
)
|
||||||
|
store = FakeInferenceStore(tasks=[task], nodes=[_node("n1")])
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_status", _status_idle)
|
||||||
|
|
||||||
|
asyncio.run(reconcile_inference_loads(store))
|
||||||
|
item = store._tasks["t1"]["load_status"]["loaded_models"][0]
|
||||||
|
assert item["status"] == "error"
|
||||||
|
assert "disappeared" in item["error"]
|
||||||
|
assert store._tasks["t1"]["status"] == "failed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_reconcile_unreachable_node_flips_to_error_after_cap(monkeypatch) -> None:
|
||||||
|
async def _raise(self) -> dict[str, Any]:
|
||||||
|
raise RuntimeError("conn refused")
|
||||||
|
|
||||||
|
task = _task(
|
||||||
|
"t1",
|
||||||
|
node_id="n1",
|
||||||
|
load_status={"loaded_models": [{"model_id": "m_1", "status": "starting", "node_id": "n1"}]},
|
||||||
|
)
|
||||||
|
store = FakeInferenceStore(tasks=[task], nodes=[_node("n1")])
|
||||||
|
monkeypatch.setattr(ComputeNodeClient, "inference_status", _raise)
|
||||||
|
|
||||||
|
# 每次轮询前重置节流时间戳,逐次推进 load_attempts 到封顶
|
||||||
|
for _ in range(MAX_STARTING_ATTEMPTS):
|
||||||
|
item = store._tasks["t1"]["load_status"]["loaded_models"][0]
|
||||||
|
item["last_polled_at"] = 0
|
||||||
|
asyncio.run(reconcile_inference_loads(store))
|
||||||
|
|
||||||
|
item = store._tasks["t1"]["load_status"]["loaded_models"][0]
|
||||||
|
assert item["status"] == "error"
|
||||||
|
assert "unreachable" in item["error"]
|
||||||
|
assert store._tasks["t1"]["status"] == "failed"
|
||||||
57
backend/tests/test_data_convert_security.py
Normal file
57
backend/tests/test_data_convert_security.py
Normal file
@@ -0,0 +1,57 @@
|
|||||||
|
"""data_convert 模块安全回归测试:输出文件名路径穿越与鉴权。
|
||||||
|
|
||||||
|
- ``output_filename`` 必须通过白名单校验,阻断 ``../``、``/``、``\\`` 及控制字符,
|
||||||
|
否则转换结果可被写出到存储根目录之外(任意文件读写/删除)。
|
||||||
|
- 所有 data_convert 路由必须挂载 ``get_current_user`` 鉴权依赖。
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from app.core.auth import get_current_user
|
||||||
|
from app.modules.data_convert.router import _safe_output_filename, router
|
||||||
|
|
||||||
|
|
||||||
|
def test_safe_output_filename_defaults() -> None:
|
||||||
|
assert _safe_output_filename(None) == "converted-data.jsonl"
|
||||||
|
assert _safe_output_filename("") == "converted-data.jsonl"
|
||||||
|
|
||||||
|
|
||||||
|
def test_safe_output_filename_valid() -> None:
|
||||||
|
assert _safe_output_filename("converted-data.jsonl") == "converted-data.jsonl"
|
||||||
|
assert _safe_output_filename("my-data.v1.jsonl") == "my-data.v1.jsonl"
|
||||||
|
assert _safe_output_filename(" 报告.jsonl ") == "报告.jsonl"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"bad",
|
||||||
|
[
|
||||||
|
"../../etc/passwd",
|
||||||
|
"../x.jsonl",
|
||||||
|
"a/b.jsonl",
|
||||||
|
r"a\b.jsonl",
|
||||||
|
"a\\b.jsonl",
|
||||||
|
"..",
|
||||||
|
".",
|
||||||
|
"x\x00.jsonl",
|
||||||
|
"x\n.jsonl",
|
||||||
|
"x\t.jsonl",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_safe_output_filename_rejects_traversal(bad: str) -> None:
|
||||||
|
with pytest.raises(HTTPException):
|
||||||
|
_safe_output_filename(bad)
|
||||||
|
|
||||||
|
|
||||||
|
def test_all_data_convert_routes_require_auth() -> None:
|
||||||
|
for route in router.routes:
|
||||||
|
node = getattr(route, "dependant", None)
|
||||||
|
assert node is not None, f"route {route.path} has no dependency graph"
|
||||||
|
stack = list(node.dependencies)
|
||||||
|
calls: list = []
|
||||||
|
while stack:
|
||||||
|
dep = stack.pop()
|
||||||
|
stack.extend(getattr(dep, "dependencies", []))
|
||||||
|
calls.append(getattr(dep, "call", None))
|
||||||
|
assert get_current_user in calls, f"route {route.path} is missing get_current_user auth"
|
||||||
1074
backend/tests/test_data_process_algorithms.py
Normal file
1074
backend/tests/test_data_process_algorithms.py
Normal file
File diff suppressed because it is too large
Load Diff
3528
backend/tests/test_data_process_api.py
Normal file
3528
backend/tests/test_data_process_api.py
Normal file
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user