95 Commits

Author SHA1 Message Date
caoxiaozhu
e486d36a80 feat(dataset): 支持数据任务批量删除 2026-07-27 11:36:30 +08:00
caoxiaozhu
d6e325fe9e fix(data-process): 按发布状态控制结果编辑 2026-07-27 11:23:17 +08:00
caoxiaozhu
b08a771a61 fix(data-process): 保留可验证的任务耗时 2026-07-27 11:16:48 +08:00
caoxiaozhu
8caaaa5bbc fix(data-process): 恢复提前中断的重新生成任务 2026-07-27 11:07:18 +08:00
caoxiaozhu
3f5fedb9ed build(frontend): 更新重新生成流程构建 2026-07-27 10:43:53 +08:00
caoxiaozhu
53844a3a09 fix(frontend): 保持重新生成退出前详情 2026-07-27 10:43:48 +08:00
caoxiaozhu
53014bb381 fix(data-process): 延迟重新生成破坏性变更 2026-07-27 10:43:42 +08:00
caoxiaozhu
03bd0b6d03 build(frontend): 更新数据任务详情构建 2026-07-27 10:03:58 +08:00
caoxiaozhu
b4927a8952 fix(frontend): 对齐数据任务详情状态 2026-07-27 10:03:52 +08:00
caoxiaozhu
b14b2ecf22 fix(data-process): 修正任务详情数据契约 2026-07-27 10:03:46 +08:00
caoxiaozhu
42c0e4f5c2 build(frontend): 更新数据任务列表构建 2026-07-27 09:50:33 +08:00
caoxiaozhu
88a82ed771 feat(frontend): 完善数据任务数量与状态 2026-07-27 09:50:26 +08:00
caoxiaozhu
895983ac20 feat(data-process): 返回任务文档数量 2026-07-27 09:50:19 +08:00
caoxiaozhu
680fa905f8 build(frontend): 更新数据集保留提示构建 2026-07-27 09:39:29 +08:00
caoxiaozhu
ce0f908d20 fix(frontend): 标明重新生成前的数据集状态 2026-07-27 09:39:23 +08:00
caoxiaozhu
e4ea1f168c fix(data-process): 保留重新生成前的已发布数据集 2026-07-27 09:39:19 +08:00
caoxiaozhu
915f994c45 build(frontend): 更新五十条生成上限构建 2026-07-27 09:12:03 +08:00
caoxiaozhu
762f866175 feat(data-process): 放宽问答生成数量上限 2026-07-27 09:11:58 +08:00
caoxiaozhu
9428c6b785 feat(data-process): 支持单项生成五十条数据 2026-07-27 09:11:51 +08:00
caoxiaozhu
25d75f40c7 build(frontend): 更新重新生成流程构建 2026-07-25 22:41:16 +08:00
caoxiaozhu
06e7455630 feat(data-process): 接入重新生成配置流程 2026-07-25 22:41:06 +08:00
caoxiaozhu
396d3f6f47 feat(data-process): 支持任务重新生成 2026-07-25 22:40:55 +08:00
caoxiaozhu
64d7414b04 build(frontend): 更新跳转图标对齐构建 2026-07-25 22:00:23 +08:00
caoxiaozhu
9a5282f39c fix(data-process): 对齐数据集跳转图标 2026-07-25 22:00:23 +08:00
caoxiaozhu
f21a4c954f build(frontend): 更新数据集名称对齐构建 2026-07-25 21:59:02 +08:00
caoxiaozhu
2e2cbb4976 fix(data-process): 对齐数据集名称起点 2026-07-25 21:59:02 +08:00
caoxiaozhu
2e8278636b build(frontend): 更新数据集对齐修复构建 2026-07-25 21:55:24 +08:00
caoxiaozhu
bbc0df29bf fix(data-process): 对齐输出数据集列表 2026-07-25 21:55:16 +08:00
caoxiaozhu
b20e7aa595 build(frontend): 更新结果提示修复构建 2026-07-25 18:25:58 +08:00
caoxiaozhu
17615aa17d fix(data-process): 限制结果提示浮层尺寸 2026-07-25 18:25:34 +08:00
caoxiaozhu
07e2999323 fix(data-process): 修复三数据集发布参数错误 2026-07-25 18:19:37 +08:00
caoxiaozhu
749c84a62b build(frontend): 更新会话续期修复构建 2026-07-25 18:16:02 +08:00
caoxiaozhu
ea0013b99c fix(auth): 长任务活跃期间续期会话 2026-07-25 18:15:49 +08:00
caoxiaozhu
4f8aff5fc4 build(frontend): 更新文档切分页面构建 2026-07-25 18:00:55 +08:00
caoxiaozhu
9193f10e3e feat(frontend): 更新三种切分模式与PDF定位 2026-07-25 18:00:44 +08:00
caoxiaozhu
ea08478a37 feat(data-process): 接入三种文档切分引擎 2026-07-25 18:00:21 +08:00
caoxiaozhu
4782981169 build(frontend): 更新三数据集发布页面构建 2026-07-25 17:04:27 +08:00
caoxiaozhu
939a7f8e8f fix(frontend): 展示三个独立发布数据集 2026-07-25 17:04:21 +08:00
caoxiaozhu
e9a121cfeb fix(data-process): 发布三个独立切分数据集 2026-07-25 17:04:14 +08:00
caoxiaozhu
d4b9a76aa5 build(frontend): 更新数据切分页面构建 2026-07-24 21:11:06 +08:00
caoxiaozhu
9114f3d4c7 fix(dataset): 展示并切换三路数据切分 2026-07-24 21:11:01 +08:00
caoxiaozhu
9cb77c251a fix(data-process): 发布精确三路数据切分 2026-07-24 20:43:47 +08:00
caoxiaozhu
e6a5a36bc0 build(frontend): 更新数据处理详情构建 2026-07-24 16:31:38 +08:00
caoxiaozhu
b2c570f607 fix(data-process): 补齐原文参照与详情统计 2026-07-24 16:30:29 +08:00
caoxiaozhu
994ec6644a fix(data-process): 修正详情统计字段契约 2026-07-24 16:28:47 +08:00
caoxiaozhu
d6d3d27b2d chore: 恢复 .gitignore 对 .env 文件的忽略规则
上次合并 rebase 时远端引入的 !.env 规则会强制追踪环境文件,
存在将含敏感信息的 .env 误提交的风险。恢复为 .env 忽略规则,
同时保留 !.env.example 以便示例文件继续入库。
2026-07-24 16:13:04 +08:00
caoxiaozhu
eb6ff93150 build(frontend): 更新智能预处理生产构建 2026-07-24 15:10:48 +08:00
caoxiaozhu
215b4074e0 build(frontend): 同步PDF预览依赖锁文件 2026-07-24 15:06:37 +08:00
caoxiaozhu
4544483fc5 feat(data-process): 启用智能文档清理预览 2026-07-24 15:06:01 +08:00
caoxiaozhu
3266a6fc09 feat(data-process): 清理PDF文档级噪声 2026-07-24 15:05:39 +08:00
caoxiaozhu
a9b06140d0 docs: 补充评估工作台设计文档
从根目录评测.txt 迁入 docs/,归档评估流程、错误诊断、
报告生成与建议有效性验证的设计说明。
2026-07-24 14:36:11 +08:00
caoxiaozhu
663b73af2e feat(frontend): 新增登录页 hero 流程图资源 2026-07-24 14:36:11 +08:00
caoxiaozhu
4a2f1f5dcd build(frontend): 纳管 pnpm 锁文件与 workspace 配置
pnpm-lock.yaml 锁定依赖版本以保证可复现安装;
pnpm-workspace.yaml 声明构建脚本白名单配置。
2026-07-24 14:36:11 +08:00
caoxiaozhu
8ac39cf007 chore: 移除已废弃的 .env.example 示例文件
配置入口已迁移至 backend/config.yaml(本地,已 gitignore)
与各环境实际的 .env 文件,示例文件不再需要维护。
2026-07-24 14:36:11 +08:00
caoxiaozhu
faad88dfcd chore: 忽略本地敏感配置与工具产物
- backend/config.yaml 含数据库账号密码等敏感信息
- .codex-backups/、.pnpm-store/、.zcode/ 为工具产物,不应进版本库
2026-07-24 14:36:11 +08:00
caoxiaozhu
d3a25f3a4b build(frontend): 更新切分缓存版本构建 2026-07-24 14:36:11 +08:00
caoxiaozhu
b801bd314b fix(data-process): 升级预览切分缓存版本 2026-07-24 14:36:11 +08:00
caoxiaozhu
476502fc0d build(frontend): 更新数据预览生产构建 2026-07-24 14:36:11 +08:00
caoxiaozhu
975f55d06c fix(data-process): 优化预览分页与PDF兼容 2026-07-24 14:36:11 +08:00
caoxiaozhu
93373bc61f fix(data-process): 合并文档结构短切片 2026-07-24 14:36:11 +08:00
caoxiaozhu
b79a8e1499 build(frontend): 更新生产构建产物 2026-07-24 14:36:11 +08:00
caoxiaozhu
0124e28d77 docs(data-process): 更新文件处理与切分设计 2026-07-24 14:35:03 +08:00
caoxiaozhu
ad64e44860 feat(data-process): 优化上传切分流程与PDF高亮预览 2026-07-24 14:35:03 +08:00
caoxiaozhu
33d0ed2e01 feat(data-process): 完善文件解析与切分存储链路 2026-07-24 14:35:03 +08:00
wuyongtao
6d4bf85284 feat: 提交 Docker 环境配置文件,更新 .gitignore 规则
- 新增 docker/app/.env 应用环境配置
- 新增 docker/compute/.env 计算节点环境配置
- 更新 .gitignore 允许提交 .env 配置文件

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-24 10:45:10 +08:00
wuyongtao
2b10c013ce feat: 更新后端平台模块、前端组件及构建产物,新增工作计划文档
- 更新 backend 平台 API endpoints 及 platform_store
- 更新前端 ComputeNodesView、DataProcessCreateView、FineTuneCreateView 等组件
- 更新前端 API 模块(compute、fineTune)
- 重构 frontend/dist 构建产物(新 hash)
- 新增 docs/2026-07-24-work-plan.md 工作计划文档

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-24 10:27:52 +08:00
wuyongtao
b28cfbc6fa feat: 更新后端平台模块、Compute引擎、前端组件及构建产物
- 更新 backend 平台 API、platform_store、compute_gateway sync
- 更新 compute agent/engine/adapter 及 API
- 更新 Docker 部署配置(app/compute)
- 新增 frontend/src/utils/ 工具模块
- 新增 scripts/ops_diagnostics.py 运维诊断脚本
- 新增 docs/2026-07-23-development-summary.md 开发总结
- 重构 frontend/dist 构建产物(新 hash)
- 更新前端多个视图组件及 API 模块

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-23 19:32:42 +08:00
caoxiaozhu
f04dc479bb feat: 完成数据处理接口与前端接入 2026-07-23 15:10:13 +08:00
caoxiaozhu
f453234057 feat: 增加前后端一键启动脚本 2026-07-23 11:09:04 +08:00
wuyongtao
a6868ec2e5 fix: 恢复 gitignore 全局 logs/ 规则,通过 !docker/compute/data/yg-ft/logs/ 放行
Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-22 20:50:37 +08:00
wuyongtao
6cd1e46e86 feat: 添加 compute 数据目录结构,配置 gitignore 规则
- docker/compute/data/yg-ft/ 下创建 models、datasets、outputs、logs 目录
- 保留目录结构和 README,忽略子目录实际内容(日志、模型、数据集等)
- 修正 gitignore 全局 logs/ 规则为 /logs/,避免误伤

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-22 20:47:05 +08:00
wuyongtao
836343b29e feat: 新增 compute_gateway、compute_poller、agent 模块,重构前端 dist
- 新增 backend/app/modules/compute_gateway(client/sync)计算网关模块
- 新增 backend/app/workers/compute_poller 计算轮询 worker
- 新增 compute/agent/process_manager 进程管理器
- 新增 scripts/ 脚本目录
- 更新 Docker 部署配置(app/compute/nginx)
- 更新后端平台 API、数据库 SQL、core 配置
- 更新前端多个视图组件及 API 模块
- 重构 frontend/dist 构建产物(新 hash)
- 更新多项文档

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-22 20:26:16 +08:00
wuyongtao
1e438164c1 fix: 更新 Dockerfile.backend 配置
Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 12:45:47 +08:00
wuyongtao
f4864fafd0 docs: 更新 README.md
Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 12:41:30 +08:00
wuyongtao
9798b34717 feat: 重构前端 dist 构建产物,更新 Docker 配置及文档
- 重新构建 frontend/dist(新版 hash 替换旧版)
- 更新 docker 前端 Dockerfile 及 docker-compose 配置
- 新增 docs/team-development-plan.md 团队开发计划文档
- 更新 README、系统开发计划等文档
- 更新 LoginView 组件

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 12:36:33 +08:00
wuyongtao
284995d79c chore: 更新 Docker 部署配置及文档
Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 11:36:43 +08:00
wuyongtao
e18a367abb feat: 添加前端构建产物 frontend/dist,更新 gitignore 规则
Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 11:11:04 +08:00
wuyongtao
817d13c8f7 chore: 更新后端配置、Docker部署及前端API请求配置
Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 11:06:34 +08:00
wuyongtao
a72b8f1e4b feat: 更新后端平台模块、数据库、Compute引擎及多项配置文档
- 更新 backend 平台 API、platform_store、session 数据库模块
- 新增 backend SQL 初始化脚本
- 更新 compute 引擎适配器及 README
- 更新 Docker 部署配置(app/compute)
- 更新前端入口、环境类型声明及 README
- 新增 docs/menu-functional-requirements.md 菜单功能需求文档
- 更新多项项目文档

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 10:55:44 +08:00
wuyongtao
bccd3bf448 feat: 更新后端配置、Docker部署、API模块及多项文档
- 更新后端 main.py、config.py 核心配置
- 更新 compute API 模块
- 更新 Docker 部署配置(app/compute docker-compose、nginx、环境变量)
- 更新前端 API 模块(dataset、model、request)及 vite 配置
- 更新多项项目文档(架构、部署、开发计划、日志等)

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 10:09:36 +08:00
wuyongtao
a67ca2c19c feat: 添加平台管理、计算模块适配器及前端页面更新
- 新增 platform API 端点和存储
- 新增 llama_factory 适配器
- 新增前端 compute、guide、system 等视图页面
- 新增 echarts 插件和 mock 数据
- 更新 Docker 配置、后端配置及文档
- 更新前端路由、API、侧边栏等组件

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-21 09:23:43 +08:00
wuyongtao
2c1e08a271 feat: 重构 Docker 配置结构,添加 compute 模块及新增文档
- 将 Dockerfile 和 docker-compose.yml 迁移至 docker/ 目录下统一管理
- 新增 compute 计算模块(API 入口、依赖配置)
- 新增 docker/app 和 docker/compute 部署配置
- 新增 demo-development-plan.md 演示开发计划文档
- 更新后端 API 设计、部署计划、架构需求等文档
- 更新 postgres 数据库 schema

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-20 14:59:31 +08:00
wuyongtao
ba4059fe3b feat: 添加后端架构、计算模块及部署文档 2026-07-16 13:47:37 +08:00
wuyongtao
4050c120d5 feat: 添加 Docker 支持与项目文档
- 添加 Dockerfile, docker-compose.yml, .dockerignore, nginx 配置
- 添加后端 API 设计文档、平台架构需求文档、系统开发计划
- 添加 PostgreSQL schema 设计
- 更新 README.md 和 design-qa.md
2026-07-16 11:52:05 +08:00
156a952b47 Merge pull request 'dev' (#4) from dev into main
Reviewed-on: #4
2026-07-16 11:04:35 +08:00
caoxiaozhu
4173b53b1b refactor: 前端构建按需化与 Mock 懒加载
移除全量 Element Plus 与全局 VChart 注册,改为按需引入样式与组件内局部图表;Mock 适配器改为按 VITE_ENABLE_MOCK 环境变量懒加载,生产默认不拦截请求;ECharts 精简为看板所需图表,各视图与 stores 同步适配。
2026-07-16 11:03:54 +08:00
caoxiaozhu
ab9e87f948 refactor: 推理对比超时与打字机展示改进
抽取 withTimeout 替代 Promise.race 超时控制,对比结果新增打字机逐字渲染与清理,推理聊天参数与列表类型同步收敛。
2026-07-16 11:03:25 +08:00
caoxiaozhu
5a040366da refactor: usePolling 改为串行轮询
上一轮完成后才安排下一轮避免慢请求重叠,页面不可见时自动暂停、恢复后立即刷新,支持动态间隔与错误回调。
2026-07-16 11:03:15 +08:00
8789019db2 Merge pull request 'dev' (#3) from dev into main
Reviewed-on: #3
2026-07-16 10:18:12 +08:00
caoxiaozhu
3cb20a4a28 style: 登录时长图表标签回到柱条右侧
标签位置由 insideRight 改回 right、颜色改回灰色,避免短柱条内文字溢出。
2026-07-16 09:38:23 +08:00
caoxiaozhu
a6085a2612 chore: 页面表面回归与数据处理文案适配
页面表面回归选择器更新为 task-overview,任务概览补充表面背景变量,数据处理结果步骤标题文案精简。
2026-07-16 09:36:39 +08:00
caoxiaozhu
4899bc8779 style: Dashboard 看板表格行高与图表微调
训练任务表格行高加大提升可读性,登录时长图表留白与标签位置调整避免溢出,回归断言同步。
2026-07-16 09:36:39 +08:00
caoxiaozhu
e70538e64d feat: 顶部栏新增使用文档入口
AppHeader 增加使用文档按钮,新窗口打开 /guide 路由,补充响应式窄屏适配。
2026-07-16 09:36:39 +08:00
39a5390ecd Merge pull request 'dev' (#2) from dev into main
Reviewed-on: #2
2026-07-14 16:20:08 +08:00
cd354f52e6 Merge pull request 'dev' (#1) from dev into main
Reviewed-on: #1
2026-07-14 11:26:14 +08:00
501 changed files with 39380 additions and 2046 deletions

11
.dockerignore Normal file
View 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

34
.gitignore vendored
View File

@@ -12,6 +12,10 @@ __pycache__/
build/ build/
develop-eggs/ develop-eggs/
dist/ dist/
!frontend/dist/
!frontend/dist/**
node_modules/
*.tsbuildinfo
downloads/ downloads/
eggs/ eggs/
.eggs/ .eggs/
@@ -37,6 +41,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 +144,7 @@ celerybeat.pid
# Environments # Environments
.env .env
!.env.example
.venv .venv
env/ env/
venv/ venv/
@@ -137,6 +152,14 @@ ENV/
env.bak/ env.bak/
venv.bak/ venv.bak/
# Local backend config (含数据库账号密码等敏感信息,勿提交)
backend/config.yaml
# Agent / IDE 工具产物,不应进版本库
.codex-backups/
.pnpm-store/
.zcode/
# Spyder project settings # Spyder project settings
.spyderproject .spyderproject
.spyproject .spyproject
@@ -174,3 +197,14 @@ 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

271
README.md
View File

@@ -1,133 +1,198 @@
# YG_FT # YG_FT 模型微调平台
远光微调平台 - 面向大语言模型微调、评测、推理与对比一体化前端 YG_FT 是一个面向企业治理场景的模型微调平台,覆盖用户中心、多租户、项目隔离、数据集管理、模型管理、训练任务、评测、推理、审批流、审计留存、算力调度和训练引擎适配
## 技术栈 当前前端已有基础页面,后端与算力平台已按多人协作开发方式建立工程骨架,并开始实现正式系统主链路能力。当前代码和 SQL 均作为后续生产演进基线维护,不再以一次性演示或静态 Mock 为开发准则。
| 类别 | 技术 | 版本 | ## 总体架构
|------|------|------|
| 框架 | Vue 3 | ^3.5.13 |
| 语言 | TypeScript | ~5.7.2 |
| 构建工具 | Vite | ^6.0.7 |
| 路由 | Vue Router | ^4.5.0 |
| 状态管理 | Pinia | ^2.3.0 |
| UI 组件库 | Element Plus | ^2.9.1 |
| 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 ```text
YG_FT/
frontend/ # 前端控制台
backend/ # FastAPI 应用平台后端
app/
api/v1/ # 对前端暴露的 REST API
core/ # 配置、日志、中间件、权限等基础能力
db/ # 数据库连接、迁移、事务工具
modules/ # 业务模块目录
schemas/ # Pydantic 入参/出参模型
services/ # 跨模块应用服务
workers/ # 后台任务入口
requirements.txt # 后端 Python 第三方依赖
compute/ # 算力平台与训练框架适配层
api/ # 内部 Compute API
agent/ # 单机多 GPU 调度与进程管理
engines/llama_factory/ # LLaMA-Factory 适配器
file_gateway/ # 本地文件上传、下载、导入、产物管理
docs/ # 需求、接口、数据库、开发计划和部署文档
docker/ # 容器化配置
```
## 环境要求 ## 平台分层
- **Node.js** >= 18推荐 20 LTS | 层级 | 职责 | 主要目录 |
- **npm** >= 9 | --- | --- | --- |
- 后端服务运行于 `http://localhost:7861`(前端通过代理转发,见下文) | 前端控制台 | 用户操作入口、任务看板、项目/模型/数据集/训练/审批/审计页面 | `frontend/` |
| 应用平台后端 | 用户中心、多租户、RBAC/ABAC、项目隔离、元数据、审批流、审计、API 编排 | `backend/` |
| 算力平台 | GPU 发现、资源锁定、训练进程管理、日志采集、产物归档、任务状态同步 | `compute/` |
| 训练引擎 | 当前固定接入 LLaMA-Factory预留其他训练平台适配标准 | `compute/engines/` |
| 数据层 | PostgreSQL、Redis、本地文件存储、日志归档 | `docs/postgres-schema.sql` |
## 快速开始 ## 当前开发基线
### 1. 安装依赖 - 使用 FastAPI 提供统一 API 响应结构 `{ code, message, data }`
- 本地运行阶段统一使用 PostgreSQL后端启动时会在 PG 中初始化当前运行表和系统内置账号模型、数据集、算力节点、GPU、微调任务等业务数据必须通过页面、接口或正式导入流程产生。
- 支持登录、模型管理、数据集管理、微调任务创建/启动/停止/进度轮询。
- 支持训练日志、loss 指标、checkpoint 和训练产物接口;真实训练执行器接入前,联调状态机必须通过显式环境变量开启。
- 支持多算力节点、GPU、任务队列、资源副本和资源同步状态接口。
- 前端新增 `/compute` 算力节点页面展示节点地址、权重、标签、启用状态、GPU、队列和资源副本。
- `compute/engines/llama_factory/adapter.py` 提供 LLaMA-Factory 参数校验、命令生成和日志解析基础能力。
## 前后端一键启动
首次使用前,请先按下方“后端启动”和“前端启动”说明安装依赖,并确保
PostgreSQL 已可用。之后在项目根目录执行:
```bash
bash ./start.sh
```
脚本会同时启动前端 `http://localhost:16801` 和后端
`http://127.0.0.1:17861`,按 `Ctrl+C` 会同时停止两个服务。脚本只负责
启动前后端,不会自动安装依赖,也不会启动 PostgreSQL、Redis 或算力服务。
仅检查依赖和端口而不启动服务:
```bash
bash ./start.sh --check
```
本地启动推荐只配置数据库主机。脚本会复用 `docker/app/.env` 中已有的
`POSTGRES_USER``POSTGRES_PASSWORD``POSTGRES_DB`,端口默认使用
PostgreSQL 标准端口 `5432`
```bash
DATABASE_HOST='www.caoxiaozhu.com' bash ./start.sh
```
也可以在 `docker/app/.env` 中增加:
```env
DATABASE_HOST=www.caoxiaozhu.com
```
需要使用非标准端口时再设置 `DATABASE_PORT``DATABASE_URL` 仍可作为完整连接串
高级覆盖项;终端环境变量优先级最高。脚本不会输出数据库密码。
## 后端启动
```bash
cd backend
python -m venv .venv
.venv\Scripts\activate
pip install -r requirements.txt
uvicorn app.main:app --reload --port 17861
```
默认接口前缀为 `/modelTF`,例如:
```text
GET /modelTF/health
POST /modelTF/login
GET /modelTF/model-manage
GET /modelTF/dataset-manage
GET /modelTF/fine-tune
GET /modelTF/compute/nodes
```
本地运行时默认 PostgreSQL 连接:
```text
DATABASE_URL=postgresql+psycopg://yg_ft:change_me@localhost:15432/yg_ft
```
本地启动前需要确保 PostgreSQL 已监听 `localhost:15432`,并已创建 `yg_ft` 数据库和 `yg_ft` 用户。后端启动后会自动创建当前运行表并写入内置管理员账号,运行数据统一写入 PostgreSQL。
开发阶段内置登录账号:
| 角色 | 账号 | 密码 | 说明 |
| --- | --- | --- | --- |
| 超级管理员 | `admin` | `admin123` | 拥有当前全部页面权限 |
| 操作员 | `operator` | `operator123` | 拥有业务操作相关页面权限 |
以上账号仅用于本地开发和联调。生产环境初始化后应立即修改密码,或改为企业统一身份认证/管理员初始化流程。
## 前端启动
```bash ```bash
cd frontend cd frontend
npm install npm install
```
### 2. 启动开发服务器
```bash
npm run dev npm run dev
``` ```
开发服务默认运行在 `http://localhost:6801` 前端开发服务默认运行在 `http://localhost:16801`,并通过 Vite proxy 将 `/modelTF` 转发到 `http://localhost:17861`
### 3. 构建生产包 ## 算力服务启动
```bash ```bash
npm run build # 类型检查 + 生产构建,产物输出到 dist/ cd compute
npm run preview # 本地预览构建产物 uvicorn api.main:app --reload --port 19100
``` ```
### 4. 类型检查 默认 `COMPUTE_MODE=real`。真实 GPU 接入时,在每台算力服务器上部署 Compute API、Agent、File Gateway 和 LLaMA-Factory应用平台通过 `compute_nodes.api_base_url``compute_nodes.file_gateway_url` 主动轮询。仅在隔离联调环境可显式设置 `COMPUTE_MODE=simulator``COMPUTE_EXECUTION_MODE=simulator`
## 日志
后端日志模块位于 `backend/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 天。
## 主要文档
- `docs/platform-architecture-requirements.md`:平台需求、功能模块、页面补全建议。
- `docs/menu-functional-requirements.md`:当前菜单、二级路由、规划菜单、功能需求、接口和数据库映射。
- `docs/backend-api-design.md`FastAPI 接口分组、参数定义、权限说明。
- `docs/postgres-schema.sql`PostgreSQL 数据库脚本,包含权限、用户中心、多租户、审批、审计等模型。
- `docs/system-development-plan.md`多人协作开发计划按前端、后端、DB、部署拆分。
- `docs/team-development-plan.md`3-4 人并行开发分工计划,按人员边界标注页面、接口、数据库和交付节奏。
- `docs/first-version-development-plan.md`当前系统主链路开发计划覆盖前端、后端、DB、Compute API、GPU 和 LLaMA-Factory 适配。
- `docs/backend-logging.md`:后端日志模块使用说明。
- `docs/deployment-plan.md`:后期部署方案,覆盖单机算力服务器部署与应用/算力分离部署。
- `docker/README.md`Docker 部署入口,包含应用服务器和算力服务器两套 Compose 使用方式。
## Docker 部署入口
应用服务器:
```bash ```bash
npm run type-check cd docker/app
cp .env.example .env
docker compose up -d
``` ```
## 测试 算力服务器:
内置基于 Playwright 的 UI 回归脚本,首次运行前需安装浏览器:
```bash ```bash
npx playwright install chromium cd docker/compute
cp .env.example .env
docker compose up -d
``` ```
执行已注册的回归脚本: 两套 Compose 均采用代码外挂方式运行,镜像只包含运行时环境和第三方依赖。项目根目录不再保留 `Dockerfile``docker-compose.yml`,部署时统一进入 `docker/app``docker/compute` 目录执行。
```bash ## 后续开发原则
npm run test:data-process-wizard # 数据处理向导
npm run test:model-manage # 模型管理
npm run test:training-log-layout # 训练日志布局
npm run test:page-surface # 页面表层级
```
其余脚本可直接运行: - 接口实现优先遵循 `docs/backend-api-design.md`
- 数据库实现优先遵循 `docs/postgres-schema.sql`,后续通过 Alembic 迁移管理变更。
```bash - 前端页面与后端接口、数据库表之间的映射以文档中的“对应页面/功能模块”为准。
node scripts/regression-back-navigation.mjs # 返回导航 - 训练引擎适配必须通过 `compute/engines/` 下的标准接口,不在应用平台后端直接拼接训练命令。
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
View 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/日志平台采集。

1
backend/app/__init__.py Normal file
View File

@@ -0,0 +1 @@
"""Application package."""

View File

@@ -0,0 +1 @@
"""API package."""

View File

@@ -0,0 +1 @@
"""Versioned API package."""

View File

@@ -0,0 +1 @@
"""API endpoint modules."""

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,14 @@
from fastapi import APIRouter
from app.core.logging import get_logger
from app.db.platform_store import get_platform_store
router = APIRouter()
logger = get_logger(__name__)
@router.get("/health")
async def health_check() -> dict[str, object]:
logger.info("health check requested")
return {"code": 0, "message": "ok", "data": get_platform_store().health_metrics()}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,10 @@
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
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"])

View File

@@ -0,0 +1 @@
"""Core infrastructure modules."""

View File

@@ -0,0 +1,59 @@
from dataclasses import dataclass
from functools import lru_cache
import os
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()]
@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)
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)
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",
],
),
)
@lru_cache
def get_settings() -> Settings:
return Settings()

253
backend/app/core/logging.py Normal file
View File

@@ -0,0 +1,253 @@
from __future__ import annotations
from contextvars import ContextVar
from datetime import date, datetime, timedelta
import json
import logging
from logging import Handler, LogRecord
from pathlib import Path
import re
import time
from typing import Any
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="-")
class RequestIdFilter(logging.Filter):
def filter(self, record: LogRecord) -> bool:
record.request_id = request_id_var.get()
return True
class JsonLogFormatter(logging.Formatter):
"""Format one JSON object per line for ELK/Filebeat collection."""
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", "-"),
}
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=(",", ":"))
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)
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
def get_logger(name: str) -> logging.Logger:
return logging.getLogger(name)
def set_request_id(request_id: str) -> None:
request_id_var.set(request_id)
def setup_request_logging(app: FastAPI) -> None:
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
logger.info(
"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 "-",
)
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 "-",
)
raise
finally:
request_id_var.reset(token)

View File

@@ -0,0 +1 @@
"""Database infrastructure package."""

File diff suppressed because it is too large Load Diff

40
backend/app/db/session.py Normal file
View 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()

View File

@@ -0,0 +1,284 @@
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
);
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 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);

View File

@@ -0,0 +1,235 @@
-- 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,
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;
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,
original_instruction TEXT,
original_input TEXT,
original_output 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()
);
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;

43
backend/app/main.py Normal file
View 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 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)
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()

View 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` 的页面模块开发工作包为准。

View File

@@ -0,0 +1 @@
"""Approval workflow module."""

View File

@@ -0,0 +1 @@
"""Audit log module."""

View File

@@ -0,0 +1 @@
"""Authentication and user session module."""

View File

@@ -0,0 +1 @@
"""Application-side compute platform gateway module."""

View File

@@ -0,0 +1,206 @@
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 {}
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 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)}
async with httpx.AsyncClient(timeout=max(self.timeout, 60), 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())

View File

@@ -0,0 +1,44 @@
from __future__ import annotations
from typing import Any
from app.db.platform_store import get_platform_store
from app.modules.compute_gateway.client import ComputeNodeClient
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)
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
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)})
return {"synced": len(synced) + len(standalone_synced), "failed": failed, "items": synced, "standalone": standalone_synced}

View File

@@ -0,0 +1 @@
"""Data processing module."""

File diff suppressed because it is too large Load Diff

View 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"]

View File

@@ -0,0 +1,444 @@
"""基于 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:
return tiktoken.get_encoding("cl100k_base")
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.document_converter import DocumentConverter
return DocumentConverter()
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,
)

View File

@@ -0,0 +1,316 @@
"""数据处理任务的大模型生成适配器。"""
from __future__ import annotations
import hashlib
import json
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):
"""模型配置、响应或调用失败。"""
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 _message_content(payload: Mapping[str, Any]) -> str:
try:
content = payload["choices"][0]["message"]["content"]
except (KeyError, IndexError, TypeError) as exc:
raise ModelGenerationError("model response does not contain choices[0].message.content") from exc
if isinstance(content, str):
return content
if 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"}
]
if parts:
return "".join(parts)
raise ModelGenerationError("model response content must be text")
def _json_payload(content: str) -> Any:
cleaned = re.sub(r"<think>[\s\S]*?</think>", "", content, 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 exc:
raise ModelGenerationError(
f"model response is not valid JSON at line {exc.lineno}, column {exc.colno}"
) from exc
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,
) -> list[dict[str, str]]:
end_index = start_index + count - 1
schema_instruction = (
f"必须只返回 JSON 对象,格式为 {{\"items\":[{{\"instruction\":\"...\","
f"\"input\":\"...\",\"output\":\"...\"}}]}}items 必须包含 {count} 条。"
f"这是总计 {total_count} 条中的第 {start_index}-{end_index} 条,"
"各条必须使用不同的提问角度和表述,避免重复。"
"instruction 和 output 不得为空,不要输出 Markdown 代码围栏或分析过程。"
)
base_prompt = (
normalize_text(prompt)
or "请根据来源内容生成可用于监督微调的问答数据。"
)
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}]"
)
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")
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,
),
"temperature": temperature,
"max_tokens": max_tokens,
}
if bool(config.get("json_mode", False)):
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"
)
candidate_items = _result_items(
_json_payload(_message_content(body))
)
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 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": "",
"original_instruction": failure_instruction,
"original_input": content,
"original_output": "",
"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 "")
)
output = normalize_text(
str(
value.get("output")
or value.get("answer")
or value.get("response")
or ""
)
)
raw_id = f"{preview_id}:{variant_index + 1}:{instruction}:{output}"
result_id = (
f"result_{hashlib.sha256(raw_id.encode()).hexdigest()[:16]}"
)
valid = bool(instruction and output)
results.append(
{
"id": result_id,
"preview_item_id": preview_id,
"instruction": instruction,
"input": input_text,
"output": output,
"original_instruction": instruction,
"original_input": input_text,
"original_output": output,
"status": "valid" if valid else "invalid",
"error": (
None
if valid
else "model result is missing instruction or output"
),
"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"]

View File

@@ -0,0 +1,65 @@
"""数据处理运行表的显式检查与安装命令。"""
from __future__ import annotations
import argparse
from urllib.parse import urlsplit
from app.modules.data_process.store import DataProcessStore
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 EXISTS (
SELECT 1
FROM information_schema.columns
WHERE table_schema=current_schema()
AND table_name='data_process_tasks'
AND column_name='generation_run_id'
) AS ready
"""
).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())

View File

@@ -0,0 +1,499 @@
"""数据处理原始源文件的受控本地对象存储。"""
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:
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 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:
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",
]

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1 @@
"""Dataset management module."""

View File

@@ -0,0 +1 @@
"""Training engine registry module."""

View File

@@ -0,0 +1 @@
"""Evaluation module."""

View File

@@ -0,0 +1 @@
"""Application-side file gateway module."""

View File

@@ -0,0 +1 @@
"""Fine-tuning task module."""

View File

@@ -0,0 +1 @@
"""Inference and compare module."""

View File

@@ -0,0 +1 @@
"""Model registry module."""

View File

@@ -0,0 +1 @@
"""Project workspace and member module."""

View File

@@ -0,0 +1 @@
"""Retention policy and cleanup module."""

View File

@@ -0,0 +1 @@
"""System health, metrics and logs module."""

View File

@@ -0,0 +1 @@
"""Tenant management module."""

View File

@@ -0,0 +1 @@
"""Shared schemas package."""

View File

@@ -0,0 +1,321 @@
from __future__ import annotations
from enum import StrEnum
from typing import Any, Literal
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:
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 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 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 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)
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
expected_updated_at: str | None = None
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"] = "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

View File

@@ -0,0 +1 @@
"""Cross-module services package."""

View File

@@ -0,0 +1 @@
"""Background workers package."""

View File

@@ -0,0 +1,31 @@
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["synced"] or result["failed"]:
logger.info("compute jobs polled", 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)

40
backend/pyproject.toml Normal file
View File

@@ -0,0 +1,40 @@
[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",
"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"]

20
backend/requirements.txt Normal file
View File

@@ -0,0 +1,20 @@
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
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

View File

@@ -0,0 +1,772 @@
from __future__ import annotations
import io
import json
import xml.etree.ElementTree as ET
import zipfile
from datetime import datetime
import pytest
from docx import Document
from openpyxl import Workbook
from pptx import Presentation
from pptx.util import Inches
from pypdf import PdfWriter
from app.modules.data_process.algorithms import (
PdfPageText,
content_quality_flags,
desensitize_pii,
desensitize_structured_record,
detect_document_structure,
detect_pdf_document_noise,
detect_text_format,
extract_pdf_page_texts,
extract_structured_records,
generate_standard_records,
is_near_duplicate,
merge_short_blocks,
normalize_text,
parse_text_content,
preprocess_structured_records,
record_fingerprint,
remove_document_noise,
score_quality,
stable_split,
stable_split_assignments,
)
def _pdf_page_texts(*texts: str) -> tuple[PdfPageText, ...]:
pages: list[PdfPageText] = []
offset = 0
for page_number, text in enumerate(texts, start=1):
normalized = normalize_text(text)
if pages:
offset += 2
start = offset
offset += len(normalized)
pages.append(
PdfPageText(
page_number=page_number,
text=normalized,
source_start=start,
source_end=offset,
)
)
return tuple(pages)
def _minimal_pdf(text: str = "Hello PDF") -> bytes:
stream = f"BT /F1 12 Tf 72 720 Td ({text}) Tj ET".encode("ascii")
objects = [
b"<< /Type /Catalog /Pages 2 0 R >>",
b"<< /Type /Pages /Kids [3 0 R] /Count 1 >>",
(
b"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 612 792] "
b"/Resources << /Font << /F1 5 0 R >> >> /Contents 4 0 R >>"
),
b"<< /Length " + str(len(stream)).encode("ascii") + b" >>\nstream\n"
+ stream
+ b"\nendstream",
b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>",
]
result = bytearray(b"%PDF-1.4\n%\xe2\xe3\xcf\xd3\n")
offsets = [0]
for object_number, value in enumerate(objects, start=1):
offsets.append(len(result))
result.extend(f"{object_number} 0 obj\n".encode("ascii"))
result.extend(value)
result.extend(b"\nendobj\n")
xref_offset = len(result)
result.extend(f"xref\n0 {len(objects) + 1}\n".encode("ascii"))
result.extend(b"0000000000 65535 f \n")
for offset in offsets[1:]:
result.extend(f"{offset:010d} 00000 n \n".encode("ascii"))
result.extend(
(
f"trailer\n<< /Size {len(objects) + 1} /Root 1 0 R >>\n"
f"startxref\n{xref_offset}\n%%EOF\n"
).encode("ascii")
)
return bytes(result)
def _aes_encrypted_pdf(*, user_password: str) -> bytes:
writer = PdfWriter(clone_from=io.BytesIO(_minimal_pdf()))
writer.encrypt(
user_password=user_password,
owner_password="owner-secret",
algorithm="AES-256",
)
output = io.BytesIO()
writer.write(output)
return output.getvalue()
def _docx_bytes() -> bytes:
document = Document()
document.add_heading("服务说明", level=1)
document.add_paragraph("这是 DOCX 正文。")
table = document.add_table(rows=1, cols=2)
table.cell(0, 0).text = "字段"
table.cell(0, 1).text = "内容"
output = io.BytesIO()
document.save(output)
return output.getvalue()
def _xlsx_bytes() -> bytes:
workbook = Workbook()
worksheet = workbook.active
worksheet.title = "数据"
worksheet.append(["name", "score", "created_at"])
worksheet.append(["Alice", 95, datetime(2026, 7, 23, 10, 30)])
worksheet.append(["Bob", 88, datetime(2026, 7, 24, 9, 0)])
output = io.BytesIO()
workbook.save(output)
workbook.close()
return output.getvalue()
def _xlsx_with_worksheet_relationship(
raw: bytes,
target: str,
*,
target_mode: str | None = None,
) -> bytes:
member_name = "xl/_rels/workbook.xml.rels"
source = io.BytesIO(raw)
output = io.BytesIO()
with zipfile.ZipFile(source) as original, zipfile.ZipFile(output, "w") as rewritten:
for member in original.infolist():
content = original.read(member.filename)
if member.filename == member_name:
root = ET.fromstring(content)
worksheet_relationship = next(
element
for element in root
if element.attrib.get("Type", "").endswith("/worksheet")
)
worksheet_relationship.set("Target", target)
if target_mode is None:
worksheet_relationship.attrib.pop("TargetMode", None)
else:
worksheet_relationship.set("TargetMode", target_mode)
content = ET.tostring(root, encoding="utf-8", xml_declaration=True)
rewritten.writestr(member, content)
return output.getvalue()
def _pptx_bytes() -> bytes:
presentation = Presentation()
slide = presentation.slides.add_slide(presentation.slide_layouts[6])
text_box = slide.shapes.add_textbox(Inches(1), Inches(1), Inches(6), Inches(1))
text_box.text = "PPTX 页面正文"
output = io.BytesIO()
presentation.save(output)
return output.getvalue()
def test_parse_utf8_json_jsonl_csv_markdown_and_txt() -> None:
parsed_json = parse_text_content(
b'\xef\xbb\xbf{"data":[{"name":"\xe5\xbc\xa0\xe4\xb8\x89"}]}',
filename="records.json",
)
assert parsed_json.format == "json"
assert parsed_json.records == ({"name": "张三"},)
parsed_jsonl = parse_text_content('{"id":1}\n\n{"id":2}\n', filename="records.jsonl")
assert parsed_jsonl.format == "jsonl"
assert parsed_jsonl.records == ({"id": 1}, {"id": 2})
parsed_csv = parse_text_content("name,answer\r\nAlice,yes\r\nBob,no", filename="records.csv")
assert parsed_csv.format == "csv"
assert parsed_csv.text == "name,answer\nAlice,yes\nBob,no"
assert parsed_csv.records[1] == {"name": "Bob", "answer": "no"}
parsed_markdown = parse_text_content("# 标题\n\n正文", filename="README.md")
assert parsed_markdown.format == "markdown"
assert parsed_markdown.records == ()
parsed_txt = parse_text_content("普通文本", filename="note.txt")
assert parsed_txt.format == "txt"
assert parsed_txt.text == "普通文本"
def test_parse_pdf_docx_xlsx_and_pptx() -> None:
parsed_pdf = parse_text_content(_minimal_pdf(), filename="manual.pdf")
assert parsed_pdf.format == "pdf"
assert "Hello PDF" in parsed_pdf.text
assert parsed_pdf.records == ()
pdf_pages = extract_pdf_page_texts(_minimal_pdf())
assert len(pdf_pages) == 1
assert pdf_pages[0].page_number == 1
assert pdf_pages[0].text == "Hello PDF"
assert pdf_pages[0].source_start == 0
assert pdf_pages[0].source_end == len(parsed_pdf.text)
parsed_docx = parse_text_content(_docx_bytes(), filename="manual.docx")
assert parsed_docx.format == "docx"
assert "服务说明" in parsed_docx.text
assert "这是 DOCX 正文。" in parsed_docx.text
assert "字段\t内容" in parsed_docx.text
assert parsed_docx.records == ()
parsed_xlsx = parse_text_content(_xlsx_bytes(), filename="records.xlsx")
assert parsed_xlsx.format == "xlsx"
assert parsed_xlsx.records == (
{"name": "Alice", "score": 95, "created_at": "2026-07-23T10:30:00"},
{"name": "Bob", "score": 88, "created_at": "2026-07-24T09:00:00"},
)
assert json.loads(parsed_xlsx.text.splitlines()[0]) == parsed_xlsx.records[0]
parsed_pptx = parse_text_content(_pptx_bytes(), filename="slides.pptx")
assert parsed_pptx.format == "pptx"
assert parsed_pptx.text == "PPTX 页面正文"
assert parsed_pptx.records == ()
def test_pdf_document_noise_removes_headers_page_numbers_and_toc_safely() -> None:
pages = _pdf_page_texts(
"""
远光制度文件 文件编码 2024
秘密等级 商密【中】
第 1 页 共 5 页
正文第一页,关于适用范围的说明。
业务提示保留
第一页补充说明甲
第一页补充说明乙
第一页补充说明丙
""",
"""
远光制度文件 文件编码 2024
秘密等级 商密【中】
第 2 页 共 5 页
目 录
第一章 总则........3
第二章 报销申请........4
第三章 附则........5
""",
"""
远光制度文件 文件编码 2024
秘密等级 商密【中】
第 3 页 共 5 页
1.1 管理要求........6
1.2 审批职责 7
1.3 费用标准........8
1.4 例外处理........9
""",
"""
远光制度文件 文件编码 2024
秘密等级 商密【中】
第 4 页 共 5 页
正文中可以说“请参见第 3 页说明”,不应误删。
第 99 页 共 100 页
系统可用率........99.9%
业务提示保留
第四页补充说明甲
第四页补充说明乙
第四页补充说明丙
""",
"""
远光制度文件 文件编码 2024
秘密等级 商密【中】
第 5 页 共 5 页
本办法自发布之日起施行。
业务提示保留
第五页补充说明甲
第五页补充说明乙
第五页补充说明丙
""",
)
source = "\n\n".join(page.text for page in pages)
spans = detect_pdf_document_noise(pages)
cleaned = remove_document_noise(source, spans)
assert {span.kind for span in spans} == {
"page_number",
"repeated_margin",
"table_of_contents",
}
assert "远光制度文件" not in cleaned
assert "商密【中】" not in cleaned
assert "第 1 页 共 5 页" not in cleaned
assert "第一章 总则" not in cleaned
assert "1.2 审批职责 7" not in cleaned
assert "请参见第 3 页说明" in cleaned
assert "第 99 页 共 100 页" in cleaned
assert "系统可用率........99.9%" in cleaned
assert cleaned.count("业务提示保留") == 3
def test_pdf_document_noise_does_not_infer_repeated_margins_for_short_documents() -> None:
pages = _pdf_page_texts(
"公司内部文件\n正文 A",
"公司内部文件\n正文 B",
)
spans = detect_pdf_document_noise(pages)
assert not any(span.kind == "repeated_margin" for span in spans)
def test_xlsx_merged_multilevel_headers_are_flattened_without_losing_columns() -> None:
workbook = Workbook()
worksheet = workbook.active
worksheet.merge_cells("A1:A2")
worksheet.merge_cells("B1:C1")
worksheet["A1"] = "地区"
worksheet["B1"] = "销售"
worksheet["B2"] = "Q1"
worksheet["C2"] = "Q2"
worksheet.append(["华东", 100, 120])
output = io.BytesIO()
workbook.save(output)
workbook.close()
parsed = parse_text_content(output.getvalue(), filename="sales.xlsx")
assert parsed.records == ({"地区": "华东", "销售.Q1": 100, "销售.Q2": 120},)
def test_xlsx_header_inference_skips_more_than_eight_merged_report_titles() -> None:
workbook = Workbook()
worksheet = workbook.active
for row_number in range(1, 13):
worksheet.merge_cells(
start_row=row_number,
start_column=1,
end_row=row_number,
end_column=4,
)
worksheet.cell(row_number, 1, f"报表说明 {row_number}")
worksheet.append(["姓名", "部门", "得分", "日期"])
worksheet.append(["张三", "研发", 95, "2026-07-23"])
output = io.BytesIO()
workbook.save(output)
workbook.close()
parsed = parse_text_content(output.getvalue(), filename="report.xlsx")
assert parsed.records == (
{"姓名": "张三", "部门": "研发", "得分": 95, "日期": "2026-07-23"},
)
def test_xlsx_header_inference_ignores_continuous_body_merges() -> None:
workbook = Workbook()
worksheet = workbook.active
worksheet.append(["类别", "名称", "数量"])
worksheet.append(["水果", "苹果", 10])
worksheet.append([None, "香蕉", 12])
worksheet.append(["蔬菜", "白菜", 8])
worksheet.append([None, "萝卜", 9])
worksheet.merge_cells("A2:A3")
worksheet.merge_cells("A4:A5")
output = io.BytesIO()
workbook.save(output)
workbook.close()
parsed = parse_text_content(output.getvalue(), filename="inventory.xlsx")
assert parsed.records == (
{"类别": "水果", "名称": "苹果", "数量": 10},
{"类别": "", "名称": "香蕉", "数量": 12},
{"类别": "蔬菜", "名称": "白菜", "数量": 8},
{"类别": "", "名称": "萝卜", "数量": 9},
)
def test_xlsx_header_inference_supports_title_and_two_header_levels() -> None:
workbook = Workbook()
worksheet = workbook.active
worksheet.merge_cells("A1:C1")
worksheet["A1"] = "区域销售报表"
worksheet["A2"] = "统计日期"
worksheet["B2"] = "2026-07-23"
worksheet.merge_cells("A4:A5")
worksheet.merge_cells("B4:C4")
worksheet["A4"] = "地区"
worksheet["B4"] = "销售"
worksheet["B5"] = "Q1"
worksheet["C5"] = "Q2"
worksheet.append(["华南", 88, 92])
output = io.BytesIO()
workbook.save(output)
workbook.close()
parsed = parse_text_content(output.getvalue(), filename="two-level.xlsx")
assert parsed.records == (
{"地区": "华南", "销售.Q1": 88, "销售.Q2": 92},
)
def test_xlsx_header_inference_supports_title_and_three_header_levels() -> None:
workbook = Workbook()
worksheet = workbook.active
worksheet.merge_cells("A1:D1")
worksheet["A1"] = "年度销售分析报告"
worksheet["A2"] = "统计日期"
worksheet["B2"] = "2026-07-23"
worksheet.merge_cells("A4:A6")
worksheet.merge_cells("B4:D4")
worksheet.merge_cells("B5:C5")
worksheet.merge_cells("D5:D6")
worksheet["A4"] = "地区"
worksheet["B4"] = "销售"
worksheet["B5"] = "国内"
worksheet["D5"] = "海外"
worksheet["B6"] = "Q1"
worksheet["C6"] = "Q2"
worksheet.append(["华东", 100, 120, 80])
output = io.BytesIO()
workbook.save(output)
workbook.close()
parsed = parse_text_content(output.getvalue(), filename="three-level.xlsx")
assert parsed.records == (
{
"地区": "华东",
"销售.国内.Q1": 100,
"销售.国内.Q2": 120,
"销售.海外": 80,
},
)
def test_xlsx_header_inference_keeps_an_ordinary_single_header_row() -> None:
parsed = parse_text_content(_xlsx_bytes(), filename="ordinary.xlsx")
assert tuple(parsed.records[0]) == ("name", "score", "created_at")
assert len(parsed.records) == 2
@pytest.mark.parametrize(
"target",
[
"./worksheets/../worksheets/sheet1.xml",
"./worksheets/%2e%2e/worksheets/sheet1.xml",
"../xl/worksheets/sheet1.xml",
"/xl/worksheets/./sheet1.xml",
],
)
def test_xlsx_worksheet_relationship_allows_safe_dot_segments(target: str) -> None:
raw = _xlsx_with_worksheet_relationship(_xlsx_bytes(), target)
parsed = parse_text_content(raw, filename="records.xlsx")
assert parsed.records[0]["name"] == "Alice"
@pytest.mark.parametrize(
"target",
[
"../../outside.xml",
"worksheets\\sheet1.xml",
"%2e%2e/%2e%2e/outside.xml",
"%252e%252e/%252e%252e/outside.xml",
"https://example.com/sheet1.xml",
],
)
def test_xlsx_worksheet_relationship_rejects_path_traversal(target: str) -> None:
raw = _xlsx_with_worksheet_relationship(_xlsx_bytes(), target)
with pytest.raises(ValueError, match="unsafe worksheet path"):
parse_text_content(raw, filename="unsafe.xlsx")
def test_xlsx_worksheet_relationship_rejects_external_and_missing_targets() -> None:
external = _xlsx_with_worksheet_relationship(
_xlsx_bytes(),
"https://example.com/sheet1.xml",
target_mode="External",
)
with pytest.raises(ValueError, match="external relationship"):
parse_text_content(external, filename="external.xlsx")
missing = _xlsx_with_worksheet_relationship(
_xlsx_bytes(),
"worksheets/missing.xml",
)
with pytest.raises(ValueError, match="target does not exist"):
parse_text_content(missing, filename="missing.xlsx")
@pytest.mark.parametrize(
("filename", "replacement"),
[
("legacy.doc", ".docx"),
("legacy.xls", ".xlsx"),
("legacy.ppt", ".pptx"),
],
)
def test_legacy_office_formats_require_conversion(filename: str, replacement: str) -> None:
with pytest.raises(ValueError, match=rf"convert the file to \{replacement}"):
parse_text_content(b"legacy", filename=filename)
def test_office_zip_bomb_and_invalid_pdf_are_rejected_before_parsing() -> None:
archive = io.BytesIO()
with zipfile.ZipFile(archive, "w", compression=zipfile.ZIP_DEFLATED) as package:
package.writestr("[Content_Types].xml", "<Types/>")
package.writestr("word/document.xml", b"A" * (2 * 1024 * 1024))
with pytest.raises(ValueError, match="unsafe compression ratio"):
parse_text_content(archive.getvalue(), filename="unsafe.docx")
active_xml = io.BytesIO()
with zipfile.ZipFile(active_xml, "w") as package:
package.writestr("[Content_Types].xml", "<Types/>")
package.writestr(
"word/document.xml",
'<!DOCTYPE document [<!ENTITY xxe SYSTEM "file:///etc/passwd">]><document/>',
)
with pytest.raises(ValueError, match="unsupported active XML"):
parse_text_content(active_xml.getvalue(), filename="active.docx")
with pytest.raises(ValueError, match="missing PDF header"):
parse_text_content(b"not a pdf", filename="broken.pdf")
blank_pdf = io.BytesIO()
blank_writer = PdfWriter()
blank_writer.add_blank_page(width=612, height=792)
blank_writer.write(blank_pdf)
with pytest.raises(ValueError, match="scanned PDF requires OCR"):
parse_text_content(blank_pdf.getvalue(), filename="scanned.pdf")
aes_pdf_without_open_password = parse_text_content(
_aes_encrypted_pdf(user_password=""),
filename="aes-no-password.pdf",
)
assert "Hello PDF" in aes_pdf_without_open_password.text
with pytest.raises(ValueError, match="password-protected PDF files are not supported"):
parse_text_content(
_aes_encrypted_pdf(user_password="secret"),
filename="aes-password.pdf",
)
def test_invalid_utf8_and_malformed_structured_content_fail_loudly() -> None:
with pytest.raises(ValueError, match="not valid UTF-8"):
parse_text_content(b"\xff\xfe", filename="broken.txt")
with pytest.raises(ValueError, match="invalid JSONL at line 2"):
extract_structured_records('{"id":1}\nnot-json', "jsonl")
with pytest.raises(ValueError, match="more fields"):
extract_structured_records("a,b\n1,2,3", "csv")
def test_detect_format_from_content_and_normalize() -> None:
assert detect_text_format(text='{"id":1}\n{"id":2}') == "jsonl"
assert detect_text_format(text="# Heading\ntext") == "markdown"
assert detect_text_format(text="a,b\n1,2") == "csv"
assert normalize_text("\ufeff \r\n第二\x00\u200b\t \r\n") == "ABC\n第二行"
def test_extract_json_scalar_and_nested_values_are_stable() -> None:
assert extract_structured_records("[1, true, null]", "json") == [
{"value": 1},
{"value": True},
{"value": None},
]
result = extract_structured_records(
json.dumps({"items": [{"text": " 内容 "}], "ignored": 1}, ensure_ascii=False),
"json",
)
assert result == [{"text": "内容"}]
def test_desensitize_pii_returns_masked_text_and_counts() -> None:
source = "邮箱 a.user+tag@example.com手机 +86 13800138000身份证 11010519491231002X。"
masked, counts = desensitize_pii(source)
assert masked == "邮箱 [EMAIL],手机 [PHONE],身份证 [ID_CARD]。"
assert counts == {"email": 1, "phone": 1, "id_card": 1, "total": 3}
def test_every_structured_preprocess_option_has_independent_behavior() -> None:
clean_source = [
{"id": "1", "name": "有效", "empty_column": ""},
{"id": "", "name": "缺少关键字段", "empty_column": ""},
{"id": "2", "name": "有效", "empty_column": ""},
]
assert preprocess_structured_records(clean_source, []) == clean_source
assert preprocess_structured_records(clean_source, ["clean_invalid"]) == [
{"id": "1", "name": "有效"},
{"id": "2", "name": "有效"},
]
nested = [{"id": 1, "profile": {"name": "张三", "level": 2}}]
assert "profile" in preprocess_structured_records(nested, [])[0]
assert preprocess_structured_records(nested, ["detect_structure"])[0] == {
"id": 1,
"profile.name": "张三",
"profile.level": 2,
}
duplicates = [
{"customer_id": "C-1", "value": "first"},
{"customer_id": "C-1", "value": "updated"},
{"customer_id": "", "value": "blank-one"},
{"customer_id": "", "value": "blank-two"},
]
assert len(preprocess_structured_records(duplicates, [])) == 4
deduplicated = preprocess_structured_records(duplicates, ["deduplicate"])
assert [record["value"] for record in deduplicated] == [
"first",
"blank-one",
"blank-two",
]
unnormalized = [{" User Name ": "\r\n第二行"}]
assert preprocess_structured_records(unnormalized, []) == unnormalized
assert preprocess_structured_records(unnormalized, ["normalize_format"]) == [
{"user_name": "ABC\n第二行"}
]
anomaly_source = [
{"id": 10_000 + index, "amount": amount, "text": "正常内容"}
for index, amount in enumerate((10, 10, 11, 11, 12, 12, 13, 1000))
]
assert len(preprocess_structured_records(anomaly_source, [])) == 8
filtered = preprocess_structured_records(anomaly_source, ["filter_anomaly"])
assert len(filtered) == 7
assert all(record["amount"] != 1000 for record in filtered)
assert max(record["id"] for record in filtered) > 10_000
sensitive = [{"姓名": "张三", "phone": "13800138000", "email": "a@b.com"}]
assert preprocess_structured_records(sensitive, []) == sensitive
masked = preprocess_structured_records(sensitive, ["desensitize"])[0]
assert masked == {"姓名": "[NAME]", "phone": "[PHONE]", "email": "[EMAIL]"}
def test_structured_desensitization_counts_and_document_helpers() -> None:
masked, counts = desensitize_structured_record(
{"联系人姓名": "李四", "说明": "邮箱 user@example.com手机 13900139000"}
)
assert masked == {
"联系人姓名": "[NAME]",
"说明": "邮箱 [EMAIL],手机 [PHONE]",
}
assert counts == {"email": 1, "phone": 1, "id_card": 0, "name": 1, "total": 3}
structure = detect_document_structure(
"# 第一章\n正文\n\n## 细节\n- 项目一\n- 项目二\n\n```python\nprint(1)\n```"
)
assert [heading.title for heading in structure.headings] == ["第一章", "细节"]
assert structure.list_block_count == 1
assert structure.code_block_count == 1
assert merge_short_blocks(["短一", "短二", "这是一段足够长的正文内容"], min_token_count=4)
assert "mojibake" in content_quality_flags("正常文字锟斤拷内容", min_chars=0, min_tokens=0)
assert is_near_duplicate(
"alpha beta gamma delta epsilon zeta eta theta iota kappa lambda mu nu xi omicron",
"alpha beta gamma, delta epsilon zeta eta theta iota kappa lambda mu nu xi omicron",
similarity_threshold=0.92,
max_hamming_distance=2,
)
def test_quality_scoring_covers_all_dimensions_and_duplicates() -> None:
valid = {
"instruction": "如何修改收货地址?",
"input": "订单尚未发货",
"output": "可以在订单详情页申请修改收货地址。",
}
source = "订单尚未发货时,可以在订单详情页申请修改收货地址。"
first_score = score_quality(valid, min_output_length=10, source_content=source)
assert first_score.is_valid
assert first_score.completeness == 100
assert first_score.length == 100
assert first_score.readability >= 90
assert first_score.relevance >= 70
assert first_score.duplicate == 100
duplicate_score = score_quality(valid, known_fingerprints={first_score.fingerprint})
assert duplicate_score.duplicate == 0
assert "duplicate_record" in duplicate_score.flags
unrelated_score = score_quality(
valid,
min_output_length=10,
source_content="量子计算使用量子比特处理信息。",
)
assert unrelated_score.relevance < first_score.relevance
assert "low_source_relevance" in unrelated_score.flags
invalid_score = score_quality({"instruction": "", "output": ""}, min_output_length=10)
assert not invalid_score.is_valid
assert {"missing_instruction", "output_too_short"}.issubset(invalid_score.flags)
assert record_fingerprint(valid) == record_fingerprint(dict(reversed(list(valid.items()))))
def test_stable_split_is_reproducible_and_validates_ratios() -> None:
first = stable_split("record-42", seed="task-1")
assert stable_split("record-42", seed="task-1") == first
assert first in {"train", "validation", "test"}
assert stable_split("record-42", {"train": 100, "validation": 0, "test": 0}) == "train"
with pytest.raises(ValueError, match="sum to 100"):
stable_split("record", {"train": 80, "validation": 10, "test": 9})
def test_stable_split_assignments_use_exact_deterministic_quotas() -> None:
values = [f"record-{index}" for index in range(28)]
first = stable_split_assignments(values, seed="task-1")
second = stable_split_assignments(values, seed="task-1")
assert first == second
assert first.count("train") == 22
assert first.count("validation") == 3
assert first.count("test") == 3
def test_generate_standard_records_supports_json_qa_and_stable_variants() -> None:
previews = [
{
"id": "preview-json",
"edited_content": json.dumps(
{"instruction": "问题", "input": "上下文", "output": "答案"},
ensure_ascii=False,
),
},
{"id": "preview-qa", "editedContent": "问:如何操作?\n答:按步骤操作。"},
]
records = generate_standard_records(
previews,
qa_pairs_per_item=2,
semantic_enrichment=True,
split={"train": 100, "validation": 0, "test": 0},
split_seed="task-1",
)
assert len(records) == 4
assert records[0]["instruction"] == "问题"
assert records[0]["input"] == "上下文"
assert records[0]["output"] == "答案"
assert records[1]["instruction"].endswith("问题")
assert records[2]["instruction"] == "如何操作?"
assert records[2]["output"] == "按步骤操作。"
assert all(record["status"] == "valid" for record in records)
assert all(record["split"] == "train" for record in records)
assert records == generate_standard_records(
previews,
qa_pairs_per_item=2,
semantic_enrichment=True,
split={"train": 100, "validation": 0, "test": 0},
split_seed="task-1",
)
def test_generate_standard_records_supports_fifty_unique_semantic_variants() -> None:
records = generate_standard_records(
[{"id": "preview-50", "edited_content": "问:如何操作?\n答:按步骤操作。"}],
qa_pairs_per_item=50,
semantic_enrichment=True,
split={"train": 100, "validation": 0, "test": 0},
split_seed="task-50",
)
assert len(records) == 50
assert len({record["id"] for record in records}) == 50
assert len({record["instruction"] for record in records}) == 50
assert all(record["status"] == "valid" for record in records)
@pytest.mark.parametrize("qa_pairs_per_item", [0, 51])
def test_generate_standard_records_rejects_out_of_range_count(
qa_pairs_per_item: int,
) -> None:
with pytest.raises(ValueError, match=r"\[1, 50\]"):
generate_standard_records([], qa_pairs_per_item=qa_pairs_per_item)

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,269 @@
from __future__ import annotations
import json
import httpx
import pytest
from app.modules.data_process.generation import (
ModelGenerationError,
chat_completions_url,
generate_model_records,
)
def test_chat_completions_url_accepts_host_base_and_complete_url() -> None:
assert chat_completions_url("www.caoxiaozhu.com") == (
"https://www.caoxiaozhu.com/v1/chat/completions"
)
assert chat_completions_url("https://model.example/v1") == (
"https://model.example/v1/chat/completions"
)
complete = "https://model.example/openai/v1/chat/completions"
assert chat_completions_url(complete) == complete
def test_generate_model_records_uses_prompt_auth_and_stable_split() -> None:
requests: list[httpx.Request] = []
progress_updates: list[tuple[int, int]] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
payload = json.loads(request.content)
assert payload["model"] == "qwen-plus"
assert payload["response_format"] == {"type": "json_object"}
assert "客户反馈页面加载慢" in payload["messages"][1]["content"]
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "请生成简洁客服回复",
"input": "客户反馈页面加载慢",
"output": "已收到反馈,我们正在排查。",
}
]
},
ensure_ascii=False,
)
}
}
]
},
)
client = httpx.Client(transport=httpx.MockTransport(handler))
records = generate_model_records(
[{"id": "preview-1", "edited_content": "客户反馈页面加载慢"}],
model={
"name": "Qwen",
"online_model_name": "qwen-plus",
"api_url": "model.example",
"api_key": "test-secret",
},
config={
"generation_prompt": "请处理:{{ content }}",
"json_mode": True,
"temperature": 0.2,
"max_tokens": 512,
},
task_id="task-1",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=1,
client=client,
on_progress=lambda processed, total: progress_updates.append((processed, total)),
)
assert len(records) == 1
assert records[0]["status"] == "valid"
assert records[0]["split"] == "train"
assert requests[0].headers["Authorization"] == "Bearer test-secret"
assert progress_updates == [(1, 1)]
def test_generate_model_records_keeps_partial_failure_for_manual_repair() -> None:
client = httpx.Client(
transport=httpx.MockTransport(
lambda _: httpx.Response(200, json={"choices": [{"message": {"content": "not-json"}}]})
)
)
records = generate_model_records(
[{"id": "preview-1", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-1",
split={"train": 80, "validation": 10, "test": 10},
qa_pairs_per_item=1,
client=client,
)
assert len(records) == 1
assert records[0]["status"] == "invalid"
assert records[0]["error"]
def test_generate_model_records_batches_fifty_results_with_unique_ids() -> None:
requests: list[httpx.Request] = []
def handler(request: httpx.Request) -> httpx.Response:
requests.append(request)
batch_start = (len(requests) - 1) * 10 + 1
batch_end = batch_start + 9
payload = json.loads(request.content)
system_prompt = payload["messages"][0]["content"]
assert "items 必须包含 10 条" in system_prompt
assert f"{batch_start}-{batch_end}" in system_prompt
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "同一问题",
"input": "来源正文",
"output": "同一答案",
}
for _ in range(batch_start, batch_end + 1)
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-50", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={},
task_id="task-50",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=50,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert len(requests) == 5
assert len(records) == 50
assert len({record["id"] for record in records}) == 50
assert {record["instruction"] for record in records} == {"同一问题"}
assert all(record["status"] == "valid" for record in records)
def test_generate_model_records_preserves_successful_batches_when_one_fails() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
if request_count == 2:
return httpx.Response(500)
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": f"问题 {index}",
"output": f"答案 {index}",
}
for index in range(1, 11)
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-partial", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 0},
task_id="task-partial",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=20,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert len(records) == 11
assert sum(record["status"] == "valid" for record in records) == 10
failed = next(record for record in records if record["status"] == "invalid")
assert "第 11-20 条" in failed["instruction"]
assert len({record["id"] for record in records}) == len(records)
def test_generate_model_records_retries_short_batch_then_marks_it_invalid() -> None:
request_count = 0
def handler(_: httpx.Request) -> httpx.Response:
nonlocal request_count
request_count += 1
return httpx.Response(
200,
json={
"choices": [
{
"message": {
"content": json.dumps(
{
"items": [
{
"instruction": "只有一条",
"output": "不足本批要求数量",
}
]
},
ensure_ascii=False,
)
}
}
]
},
)
records = generate_model_records(
[{"id": "preview-short", "edited_content": "来源正文"}],
model={"name": "model", "api_url": "https://model.example/v1"},
config={"generation_retries": 1},
task_id="task-short",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=10,
client=httpx.Client(transport=httpx.MockTransport(handler)),
)
assert request_count == 2
assert len(records) == 1
assert records[0]["status"] == "invalid"
assert "expected 10, got 1" in records[0]["error"]
@pytest.mark.parametrize("qa_pairs_per_item", [0, 51])
def test_generate_model_records_rejects_out_of_range_count(
qa_pairs_per_item: int,
) -> None:
with pytest.raises(ModelGenerationError, match=r"\[1, 50\]"):
generate_model_records(
[],
model={"name": "model", "api_url": "https://model.example/v1"},
config={},
task_id="task-invalid",
split={"train": 100, "validation": 0, "test": 0},
qa_pairs_per_item=qa_pairs_per_item,
)

View File

@@ -0,0 +1,29 @@
from __future__ import annotations
from pathlib import Path
from app.modules.data_process.schema_cli import _target_label
def test_runtime_migration_fails_fast_on_incompatible_schema() -> None:
sql_path = (
Path(__file__).resolve().parents[1]
/ "app"
/ "db"
/ "sql"
/ "002_data_process.sql"
)
sql = sql_path.read_text(encoding="utf-8")
assert "requires 001_platform_runtime.sql first" in sql
assert "supports only the current TEXT runtime schema" in sql
assert "generation_run_id" in sql
assert "CREATE TABLE IF NOT EXISTS data_process_results" in sql
assert sql.count("BEGIN;") == 1
assert sql.rstrip().endswith("COMMIT;")
def test_schema_cli_target_label_never_contains_credentials() -> None:
label = _target_label("postgresql://secret-user:secret-password@db.example:5433/yg_ft")
assert label == "db.example:5433/yg_ft"
assert "secret" not in label

View File

@@ -0,0 +1,242 @@
from __future__ import annotations
from pathlib import Path, PurePosixPath
import pytest
from app.modules.data_process import storage as storage_module
from app.modules.data_process.storage import (
DataProcessStorageError,
LocalDataProcessStorage,
StagedSourceObject,
)
def _stage(
storage: LocalDataProcessStorage,
*,
batch_id: str = "batch-main",
task_id: str = "task-1",
source_file_id: str = "source-1",
version: int = 1,
name: str = "source.txt",
content: bytes = b"payload",
) -> StagedSourceObject:
return storage.stage_bytes(
batch_id=batch_id,
task_id=task_id,
source_file_id=source_file_id,
version=version,
name=name,
content=content,
)
def _create_symlink(link: Path, target: Path, *, target_is_directory: bool = False) -> None:
try:
link.symlink_to(target, target_is_directory=target_is_directory)
except (NotImplementedError, OSError) as exc:
pytest.skip(f"当前平台不支持创建测试所需的符号链接: {exc}")
def _assert_staging_empty(storage: LocalDataProcessStorage) -> None:
assert list((storage.root / ".staging").iterdir()) == []
def test_stage_publish_read_delete_roundtrip_with_unicode_filename(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
content = "第一行\n第二行100% 完成".encode()
staged = _stage(
storage,
name="中文 数据 100%.csv",
content=content,
)
assert "%20" in staged.reference
assert "%25" in staged.reference
storage.publish([staged])
assert storage.read(staged.reference) == content
assert storage.delete(staged.reference) is True
assert storage.delete(staged.reference) is False
_assert_staging_empty(storage)
def test_db_reference_is_left_to_database_storage(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
assert storage.read("db://source-files/source-1") is None
assert storage.delete("db://source-files/source-1") is False
def test_owned_source_can_be_streamed_by_byte_range(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
content = b"0123456789abcdef"
staged = _stage(storage, content=content)
storage.publish([staged])
assert storage.file_size(
staged.reference,
expected_task_id="task-1",
expected_source_file_id="source-1",
) == len(content)
assert b"".join(storage.iter_bytes(
staged.reference,
expected_task_id="task-1",
expected_source_file_id="source-1",
expected_size=len(content),
start=4,
length=6,
chunk_size=2,
)) == b"456789"
with pytest.raises(DataProcessStorageError, match="owner mismatch"):
storage.file_size(
staged.reference,
expected_task_id="another-task",
expected_source_file_id="source-1",
)
with pytest.raises(DataProcessStorageError, match="does not match metadata"):
b"".join(storage.iter_bytes(
staged.reference,
expected_task_id="task-1",
expected_source_file_id="source-1",
expected_size=len(content) + 1,
))
@pytest.mark.parametrize(
"reference",
[
"local://data-process/../source-1/v1/file.txt",
"local://data-process/task-1/source-1/v1/file%2Fname.txt",
"local://data-process/task-1/source-1/v1/file.txt?download=1",
"local://data-process/task-1/source-1/v1/file.txt#fragment",
"https://data-process/task-1/source-1/v1/file.txt",
],
ids=[
"parent-traversal",
"percent-encoded-slash",
"query",
"fragment",
"wrong-scheme",
],
)
def test_unsafe_references_are_rejected(tmp_path: Path, reference: str) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
with pytest.raises(DataProcessStorageError):
storage.read(reference)
with pytest.raises(DataProcessStorageError):
storage.delete(reference)
def test_publish_rejects_intermediate_directory_symlink(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
outside = tmp_path / "outside"
outside.mkdir()
staged = _stage(storage, task_id="linked-task")
_create_symlink(
storage.root / "linked-task",
outside,
target_is_directory=True,
)
with pytest.raises(DataProcessStorageError, match="symlink|non-directory"):
storage.publish([staged])
assert list(outside.iterdir()) == []
_assert_staging_empty(storage)
def test_target_symlink_is_never_followed_or_deleted(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
staged = _stage(storage, task_id="task-link", source_file_id="source-link")
outside_file = tmp_path / "outside.txt"
outside_file.write_bytes(b"outside sentinel")
final_path = storage.root.joinpath(*staged._relative_path.parts)
final_path.parent.mkdir(parents=True)
_create_symlink(final_path, outside_file)
with pytest.raises(DataProcessStorageError, match="already exists"):
storage.publish([staged])
with pytest.raises(DataProcessStorageError, match="regular file"):
storage.read(staged.reference)
with pytest.raises(DataProcessStorageError, match="non-regular"):
storage.delete(staged.reference)
assert final_path.is_symlink()
assert outside_file.read_bytes() == b"outside sentinel"
_assert_staging_empty(storage)
def test_publish_rolls_back_first_object_when_second_target_collides(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
existing = _stage(
storage,
batch_id="batch-existing",
source_file_id="source-existing",
content=b"existing content",
)
storage.publish([existing])
first = _stage(
storage,
batch_id="batch-new",
source_file_id="source-new",
content=b"must be rolled back",
)
colliding_second = _stage(
storage,
batch_id="batch-new",
source_file_id="source-existing",
content=b"must not replace existing content",
)
with pytest.raises(DataProcessStorageError, match="already exists"):
storage.publish([first, colliding_second])
with pytest.raises(DataProcessStorageError, match="does not exist"):
storage.read(first.reference)
assert storage.read(existing.reference) == b"existing content"
_assert_staging_empty(storage)
def test_publish_rejects_manually_forged_staged_object(tmp_path: Path) -> None:
storage = LocalDataProcessStorage(tmp_path / "storage")
temporary_path = storage.root / ".staging" / "batch-forged" / "forged.tmp"
temporary_path.parent.mkdir()
temporary_path.write_bytes(b"forged content")
relative_path = PurePosixPath("task-forged", "source-forged", "v1", "forged.txt")
forged = StagedSourceObject(
reference="local://data-process/task-forged/source-forged/v1/forged.txt",
_temporary_path=temporary_path,
_relative_path=relative_path,
)
with pytest.raises(DataProcessStorageError, match="was not issued"):
storage.publish([forged])
with pytest.raises(DataProcessStorageError, match="was not issued"):
storage.discard([forged])
assert temporary_path.read_bytes() == b"forged content"
assert not storage.root.joinpath(*relative_path.parts).exists()
def test_relative_storage_configuration_is_anchored_to_backend_root(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
relative_configuration = Path("relative-storage") / tmp_path.name
backend_root = Path(storage_module.__file__).resolve().parents[3]
monkeypatch.chdir(tmp_path)
monkeypatch.setenv("DATA_PROCESS_STORAGE_DIR", str(relative_configuration))
storage_module.get_data_process_storage.cache_clear()
try:
configured_root = storage_module._configured_storage_root()
assert configured_root == backend_root / relative_configuration
assert not configured_root.exists()
finally:
storage_module.get_data_process_storage.cache_clear()

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,85 @@
from __future__ import annotations
from llama_index.core.embeddings import MockEmbedding
from app.modules.data_process.document_chunking import (
DocumentChunk,
_compact_with_offsets,
_project_layout_span,
chunk_fixed_text,
chunk_semantic_text,
merge_short_chunks,
)
def test_fixed_splitter_preserves_offsets_and_token_limit() -> None:
text = "第一段说明苹果。第二段说明香蕉。\n第三段说明数据库。第四段说明索引。"
chunks = chunk_fixed_text(text, chunk_size=20, chunk_overlap=0)
assert len(chunks) > 1
assert all(chunk.source_start is not None for chunk in chunks)
assert all(chunk.source_end is not None for chunk in chunks)
assert all(
chunk.original_content == text[chunk.source_start : chunk.source_end]
for chunk in chunks
if chunk.source_start is not None and chunk.source_end is not None
)
assert all(chunk.token_count <= 20 for chunk in chunks)
def test_semantic_splitter_uses_llamaindex_and_reapplies_maximum_size() -> None:
text = "第一段讨论水果。第二段继续讨论香蕉。第三段讨论数据库。第四段讨论索引。"
chunks = chunk_semantic_text(
text,
chunk_size=30,
chunk_overlap=0,
breakpoint_percentile_threshold=95,
embed_model=MockEmbedding(embed_dim=8),
)
assert len(chunks) >= 2
assert all(chunk.token_count <= 30 for chunk in chunks)
assert "".join(chunk.original_content for chunk in chunks) == text
def test_layout_projection_ignores_layout_whitespace_but_keeps_source_lines() -> None:
source = "标题\n第一条 这是正文。\n第二条 后续正文。"
compact_source, offsets = _compact_with_offsets(source)
start, end, cursor = _project_layout_span(
source,
"第一条\n这是正文。",
compact_source=compact_source,
source_offsets=offsets,
compact_start=0,
)
assert source[start:end] == "第一条 这是正文。"
assert cursor > 0
def test_short_layout_chunk_merges_with_neighbor_and_keeps_page_provenance() -> None:
source = "短标题\n这是一段足够长的正文内容,用于测试相邻切片合并。"
chunks = [
DocumentChunk("短标题", "短标题", 0, 3, 1, 1, 2, source_pages=(1,)),
DocumentChunk(
"这是一段足够长的正文内容,用于测试相邻切片合并。",
"这是一段足够长的正文内容,用于测试相邻切片合并。",
4,
len(source),
2,
2,
20,
source_pages=(1, 2),
),
]
merged = merge_short_chunks(
chunks,
source_text=source,
min_token_count=10,
max_token_count=100,
)
assert len(merged) == 1
assert merged[0].original_content == source
assert merged[0].source_pages == (1, 2)

81
compute/README.md Normal file
View File

@@ -0,0 +1,81 @@
# Compute Platform
算力平台与应用平台分开部署,本目录用于后续实现单机多 GPU 调度、文件网关和训练引擎适配。
## 目录结构
```text
compute/
api/ # 只允许应用平台访问的内部 Compute API
agent/ # 单机 Agent负责 GPU、进程、工作区管理
engines/
llama_factory/ # LLaMA-Factory 训练引擎适配器
file_gateway/ # 本地磁盘上传、下载、预览、离线导入
tests/
```
## 开发职责
- GPU 发现、状态上报、锁定和释放。
- 本地磁盘工作区管理。
- 创建、停止、查询训练/评测/推理/合并任务。
- LLaMA-Factory 命令生成、日志解析、产物收集。
- 分片上传、短时下载、离线导入。
- 通过服务间 token 接受应用平台调用。
## 运行模式
- 默认 `COMPUTE_EXECUTION_MODE=real`Compute API 会通过 `compute.agent.process_manager.ProcessManager` 启动真实 `llamafactory-cli train` 子进程,并将日志写入 `TRAINING_LOG_ROOT`
- 真实模式下 GPU 发现优先使用宿主机 `nvidia-smi`。如果部署环境暂时无法调用 `nvidia-smi`,可通过 `COMPUTE_GPU_COUNT``COMPUTE_GPU_NAME``COMPUTE_GPU_MEMORY_GB``COMPUTE_GPU_POWER_LIMIT_W` 声明兼容 GPU 清单,便于应用侧先完成节点登记和联调。
- 仅隔离联调时可设置 `COMPUTE_EXECUTION_MODE=simulator`,启用内存状态机和合成 GPU/日志数据。该模式不得作为生产运行路径。
- 服务间鉴权默认开启:设置 `COMPUTE_AUTH_ENABLED=true` 和一致的 `COMPUTE_SERVICE_TOKEN`,应用侧会通过 `X-Compute-Token` 调用 Compute API。
- 真实训练作业会登记到 `TRAINING_LOG_ROOT/compute-jobs.json`。Compute API 重启后会恢复作业索引,继续提供状态、停止和日志查询。
- 同一算力节点内按 GPU ID 做轻量锁定;已有运行中作业占用的 GPU 不允许再次提交,避免同机多 GPU 场景下误复用。
真实执行前提:
- 镜像或宿主机环境中 `llamafactory-cli` 可执行。
- `LLAMA_FACTORY_HOME` 指向 LLaMA-Factory 工作目录。
- 基座模型路径和数据集名称/目录已经在算力服务器本地可访问。
- 应用侧训练任务中的 GPU、模型、数据集配置能映射到当前节点本地路径。
## 应用侧接入
应用平台通过“算力节点”页面维护每台 GPU 服务器的 `Compute API``File Gateway` 地址。点击连接测试时Backend API 会主动调用:
```text
GET /modelTF/v1/compute/health
GET /modelTF/compute/resources/gpus
```
连接成功后,应用侧会同步节点健康信息、能力标签和 GPU 清单到 PostgreSQL。多节点阶段仍按“每台算力服务器 = 单机多 GPU 节点”管理,每台服务器都部署 Compute API、Agent、File Gateway 契约和 LLaMA-Factory。
训练闭环:
```text
Frontend 创建/启动训练
-> Backend API 选择 compute_nodes 节点
-> Backend API POST /modelTF/compute/jobs 到目标 Compute API
-> Compute API 启动 llamafactory-cli 子进程
-> Backend Worker 定时 GET /modelTF/compute/jobs/{id}
-> Backend API 同步 fine_tune_tasks 状态、进度、PID、日志路径和产物索引
```
## 当前接口能力
日志接口:
```text
GET /modelTF/compute/jobs/{job_id}/logs?tail_lines=200
GET /modelTF/compute/jobs/{job_id}/logs?offset=0&limit=500
```
返回 `content``metrics``total_lines``offset``limit``has_more``next_offset`,用于前端增量刷新和日志平台采集。
文件导入:
```text
POST /modelTF/compute/files/import-local
```
该接口用于应用侧调度前把算力服务器本地可访问的模型/数据集路径导入到 `YG_FT_DATA_ROOT` 内部。目标路径会校验不能逃逸出 `YG_FT_DATA_ROOT`,源路径必须已存在于算力服务器本地或挂载目录。

1
compute/__init__.py Normal file
View File

@@ -0,0 +1 @@
"""Compute platform package."""

View File

@@ -0,0 +1 @@
"""Compute agent package."""

View File

@@ -0,0 +1,281 @@
from __future__ import annotations
import os
import json
import contextlib
import hashlib
import signal
import subprocess
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
TERMINAL_STATUSES = {"completed", "failed", "stopped"}
@dataclass
class ManagedProcess:
id: str
name: str
command: list[str]
work_dir: str
log_path: Path
output_dir: str
gpus: list[int]
process: subprocess.Popen[Any] | None
created_at: float
pid: int | None = None
status: str = "running"
progress: int = 5
artifacts: list[dict[str, Any]] = field(default_factory=list)
class ProcessManager:
def __init__(self, log_root: str) -> None:
self.log_root = Path(log_root)
self.log_root.mkdir(parents=True, exist_ok=True)
self.registry_path = self.log_root / "compute-jobs.json"
self.jobs: dict[str, ManagedProcess] = {}
self._load_registry()
def create_job(self, payload: dict[str, Any], command: list[str], work_dir: str) -> dict[str, Any]:
job_id = str(payload.get("id") or f"job_{int(time.time() * 1000)}")
if job_id in self.jobs and self.jobs[job_id].status not in TERMINAL_STATUSES:
raise ValueError(f"job {job_id} is already running")
output_dir = str(payload.get("output_dir") or f"/data/yg-ft/outputs/{payload.get('name', job_id)}")
Path(output_dir).mkdir(parents=True, exist_ok=True)
log_path = self.log_root / f"{job_id}.log"
env = os.environ.copy()
gpus = [int(item) for item in payload.get("gpus") or []]
locked = self.locked_gpus()
conflict = sorted(set(gpus).intersection(locked))
if conflict:
raise ValueError(f"gpu already locked: {conflict}")
if gpus:
env["CUDA_VISIBLE_DEVICES"] = ",".join(str(item) for item in gpus)
env.update({str(k): str(v) for k, v in payload.get("env", {}).items()})
cwd = work_dir if Path(work_dir).exists() else None
with log_path.open("ab") as log_file:
log_file.write(f"[INFO] starting job_id={job_id} command={' '.join(command)}\n".encode("utf-8"))
process = subprocess.Popen(
command,
cwd=cwd,
env=env,
stdout=log_file,
stderr=subprocess.STDOUT,
)
managed = ManagedProcess(
id=job_id,
name=str(payload.get("name") or job_id),
command=command,
work_dir=work_dir,
log_path=log_path,
output_dir=output_dir,
gpus=gpus,
process=process,
created_at=time.time(),
pid=process.pid,
progress=10,
)
self.jobs[job_id] = managed
data = self.serialize(managed)
self._save_registry()
return data
def get_job(self, job_id: str) -> dict[str, Any] | None:
job = self.jobs.get(job_id)
if not job:
return None
return self.serialize(job)
def list_jobs(self) -> list[dict[str, Any]]:
return [self.serialize(job) for job in self.jobs.values()]
def stop_job(self, job_id: str) -> dict[str, Any] | None:
job = self.jobs.get(job_id)
if not job:
return None
if job.status not in TERMINAL_STATUSES:
try:
if job.process is not None and os.name == "nt":
job.process.terminate()
elif job.pid is not None:
os.kill(job.pid, signal.SIGTERM)
if job.process is not None:
job.process.wait(timeout=10)
except Exception:
if job.process is not None:
job.process.kill()
elif job.pid is not None:
with contextlib.suppress(Exception):
os.kill(job.pid, signal.SIGKILL)
job.status = "stopped"
job.progress = min(job.progress, 99)
data = self.serialize(job)
self._save_registry()
return data
def logs(self, job_id: str) -> str:
job = self.jobs.get(job_id)
if not job or not job.log_path.exists():
return ""
return job.log_path.read_text(encoding="utf-8", errors="replace")
def serialize(self, job: ManagedProcess) -> dict[str, Any]:
code = job.process.poll() if job.process is not None else None
checkpoints = self._collect_checkpoints(job.output_dir)
if job.status not in TERMINAL_STATUSES:
if job.process is None and job.pid is not None and not self._pid_alive(job.pid):
job.status = "failed"
job.progress = min(job.progress, 99)
code = -1
elif code is None:
job.status = "running"
elapsed = max(0, int(time.time() - job.created_at))
job.progress = min(95, max(job.progress, 10 + elapsed // 6))
elif code == 0:
job.status = "completed"
job.progress = 100
job.artifacts = self._collect_artifacts(job.output_dir)
else:
job.status = "failed"
job.progress = min(job.progress, 99)
self._save_registry()
return {
"id": job.id,
"name": job.name,
"status": job.status,
"progress": job.progress,
"pid": job.pid,
"gpus": job.gpus,
"created_at": job.created_at,
"command": job.command,
"work_dir": job.work_dir,
"output_dir": job.output_dir,
"log_file": str(job.log_path),
"artifacts": job.artifacts,
"checkpoints": checkpoints,
"return_code": code,
}
def locked_gpus(self) -> set[int]:
locked: set[int] = set()
for job in self.jobs.values():
status = self.serialize(job)["status"]
if status in {"queued", "running"}:
locked.update(job.gpus)
return locked
def _collect_artifacts(self, output_dir: str) -> list[dict[str, Any]]:
root = Path(output_dir)
if not root.exists():
return []
artifacts: list[dict[str, Any]] = []
for path in root.rglob("*"):
if path.is_file():
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
size = path.stat().st_size
artifacts.append(
{
"path": str(path),
"name": path.name,
"size": size,
"size_bytes": size,
"checksum_sha256": digest.hexdigest(),
}
)
return artifacts[:200]
def _collect_checkpoints(self, output_dir: str) -> list[dict[str, Any]]:
root = Path(output_dir)
if not root.exists():
return []
checkpoints: list[dict[str, Any]] = []
for path in root.glob("checkpoint-*"):
if not path.is_dir():
continue
step = 0
try:
step = int(path.name.rsplit("-", 1)[-1])
except ValueError:
step = 0
size_bytes = sum(item.stat().st_size for item in path.rglob("*") if item.is_file())
checkpoints.append(
{
"step": step,
"name": path.name,
"path": str(path),
"size_bytes": size_bytes,
"create_time": path.stat().st_mtime,
}
)
return sorted(checkpoints, key=lambda item: (int(item.get("step") or 0), str(item.get("name") or "")))
def _save_registry(self) -> None:
items = []
for job in self.jobs.values():
items.append(
{
"id": job.id,
"name": job.name,
"command": job.command,
"work_dir": job.work_dir,
"log_path": str(job.log_path),
"output_dir": job.output_dir,
"gpus": job.gpus,
"pid": job.pid,
"created_at": job.created_at,
"status": job.status,
"progress": job.progress,
"artifacts": job.artifacts,
}
)
self.registry_path.write_text(json.dumps(items, ensure_ascii=False, indent=2), encoding="utf-8")
def _load_registry(self) -> None:
if not self.registry_path.exists():
return
try:
items = json.loads(self.registry_path.read_text(encoding="utf-8"))
except json.JSONDecodeError:
return
for item in items if isinstance(items, list) else []:
if not isinstance(item, dict):
continue
pid = item.get("pid")
status = item.get("status", "failed")
if status not in TERMINAL_STATUSES and pid and not self._pid_alive(int(pid)):
status = "failed"
job = ManagedProcess(
id=str(item["id"]),
name=str(item.get("name") or item["id"]),
command=[str(part) for part in item.get("command") or []],
work_dir=str(item.get("work_dir") or ""),
log_path=Path(item.get("log_path") or self.log_root / f"{item['id']}.log"),
output_dir=str(item.get("output_dir") or ""),
gpus=[int(gpu) for gpu in item.get("gpus") or []],
process=None,
pid=int(pid) if pid else None,
created_at=float(item.get("created_at") or time.time()),
status=status,
progress=int(item.get("progress") or 0),
artifacts=item.get("artifacts") or [],
)
self.jobs[job.id] = job
def _pid_alive(self, pid: int) -> bool:
if pid <= 0:
return False
try:
os.kill(pid, 0)
return True
except OSError:
return False

1
compute/api/__init__.py Normal file
View File

@@ -0,0 +1 @@
"""Compute API package."""

747
compute/api/main.py Normal file
View File

@@ -0,0 +1,747 @@
from __future__ import annotations
import os
import math
import hashlib
import shutil
import subprocess
import time
from pathlib import Path
from typing import Any
from fastapi import FastAPI, File, Form, HTTPException, Query, Request, UploadFile
from fastapi.responses import FileResponse, JSONResponse
from compute.agent.process_manager import ProcessManager
from compute.engines.llama_factory.adapter import build_command, parse_log_line, prepare_runtime_files
def create_app() -> FastAPI:
app = FastAPI(title="YG Fine-Tune Compute API")
jobs: dict[str, dict[str, Any]] = {}
route_prefix = os.getenv("MODELTF_ROUTE_PREFIX", "/modelTF").rstrip("/") or "/modelTF"
process_manager = ProcessManager(os.getenv("TRAINING_LOG_ROOT", "/opt/yg-ft/logs/training"))
@app.middleware("http")
async def compute_token_auth(request: Request, call_next):
token = os.getenv("COMPUTE_SERVICE_TOKEN", "")
auth_enabled = os.getenv("COMPUTE_AUTH_ENABLED", "true").lower() == "true"
public_paths = {f"{route_prefix}/health", "/health"}
if auth_enabled and token and request.url.path not in public_paths:
header_token = request.headers.get("x-compute-token", "")
auth_header = request.headers.get("authorization", "")
bearer_token = auth_header.removeprefix("Bearer ").strip() if auth_header.startswith("Bearer ") else ""
if header_token != token and bearer_token != token:
return JSONResponse({"detail": "invalid compute service token"}, status_code=401)
return await call_next(request)
def now() -> float:
return time.time()
def host_id() -> str:
return os.getenv("COMPUTE_HOST_ID", "gpu-node-01")
def execution_mode() -> str:
return os.getenv("COMPUTE_EXECUTION_MODE", os.getenv("COMPUTE_MODE", "real")).lower()
def _int_env(name: str, default: int) -> int:
raw = os.getenv(name)
if raw is None or raw == "":
return default
return int(raw)
def _float_env(name: str, default: float) -> float:
raw = os.getenv(name)
if raw is None or raw == "":
return default
return float(raw)
def _path_inside(root: Path, candidate: Path) -> bool:
try:
candidate.resolve().relative_to(root.resolve())
return True
except ValueError:
return False
def _llama_factory_version() -> str:
for command in (["llamafactory-cli", "version"], ["llamafactory-cli", "--version"]):
try:
result = subprocess.run(command, capture_output=True, text=True, timeout=5)
except Exception:
continue
output = (result.stdout or result.stderr).strip()
if result.returncode == 0 and output:
return output.splitlines()[0][:120]
return ""
def torch_cuda_status() -> dict[str, Any]:
try:
import torch # type: ignore[import-not-found]
except Exception as exc: # noqa: BLE001 - keep health endpoint resilient
return {
"available": False,
"device_count": 0,
"torch_version": "",
"torch_cuda_version": "",
"error": f"torch import failed: {exc}",
}
try:
available = bool(torch.cuda.is_available())
device_count = int(torch.cuda.device_count())
devices = []
for index in range(device_count):
props = torch.cuda.get_device_properties(index)
devices.append(
{
"index": index,
"name": props.name,
"memory_total_gb": round(props.total_memory / 1024 / 1024 / 1024, 2),
}
)
return {
"available": available,
"device_count": device_count,
"torch_version": str(torch.__version__),
"torch_cuda_version": str(torch.version.cuda or ""),
"devices": devices,
"error": "" if available else "torch cuda is not available",
}
except Exception as exc: # noqa: BLE001 - expose CUDA initialization failures
return {
"available": False,
"device_count": 0,
"torch_version": str(getattr(torch, "__version__", "")),
"torch_cuda_version": str(getattr(torch.version, "cuda", "") or ""),
"devices": [],
"error": str(exc),
}
def _slice_log_content(
content: str,
tail_lines: int | None = None,
offset: int | None = None,
limit: int | None = None,
) -> dict[str, Any]:
lines = content.splitlines()
total = len(lines)
if offset is not None or limit is not None:
start = max(0, offset or 0)
end = start + limit if limit else total
selected = lines[start:end]
else:
tail = tail_lines or 200
start = max(0, total - tail)
selected = lines[start:]
next_offset = start + len(selected)
return {
"content": "\n".join(selected),
"total_lines": total,
"offset": start,
"limit": len(selected),
"has_more": next_offset < total,
"next_offset": next_offset if next_offset < total else None,
}
def _safe_float(value: Any, default: float = 0) -> float:
try:
return float(str(value).replace("[N/A]", "").strip() or default)
except (TypeError, ValueError):
return default
def job_status(job: dict[str, Any]) -> dict[str, Any]:
if execution_mode() != "simulator":
return job
elapsed = max(0, int(now() - job["created_at"]))
if job["status"] not in {"stopped", "failed", "completed"}:
if elapsed < 5:
job["status"] = "queued"
job["progress"] = 12 + elapsed * 3
elif elapsed < 60:
job["status"] = "running"
job["progress"] = min(96, 25 + int((elapsed - 5) / 55 * 70))
else:
job["status"] = "completed"
job["progress"] = 100
job["logs"] = generate_logs(job)
return job
def generate_logs(job: dict[str, Any]) -> str:
progress = int(job.get("progress", 0) or 0)
points = max(1, min(80, progress))
lines = [
f"[INFO] compute_host_id={host_id()} job_id={job['id']} engine=llama_factory",
f"[INFO] command={' '.join(job['command'])}",
]
for step in range(1, points + 1):
if step % 4 != 0 and step != points:
continue
loss = max(0.11, 2.5 * math.exp(-step / 40))
grad_norm = 0.4 + (step % 5) * 0.04
lr = 0.0002 * max(0.05, 1 - step / 100)
epoch = round(step / points * 3, 4)
lines.append(
"{"
f"'loss': {loss:.4f}, 'grad_norm': {grad_norm:.4f}, "
f"'learning_rate': {lr:.8f}, 'epoch': {epoch:.4f}"
"}"
)
if job.get("status") == "completed":
lines.extend(
[
"***** train metrics *****",
"epoch = 3",
"train_loss = 0.1181",
"train_runtime = 1m 0s",
"***** train metrics end *****",
]
)
return "\n".join(lines)
def real_gpu_resources() -> list[dict[str, Any]]:
query = (
"index,uuid,name,memory.total,memory.used,utilization.gpu,"
"temperature.gpu,power.draw,power.limit"
)
try:
result = subprocess.run(
["nvidia-smi", f"--query-gpu={query}", "--format=csv,noheader,nounits"],
check=True,
capture_output=True,
text=True,
timeout=5,
)
except Exception:
return fallback_gpu_resources()
items: list[dict[str, Any]] = []
for line in result.stdout.splitlines():
parts = [part.strip() for part in line.split(",")]
if len(parts) < 9:
continue
idx, uuid, name, mem_total, mem_used, util, temp, power, power_limit = parts[:9]
total_gb = round(_safe_float(mem_total) / 1024, 2)
used_gb = round(_safe_float(mem_used) / 1024, 2)
memory_percent = round(used_gb / total_gb * 100, 1) if total_gb else 0
gpu_percent = int(_safe_float(util))
items.append(
{
"id": int(idx),
"gpu_index": int(idx),
"uuid": uuid,
"name": name,
"status": "busy" if gpu_percent >= 5 or used_gb > 1 else "idle",
"gpu_percent": gpu_percent,
"memory_used_gb": used_gb,
"memory_total_gb": total_gb,
"memory_percent": memory_percent,
"temperature": int(_safe_float(temp)),
"power_w": round(_safe_float(power), 1),
"power_limit_w": round(_safe_float(power_limit), 1),
"processes": [],
}
)
return items
def fallback_gpu_resources() -> list[dict[str, Any]]:
count = _int_env("COMPUTE_GPU_COUNT", 0)
if count <= 0:
return []
name = os.getenv("COMPUTE_GPU_NAME", "Configured GPU")
memory_total = _float_env("COMPUTE_GPU_MEMORY_GB", 80.0)
power_limit = _float_env("COMPUTE_GPU_POWER_LIMIT_W", 300.0)
return [
{
"id": idx,
"gpu_index": idx,
"uuid": f"GPU-{host_id().upper()}-{idx}",
"name": name,
"status": "idle",
"gpu_percent": 0,
"memory_used_gb": 0,
"memory_total_gb": memory_total,
"memory_percent": 0,
"temperature": _int_env("COMPUTE_GPU_BASE_TEMPERATURE", 35),
"power_w": 0,
"power_limit_w": power_limit,
"processes": [],
}
for idx in range(count)
]
def gpu_resources() -> list[dict[str, Any]]:
if execution_mode() != "simulator":
return real_gpu_resources()
active_jobs = [job_status(job) for job in jobs.values() if job["status"] in {"queued", "running"}]
gpus: list[dict[str, Any]] = []
for idx in range(4):
task = next((job for job in active_jobs if idx in job.get("gpus", [])), None)
busy = task is not None and task["status"] == "running"
reserved = task is not None and task["status"] == "queued"
gpus.append(
{
"id": idx,
"uuid": f"GPU-{host_id().upper()}-{idx}",
"name": os.getenv("COMPUTE_GPU_NAME", "NVIDIA A800-SXM4-80GB"),
"status": "busy" if busy else "reserved" if reserved else "idle",
"gpu_percent": 88 if busy else 25 if reserved else 4,
"memory_used_gb": 58 if busy else 12 if reserved else 2,
"memory_total_gb": 80,
"temperature": 61 if busy else 45 if reserved else 36,
"power_w": 215 if busy else 80 if reserved else 25,
"power_limit_w": 300,
"processes": [
{
"pid": task["pid"],
"name": "llamafactory-cli",
"task_name": task["name"],
"memory_used_gb": 58 if busy else 12,
}
]
if task
else [],
}
)
return gpus
def _validate_training_accelerator(payload: dict[str, Any]) -> tuple[list[str], list[str], dict[str, Any]]:
errors: list[str] = []
warnings: list[str] = []
if str(payload.get("engine") or payload.get("training_engine") or "llama_factory") == "smoke":
return errors, warnings, {}
requested_gpus = [int(item) for item in payload.get("gpus") or []]
if not requested_gpus:
warnings.append("no gpu selected; training will run on CPU")
return errors, warnings, {}
cuda = torch_cuda_status()
if not cuda.get("available"):
errors.append(f"torch cuda unavailable on compute node: {cuda.get('error') or 'unknown error'}")
device_count = int(cuda.get("device_count") or 0)
if device_count and max(requested_gpus) >= device_count:
errors.append(f"requested gpu index out of torch device range: requested={requested_gpus}, device_count={device_count}")
min_memory_gb = _float_env("MIN_TRAINING_GPU_MEMORY_GB", 4.0)
gpus = {int(item["gpu_index"]): item for item in gpu_resources() if "gpu_index" in item}
for gpu_index in requested_gpus:
gpu = gpus.get(gpu_index)
if not gpu:
errors.append(f"requested gpu not found by nvidia-smi: {gpu_index}")
continue
memory_total = float(gpu.get("memory_total_gb") or 0)
if memory_total and memory_total < min_memory_gb:
errors.append(
f"gpu {gpu_index} memory too small: {memory_total}GB < required {min_memory_gb}GB"
)
return errors, warnings, cuda
def _check_path_item(item: dict[str, Any]) -> dict[str, Any]:
path = Path(str(item.get("path") or ""))
exists = path.exists()
expected_type = str(item.get("type") or "any")
ok = exists
if exists and expected_type == "dir":
ok = path.is_dir()
if exists and expected_type == "file":
ok = path.is_file()
return {
"name": item.get("name") or "",
"path": str(path),
"type": expected_type,
"required": bool(item.get("required", True)),
"exists": exists,
"is_dir": path.is_dir() if exists else False,
"is_file": path.is_file() if exists else False,
"byte_size": sum(child.stat().st_size for child in path.rglob("*") if child.is_file()) if exists and path.is_dir() else path.stat().st_size if exists and path.is_file() else 0,
"ok": ok or not item.get("required", True),
}
def _job_preview(payload: dict[str, Any], check_paths: bool) -> dict[str, Any]:
warnings: list[str] = []
runtime_files: list[dict[str, str]] = []
command_payload = {**payload, "require_dataset_files": check_paths}
if check_paths:
try:
runtime_files = prepare_runtime_files(command_payload)
except OSError as exc:
return {
"valid": False,
"errors": [f"prepare runtime files failed: {exc}"],
"warnings": warnings,
"engine": str(payload.get("engine") or payload.get("training_engine") or "llama_factory"),
"command": [],
"command_text": "",
"work_dir": os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"),
"env": {},
"runtime_files": [],
"path_checks": [],
}
try:
command = build_command(command_payload, os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"))
except ValueError as exc:
return {
"valid": False,
"errors": [part.strip() for part in str(exc).split(";") if part.strip()],
"warnings": warnings,
"engine": str(payload.get("engine") or payload.get("training_engine") or "llama_factory"),
"command": [],
"command_text": "",
"work_dir": os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"),
"env": {},
"runtime_files": runtime_files,
"path_checks": [],
}
errors: list[str] = []
engine = str(payload.get("engine") or payload.get("training_engine") or "llama_factory")
path_checks: list[dict[str, Any]] = []
accelerator: dict[str, Any] = {}
if check_paths and engine != "smoke":
path_checks = [
_check_path_item(
{
"name": "model_name_or_path",
"path": payload.get("model_name_or_path") or payload.get("base_model") or payload.get("base_model_path") or "",
"type": "any",
"required": True,
}
)
]
if engine in {"merge", "export", "llama_factory_export"} and payload.get("adapter_name_or_path"):
path_checks.append(
_check_path_item(
{
"name": "adapter_name_or_path",
"path": payload.get("adapter_name_or_path"),
"type": "any",
"required": True,
}
)
)
if payload.get("dataset_dir"):
path_checks.append(
_check_path_item(
{
"name": "dataset_dir",
"path": payload.get("dataset_dir"),
"type": "dir",
"required": True,
}
)
)
output_dir = Path(str(payload.get("output_dir") or "/data/yg-ft/outputs/training-job"))
path_checks.append(
_check_path_item(
{
"name": "output_parent",
"path": str(output_dir.parent),
"type": "dir",
"required": False,
}
)
)
errors.extend(
[f"{item['name']} path not available: {item['path']}" for item in path_checks if not item["ok"] and item["required"]]
)
if shutil.which(command.command[0]) is None:
errors.append(f"training command not found: {command.command[0]}")
if not Path(command.work_dir).exists():
errors.append(f"llama_factory_home not found: {command.work_dir}")
if engine not in {"merge", "export", "llama_factory_export"}:
accelerator_errors, accelerator_warnings, accelerator = _validate_training_accelerator(payload)
errors.extend(accelerator_errors)
warnings.extend(accelerator_warnings)
elif engine == "smoke":
warnings.append("smoke engine skips model and dataset path checks")
return {
"valid": not errors,
"errors": errors,
"warnings": warnings,
"engine": engine,
"command": command.command,
"command_text": " ".join(command.command),
"work_dir": command.work_dir,
"env": command.env,
"runtime_files": runtime_files,
"accelerator": accelerator,
"path_checks": path_checks,
}
@app.get(f"{route_prefix}/health")
async def health_check() -> dict[str, str]:
return {
"status": "ok",
"compute_host_id": os.getenv("COMPUTE_HOST_ID", "unknown"),
}
@app.get("/health")
async def health_check_root() -> dict[str, str]:
return await health_check()
@app.get(f"{route_prefix}/v1/compute/health")
async def compute_health_check() -> dict[str, Any]:
data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft"))
dataset_root = Path(os.getenv("YG_FT_DATASET_ROOT", str(data_root / "datasets")))
output_root = Path(os.getenv("YG_FT_OUTPUT_ROOT", str(data_root / "outputs")))
llama_factory_home = Path(os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"))
gpu_items = gpu_resources()
torch_cuda = torch_cuda_status()
return {
"status": "ok",
"api_version": "v1",
"compute_host_id": os.getenv("COMPUTE_HOST_ID", "unknown"),
"app_callback_enabled": os.getenv("ENABLE_APP_CALLBACK", "false").lower() == "true",
"data_root": str(data_root),
"data_root_exists": data_root.exists(),
"model_root": os.getenv("YG_FT_MODEL_ROOT", str(data_root / "models")),
"dataset_root": str(dataset_root),
"dataset_root_exists": dataset_root.exists(),
"output_root": str(output_root),
"output_root_exists": output_root.exists(),
"log_root": os.getenv("TRAINING_LOG_ROOT", "/opt/yg-ft/logs/training"),
"llama_factory_home": str(llama_factory_home),
"llama_factory_home_exists": llama_factory_home.exists(),
"llama_factory_version": os.getenv("LLAMA_FACTORY_VERSION", ""),
"execution_mode": execution_mode(),
"gpu_count": _int_env("COMPUTE_GPU_COUNT", 0),
"nvidia_gpu_count": len(gpu_items),
"torch_cuda": torch_cuda,
"gpu_discovery_endpoint": f"{route_prefix}/compute/resources/gpus",
"capabilities": ["gpu_discovery", "torch_cuda_diagnostics", "llama_factory", "file_gateway", "job_polling"],
}
@app.get(f"{route_prefix}/v1/compute/jobs")
async def list_jobs_alias() -> dict[str, list[dict[str, Any]]]:
items = process_manager.list_jobs() if execution_mode() != "simulator" else [job_status(job) for job in jobs.values()]
return {"items": items}
@app.get(f"{route_prefix}/compute/resources/gpus")
async def list_gpus() -> dict[str, Any]:
return {"items": gpu_resources(), "compute_host_id": host_id()}
@app.get(f"{route_prefix}/v1/compute/resources/gpus")
async def list_gpus_v1() -> dict[str, Any]:
return {"items": gpu_resources(), "compute_host_id": host_id()}
@app.post(f"{route_prefix}/compute/jobs/preview")
async def preview_job(payload: dict[str, Any]) -> dict[str, Any]:
return _job_preview(payload, check_paths=False)
@app.post(f"{route_prefix}/compute/jobs/validate")
async def validate_job(payload: dict[str, Any]) -> dict[str, Any]:
return _job_preview(payload, check_paths=True)
@app.post(f"{route_prefix}/v1/compute/jobs/preview")
async def preview_job_v1(payload: dict[str, Any]) -> dict[str, Any]:
return await preview_job(payload)
@app.post(f"{route_prefix}/v1/compute/jobs/validate")
async def validate_job_v1(payload: dict[str, Any]) -> dict[str, Any]:
return await validate_job(payload)
@app.post(f"{route_prefix}/compute/files/check-paths")
async def check_paths(payload: dict[str, Any]) -> dict[str, Any]:
items = [_check_path_item(item) for item in payload.get("paths", []) if isinstance(item, dict)]
return {"valid": all(item["ok"] for item in items), "items": items}
@app.get(f"{route_prefix}/compute/files/list")
async def list_files(
root: str = Query(default="data"),
relative_path: str = Query(default=""),
directories_only: bool = Query(default=False),
) -> dict[str, Any]:
roots = {
"data": Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft")),
"models": Path(os.getenv("YG_FT_MODEL_ROOT", os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft") + "/models")),
"datasets": Path(os.getenv("YG_FT_DATASET_ROOT", os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft") + "/datasets")),
"outputs": Path(os.getenv("YG_FT_OUTPUT_ROOT", os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft") + "/outputs")),
}
base = roots.get(root)
if base is None:
raise HTTPException(status_code=400, detail="invalid root")
target = (base / relative_path.lstrip("/\\")).resolve()
if not _path_inside(base, target):
raise HTTPException(status_code=400, detail="path must stay inside selected root")
if not target.exists():
return {"root": root, "base_path": str(base), "relative_path": relative_path, "items": []}
items = []
for child in sorted(target.iterdir(), key=lambda path: (not path.is_dir(), path.name.lower())):
if directories_only and not child.is_dir():
continue
items.append(
{
"name": child.name,
"path": str(child),
"relative_path": str(child.relative_to(base)).replace("\\", "/"),
"type": "directory" if child.is_dir() else "file",
"byte_size": child.stat().st_size if child.is_file() else 0,
}
)
return {"root": root, "base_path": str(base), "relative_path": relative_path, "items": items}
@app.post(f"{route_prefix}/compute/jobs")
async def create_job(payload: dict[str, Any]) -> dict[str, Any]:
payload = {**payload, "require_dataset_files": True}
try:
prepare_runtime_files(payload)
except OSError as exc:
raise HTTPException(status_code=400, detail=f"prepare runtime files failed: {exc}")
try:
command = build_command(payload, os.getenv("LLAMA_FACTORY_HOME", "/app/LLaMA-Factory"))
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc))
job_id = str(payload.get("id") or f"job_{int(now() * 1000)}")
if execution_mode() != "simulator":
try:
return process_manager.create_job({**payload, "id": job_id}, command.command, command.work_dir)
except FileNotFoundError as exc:
raise HTTPException(status_code=500, detail=f"training command not found: {exc.filename}")
except ValueError as exc:
raise HTTPException(status_code=409, detail=str(exc))
job = {
"id": job_id,
"name": payload.get("name", job_id),
"status": "queued",
"progress": 10,
"pid": int(52000 + now() % 10000),
"gpus": payload.get("gpus") or [0],
"created_at": now(),
"command": command.command,
"work_dir": command.work_dir,
"artifacts": [],
"logs": "",
}
jobs[job_id] = job
return job_status(job)
@app.get(f"{route_prefix}/compute/jobs")
async def list_jobs() -> dict[str, Any]:
items = process_manager.list_jobs() if execution_mode() != "simulator" else [job_status(job) for job in jobs.values()]
return {"items": items}
@app.get(f"{route_prefix}/compute/jobs/{{job_id}}")
async def get_job(job_id: str) -> dict[str, Any]:
job = jobs.get(job_id)
if execution_mode() != "simulator":
job = process_manager.get_job(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
return job
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
return job_status(job)
@app.post(f"{route_prefix}/compute/jobs/{{job_id}}/stop")
async def stop_job(job_id: str) -> dict[str, Any]:
if execution_mode() != "simulator":
job = process_manager.stop_job(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
return job
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
job["status"] = "stopped"
job["progress"] = min(job.get("progress", 0), 99)
return job
@app.get(f"{route_prefix}/compute/jobs/{{job_id}}/logs")
async def job_logs(
job_id: str,
tail_lines: int | None = Query(default=200, ge=1, le=5000),
offset: int | None = Query(default=None, ge=0),
limit: int | None = Query(default=None, ge=1, le=5000),
) -> dict[str, Any]:
if execution_mode() != "simulator":
job = process_manager.get_job(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
content = process_manager.logs(job_id)
else:
job = jobs.get(job_id)
if not job:
raise HTTPException(status_code=404, detail="job not found")
job = job_status(job)
content = job["logs"]
window = _slice_log_content(content, tail_lines, offset, limit)
metrics = [parse_log_line(line) for line in window["content"].splitlines()]
return {"job_id": job_id, **window, "metrics": [m for m in metrics if m]}
@app.post(f"{route_prefix}/compute/files/upload")
async def upload_file(
file: UploadFile | None = File(default=None),
target_relative_path: str | None = Form(default=None),
resource_type: str | None = Form(default=None),
resource_id: str | None = Form(default=None),
) -> dict[str, Any]:
file_id = f"file_{int(now() * 1000)}"
data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft"))
data_root.mkdir(parents=True, exist_ok=True)
filename = Path(file.filename if file else file_id).name
if target_relative_path:
target = (data_root / target_relative_path.lstrip("/\\")).resolve()
if not _path_inside(data_root, target):
raise HTTPException(status_code=400, detail="target path must stay inside YG_FT_DATA_ROOT")
else:
target = data_root / "uploads" / f"{file_id}_{filename}"
if file:
target.parent.mkdir(parents=True, exist_ok=True)
with target.open("wb") as output:
while chunk := await file.read(1024 * 1024):
output.write(chunk)
else:
target.parent.mkdir(parents=True, exist_ok=True)
target.write_text("", encoding="utf-8")
return {
"id": file_id,
"resource_type": resource_type,
"resource_id": resource_id,
"status": "available",
"local_path": str(target),
"byte_size": target.stat().st_size,
"checksum_sha256": hashlib.sha256(target.read_bytes()).hexdigest() if target.is_file() else "",
}
@app.post(f"{route_prefix}/compute/files/import-local")
async def import_local_file(payload: dict[str, Any]) -> dict[str, Any]:
source = Path(str(payload.get("source_path") or ""))
if not source.exists():
raise HTTPException(status_code=404, detail="source path not found")
data_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft"))
data_root.mkdir(parents=True, exist_ok=True)
relative = str(payload.get("target_relative_path") or f"imports/{source.name}").lstrip("/\\")
target = (data_root / relative).resolve()
if not _path_inside(data_root, target):
raise HTTPException(status_code=400, detail="target path must stay inside YG_FT_DATA_ROOT")
target.parent.mkdir(parents=True, exist_ok=True)
if source.is_dir():
if target.exists():
shutil.rmtree(target)
shutil.copytree(source, target)
byte_size = sum(path.stat().st_size for path in target.rglob("*") if path.is_file())
checksum = ""
else:
shutil.copy2(source, target)
byte_size = target.stat().st_size
checksum = hashlib.sha256(target.read_bytes()).hexdigest()
return {
"id": str(payload.get("id") or f"file_{int(now() * 1000)}"),
"resource_type": payload.get("resource_type"),
"resource_id": payload.get("resource_id"),
"status": "available",
"local_path": str(target),
"byte_size": byte_size,
"checksum_sha256": checksum,
}
@app.get(f"{route_prefix}/compute/files/{{file_id}}/download")
async def download_file(file_id: str) -> FileResponse:
upload_root = Path(os.getenv("YG_FT_DATA_ROOT", "/data/yg-ft")) / "uploads"
matches = list(upload_root.glob(f"{file_id}_*"))
if not matches:
raise HTTPException(status_code=404, detail="file not found")
return FileResponse(matches[0])
return app
app = create_app()

View File

@@ -0,0 +1 @@
"""Training engine adapters package."""

View File

@@ -0,0 +1 @@
"""LLaMA-Factory engine adapter package."""

View File

@@ -0,0 +1,299 @@
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Any
@dataclass(frozen=True)
class LlamaFactoryCommand:
command: list[str]
work_dir: str
env: dict[str, str]
def _load_dataset_preview(path: Path) -> list[dict[str, Any]]:
if not path.exists():
return []
text = path.read_text(encoding="utf-8", errors="replace").strip()
if not text:
return []
if path.suffix.lower() == ".jsonl":
items: list[dict[str, Any]] = []
for line in text.splitlines()[:20]:
line = line.strip()
if not line:
continue
value = json.loads(line)
if isinstance(value, dict):
items.append(value)
return items
value = json.loads(text)
if isinstance(value, list):
return [item for item in value[:20] if isinstance(item, dict)]
if isinstance(value, dict):
return [value]
return []
def _validate_dataset_columns(config: dict[str, Any]) -> list[str]:
dataset_dir = config.get("dataset_dir")
dataset_info = config.get("dataset_info")
if not dataset_dir or not isinstance(dataset_info, dict):
return []
root = Path(str(dataset_dir))
errors: list[str] = []
for dataset_key, item in dataset_info.items():
if not isinstance(item, dict):
continue
file_name = item.get("file_name")
file_names = file_name if isinstance(file_name, list) else [file_name]
columns = item.get("columns") if isinstance(item.get("columns"), dict) else {}
required_columns = [str(value) for value in columns.values() if value]
for name in file_names:
if not name:
continue
path = root / str(name).lstrip("/\\")
if not path.exists():
continue
try:
preview_rows = _load_dataset_preview(path)
except Exception as exc: # noqa: BLE001 - expose malformed data as validation error
errors.append(f"dataset file parse failed: {path}: {exc}")
continue
if not preview_rows:
errors.append(f"dataset file has no valid object records: {path}")
continue
available = set().union(*(row.keys() for row in preview_rows))
missing = [column for column in required_columns if column not in available]
if missing:
errors.append(
f"dataset columns missing in {path.name} for {dataset_key}: {', '.join(sorted(set(missing)))}"
)
return errors
def validate_config(config: dict[str, Any]) -> list[str]:
errors: list[str] = []
if not config.get("base_model") and not config.get("model_name_or_path"):
errors.append("base_model or model_name_or_path is required")
if not config.get("dataset") and not config.get("dataset_dir"):
errors.append("dataset or dataset_dir is required")
try:
learning_rate = float(config.get("learning_rate", 0.0002))
except (TypeError, ValueError):
learning_rate = 0
if learning_rate <= 0:
errors.append("learning_rate must be greater than zero")
try:
epochs = int(config.get("n_epochs", config.get("num_train_epochs", 1)))
except (TypeError, ValueError):
epochs = 0
if epochs <= 0:
errors.append("n_epochs must be greater than zero")
dataset_dir = config.get("dataset_dir")
dataset_info = config.get("dataset_info")
if config.get("require_dataset_files") and dataset_dir and isinstance(dataset_info, dict):
root = Path(str(dataset_dir))
for dataset_key, item in dataset_info.items():
if not isinstance(item, dict):
errors.append(f"dataset_info entry must be object: {dataset_key}")
continue
file_name = item.get("file_name")
file_names = file_name if isinstance(file_name, list) else [file_name]
for name in file_names:
if not name:
errors.append(f"dataset_info file_name is required: {dataset_key}")
continue
path = root / str(name).lstrip("/\\")
if not path.exists():
errors.append(f"dataset file not found: {path}")
errors.extend(_validate_dataset_columns(config))
return errors
def _optional_arg(config: dict[str, Any], command: list[str], option: str, *keys: str) -> None:
for key in keys:
value = config.get(key)
if value is not None and value != "":
command.extend([option, str(value)])
return
def _optional_bool_arg(config: dict[str, Any], command: list[str], option: str, *keys: str) -> None:
for key in keys:
value = config.get(key)
if value is True or str(value).lower() == "true":
command.extend([option, "true"])
return
def _normalize_stage(config: dict[str, Any]) -> str:
raw = str(config.get("stage") or config.get("train_type") or "sft").strip().lower()
return {
"sft": "sft",
"dpo": "dpo",
"cpt": "pt",
"pt": "pt",
"pretrain": "pt",
"rm": "rm",
"ppo": "ppo",
"kto": "kto",
}.get(raw, raw or "sft")
def prepare_runtime_files(config: dict[str, Any]) -> list[dict[str, str]]:
dataset_dir = config.get("dataset_dir")
dataset_info = config.get("dataset_info")
if not dataset_dir or not isinstance(dataset_info, dict):
return []
root = Path(str(dataset_dir))
root.mkdir(parents=True, exist_ok=True)
path = root / "dataset_info.json"
existing: dict[str, Any] = {}
if path.exists():
try:
loaded = json.loads(path.read_text(encoding="utf-8"))
existing = loaded if isinstance(loaded, dict) else {}
except json.JSONDecodeError:
existing = {}
existing.update(dataset_info)
path.write_text(json.dumps(existing, ensure_ascii=False, indent=2), encoding="utf-8")
return [{"name": "dataset_info", "path": str(path)}]
def build_command(config: dict[str, Any], llama_factory_home: str = "/app/LLaMA-Factory") -> LlamaFactoryCommand:
engine = str(config.get("engine") or config.get("training_engine") or "llama_factory")
if engine in {"merge", "export", "llama_factory_export"}:
model_path = config.get("base_model") or config.get("model_name_or_path") or config.get("base_model_path")
adapter_path = config.get("adapter_name_or_path") or config.get("adapter_path") or config.get("lora_path")
output_dir = config.get("output_dir") or config.get("export_dir")
errors: list[str] = []
if not model_path:
errors.append("base_model or model_name_or_path is required")
if not adapter_path and engine == "merge":
errors.append("adapter_name_or_path or adapter_path is required")
if not output_dir:
errors.append("output_dir or export_dir is required")
if errors:
raise ValueError("; ".join(errors))
command = [
"llamafactory-cli",
"export",
"--model_name_or_path",
str(model_path),
"--template",
str(config.get("template", "qwen")),
"--finetuning_type",
str(config.get("train_method", config.get("finetuning_type", "lora"))),
"--export_dir",
str(output_dir),
"--export_size",
str(config.get("export_size", 2)),
"--export_device",
str(config.get("export_device", "cpu")),
"--export_legacy_format",
str(config.get("export_legacy_format", False)).lower(),
]
if adapter_path:
command.extend(["--adapter_name_or_path", str(adapter_path)])
quantization_bit = int(config.get("export_quantization_bit", config.get("quantization_bit", 0)) or 0)
if quantization_bit in {4, 8}:
command.extend(["--quantization_bit", str(quantization_bit)])
return LlamaFactoryCommand(command=command, work_dir=str(Path(llama_factory_home)), env={})
errors = validate_config(config)
if errors:
raise ValueError("; ".join(errors))
if engine == "smoke":
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'training-smoke')}"
script = (
"import json, os, time; "
f"out={str(output_dir)!r}; "
"os.makedirs(out, exist_ok=True); "
"print('[INFO] smoke training started', flush=True); "
"\nfor step in range(1, 7):\n"
" loss=round(1.8/(step+1), 4)\n"
" lr=round(0.0002*(1-step/10), 8)\n"
" print({'loss': loss, 'grad_norm': round(0.4 + step*0.03, 4), 'learning_rate': lr, 'epoch': round(step/6, 4)}, flush=True)\n"
" time.sleep(0.4)\n"
"\nopen(os.path.join(out, 'adapter_config.json'), 'w', encoding='utf-8').write(json.dumps({'engine':'smoke','status':'completed'})); "
"print('***** train metrics *****', flush=True); "
"print('train_loss = 0.12', flush=True); "
"print('***** train metrics end *****', flush=True)"
)
return LlamaFactoryCommand(command=["python", "-u", "-c", script], work_dir="/app", env={})
model_path = config.get("base_model") or config.get("model_name_or_path")
dataset = config.get("dataset") or config.get("dataset_name")
dataset_dir = config.get("dataset_dir")
output_dir = config.get("output_dir") or f"/data/yg-ft/outputs/{config.get('name', 'training-job')}"
command = [
"llamafactory-cli",
"train",
"--stage",
_normalize_stage(config),
"--do_train",
"true",
"--model_name_or_path",
str(model_path),
"--dataset",
str(dataset or "default"),
"--template",
str(config.get("template", "qwen")),
"--finetuning_type",
str(config.get("train_method", config.get("finetuning_type", "lora"))),
"--output_dir",
str(output_dir),
"--per_device_train_batch_size",
str(config.get("batch_size", 2)),
"--learning_rate",
str(config.get("learning_rate", 0.0002)),
"--num_train_epochs",
str(config.get("n_epochs", 3)),
"--save_steps",
str(config.get("save_steps", 50)),
"--logging_steps",
str(config.get("logging_steps", 10)),
"--overwrite_output_dir",
"true",
"--plot_loss",
"true",
]
if dataset_dir:
command.extend(["--dataset_dir", str(dataset_dir)])
eval_dataset = config.get("eval_dataset")
if eval_dataset:
command.extend(["--eval_dataset", str(eval_dataset), "--do_eval", "true"])
_optional_arg(config, command, "--cutoff_len", "max_length", "cutoff_len")
_optional_arg(config, command, "--lr_scheduler_type", "lr_scheduler_type")
_optional_arg(config, command, "--warmup_ratio", "warmup_ratio")
_optional_arg(config, command, "--weight_decay", "weight_decay")
_optional_arg(config, command, "--lora_rank", "lora_rank", "rank")
_optional_arg(config, command, "--lora_alpha", "lora_alpha")
_optional_arg(config, command, "--lora_dropout", "lora_dropout")
_optional_arg(config, command, "--gradient_accumulation_steps", "gradient_accumulation_steps")
if not eval_dataset:
_optional_arg(config, command, "--val_size", "val_size")
_optional_arg(config, command, "--max_samples", "max_samples")
_optional_arg(config, command, "--preprocessing_num_workers", "preprocessing_num_workers")
_optional_bool_arg(config, command, "--fp16", "fp16")
_optional_bool_arg(config, command, "--bf16", "bf16")
quantization_bit = int(config.get("quantization_bit", 0) or 0)
if quantization_bit in {4, 8}:
command.extend(["--quantization_bit", str(quantization_bit)])
return LlamaFactoryCommand(command=command, work_dir=str(Path(llama_factory_home)), env={})
def parse_log_line(line: str) -> dict[str, float] | None:
if "loss" not in line or "learning_rate" not in line:
return None
result: dict[str, float] = {}
for key in ["loss", "grad_norm", "learning_rate", "epoch"]:
match = re.search(rf"['\"]?{key}['\"]?\s*:\s*([-+]?\d+(?:\.\d+)?(?:[eE][-+]?\d+)?)", line)
if match:
result[key] = float(match.group(1))
return result or None

View File

@@ -0,0 +1 @@
"""Local file gateway package."""

6
compute/requirements.txt Normal file
View File

@@ -0,0 +1,6 @@
fastapi>=0.111.0
uvicorn[standard]>=0.30.0
python-multipart>=0.0.9
pydantic>=2.7.0
python-dotenv>=1.0.1
httpx>=0.27.0

View File

@@ -0,0 +1 @@
"""Compute platform tests package."""

View File

@@ -0,0 +1,23 @@
from __future__ import annotations
from compute.engines.llama_factory.adapter import build_command
def test_build_command_uses_explicit_validation_dataset_without_resplitting() -> None:
result = build_command(
{
"base_model": "/models/qwen",
"dataset": "ygft_dataset_train",
"eval_dataset": "ygft_dataset_validation",
"dataset_dir": "/datasets/example",
"output_dir": "/outputs/example",
"val_size": 0.1,
}
)
assert result.command[result.command.index("--dataset") + 1] == "ygft_dataset_train"
assert result.command[result.command.index("--eval_dataset") + 1] == (
"ygft_dataset_validation"
)
assert "--do_eval" in result.command
assert "--val_size" not in result.command

299
docker/README.md Normal file
View File

@@ -0,0 +1,299 @@
# Docker 部署说明
本目录按应用服务器和算力服务器拆分 Dockerfile 与 Docker Compose 文件。Compose 文件不包含 `build:`,不会在 `docker compose up` 时自动构建业务镜像。所有业务镜像需要先通过手动 `docker build` 构建,再由 Compose 启动。
## 基础镜像
| 镜像 | 用途 |
| --- | --- |
| `python:3.12-slim` | 应用后端基础镜像,后端运行环境要求 Python 3.12 及以上 |
| `nginx:1.27-alpine` | 前端静态资源与 `/modelTF` 反向代理运行镜像 |
| `hiyouga/llamafactory:latest` | 算力服务基础镜像,基于 LLaMA-Factory 官方镜像扩展 Compute API |
| `postgres:16-alpine` | 开发阶段内置 PostgreSQL |
| `redis:7-alpine` | 开发阶段内置 Redis |
一键拉取基础镜像:
```bash
docker pull python:3.12-slim && \
docker pull nginx:1.27-alpine && \
docker pull hiyouga/llamafactory:latest && \
docker pull postgres:16-alpine && \
docker pull redis:7-alpine
```
Windows PowerShell
```powershell
$images = @(
"python:3.12-slim",
"nginx:1.27-alpine",
"hiyouga/llamafactory:latest",
"postgres:16-alpine",
"redis:7-alpine"
)
$images | ForEach-Object { docker pull $_ }
```
如果部署环境不能访问外网,需要提前在可联网环境执行上述拉取命令,再用 `docker save` / `docker load` 导出导入。
## 业务镜像
| 镜像 | Dockerfile | 构建命令 |
| --- | --- | --- |
| `yg-ft-backend-api:latest` | `docker/app/Dockerfile.backend` | `docker build -f docker/app/Dockerfile.backend -t yg-ft-backend-api:latest .` |
| `yg-ft-frontend-runtime:latest` | `docker/app/Dockerfile.frontend` | `docker build -f docker/app/Dockerfile.frontend -t yg-ft-frontend-runtime:latest .` |
| `yg-ft-compute-api:latest` | `docker/compute/Dockerfile.compute` | `docker build -f docker/compute/Dockerfile.compute -t yg-ft-compute-api:latest .` |
## 对外端口
所有宿主机对外端口统一使用 5 位端口。容器内部端口保持镜像默认端口,便于容器内服务和健康检查稳定。
| 服务 | 宿主机对外端口 | 容器内部端口 | 说明 |
| --- | --- | --- | --- |
| 前端 Nginx | `16801` | `80` | 前端页面入口 |
| 后端 API | `17861` | `8000` | FastAPI 服务 |
| PostgreSQL | `15432` | `5432` | 开发阶段内置数据库 |
| Redis | `16379` | `6379` | 开发阶段内置缓存 |
| Compute API | `19100` | `9100` | 算力服务器 API |
| File Gateway | `19101` | `9100` | 当前由 Compute API 暴露文件网关契约,后续可拆为独立服务 |
注意:`8000` 是后端容器内部端口,不作为宿主机对外访问端口。宿主机或浏览器应访问 `http://<app-server-ip>:17861/modelTF/health`;前端 Nginx 容器在 Docker 网络内部访问 `http://backend-api:8000/modelTF/...`
对应配置文件:
- `docker/app/.env.example`
- `FRONTEND_PORT=16801`
- `BACKEND_API_PORT=17861`
- `POSTGRES_PORT=15432`
- `REDIS_PORT=16379`
- `docker/compute/.env.example`
- `COMPUTE_API_PORT=19100`
- `FILE_GATEWAY_PORT=19101`
## 运行模式
- 应用侧默认 `COMPUTE_MODE=real`,任务状态必须由真实算力同步逻辑更新。
- 算力侧默认 `COMPUTE_EXECUTION_MODE=real`,真实执行器未完成前不会伪造训练作业。
- 仅隔离联调时可显式设置 `COMPUTE_MODE=simulator``COMPUTE_EXECUTION_MODE=simulator`,该模式不得用于测试环境、生产环境或生产升级基线。
## 应用服务器部署
应用服务器包含前端 Nginx、Backend API、PostgreSQL、Redis。
当前 Compose 内置 PostgreSQL 使用 `backend/app/db/sql/001_platform_runtime.sql` 初始化运行库。`docs/postgres-schema.sql` 是完整目标架构设计,不应直接挂载为当前运行库初始化脚本,否则会与当前后端代码的运行表结构不兼容。
首次部署:
```bash
cd <repo-root>
# 1. 使用当前 Windows/宿主机 npm 构建前端静态产物
cd frontend
npm ci
npm run build
cd ..
# 2. 手动构建业务镜像
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 .
# 3. 启动应用服务
cd docker/app
cp .env.example .env
docker compose up -d
```
后端镜像构建过程中会执行依赖导入自检,确认 `fastapi``uvicorn``psycopg``sqlalchemy``redis` 等运行依赖已安装。构建后也可以手动检查:
```bash
docker run --rm yg-ft-backend-api:latest python -c "import psycopg; print(psycopg.__version__)"
```
默认访问地址:
```text
http://<app-server-ip>:16801
```
应用侧代码和数据外挂:
```text
../../backend -> /app
../../frontend/dist -> /usr/share/nginx/html
../../runtime/app/logs/backend -> /opt/yg-ft/logs/backend
../../runtime/app/data -> /data/yg-ft
```
前端容器启动前必须确保 `../../frontend/dist/index.html` 已存在。若前端 Nginx 日志出现 `directory index of "/usr/share/nginx/html/" is forbidden``rewrite or internal redirection cycle while internally redirecting to "/index.html"`,通常表示当前执行 `docker compose` 的项目目录下没有构建好的 `frontend/dist`,或挂载路径不是同一份代码目录。
```bash
# 在执行 docker compose 的同一份代码目录中检查
cd <repo-root>/frontend
npm run build
test -f dist/index.html && ls -lh dist/index.html
cd ../docker/app
docker compose up -d --force-recreate frontend
docker compose logs --tail=80 frontend
```
如果使用 Windows npm 构建、WSL 中运行 Docker Compose需要确认 Windows 路径和 WSL 路径指向同一份仓库。例如在 `D:\...\YG_FT\frontend` 构建不会自动生成 `/mnt/d/wuyongtao/Code/YG_FT/frontend/dist` 下的产物,除非二者本就是同一个目录。
如果使用企业统一 PostgreSQL/Redis修改 `docker/app/.env`
如果前端 Nginx 日志出现 `open() "/usr/share/nginx/html/modelTF/login" failed``open() "/usr/share/nginx/html/login" failed`,说明当前容器没有加载项目的 Nginx 代理配置,`/modelTF/*` 被当成静态文件查找。处理方式:
```bash
cd <repo-root>/docker/app
docker compose up -d --force-recreate frontend
docker compose exec frontend nginx -T | grep -n "location.*modelTF" -A12
```
正常配置中应存在 `location ^~ /modelTF/`,并代理到 `BACKEND_PROXY_PASS`,默认是 `http://backend-api:8000`
```env
DATABASE_URL=postgresql+psycopg://<user>:<password>@<postgres-host>:15432/<db>
REDIS_URL=redis://<redis-host>:16379/0
USE_BUILTIN_POSTGRES=false
USE_BUILTIN_REDIS=false
```
生产环境如完全使用外部基础设施,可以删除或注释 Compose 中的 `postgres``redis` 服务及 `backend-api.depends_on` 中对应依赖。
## 算力服务器部署
算力服务器包含 Compute API、后续 Compute Agent、File Gateway、GPU runtime、本地训练数据目录和 LLaMA-Factory。`Dockerfile.compute` 基于 LLaMA-Factory 官方镜像:
```dockerfile
FROM hiyouga/llamafactory:latest
```
部署前需要安装:
- NVIDIA Driver
- NVIDIA Container Toolkit
- Docker Engine 和 Docker Compose Plugin
- 本地训练数据目录,默认 `/data/yg-ft`
首次部署:
```bash
cd <repo-root>
# 手动构建算力业务镜像
docker build -f docker/compute/Dockerfile.compute -t yg-ft-compute-api:latest .
# 启动算力服务
cd docker/compute
cp .env.example .env
docker compose up -d
```
健康检查:
```text
GET http://<compute-server-ip>:19100/modelTF/health
GET http://<compute-server-ip>:19100/modelTF/v1/compute/health
```
算力侧代码和数据外挂:
```text
../../compute -> /app/compute
${YG_FT_DATA_ROOT_HOST} -> /data/yg-ft
${YG_FT_MODEL_ROOT_HOST} -> /data/yg-ft/models
${YG_FT_DATASET_ROOT_HOST} -> /data/yg-ft/datasets
${YG_FT_OUTPUT_ROOT_HOST} -> /data/yg-ft/outputs
${COMPUTE_LOG_ROOT_HOST} -> /opt/yg-ft/logs/compute
${TRAINING_LOG_ROOT_HOST} -> /opt/yg-ft/logs/training
```
算力服务器启动前必须先在宿主机创建持久化目录,基座模型、训练数据、训练产物和训练日志都应落在宿主机磁盘上,不能只写入容器层。推荐默认目录:
```bash
cd <repo-root>/docker/compute
mkdir -p data/yg-ft/models \
data/yg-ft/datasets \
data/yg-ft/outputs \
data/yg-ft/logs/compute \
data/yg-ft/logs/training
```
默认 `docker/compute/.env.example` 使用 `./data/yg-ft`,该相对路径以 `docker/compute/docker-compose.yml` 所在目录为基准,因此实际宿主机目录是 `<repo-root>/docker/compute/data/yg-ft`。如企业环境模型盘、数据盘、产物盘分盘挂载,可在 `docker/compute/.env` 中分别调整 `YG_FT_MODEL_ROOT_HOST``YG_FT_DATASET_ROOT_HOST``YG_FT_OUTPUT_ROOT_HOST``COMPUTE_LOG_ROOT_HOST``TRAINING_LOG_ROOT_HOST`,容器内路径建议保持 `/data/yg-ft/models``/data/yg-ft/datasets``/data/yg-ft/outputs`,避免训练参数和节点配置复杂化。
页面上传数据集时,文件先进入 Backend API再由 Backend API 调用目标算力节点的 `POST /modelTF/compute/files/upload`,写入容器内 `/data/yg-ft/datasets/{dataset_id}/`。在默认开发配置下,宿主机可在 `<repo-root>/docker/compute/data/yg-ft/datasets/{dataset_id}/` 看到对应文件。仅创建 bind mount 不会自动让应用侧上传文件出现在算力目录,必须通过这条 File Gateway 链路同步。
## 应用与算力分离部署
应用服务器只需要主动访问算力服务器,不要求算力服务器回调应用服务器。
`docker/app/.env` 中配置:
```env
COMPUTE_API_BASE_URL=http://<compute-server-ip>:19100
FILE_GATEWAY_BASE_URL=http://<compute-server-ip>:19101
COMPUTE_SERVICE_TOKEN=change_me
COMPUTE_STATUS_SYNC_MODE=polling
COMPUTE_POLL_INTERVAL_SECONDS=3
COMPUTE_POLL_BATCH_SIZE=100
```
交互链路:
```text
Frontend
-> Backend API
-> Compute API
-> Compute Agent / LLaMA-Factory
-> 本地数据目录 / 模型目录 / 训练产物
<- Backend Worker 定时轮询 Compute API
```
算力服务默认开启服务间鉴权。`docker/compute/.env` 中保持 `COMPUTE_AUTH_ENABLED=true`,并确保 `COMPUTE_SERVICE_TOKEN``docker/app/.env` 一致;健康检查路径仍可用于容器探活。
## 多算力节点部署
多算力节点仍按“单机多 GPU 节点”部署。每台 GPU 服务器都独立部署一套 `docker/compute`
```text
gpu-node-01: docker/compute + /data/yg-ft + 19100/19101
gpu-node-02: docker/compute + /data/yg-ft + 19100/19101
gpu-node-03: docker/compute + /data/yg-ft + 19100/19101
```
节点之间默认不互访。应用平台主动访问每个节点的 Compute API/File Gateway并通过 `compute_nodes``resource_replicas``resource_sync_jobs` 统一调度和同步。
节点地址、权重、标签、启用状态和本地路径在前端“算力节点”页面动态维护。新增或编辑节点后,点击“测试”会由 Backend API 主动访问该节点的 `GET /modelTF/v1/compute/health``GET /modelTF/compute/resources/gpus`,并把健康信息与 GPU 清单同步到 PostgreSQL。
## 常用命令
重新构建应用镜像:
```bash
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 .
```
重新构建算力镜像:
```bash
docker build -f docker/compute/Dockerfile.compute -t yg-ft-compute-api:latest .
```
启动服务:
```bash
cd docker/app
docker compose up -d
cd ../compute
docker compose up -d
```
查看服务:
```bash
docker compose ps
docker compose logs -f
```

46
docker/app/.env Normal file
View File

@@ -0,0 +1,46 @@
APP_ENV=prod
APP_NAME=YG Fine-Tune Platform API
MODELTF_ROUTE_PREFIX=/modelTF
CORS_ALLOW_ORIGINS=http://localhost:16801,http://127.0.0.1:16801
FRONTEND_IMAGE=yg-ft-frontend-runtime:latest
BACKEND_API_IMAGE=yg-ft-backend-api:latest
# Five-digit host ports exposed outside the application server.
FRONTEND_PORT=16801
BACKEND_API_PORT=17861
POSTGRES_PORT=15432
REDIS_PORT=16379
POSTGRES_DB=yg_ft
POSTGRES_USER=root
POSTGRES_PASSWORD=8811614287327Leo
DATABASE_URL=postgresql+psycopg://root:8811614287327Leo@www.caoxiaozhu.com:5432/yg_ft
REDIS_URL=redis://redis:6379/0
# PostgreSQL uses the shared external database. The local postgres service is disabled in docker-compose.yml.
# Redis still uses the built-in service during current development.
USE_BUILTIN_POSTGRES=false
USE_BUILTIN_REDIS=true
LOG_LEVEL=INFO
LOG_DIR=/opt/yg-ft/logs/backend
LOG_FILE_PREFIX=backend
LOG_ERROR_FILE_PREFIX=error
LOG_MAX_BYTES=20971520
LOG_RETENTION_DAYS=10
BACKEND_PROXY_PASS=http://backend-api:8000
# Split deployment: set these to the compute server address, for example http://10.10.20.31:19100.
COMPUTE_API_BASE_URL=http://compute-api:9100
COMPUTE_SERVICE_TOKEN=change_me
FILE_GATEWAY_BASE_URL=http://compute-api:9101
# The application side polls Compute API for job state to avoid opening reverse network access.
COMPUTE_MODE=real
COMPUTE_STATUS_SYNC_MODE=polling
COMPUTE_POLL_INTERVAL_SECONDS=10
COMPUTE_POLL_BATCH_SIZE=100
COMPUTE_REQUEST_TIMEOUT_SECONDS=5

View File

@@ -0,0 +1,21 @@
FROM python:3.12-slim
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1
WORKDIR /app
COPY backend/requirements.txt /tmp/requirements.txt
RUN pip install --upgrade pip -i https://pypi.tuna.tsinghua.edu.cn/simple \
&& pip install -r /tmp/requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple \
&& rm -f /tmp/requirements.txt
RUN python -c "import fastapi, uvicorn, psycopg, sqlalchemy, redis, jwt, passlib, httpx, alembic; print('backend dependency check ok')"
RUN mkdir -p /opt/yg-ft/logs/backend /data/yg-ft \
&& chmod -R 0775 /opt/yg-ft /data/yg-ft
EXPOSE 8000
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]

View File

@@ -0,0 +1,9 @@
FROM nginx:1.27-alpine
RUN mkdir -p /usr/share/nginx/html
EXPOSE 80
HEALTHCHECK --interval=30s --timeout=3s --start-period=10s --retries=3 \
CMD test -f /usr/share/nginx/html/index.html && wget -qO- http://127.0.0.1/index.html >/dev/null || exit 1

View File

@@ -0,0 +1,126 @@
services:
frontend:
image: ${FRONTEND_IMAGE:-yg-ft-frontend-runtime:latest}
container_name: yg-ft-frontend
depends_on:
backend-api:
condition: service_started
ports:
- "${FRONTEND_PORT:-16801}:80"
environment:
BACKEND_PROXY_PASS: ${BACKEND_PROXY_PASS:-http://backend-api:8000}
volumes:
- ../../frontend/dist:/usr/share/nginx/html:ro
- ../nginx.conf.template:/etc/nginx/templates/default.conf.template:ro
command:
- /bin/sh
- -c
- |
if [ ! -f /usr/share/nginx/html/index.html ]; then
echo "frontend dist is missing: build frontend first and ensure ../../frontend/dist is mounted";
ls -la /usr/share/nginx/html;
exit 1;
fi;
envsubst '$$BACKEND_PROXY_PASS' < /etc/nginx/templates/default.conf.template > /etc/nginx/conf.d/default.conf;
nginx -t;
nginx -g 'daemon off;'
networks:
- yg-ft-app
restart: unless-stopped
backend-api:
image: ${BACKEND_API_IMAGE:-yg-ft-backend-api:latest}
container_name: yg-ft-backend-api
depends_on:
redis:
condition: service_healthy
expose:
- "8000"
ports:
- "${BACKEND_API_PORT:-17861}:8000"
environment:
APP_ENV: ${APP_ENV:-prod}
APP_NAME: ${APP_NAME:-YG Fine-Tune Platform API}
MODELTF_ROUTE_PREFIX: ${MODELTF_ROUTE_PREFIX:-/modelTF}
CORS_ALLOW_ORIGINS: ${CORS_ALLOW_ORIGINS:-http://localhost:16801,http://127.0.0.1:16801}
DATABASE_URL: ${DATABASE_URL:-postgresql+psycopg://root:8811614287327Leo@www.caoxiaozhu.com:5432/yg_ft}
REDIS_URL: ${REDIS_URL:-redis://redis:6379/0}
USE_BUILTIN_POSTGRES: ${USE_BUILTIN_POSTGRES:-false}
USE_BUILTIN_REDIS: ${USE_BUILTIN_REDIS:-true}
LOG_LEVEL: ${LOG_LEVEL:-INFO}
LOG_DIR: ${LOG_DIR:-/opt/yg-ft/logs/backend}
LOG_FILE_PREFIX: ${LOG_FILE_PREFIX:-backend}
LOG_ERROR_FILE_PREFIX: ${LOG_ERROR_FILE_PREFIX:-error}
LOG_MAX_BYTES: ${LOG_MAX_BYTES:-20971520}
LOG_RETENTION_DAYS: ${LOG_RETENTION_DAYS:-10}
COMPUTE_API_BASE_URL: ${COMPUTE_API_BASE_URL:-http://compute-api:9100}
COMPUTE_SERVICE_TOKEN: ${COMPUTE_SERVICE_TOKEN:-change_me}
FILE_GATEWAY_BASE_URL: ${FILE_GATEWAY_BASE_URL:-http://compute-api:9101}
COMPUTE_MODE: ${COMPUTE_MODE:-real}
COMPUTE_STATUS_SYNC_MODE: ${COMPUTE_STATUS_SYNC_MODE:-polling}
COMPUTE_POLL_INTERVAL_SECONDS: ${COMPUTE_POLL_INTERVAL_SECONDS:-3}
COMPUTE_POLL_BATCH_SIZE: ${COMPUTE_POLL_BATCH_SIZE:-100}
COMPUTE_REQUEST_TIMEOUT_SECONDS: ${COMPUTE_REQUEST_TIMEOUT_SECONDS:-5}
PYTHONPATH: /app
volumes:
- ../../backend:/app:ro
- ../../runtime/app/logs/backend:/opt/yg-ft/logs/backend
- ../../runtime/app/data:/data/yg-ft
networks:
- yg-ft-app
healthcheck:
test: ["CMD-SHELL", "python -c \"import urllib.request; urllib.request.urlopen('http://127.0.0.1:8000/modelTF/health', timeout=3).read()\""]
interval: 30s
timeout: 5s
retries: 3
start_period: 20s
restart: unless-stopped
# PostgreSQL uses the shared external database configured by DATABASE_URL in docker/app/.env.
# Keep this local service commented out unless development needs an isolated database again.
# postgres:
# image: postgres:16-alpine
# container_name: yg-ft-postgres
# environment:
# POSTGRES_DB: ${POSTGRES_DB:-yg_ft}
# POSTGRES_USER: ${POSTGRES_USER:-yg_ft}
# POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:-change_me}
# PGDATA: /var/lib/postgresql/data/pgdata
# volumes:
# - postgres_data:/var/lib/postgresql/data
# - ../../backend/app/db/sql/001_platform_runtime.sql:/docker-entrypoint-initdb.d/001-platform-runtime.sql:ro
# ports:
# - "${POSTGRES_PORT:-15432}:5432"
# networks:
# - yg-ft-app
# healthcheck:
# test: ["CMD-SHELL", "pg_isready -U $${POSTGRES_USER} -d $${POSTGRES_DB}"]
# interval: 10s
# timeout: 5s
# retries: 5
# restart: unless-stopped
redis:
image: redis:7-alpine
container_name: yg-ft-redis
command: ["redis-server", "--appendonly", "yes"]
volumes:
- redis_data:/data
ports:
- "${REDIS_PORT:-16379}:6379"
networks:
- yg-ft-app
healthcheck:
test: ["CMD", "redis-cli", "ping"]
interval: 10s
timeout: 3s
retries: 5
restart: unless-stopped
networks:
yg-ft-app:
name: yg-ft-app
volumes:
# postgres_data:
redis_data:

40
docker/compute/.env Normal file
View File

@@ -0,0 +1,40 @@
COMPUTE_ENV=prod
COMPUTE_HOST_ID=gpu-node-01
COMPUTE_EXECUTION_MODE=real
MODELTF_ROUTE_PREFIX=/modelTF
# Five-digit host ports exposed outside the compute server.
COMPUTE_API_PORT=19100
FILE_GATEWAY_PORT=19101
COMPUTE_API_IMAGE=yg-ft-compute-api:latest
# The application server actively polls Compute API; compute server does not need reverse access.
COMPUTE_AUTH_ENABLED=true
COMPUTE_SERVICE_TOKEN=change_me
ENABLE_APP_CALLBACK=false
# LLaMA-Factory is provided by the official hiyouga/llamafactory base image.
LLAMA_FACTORY_HOME=/app/LLaMA-Factory
YG_FT_DATA_ROOT=/data/yg-ft
YG_FT_DATA_ROOT_HOST=./data/yg-ft
YG_FT_MODEL_ROOT=/data/yg-ft/models
YG_FT_MODEL_ROOT_HOST=./data/yg-ft/models
YG_FT_DATASET_ROOT=/data/yg-ft/datasets
YG_FT_DATASET_ROOT_HOST=./data/yg-ft/datasets
YG_FT_OUTPUT_ROOT=/data/yg-ft/outputs
YG_FT_OUTPUT_ROOT_HOST=./data/yg-ft/outputs
TRAINING_LOG_ROOT=/opt/yg-ft/logs/training
TRAINING_LOG_ROOT_HOST=./data/yg-ft/logs/training
COMPUTE_LOG_ROOT_HOST=./data/yg-ft/logs/compute
# Optional fallback used when nvidia-smi is unavailable.
# Leave COMPUTE_GPU_COUNT=0 on real GPU servers with working NVIDIA runtime.
COMPUTE_GPU_COUNT=0
COMPUTE_GPU_NAME=NVIDIA A800-SXM4-80GB
COMPUTE_GPU_MEMORY_GB=80
COMPUTE_GPU_POWER_LIMIT_W=300
LOG_DIR=/opt/yg-ft/logs/compute
CUDA_VISIBLE_DEVICES=all
NVIDIA_VISIBLE_DEVICES=all
NVIDIA_DRIVER_CAPABILITIES=compute,utility

View File

@@ -0,0 +1,44 @@
COMPUTE_ENV=prod
COMPUTE_HOST_ID=gpu-node-01
COMPUTE_EXECUTION_MODE=real
MODELTF_ROUTE_PREFIX=/modelTF
# Five-digit host ports exposed outside the compute server.
COMPUTE_API_PORT=19100
FILE_GATEWAY_PORT=19101
COMPUTE_API_IMAGE=yg-ft-compute-api:latest
# The application server actively polls Compute API; compute server does not need reverse access.
COMPUTE_AUTH_ENABLED=true
COMPUTE_SERVICE_TOKEN=change_me
ENABLE_APP_CALLBACK=false
# LLaMA-Factory is provided by the official hiyouga/llamafactory base image.
LLAMA_FACTORY_HOME=/app/LLaMA-Factory
# Persistent host directories on the compute server.
# Create these directories before starting docker compose. They are mounted into
# the container so base models, datasets, training outputs and logs survive
# container recreation or image upgrades.
YG_FT_DATA_ROOT=/data/yg-ft
YG_FT_DATA_ROOT_HOST=./data/yg-ft
YG_FT_MODEL_ROOT=/data/yg-ft/models
YG_FT_MODEL_ROOT_HOST=./data/yg-ft/models
YG_FT_DATASET_ROOT=/data/yg-ft/datasets
YG_FT_DATASET_ROOT_HOST=./data/yg-ft/datasets
YG_FT_OUTPUT_ROOT=/data/yg-ft/outputs
YG_FT_OUTPUT_ROOT_HOST=./data/yg-ft/outputs
TRAINING_LOG_ROOT=/opt/yg-ft/logs/training
TRAINING_LOG_ROOT_HOST=./data/yg-ft/logs/training
COMPUTE_LOG_ROOT_HOST=./data/yg-ft/logs/compute
# Optional fallback used when nvidia-smi is unavailable.
# Leave COMPUTE_GPU_COUNT=0 on real GPU servers with working NVIDIA runtime.
COMPUTE_GPU_COUNT=0
COMPUTE_GPU_NAME=NVIDIA A800-SXM4-80GB
COMPUTE_GPU_MEMORY_GB=80
COMPUTE_GPU_POWER_LIMIT_W=300
LOG_DIR=/opt/yg-ft/logs/compute
CUDA_VISIBLE_DEVICES=all
NVIDIA_VISIBLE_DEVICES=all
NVIDIA_DRIVER_CAPABILITIES=compute,utility

View File

@@ -0,0 +1,27 @@
FROM hiyouga/llamafactory:latest
ENV DEBIAN_FRONTEND=noninteractive \
PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1
WORKDIR /app
RUN apt-get update \
&& apt-get install -y --no-install-recommends tini \
&& rm -rf /var/lib/apt/lists/*
COPY compute/requirements.txt /tmp/requirements.txt
RUN pip install --upgrade pip -i https://pypi.tuna.tsinghua.edu.cn/simple \
&& pip install -r /tmp/requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple \
&& rm -f /tmp/requirements.txt
RUN mkdir -p /opt/yg-ft/logs/compute /opt/yg-ft/logs/training /data/yg-ft /app/LLaMA-Factory \
&& chmod -R 0775 /opt/yg-ft /data/yg-ft /app/LLaMA-Factory
ENTRYPOINT ["/usr/bin/tini", "--"]
EXPOSE 9100
CMD ["uvicorn", "compute.api.main:app", "--host", "0.0.0.0", "--port", "9100"]

View File

@@ -0,0 +1,38 @@
# YG-FT Compute 数据目录说明
本目录挂载到 `yg-ft-compute-api` 容器的 `/data/yg-ft`,用于持久化存储训练相关的数据。
## 目录结构
```
data/yg-ft/
├── datasets/ # 数据集存储目录
├── models/ # 模型文件存储目录
├── outputs/ # 训练/推理输出结果目录
└── logs/
├── compute/ # 计算服务运行日志
└── training/ # 训练任务执行日志
```
## 各目录说明
### datasets/
训练和评估所使用的数据集文件,包括 JSON、JSONL、CSV 等格式。数据集由用户上传或通过平台创建,供 LLaMA-Factory 等训练引擎读取。
### models/
存放模型文件,包括:
- 预训练基座模型(如 LLaMA、Qwen 等)
- 微调后的自定义模型权重
- 合并后的部署模型
### outputs/
训练任务和推理任务的输出结果,包括:
- 训练过程中的 checkpoint 文件
- 评估结果和指标报告
- 推理生成的结果文本
### logs/compute/
计算服务compute-api的运行时日志用于排查服务启动、GPU 调度、健康检查等问题。
### logs/training/
各训练任务的执行日志,记录训练过程状态、报错信息等,便于追踪单个任务的运行情况。

View File

@@ -0,0 +1,2 @@
*
!.gitignore

View File

@@ -0,0 +1,2 @@
*
!.gitignore

View File

@@ -0,0 +1,2 @@
*
!.gitignore

View File

@@ -0,0 +1,2 @@
*
!.gitignore

View File

@@ -0,0 +1,2 @@
*
!.gitignore

View File

@@ -0,0 +1,53 @@
services:
compute-api:
image: ${COMPUTE_API_IMAGE:-yg-ft-compute-api:latest}
container_name: yg-ft-compute-api
gpus: all
ports:
- "${COMPUTE_API_PORT:-19100}:9100"
- "${FILE_GATEWAY_PORT:-19101}:9100"
environment:
COMPUTE_ENV: ${COMPUTE_ENV:-prod}
COMPUTE_HOST_ID: ${COMPUTE_HOST_ID:-gpu-node-01}
COMPUTE_EXECUTION_MODE: ${COMPUTE_EXECUTION_MODE:-real}
MODELTF_ROUTE_PREFIX: ${MODELTF_ROUTE_PREFIX:-/modelTF}
COMPUTE_AUTH_ENABLED: ${COMPUTE_AUTH_ENABLED:-true}
COMPUTE_SERVICE_TOKEN: ${COMPUTE_SERVICE_TOKEN:-change_me}
ENABLE_APP_CALLBACK: ${ENABLE_APP_CALLBACK:-false}
LLAMA_FACTORY_HOME: ${LLAMA_FACTORY_HOME:-/app/LLaMA-Factory}
YG_FT_DATA_ROOT: ${YG_FT_DATA_ROOT:-/data/yg-ft}
YG_FT_MODEL_ROOT: ${YG_FT_MODEL_ROOT:-/data/yg-ft/models}
YG_FT_DATASET_ROOT: ${YG_FT_DATASET_ROOT:-/data/yg-ft/datasets}
YG_FT_OUTPUT_ROOT: ${YG_FT_OUTPUT_ROOT:-/data/yg-ft/outputs}
TRAINING_LOG_ROOT: ${TRAINING_LOG_ROOT:-/opt/yg-ft/logs/training}
COMPUTE_GPU_COUNT: ${COMPUTE_GPU_COUNT:-0}
COMPUTE_GPU_NAME: ${COMPUTE_GPU_NAME:-NVIDIA A800-SXM4-80GB}
COMPUTE_GPU_MEMORY_GB: ${COMPUTE_GPU_MEMORY_GB:-80}
COMPUTE_GPU_POWER_LIMIT_W: ${COMPUTE_GPU_POWER_LIMIT_W:-300}
MIN_TRAINING_GPU_MEMORY_GB: ${MIN_TRAINING_GPU_MEMORY_GB:-4}
LOG_DIR: ${LOG_DIR:-/opt/yg-ft/logs/compute}
CUDA_VISIBLE_DEVICES: ${CUDA_VISIBLE_DEVICES:-all}
NVIDIA_VISIBLE_DEVICES: ${NVIDIA_VISIBLE_DEVICES:-all}
NVIDIA_DRIVER_CAPABILITIES: ${NVIDIA_DRIVER_CAPABILITIES:-compute,utility}
PYTHONPATH: /app
volumes:
- ../../compute:/app/compute:ro
- ${YG_FT_DATA_ROOT_HOST:-./data/yg-ft}:${YG_FT_DATA_ROOT:-/data/yg-ft}
- ${YG_FT_MODEL_ROOT_HOST:-./data/yg-ft/models}:${YG_FT_MODEL_ROOT:-/data/yg-ft/models}
- ${YG_FT_DATASET_ROOT_HOST:-./data/yg-ft/datasets}:${YG_FT_DATASET_ROOT:-/data/yg-ft/datasets}
- ${YG_FT_OUTPUT_ROOT_HOST:-./data/yg-ft/outputs}:${YG_FT_OUTPUT_ROOT:-/data/yg-ft/outputs}
- ${COMPUTE_LOG_ROOT_HOST:-./data/yg-ft/logs/compute}:${LOG_DIR:-/opt/yg-ft/logs/compute}
- ${TRAINING_LOG_ROOT_HOST:-./data/yg-ft/logs/training}:${TRAINING_LOG_ROOT:-/opt/yg-ft/logs/training}
networks:
- yg-ft-compute
healthcheck:
test: ["CMD-SHELL", "python -c \"import urllib.request; urllib.request.urlopen('http://127.0.0.1:9100/modelTF/health', timeout=3).read()\""]
interval: 30s
timeout: 5s
retries: 3
start_period: 20s
restart: unless-stopped
networks:
yg-ft-compute:
name: yg-ft-compute

View File

@@ -0,0 +1,41 @@
server {
listen 80;
server_name _;
root /usr/share/nginx/html;
index index.html;
client_max_body_size 200m;
location ^~ /modelTF/ {
proxy_pass ${BACKEND_PROXY_PASS};
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_read_timeout 300s;
proxy_send_timeout 300s;
}
location = /modelTF {
proxy_pass ${BACKEND_PROXY_PASS};
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_read_timeout 300s;
proxy_send_timeout 300s;
}
location ~* \.(?:js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf)$ {
try_files $uri =404;
expires 30d;
add_header Cache-Control "public, immutable";
}
location / {
try_files $uri $uri/ /index.html;
}
}

View File

@@ -0,0 +1,153 @@
# 2026-07-23 开发总结
本文档汇总 2026-07-23 当天围绕模型微调平台完成的系统开发内容。当前开发不再按临时 Demo 口径处理,以下能力均按后续可继续演进到生产系统的代码和数据结构推进。
## 1. 今日完成内容
### 1.1 模型训练最小闭环增强
| 功能点 | 作用 | 对应页面/模块 | 用户如何操作 |
| --- | --- | --- | --- |
| 训练任务创建前预检 | 在任务进入运行前校验模型、数据集、算力节点、路径和 LLaMA-Factory 参数,减少任务启动后才失败的情况 | 模型微调列表、创建模型微调、训练日志 | 在“模型微调”页面创建任务并提交,系统自动执行预检;失败时页面提示具体原因 |
| 训练数据同步到算力节点 | 将应用侧选择的数据集文件同步到目标 Compute 节点本地磁盘,保证 LLaMA-Factory 能读取到真实训练文件 | 创建模型微调、训练日志、算力节点/资源副本 | 用户选择数据集后提交训练任务,系统自动同步数据集到算力节点,无需手工进入容器复制 |
| 动态生成训练参数 | 根据用户在页面选择的训练数据、基座模型、训练方法、输出目录等动态生成 LLaMA-Factory 训练参数 | 创建模型微调、训练日志 | 用户在创建任务表单中选择模型和数据集,启动训练后可在训练日志中查看实际执行命令 |
| 训练日志和状态轮询 | 应用侧主动轮询 Compute API 获取训练状态、进度、日志、指标和 checkpoint | 训练日志详情页 | 用户进入训练日志页,可查看任务状态、日志输出、指标和 checkpoint 信息 |
| 训练失败信息回传 | 训练命令失败后将失败状态和日志内容回传应用侧,便于定位数据格式、模型路径或参数问题 | 模型微调列表、训练日志详情页 | 任务失败后,用户进入训练日志页查看失败日志和错误堆栈 |
### 1.2 B. 模型资产、训练与 LLaMA-Factory 任务能力
| 功能点 | 作用 | 对应页面/模块 | 用户如何操作 |
| --- | --- | --- | --- |
| 训练模型产物 artifact 展示 | 展示训练/合并产生的 Adapter、合并模型、量化模型等产物路径、大小、checksum | 模型管理 -> 训练模型列表 | 用户进入“模型管理”,切换到“训练模型”,展开某个模型行查看“模型产物” |
| 模型血缘 lineage 展示 | 记录和展示基座模型、训练模型、合并产物之间的来源关系 | 模型管理 -> 训练模型列表 | 用户展开训练模型行,在“模型血缘”区域查看父资源、子资源和对应 Compute Job |
| 模型导出任务 export job 状态展示 | 展示模型合并/导出的任务状态、输出目录和创建时间 | 模型管理 -> 训练模型列表 | 用户展开训练模型行,在“导出任务”区域查看导出任务状态 |
| export 完成后回填 checksum | Compute export 完成后扫描产物文件,回填大小和 `checksum_sha256`,为后续一致性校验和制品治理做准备 | 后端模型资产模块、模型管理页面 | 用户无需手工操作,任务完成后页面自动展示产物大小和 checksum |
| 合并模型任务接入 Compute | 模型合并不再只是生成模拟任务,而是通过 Compute 节点调用模型合并任务并记录导出任务 | 模型管理 -> 合并权重 | 用户在训练模型列表点击“合并权重”,提交后系统创建 Compute 合并任务 |
### 1.3 D. 算力平台、部署与运维能力
| 功能点 | 作用 | 对应页面/模块 | 用户如何操作 |
| --- | --- | --- | --- |
| 算力节点动态配置 | 支持维护 Compute API、File Gateway、权重、标签、启用状态、调度状态等信息 | 算力节点 -> 节点 | 用户进入“算力节点”,点击“新增节点”或“编辑”,填写节点地址和调度参数 |
| 节点连通性测试 | 验证应用服务器是否能访问 Compute API并返回 GPU 发现结果和延迟 | 算力节点 -> 节点 | 用户在节点列表点击“测试”,查看连接成功或失败提示 |
| GPU 状态展示 | 展示每个算力节点的 GPU 利用率、显存、温度、功耗、进程信息 | 算力节点 -> GPU | 用户进入“算力节点”切换到“GPU”页签查看 |
| 训练队列展示 | 展示 Compute 侧任务队列、状态、进度、节点和 GPU 分配 | 算力节点 -> 队列 | 用户切换到“队列”页签查看当前运行或等待任务 |
| 资源副本列表 | 展示数据集、模型、训练产物在算力节点本地磁盘上的副本路径和同步状态 | 算力节点 -> 资源副本 | 用户切换到“资源副本”页签,选择目标节点查看本地副本 |
| 资源副本漂移检测 | 检查副本路径是否仍存在、大小是否可读取,并将异常副本标记为漂移 | 算力节点 -> 资源副本 | 用户点击“漂移检测”,系统调用 Compute API 检查节点本地路径 |
| 资源副本 repair 重同步 | 对漂移、失败、待修复副本执行真实重同步,而不只是标记状态 | 算力节点 -> 资源副本 | 用户点击“修复副本”,系统从权威源重新上传或导入到目标算力节点 |
| 应用侧轮询模式 | 应用服务主动轮询 Compute API避免需要 Compute 侧反向访问应用侧 | 后端 Compute Gateway、部署架构 | 用户无感知,部署时只需保证应用侧能访问各 Compute 节点 |
| 本地磁盘挂载约定 | 基座模型、训练数据、训练产物、训练日志均通过宿主机目录挂载到 Compute 容器 | docker/compute 部署 | 运维人员在启动 Compute 服务前准备并挂载 `docker/compute/data/yg-ft/*` 目录 |
### 1.4 前端状态中文化
| 功能点 | 作用 | 对应页面/模块 | 用户如何操作 |
| --- | --- | --- | --- |
| 统一状态字典 | 将 `pending``running``completed``failed``online``synced``drifted` 等状态统一展示为中文 | 通用前端组件 `ModelStatusTag`、状态工具 `status.ts` | 用户在各列表页看到中文状态,不再直接看到英文状态值 |
| 模型管理状态中文化 | 合并状态、导出任务状态、产物状态使用中文展示 | 模型管理 | 用户查看训练模型列表和展开详情时看到“未合并/合并中/已合并/已完成/失败”等中文 |
| 算力节点状态中文化 | 节点状态、GPU 状态、队列状态、副本同步状态使用中文展示 | 算力节点 | 用户查看“节点/GPU/队列/资源副本”时看到中文状态标签 |
| 推理/对比/评测/用户状态中文化 | 将推理任务、对比模型加载、评测任务、用户状态统一为中文标签 | 模型推理、模型对比、模型评测、用户设置 | 用户查看相关页面时看到“启动中/已就绪/已完成/启用”等中文状态 |
### 1.5 前端构建产物更新
| 功能点 | 作用 | 对应页面/模块 | 用户如何操作 |
| --- | --- | --- | --- |
| 刷新 `frontend/dist` | 保证 Docker 前端 nginx 容器可以直接加载最新页面代码 | 前端部署 | 用户重新访问前端服务即可看到最新页面 |
| 重启前后端服务验证 | 验证源码挂载和 dist 挂载方式下服务可正常加载最新代码 | docker/app、docker/compute | 开发人员重启 `backend-api``frontend``compute-api` 后验证页面和接口 |
## 2. 今日涉及的主要代码位置
| 文件 | 说明 |
| --- | --- |
| `backend/app/api/v1/endpoints/platform.py` | 新增/完善训练预检、资源同步、模型 artifact/lineage/export job 接口、资源副本漂移检测和 repair |
| `backend/app/db/platform_store.py` | 新增模型产物、血缘、导出任务、副本同步结果、artifact 查询等数据访问能力 |
| `compute/agent/process_manager.py` | export/artifact 扫描时补充文件大小和 checksum |
| `frontend/src/api/modules/model.ts` | 增加训练模型 artifact、lineage、export job API |
| `frontend/src/api/modules/compute.ts` | 增加副本漂移检测、repair API 和副本字段 |
| `frontend/src/views/model/ModelManageView.vue` | 增加训练模型展开详情:产物、血缘、导出任务 |
| `frontend/src/views/compute/ComputeNodesView.vue` | 增加资源副本页签、漂移检测、repair 操作和状态中文化 |
| `frontend/src/utils/status.ts` | 新增统一状态中文映射和标签类型映射 |
| `frontend/src/components/ModelStatusTag.vue` | 改为复用统一状态字典 |
| `frontend/src/components/DataTablePage.vue` | 增加表格展开事件支持 |
| `frontend/dist/*` | 前端生产构建产物已更新 |
## 3. 今日验证结果
| 验证项 | 结果 |
| --- | --- |
| Python 编译检查 | 通过:`python -m compileall backend compute scripts` |
| 前端生产构建 | 通过:`npm run build` |
| 前端入口 | 通过:`http://127.0.0.1:16801/` 返回 200 |
| 后端健康检查 | 通过:`/modelTF/health` 返回正常 |
| Compute 健康检查 | 通过:`/modelTF/health` 返回正常 |
| 模型导出任务接口 | 通过:`/modelTF/model-manage/export-jobs` 返回正常 |
| 算力节点列表接口 | 通过:`/modelTF/compute/nodes` 返回正常 |
| 资源副本漂移检测 | 通过:当前节点副本检测返回 `drifted: 0` |
## 4. 当前仍需注意的问题
| 问题 | 影响 | 建议 |
| --- | --- | --- |
| 训练数据格式仍依赖 LLaMA-Factory 约定 | 如果用户上传的数据字段不符合模板要求,训练仍会失败 | 下一步增加数据集格式校验和模板转换预检 |
| 单机 MX350 显存较小 | 真实训练大模型时容易因显存不足失败 | 当前环境用于链路验证;真实训练应使用高显存 GPU 节点 |
| resource replica repair 对大型模型仍是同步调用 | 大模型重同步可能耗时较长 | 下一步将 repair 完整异步化,并展示 sync job 进度 |
| artifact checksum 目前在 Compute 扫描阶段计算 | 对超大目录递归扫描可能较慢 | 下一步支持分文件 checksum、manifest 文件和后台扫描 |
| 模型评测、模型推理仍未形成完整生产闭环 | 目前页面已有基础能力,但后端表结构、任务运行、日志、治理能力还需补齐 | 下一步将评测和推理纳入正式开发计划 |
## 5. 下一步开发计划
### 5.1 B. 模型资产、训练与 LLaMA-Factory 任务
| 优先级 | 开发任务 | 目标页面/模块 | 交付结果 |
| --- | --- | --- | --- |
| P0 | 数据集格式预检 | 创建模型微调、数据集管理 | 提交训练前校验 Alpaca/ShareGPT/OpenAI Messages 等格式,提示缺失字段 |
| P0 | 训练参数可视化确认 | 创建模型微调 | 提交前展示最终 LLaMA-Factory 参数预览,减少参数不一致问题 |
| P0 | 训练任务失败诊断 | 训练日志详情 | 识别常见错误:数据字段缺失、模型路径不存在、显存不足、依赖缺失,并生成中文诊断 |
| P1 | checkpoint 管理 | 训练日志详情、模型管理 | 展示 checkpoint 列表、大小、路径、保留策略,支持标记最佳 checkpoint |
| P1 | 模型产物 manifest | 模型管理 | 为每个训练/导出产物生成 manifest记录文件清单、大小、checksum、来源任务 |
| P1 | 合并/导出任务详情页 | 模型管理、训练日志 | 展示合并任务日志、状态、产物、失败原因 |
| P2 | 模型版本治理 | 模型管理 | 支持版本号、标签、发布状态、归档状态、审批状态 |
| P2 | 训练模板管理 | 系统设置或训练配置 | 将 LLaMA-Factory 模板、数据格式、默认超参做成可维护配置 |
### 5.2 D. 算力平台、部署与运维
| 优先级 | 开发任务 | 目标页面/模块 | 交付结果 |
| --- | --- | --- | --- |
| P0 | repair 异步化 | 算力节点 -> 资源副本 | repair 创建 sync job 后后台执行,页面展示进度和失败原因 |
| P0 | 多算力节点调度策略 | 算力节点、创建模型微调 | 支持按标签、权重、空闲 GPU、显存要求选择节点 |
| P1 | 节点资源水位告警 | 算力节点、硬件监控 | 展示磁盘、GPU、显存、训练日志目录水位和告警状态 |
| P1 | 节点维护窗口 | 算力节点 | 支持维护中节点不再调度新任务,已有任务可继续或迁移 |
| P1 | 文件副本治理 | 算力节点 -> 资源副本 | 支持副本清理、重建、过期策略和跨节点一致性检查 |
| P2 | 部署健康巡检脚本 | 部署运维 | 一键检查 app、backend、redis、pg、compute、GPU、挂载目录、端口连通 |
| P2 | Compute Agent 插件标准 | 算力平台 | 抽象 LLaMA-Factory 接入规范,预留其他训练框架 |
### 5.3 模型评测页面开发计划
| 优先级 | 开发任务 | 目标页面/模块 | 交付结果 |
| --- | --- | --- | --- |
| P0 | 评测任务表结构和接口补齐 | 模型评测列表、创建评测、评测详情 | 建立 `eval_tasks``eval_dimensions``eval_sample_results` 等运行表和接口 |
| P0 | 创建评测任务真实提交 | 创建评测 | 支持选择模型、数据集、评测维度、GPU/节点,提交后生成评测任务 |
| P0 | 评测任务运行闭环 | 评测详情 | 支持状态、进度、日志、样本级结果回传 |
| P1 | 评测维度管理 | 评测维度创建/编辑 | 支持规则、Prompt、评分器、权重、适用数据集配置 |
| P1 | 样本级评分展示 | 评测详情 | 展示每条样本的输入、模型输出、评分、原因、人工复核状态 |
| P1 | 综合报告生成 | 评测详情、排行榜 | 生成维度汇总、综合分、问题样本、改进建议 |
| P2 | 评测审批和审计 | 审批流、审计日志 | 评测任务创建、发布报告、删除报告纳入治理 |
### 5.4 模型推理页面开发计划
| 优先级 | 开发任务 | 目标页面/模块 | 交付结果 |
| --- | --- | --- | --- |
| P0 | 推理任务表结构和接口补齐 | 模型推理列表、新建推理 | 建立 `inference_tasks``inference_task_models``chat_sessions``chat_messages` |
| P0 | 模型加载/卸载真实闭环 | 模型推理列表 | 支持选择训练产物加载推理服务,展示加载状态、端口、进程和错误 |
| P0 | 单模型对话持久化 | 模型对话 | 保存会话、消息、参数、响应耗时、token 统计 |
| P1 | 多模型对比任务 | 模型对比 | 支持多模型同时加载、同一问题并发请求、结果对比展示 |
| P1 | 推理资源管控 | 模型推理、算力节点 | 支持 GPU 选择、并发限制、空闲自动卸载、异常进程清理 |
| P1 | 推理日志和调用审计 | 日志、审计中心 | 记录加载、卸载、对话请求、失败原因、用户和租户信息 |
| P2 | 推理服务发布 | 模型管理、模型推理 | 支持将某个训练模型发布为内部推理服务,并配置访问权限 |
## 6. 建议的下一阶段顺序
1. 先完成 B+D 的训练稳定性增强数据格式预检、训练参数预览、失败诊断、repair 异步化。
2. 再补齐模型评测的真实任务闭环:任务表、创建任务、运行状态、样本结果。
3. 然后补齐模型推理闭环:加载/卸载、对话持久化、多模型对比。
4. 最后统一治理能力:审批、审计、租户隔离、资源配额、保留策略和运维巡检。

View File

@@ -0,0 +1,281 @@
# 2026-07-24 工作计划
本文基于 `docs/2026-07-23-development-summary.md`、当前 B+D 开发进度,以及 2026-07-24 已完成的训练预检、失败诊断、资源修复异步化能力整理。后续开发仍按正式系统演进推进,不以临时演示能力作为交付标准。
## 1. 当前完成基线
### 1.1 训练创建与预检
- 对应页面:`模型微调 / 创建训练任务`
- 已完成能力:
- 创建训练前调用训练预检接口。
- 展示真实 LLaMA-Factory 命令预览。
- 展示预检节点、错误、警告和中文诊断建议。
- 后端预检不创建任务、不落库、不占用调度锁。
- 后端兼容 `base_model/train_dataset_id``model_id/dataset_id` 两套字段。
- 对应接口:
- `POST /modelTF/fine-tune/preflight`
- `POST /modelTF/fine-tune/command-preview`
### 1.2 训练日志与失败诊断
- 对应页面:`系统日志 / 训练日志`
- 已完成能力:
- 训练失败或停止后可查询诊断建议。
- 根据训练日志和失败原因识别模型路径、数据集字段、CUDA/GPU、LLaMA-Factory 命令等常见问题。
- 页面以中文展示失败诊断。
- 对应接口:
- `GET /modelTF/fine-tune/{task_id}/diagnostics`
### 1.3 算力节点与资源修复
- 对应页面:`算力平台 / 算力节点`
- 已完成能力:
- 算力节点可配置、可测试连接、可查看 GPU 与健康状态。
- 资源副本 repair 改为异步提交。
- 页面展示资源同步任务进度,避免长请求阻塞页面。
- 对应接口:
- `POST /modelTF/compute/nodes/{node_id}/replicas/repair`
- `GET /modelTF/compute/sync-jobs/{sync_id}`
## 2. 未完成任务清单
### P0真实训练成功闭环
- 对应页面:
- `模型管理 / 新增模型`
- `数据集管理 / 上传数据集`
- `模型微调 / 创建训练任务`
- `系统日志 / 训练日志`
- 未完成内容:
- 基座模型必须支持算力服务器本地路径校验,避免选择 API 模型或应用侧路径后进入训练。
- 数据集必须支持上传后格式校验,提前发现 Alpaca、ShareGPT、DPO、CPT 字段不匹配问题。
- 训练成功后需要完成模型产物扫描、产物入库、训练任务状态回填。
- 训练失败时需要强制拉取最后日志片段,保证失败原因可见。
- 后端开发:
- 增强模型路径校验,明确区分 `本地训练模型``API 模型``已训练模型`
- 增加数据集格式校验服务,支持字段级错误返回。
- 完善训练任务完成后的 artifact 回填逻辑。
- Compute 开发:
- 训练结束后扫描输出目录。
- 返回产物列表、文件大小、目录结构、训练日志路径。
- 验收标准:
- 使用算力节点可访问的本地模型路径和合法数据集,可以完成一次真实 LLaMA-Factory 训练。
- 训练成功后页面能看到完成状态、输出目录、模型产物。
- 训练失败时页面能看到中文诊断和最后错误日志。
### P0LLaMA-Factory 参数映射完善
- 对应页面:`模型微调 / 创建训练任务`
- 未完成内容:
- SFT、DPO、CPT 参数映射仍需细化。
- LoRA、Full、QLoRA、量化导出参数需要按训练方式校验。
- 不同模板与数据集格式之间的兼容关系需要预检。
- 后端开发:
- 建立训练参数标准化层。
- 建立训练方式到 LLaMA-Factory 参数的映射表。
- 对无效组合返回中文错误,例如 DPO 缺少 rejected 字段、CPT 不应使用 instruction/output 格式等。
- 验收标准:
- 页面选择不同训练方式时,预检能返回准确命令。
- 无效参数组合不能启动训练。
### P1资源副本 repair 自动重同步
- 对应页面:`算力平台 / 算力节点`
- 未完成内容:
- repair 当前已异步化,但还需要基于权威源路径自动重同步。
- 修复完成后需要重新校验副本状态。
- 后端开发:
- 为模型、数据集、训练产物定义权威源路径。
- repair job 根据权威源自动发起重传或重新扫描。
- 修复完成后更新 replica 状态、checksum、错误原因。
- Compute 开发:
- 支持接收重同步请求。
- 支持按资源类型写入目标路径并返回校验信息。
- 验收标准:
- 将副本标记为异常后,点击修复可自动完成重同步并恢复为正常。
### P1artifact checksum 与 manifest
- 对应页面:
- `模型管理 / 已训练模型`
- `模型管理 / 模型详情`
- `算力平台 / 资源副本`
- 未完成内容:
- artifact checksum 目前仍是预留字段。
- 大目录需要 manifest 文件,避免每次递归扫描成本过高。
- 后端开发:
- 增加 artifact checksum 回填逻辑。
- 增加 manifest 解析和存储字段。
- Compute 开发:
- export 或训练完成后扫描文件并生成 checksum。
- 对大模型目录生成 manifest。
- 验收标准:
- 模型产物列表能展示 checksum、大小、文件数、生成时间。
- 副本校验可以基于 checksum 判断一致性。
### P1模型导出闭环
- 对应页面:
- `模型管理 / 已训练模型`
- `模型管理 / 导出任务`
- `模型管理 / 合并权重`
- 未完成内容:
- 导出任务、量化导出、导出日志、失败重试、产物下载仍需完善。
- 后端开发:
- 完善 export job 创建、查询、取消、重试接口。
- 导出完成后登记 artifact 和 lineage。
- Compute 开发:
- 支持 LoRA 合并、GGUF/量化导出、导出日志回传。
- 验收标准:
- 已训练模型可发起导出。
- 导出状态、日志、产物可在页面查看。
### P1多算力节点调度增强
- 对应页面:`算力平台 / 算力节点`
- 未完成内容:
- 当前已有节点配置和基础调度,但生产级调度策略仍需增强。
- 需要支持节点标签、权重、启用状态、容量、手动指定节点。
- 后端开发:
- 增强调度策略标签匹配、权重、当前任务数、GPU 占用、显存约束。
- 增加任务排队和等待原因。
- 增加 GPU 分配释放的异常恢复。
- 前端开发:
- 创建训练时支持可选手动指定节点。
- 算力节点页展示容量、排队数、当前任务。
- 验收标准:
- 多节点时可以自动选择合适节点。
- 节点不可用时页面能明确展示不可调度原因。
### P2训练日志实时性优化
- 对应页面:`系统日志 / 训练日志`
- 未完成内容:
- 当前依赖应用侧轮询,日志实时性和失败最后日志仍需增强。
- 后端开发:
- 支持日志 offset/tail 增量读取。
- 任务失败时强制同步最后日志片段。
- 日志接口返回来源、偏移量、是否截断。
- 前端开发:
- 日志页按 offset 增量刷新。
- 失败时自动跳到底部并展示最后错误。
- 验收标准:
- 训练过程中日志持续刷新。
- 失败后无需手动刷新即可看到最后错误。
### P2权限、审计和治理落点补齐
- 对应页面:
- `用户中心`
- `项目管理`
- `模型管理`
- `数据集管理`
- `模型微调`
- `算力平台`
- 未完成内容:
- 训练链路中的租户、项目、用户权限校验还需要细粒度补齐。
- 审计事件需要覆盖训练创建、启动、停止、删除、导出、资源修复。
- 后端开发:
- 接口增加项目/租户上下文校验。
- 增加审计事件写入。
- 删除和高风险操作进入审批流。
- 验收标准:
- 用户只能访问授权项目内的模型、数据集和训练任务。
- 关键操作可以在审计日志中查询。
### P2模型评测页面真实闭环
- 对应页面:`模型评测`
- 未完成内容:
- 评测任务创建、运行、日志、指标、结果对比仍需接入真实后端。
- 后端开发:
- 评测任务表、评测指标表、评测日志接口。
- 支持指定模型、数据集、评测模板和指标。
- Compute 开发:
- 支持评测任务执行器。
- 返回指标结果和日志。
- 验收标准:
- 可创建评测任务并看到运行状态、指标结果和失败原因。
### P2模型推理页面真实闭环
- 对应页面:
- `模型推理 / 推理服务`
- `模型推理 / 对话测试`
- `模型对比`
- 未完成内容:
- 推理服务启动、停止、健康检查、会话请求、资源释放仍需完善。
- 后端开发:
- 推理服务实例管理接口。
- 对话请求代理接口。
- 推理日志和资源占用查询。
- Compute 开发:
- 支持启动本地模型推理服务。
- 支持停止服务和释放 GPU。
- 验收标准:
- 可从页面启动一个已训练模型的推理服务。
- 可进行对话测试并查看服务状态。
## 3. 推荐开发顺序
1. 完成真实训练成功闭环。
2. 完成数据集格式校验和 LLaMA-Factory 参数映射。
3. 完成 artifact、checksum、manifest 和模型导出闭环。
4. 完成资源副本 repair 自动重同步。
5. 完成多算力节点调度增强。
6. 完成训练日志实时性优化。
7. 补齐权限、审批、审计治理落点。
8. 启动模型评测真实闭环开发。
9. 启动模型推理真实闭环开发。
## 4. 下一轮优先执行任务
### 任务 1训练模型路径治理
- 页面:`模型管理 / 新增模型``模型微调 / 创建训练任务`
- 内容:
- 新增模型时区分是否可用于训练。
- API 模型不能作为 LLaMA-Factory 本地训练基座。
- 本地模型路径必须是算力节点可访问路径。
- 验收:
- 选择不可训练模型时,训练创建页预检直接给出中文错误。
### 任务 2数据集格式校验
- 页面:`数据集管理 / 上传数据集``模型微调 / 创建训练任务`
- 内容:
- 上传后扫描样本字段。
- 支持 Alpaca、ShareGPT、DPO、CPT 校验。
- 返回字段缺失、类型错误、空样本等问题。
- 验收:
- `111.json` 这类数据可以明确判断是否满足当前训练模板。
### 任务 3训练完成产物入库
- 页面:`模型微调 / 任务列表``模型管理 / 已训练模型`
- 内容:
- Compute 训练成功后返回输出目录。
- 应用侧轮询后创建 trained model 记录。
- 写入 artifact、lineage、export job 初始状态。
- 验收:
- 训练完成后无需手动登记,模型管理中自动出现新模型。
### 任务 4失败日志最后片段拉取
- 页面:`系统日志 / 训练日志`
- 内容:
- 任务失败时立即拉取最后 N 行日志。
- 页面展示最后错误、诊断建议和原始日志。
- 验收:
- 训练失败后页面不再只看到“失败”,可以直接看到失败原因。
## 5. 当前测试注意事项
- 当前环境中已有算力节点可访问,但 GPU 为 2GB 显存,预检会提示显存不足,这是符合预期的生产校验结果。
- 若要验证真实训练成功,需要提前准备:
- 算力节点可访问的本地基座模型目录。
- 合法训练数据集文件。
- 足够显存的 GPU。
- Compute 容器内可用的 LLaMA-Factory 和 `llamafactory-cli`
- 当前前端 `dist` 已按要求参与构建更新,后续修改前端页面后需要重新执行 `npm run build`

1390
docs/backend-api-design.md Normal file

File diff suppressed because it is too large Load Diff

109
docs/backend-logging.md Normal file
View File

@@ -0,0 +1,109 @@
# 后端日志模块说明
本文档对应页面/功能模块:全平台通用能力、系统设置、审计中心、任务详情、训练任务日志、运维监控。
## 设计目标
- 后端服务统一使用 `backend/app/core/logging.py` 初始化日志。
- 日志文件按日期命名,单个文件超过 20MB 自动滚动。
- 日志只保留最近 10 天,过期文件自动清理。
- 业务日志使用 JSON Lines 格式,便于 Filebeat、Vector、Logstash、ELK、OpenSearch 等日志平台采集。
- `ERROR` 及以上日志独立写入错误日志文件,便于告警与问题定位。
- 日志字段必须包含代码文件、行号、函数、日志内容、请求 ID、进程和线程信息。
## 文件命名
默认日志目录由 `LOG_DIR` 控制,本地默认是 `./logs`
```text
logs/
backend-2026-07-16.log # INFO/ERROR 等全部应用日志JSON Lines
backend-2026-07-16.1.log # 当天主日志超过 20MB 后滚动产生
error-2026-07-16.log # ERROR/CRITICAL 错误日志JSON Lines
error-2026-07-16.1.log # 当天错误日志超过 20MB 后滚动产生
```
## 环境变量
```env
LOG_LEVEL=INFO
LOG_DIR=./logs
LOG_FILE_PREFIX=backend
LOG_ERROR_FILE_PREFIX=error
LOG_MAX_BYTES=20971520
LOG_RETENTION_DAYS=10
```
## JSON 字段
每一行都是一个完整 JSON 对象。
```json
{
"@timestamp": "2026-07-16T13:20:10.123",
"level": "INFO",
"logger": "app.access",
"message": "request completed method=GET path=/modelTF/health status_code=200 duration_ms=3.12 client=127.0.0.1",
"module": "logging",
"function": "request_logging_middleware",
"file": "D:\\AI\\codex-code\\YG_FT\\backend\\app\\core\\logging.py",
"line": 169,
"process": 1234,
"thread": 5678,
"thread_name": "MainThread",
"request_id": "6f9d1c3c-8be0-4c8d-a5b2-18f9d41f9a0c"
}
```
异常日志会额外包含:
```json
{
"exception": "Traceback ..."
}
```
## 使用方式
业务代码中不要直接 `print`,统一使用:
```python
from app.core.logging import get_logger
logger = get_logger(__name__)
logger.info("dataset uploaded dataset_id=%s", dataset_id)
logger.warning("gpu queue is busy project_id=%s", project_id)
logger.exception("training job failed job_id=%s", job_id)
```
`logger.exception(...)` 只能在 `except` 代码块中使用,它会自动写入堆栈信息,并同时进入主日志和错误日志。
## FastAPI 接入
应用入口 `backend/app/main.py` 已完成接入:
```python
settings = get_settings()
configure_logging(settings)
setup_request_logging(app)
```
请求日志会自动生成或透传 `X-Request-ID`,并在响应头中返回同一个请求 ID方便前端、后端、算力服务、日志平台串联排障。
## ELK/日志平台采集建议
- 采集路径:`/app/logs/*.log` 或生产环境挂载后的日志目录。
- 解析方式:按行读取,每行作为 JSON 文档解析。
- 索引建议:
- 主日志:`yg-ft-backend-*`
- 错误日志:`yg-ft-backend-error-*`
- 推荐保留字段:`@timestamp``level``logger``message``file``line``function``request_id``tenant_id``project_id``job_id`
- 业务开发后续应在关键模块日志中补充 `tenant_id``project_id``job_id` 等上下文字段,便于企业审计和问题定位。
## 注意事项
- 当前日志落本地磁盘,生产环境建议把日志目录挂载到独立数据盘。
- 日志文件保留 10 天是应用侧兜底策略,企业侧长期留存应由 ELK、对象存储或归档服务承担。
- 敏感字段如 token、密码、密钥、原始用户数据内容不得写入日志。
- 算力节点和应用节点分开部署时,建议两侧都采用 JSON Lines 格式,并使用统一 `request_id/job_id` 贯穿链路。

279
docs/data-process-design.md Normal file
View File

@@ -0,0 +1,279 @@
# 数据处理接口与算法设计
本文是 `team-development-plan.md` 板块 C 的落地契约,约束
`/modelTF/data-process/*`、前端数据处理向导以及 PostgreSQL 数据模型。
## 1. 处理闭环
```text
创建草稿任务
→ 上传并登记源文件格式、SHA-256、版本
→ 预处理(标准化、无效过滤、去重、可选脱敏)
→ 构建可编辑预览(来源偏移与行号)
→ 生成标准训练记录
→ 质量评分与稳定数据集划分
→ 人工编辑/恢复
→ 幂等发布为数据集(保留完整来源链路)
```
任务只使用以下五种状态:
```text
pending ──start/generate──> running ──success──> completed
▲ │ ├──error───────> failed
│ │ └──stop────────> stopped
└────────retry────────────┴────────retry─────┘
```
- `pending` 允许修改配置、增删源文件和重建预览。
- `running` 拒绝重复启动、修改配置和删除任务。
- `failed``stopped` 可重试;重试前清理上一次未完成结果。
- `completed` 可编辑结果和发布;重复发布返回同一个数据集。
- 非法状态转换返回 HTTP 409。
- 每次生成分配独立 `generation_run_id`;停止或重试会使旧代次立即失效,
旧后台任务不能覆盖新代次的结果或状态。
## 2. 接口契约
所有路径由请求层统一添加 `/modelTF`,响应统一为
`{ "code": 0, "message": "ok", "data": ... }`
### 任务与进度
| 方法 | 路径 | 说明 |
| --- | --- | --- |
| GET | `/data-process` | 分页查询任务,支持 keyword/status/process_type |
| POST | `/data-process` | 创建 `pending` 草稿 |
| GET | `/data-process/{id}` | 查询任务详情,不内嵌全部结果 |
| PUT | `/data-process/{id}` | 更新草稿配置 |
| DELETE | `/data-process/{id}` | 软删除非运行任务 |
| POST | `/data-process/{id}/start` | 重建预览并生成的一键编排入口 |
| POST | `/data-process/{id}/generate` | 使用已确认预览生成结果 |
| POST | `/data-process/{id}/stop` | 请求停止运行任务 |
| GET | `/data-process/{id}/progress` | 查询阶段、进度与计数 |
### 源文件与预览
| 方法 | 路径 | 说明 |
| --- | --- | --- |
| POST | `/data-process/{id}/source-files` | multipart 上传,字段名 `files` |
| DELETE | `/data-process/{id}/source-files/{file_id}` | 删除源文件及其预览 |
| GET | `/data-process/{id}/source-files/{file_id}/content` | 按行窗口读取源文 |
| POST | `/data-process/{id}/preview/build` | 后端预处理并重建预览 |
| GET | `/data-process/{id}/preview` | 分页查询预览 |
| POST | `/data-process/{id}/preview` | 手工增加预览条目 |
| PUT | `/data-process/{id}/preview/{preview_id}` | 保存人工编辑 |
| DELETE | `/data-process/{id}/preview/{preview_id}` | 删除预览条目 |
上传批次先全部完成有界读取和解析,再在单个事务中登记;任一文件为空、超限、
重复或格式非法时整批不落库,暂存原件也会一并清理。响应不回传整个文件,只返回
逻辑对象引用、文件 ID、格式、原始字节数、记录数和原始 SHA-256。二进制文档必须
由对应解析器显式处理;不支持的格式返回 415绝不能静默替换成示例正文。
原始上传字节与解析正文采用双层存储:原件默认保存在
`backend/storage/data-process/<task_id>/<file_id>/v<version>/<安全文件名>`,数据库的
`storage_object_id` 只保存 `local://data-process/...` 逻辑引用,不保存或返回宿主机
绝对路径;完整解析正文继续保存在 `data_process_source_files.content`,列表摘要使用
`content_preview`,因此 PDF、Office 等文件的预览无需反复解析原始二进制。可通过
`DATA_PROCESS_STORAGE_DIR` 指定其他本地根目录;从 `start.sh` 启动时,该变量应在
当前终端导出。历史 `db://data-process/...` 记录继续从数据库正文预览。
单独删除源文件时先提交数据库软删除,再立即删除受控目录中的原件;若物理删除
失败,接口仍按数据库结果返回成功并标记 `storage_cleanup_pending=true`,软删除记录
中的逻辑引用可供运维补偿清理。任务软删除以及修改 `process_type` 导致的源文件
软删除按留存数据处理,当前版本不自动物理清除。
### 结果与发布
| 方法 | 路径 | 说明 |
| --- | --- | --- |
| GET | `/data-process/{id}/results` | 分页查询,支持 keyword/status/split |
| PUT | `/data-process/{id}/results/{result_id}` | 保存人工编辑并重评分 |
| POST | `/data-process/{id}/results/{result_id}/restore` | 恢复生成时的原值 |
| POST | `/data-process/{id}/publish` | 幂等发布为数据集 |
## 3. 配置校验
- `process_type``structured | unstructured | external`
- 数据集划分的 `train + validation + test` 必须等于 100各项为 0100。
- `chunk_size` 为 1632768 token`chunk_overlap` 必须小于
`chunk_size``min_chunk_size` 不得大于 `chunk_size`
- `temperature` 为 02`max_tokens` 为 132768。
- 任务名称在未删除任务中唯一。
- 选择 `generation_model_id` 后,启动生成时校验模型是否存在,并保存不含密钥的
模型版本快照。
- 当前运行库沿用平台现有的单租户模式,不接受客户端提交 tenant/owner/operator
字段,避免伪造隔离上下文;接入平台可信认证上下文后再启用数据库中预留的
tenant/project 字段。
## 4. 格式解析与标准化
上传格式按处理类型约束:
- 结构化数据支持 JSON、JSONL/NDJSON、CSV/TSV 和 XLSX。XLSX 能识别纵向、
横向合并单元格组成的多级表头,并稳定展平为 `销售.Q1` 一类字段;公式只读取
文件中已缓存的计算结果,不在服务端执行。
- 非结构化数据支持 UTF-8/UTF-8 BOM 的 TXT、Markdown、JSON/JSONL以及
文本型 PDF、DOCX 和 PPTX。PDF 按页抽取文本DOCX 抽取段落与表格PPTX
抽取幻灯片文本与表格,随后统一进入切片算法。
- 旧版二进制 DOC、XLS、PPT 不直接解析,返回 415 并提示分别转换为
DOCX、XLSX、PPTX。
- 扫描 PDF 没有文本层时明确提示需要 OCR当前流程不执行 OCR。加密、损坏或
超出页数/工作表/行列/解压规模限制的文件整批拒绝。
现代 Office 文件在交给解析库前检查 ZIP 成员路径、重复成员、加密标记、活动
XML、单成员大小、总解压大小和压缩比避免路径穿越、实体扩展与 ZIP bomb。
结构化选项按固定顺序执行,关闭某项时不会隐式执行对应业务变换:
1. `detect_structure`展平嵌套对象XLSX 上传解析阶段识别合并单元格和多级表头。
2. `normalize_format`:字段名转 snake_case执行 Unicode NFKC、换行和容器值规范化
输出键顺序稳定的 canonical JSON账号、邮编等字符串不会转成数值。
3. `clean_invalid`:删除全空列和全空记录;存在 `id/uuid/key/code/*_id` 身份字段时,
删除身份字段残缺的行,但不会因备注等可选字段为空误删有效记录。
4. `filter_anomaly`:仅对不少于 8 个样本的非身份数值字段使用 Tukey IQR 过滤离群行,
同时过滤明确乱码、不可打印或极端超长文本;小样本和 ID 字段不参与统计过滤。
5. `deduplicate`:先按整行 canonical JSON 精确去重,再按非空
`id/uuid/key/code/*_id` 字段稳定保留首条;空关键值互不视为重复。
6. `desensitize`:对结构化姓名字段和正文中的高置信上下文姓名、邮箱、手机号、
身份证号进行不可逆掩码,并分别记录命中数。
非结构化“智能预处理”由六个可独立执行的底层选项组成:
- `clean_invalid_content` 删除确定为空、不可读或纯重复符号的无效块。
- `detect_document_structure` 识别 Markdown、中文章节和数字标题切片不跨章节
并在预览质量详情中保存 `heading_path`
- `merge_short_content` 在同一章节中合并短块,合并后不突破 `chunk_size`
- `filter_low_quality` 在生成前过滤乱码、不可打印、重复或极端超长内容。
- `deduplicate_content` 先精确去重,再对足够长的内容进行保守近重复判断;数字或
否定含义变化时始终保留。
- `preserve_context` 才启用相邻切片 overlap关闭时切片不共享正文上下文且上下文
永不跨文件或章节。
表格、围栏代码块和连续列表保护是三个独立参数。启用时切点避开相应 Markdown
块,关闭时允许按正常长度切分。
脱敏是不可逆掩码:
- 邮箱:`[EMAIL]`
- 中国大陆手机号:`[PHONE]`
- 18 位身份证号:`[ID_CARD]`
- 高置信姓名:`[NAME]`
源文件原文与脱敏后的预览分开保存,结果不得反向覆盖源文件。
## 5. 切片算法
首阶段只提供三种切片策略:
- `structure` 先识别 Markdown、中文章节及编号标题再由 LlamaIndex
`SentenceSplitter` 在章节内按段落和中英文句界限长;章节之间不共享 overlap。
- `fixed` 使用 LlamaIndex `TokenTextSplitter` 按目标 token 窗口切分。
- `custom` 使用用户给定分隔符,在找不到合适分隔点时回退到固定窗口。
不提供 `semantic` 和旧 `heading` 配置;创建或更新任务时传入这些值会直接拒绝。
LlamaIndex 只负责通用切分,原文 offset、行号、标题路径和 Markdown 保护块仍由
项目适配层统一维护。
首版使用可替换的确定性 token 估算器,中文字符、标点和英文词分别计数;
所有偏移以 Python/JavaScript 都能稳定表达的 Unicode 文本偏移为准。
算法必须满足:
- 每轮游标严格前进,异常分隔符不能产生死循环。
- overlap 是最大重叠量,尾部过短切片合并到上一片。
- 代码块、Markdown 表格和连续列表在启用保护时不从中间切开。
- 每个预览条目记录 `source_file_id`、字符偏移、起止行、token 数和算法版本。
## 6. 生成与质量评分
结构化记录优先识别以下字段:
1. `instruction/input/output`
2. `question/context/answer`
3. `prompt/input/response`
已有标准字段时只做标准化;需要语义生成时调用所选模型的 OpenAI 兼容接口,
并固化模型 ID、模型版本、prompt、temperature、max_tokens 和 JSON mode 快照。
模型地址可输入域名、`/v1` 基础地址或完整地址:例如输入
`www.caoxiaozhu.com` 会规范为
`https://www.caoxiaozhu.com/v1/chat/completions`,无需用户手工拼接路径。
单条失败记录为 `invalid`,有限重试耗尽后继续处理下一条,避免整批丢失。
每条结果总分为 0100
```text
总分 = 完整性 35% + 长度合理性 20% + 可读性 20%
+ 来源相关性 15% + 非重复性 10%
```
- instruction 或 output 为空时格式硬失败并标记 `invalid`
- 开启短文本过滤且 output 低于 `min_output_length` 时标记过滤原因。
- 评分详情、命中规则与过滤原因必须落库并返回前端,不只返回一个总分。
## 7. 稳定划分
划分不能依赖结果插入顺序。对每条记录计算:
```text
bucket = SHA256(task_id + ":" + result_id) mod 10000
```
按万分位阈值映射为 `train/validation/test`。同一任务重试、分页或进程重启后,
同一结果仍落入相同 split。
## 8. 发布与来源链路
发布在一个数据库事务中完成:
```text
source_file
→ data_process_task
→ data_process_result
→ dataset
→ dataset_file + dataset_file_version
→ dataset_record
```
只发布 `valid/modified` 且满足质量门槛的结果。输出 JSONL 先计算 checksum
再登记文件版本和记录。发布请求中的 split 会重新进行稳定划分。任务的
`output_dataset_id` 是幂等键;重复调用返回已有数据集,目标数据集若已被外部
删除则解除断链并重新发布。当前运行库只开放 `local` 存储类型,正文保存在
当前平台的 `dataset_files.content`,不虚假宣称已上传 MinIO 或云存储。
## 9. 安全边界
- 文件名只保留 basename响应不返回宿主机绝对路径。
- 上传限制单文件、批次文件数与批次总大小,解析采用有界读取。
- 外部数据源凭据不写日志、不进入 localStorage、不在详情接口回显。
- 外部 PostgreSQL 只允许单条 `SELECT/WITH`、只读事务、5 秒连接超时、
30 秒语句超时和 50 MiB 响应上限;默认阻止回环、链路本地及私网地址。
可信内网部署必须显式设置 `DATA_PROCESS_ALLOW_PRIVATE_EXTERNAL_DB=true`
- SQL 迁移独立存放,应用启动不会隐式修改当前远程数据库。
## 10. 迁移边界
`backend/app/db/sql/002_data_process.sql` 只面向当前运行脚本
`001_platform_runtime.sql` 的 TEXT/最小表模型。它会在执行前检查
`datasets.id` 类型;若检测到 `docs/postgres-schema.sql` 的 UUID/JSONB 目标模型,
会直接失败而不是进行一半成功、一半失败的危险迁移。目标模型后续应由独立
Alembic 迁移和对应存储实现承接。
`DataProcessStore.ensure_schema()` 仅供受控管理命令显式调用API 路由和应用启动
均不会自动执行该迁移。本次开发和测试没有修改任何远程数据库。
在已加载 `DATABASE_URL` 的终端中可先只读检查:
```bash
cd backend
.venv/bin/python -m app.modules.data_process.schema_cli --check
```
确认目标主机和数据库名称无误后,才显式执行:
```bash
cd backend
.venv/bin/python -m app.modules.data_process.schema_cli --apply --yes
```
命令输出只显示主机、端口和数据库名,不显示用户名或密码。

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