第一次提交
This commit is contained in:
9
backend/.env.example
Normal file
9
backend/.env.example
Normal file
@@ -0,0 +1,9 @@
|
||||
APP_NAME=YG Fine-Tune Platform API
|
||||
APP_ENV=local
|
||||
API_PREFIX=/api
|
||||
LOG_LEVEL=INFO
|
||||
LOG_DIR=./logs
|
||||
LOG_FILE_PREFIX=backend
|
||||
LOG_ERROR_FILE_PREFIX=error
|
||||
LOG_MAX_BYTES=20971520
|
||||
LOG_RETENTION_DAYS=10
|
||||
75
backend/README.md
Normal file
75
backend/README.md
Normal file
@@ -0,0 +1,75 @@
|
||||
# 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/ # 本地开发日志目录,生产环境建议挂载到独立日志盘
|
||||
```
|
||||
|
||||
## 企业治理接口说明
|
||||
|
||||
用户中心、租户、项目、审批、资源授权等能力已通过对应模块 `router.py` 在 `/modelTF` 下统一暴露。其中:
|
||||
|
||||
- 审计接口挂在 `system` 模块:`GET /modelTF/system/audit-logs`、`/modelTF/system/audit-logs/export`。
|
||||
- 留存策略为租户资源的嵌套接口:`PUT /modelTF/tenants/{id}/retention-policy`,未独立成 `/modelTF/retention-policies` 路由。
|
||||
- 重置密码:`POST /modelTF/users/{id}/reset-password`(保护账号不可重置,空密码回退默认 `platform123`)。
|
||||
- 资源授权:`GET/PUT /modelTF/resources/{type}/{id}/acl`。
|
||||
- 服务看板聚合:`GET /modelTF/dashboard/stats`。
|
||||
|
||||
## 本地启动
|
||||
|
||||
```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
1
backend/app/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Application package."""
|
||||
1
backend/app/api/__init__.py
Normal file
1
backend/app/api/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""API package."""
|
||||
1
backend/app/api/v1/__init__.py
Normal file
1
backend/app/api/v1/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Versioned API package."""
|
||||
1
backend/app/api/v1/endpoints/__init__.py
Normal file
1
backend/app/api/v1/endpoints/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""API endpoint modules."""
|
||||
14
backend/app/api/v1/endpoints/health.py
Normal file
14
backend/app/api/v1/endpoints/health.py
Normal 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()}
|
||||
|
||||
844
backend/app/api/v1/endpoints/platform.py
Normal file
844
backend/app/api/v1/endpoints/platform.py
Normal file
@@ -0,0 +1,844 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Body, File, HTTPException, Query, UploadFile
|
||||
from fastapi.responses import PlainTextResponse, StreamingResponse
|
||||
|
||||
from app.db.platform_store import get_platform_store
|
||||
from app.modules.fine_tune.service import apply_presets
|
||||
from fastapi import Request as FastAPIRequest
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _actor(request: FastAPIRequest) -> str | None:
|
||||
auth = request.headers.get("Authorization", "")
|
||||
token = auth.replace("Bearer ", "").strip()
|
||||
return token or None
|
||||
|
||||
|
||||
def ok(data: Any = None, message: str = "ok") -> dict[str, Any]:
|
||||
return {"code": 0, "message": message, "data": data}
|
||||
|
||||
|
||||
def fail(status_code: int, message: str) -> HTTPException:
|
||||
return HTTPException(status_code=status_code, detail={"code": status_code, "message": message, "data": None})
|
||||
|
||||
|
||||
@router.get("/dashboard/overview")
|
||||
async def dashboard_overview() -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
tasks = store.tasks()
|
||||
return ok(
|
||||
{
|
||||
"models": len(store.models()),
|
||||
"datasets": len(store.datasets()),
|
||||
"fine_tune_tasks": len(tasks),
|
||||
"running_tasks": len([t for t in tasks if t["status"] in {"syncing", "queued", "running"}]),
|
||||
"compute_nodes": len(store.compute_nodes()),
|
||||
"gpus": len(store.gpus()),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/dashboard/stats")
|
||||
async def dashboard_stats() -> dict[str, Any]:
|
||||
"""看板聚合数据:基于平台真实数据;缺项做合理近似(见下)。"""
|
||||
store = get_platform_store()
|
||||
tasks = store.tasks()
|
||||
users = store.users()
|
||||
nodes = store.compute_nodes()
|
||||
datasets = store.datasets()
|
||||
|
||||
running_statuses = {"syncing", "queued", "running"}
|
||||
running_ft = [t for t in tasks if t.get("status") in running_statuses]
|
||||
failed_ft = [t for t in tasks if t.get("status") == "failed"]
|
||||
online_nodes = [n for n in nodes if n.get("scheduler_status") == "online"]
|
||||
|
||||
# 近 7 天训练统计(按创建日期分桶;准确率为 None,因任务无该字段)
|
||||
now = datetime.now(timezone.utc)
|
||||
train_by_day: dict[str, int] = {}
|
||||
for t in tasks:
|
||||
ct = t.get("create_time")
|
||||
if ct:
|
||||
train_by_day[ct[:10]] = train_by_day.get(ct[:10], 0) + 1
|
||||
training_7d = []
|
||||
for i in range(6, -1, -1):
|
||||
day = (now - timedelta(days=i)).strftime("%Y-%m-%d")
|
||||
training_7d.append(
|
||||
{
|
||||
"date": day[5:],
|
||||
"train": train_by_day.get(day, 0),
|
||||
"gpu": sum(len(t.get("gpus") or []) for t in running_ft),
|
||||
"accuracy": None,
|
||||
}
|
||||
)
|
||||
|
||||
# 服务状态:模型推理用在线计算节点近似;模型评测暂无独立数据源,置 0
|
||||
service_status = [
|
||||
{
|
||||
"type": "模型推理",
|
||||
"status": "error" if (nodes and not online_nodes) else ("busy" if (nodes and len(online_nodes) < len(nodes)) else "normal"),
|
||||
"count": len(online_nodes),
|
||||
},
|
||||
{
|
||||
"type": "模型微调",
|
||||
"status": "error" if failed_ft else ("busy" if running_ft else "normal"),
|
||||
"count": len(running_ft),
|
||||
},
|
||||
{"type": "模型评测", "status": "normal", "count": 0},
|
||||
{
|
||||
"type": "数据处理",
|
||||
"status": "normal" if not failed_ft else "busy",
|
||||
"count": len(datasets),
|
||||
},
|
||||
]
|
||||
|
||||
# 训练任务状态归一化(fine_tune 的 syncing/queued 等映射到前端已知状态)
|
||||
status_map = {
|
||||
"syncing": "running",
|
||||
"queued": "running",
|
||||
"running": "running",
|
||||
"pending": "pending",
|
||||
"paused": "pending",
|
||||
"completed": "completed",
|
||||
"failed": "failed",
|
||||
"error": "failed",
|
||||
"cancelled": "failed",
|
||||
}
|
||||
op_labels = [
|
||||
("模型训练", lambda a: "fine_tune" in a or "train" in a),
|
||||
("数据处理", lambda a: "data" in a or "dataset" in a),
|
||||
("模型评测", lambda a: "eval" in a),
|
||||
("模型推理", lambda a: "infer" in a or "serving" in a or "deploy" in a),
|
||||
("系统设置", lambda a: True),
|
||||
]
|
||||
|
||||
def _op_label(action: str) -> str:
|
||||
for label, fn in op_labels:
|
||||
if fn(action):
|
||||
return label
|
||||
return "系统设置"
|
||||
|
||||
training_tasks = [
|
||||
{
|
||||
"id": t.get("id"),
|
||||
"name": t.get("name"),
|
||||
"status": status_map.get(t.get("status"), "pending"),
|
||||
"train_type": t.get("train_type"),
|
||||
"train_method": t.get("train_method"),
|
||||
"base_model": t.get("base_model"),
|
||||
"progress": t.get("progress", 0),
|
||||
"accuracy": t.get("accuracy"),
|
||||
"started_at": (t.get("create_time") or "")[:16],
|
||||
}
|
||||
for t in tasks[:8]
|
||||
]
|
||||
|
||||
# 用户操作分布(按 audit action 归类为中文分类)
|
||||
audit = store.audit_logs(limit=500)
|
||||
op_counter: dict[str, int] = {}
|
||||
for log in audit.get("items", []):
|
||||
act = log.get("action") or "unknown"
|
||||
op_counter[_op_label(act)] = op_counter.get(_op_label(act), 0) + 1
|
||||
operation_distribution = [{"name": k, "value": v} for k, v in op_counter.items()]
|
||||
|
||||
# 最近登录用户:后端有 last_login 字段,返回真实数据
|
||||
recent = sorted(
|
||||
[u for u in users if u.get("last_login")],
|
||||
key=lambda u: u["last_login"],
|
||||
reverse=True,
|
||||
)[:5]
|
||||
recent_login_users = [
|
||||
{
|
||||
"user": u.get("display_name") or u.get("username"),
|
||||
"role": u.get("role"),
|
||||
"last_login": (u.get("last_login") or "")[:16],
|
||||
}
|
||||
for u in recent
|
||||
]
|
||||
# 登录时长:后端暂无该数据源,先留空,待接入后补充
|
||||
login_duration_rank: list = []
|
||||
|
||||
return ok(
|
||||
{
|
||||
"online_services": sum(s["count"] for s in service_status),
|
||||
"running_tasks": len(running_ft),
|
||||
"pending_alerts": 0, # 平台暂无独立告警数据源,先置 0,待接入后补充
|
||||
"training_7d": training_7d,
|
||||
"service_status": service_status,
|
||||
"training_tasks": training_tasks,
|
||||
"operation_distribution": operation_distribution,
|
||||
"login_duration_rank": login_duration_rank,
|
||||
"recent_login_users": recent_login_users,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/system-info")
|
||||
async def system_info() -> dict[str, Any]:
|
||||
return ok(get_platform_store().system_info())
|
||||
|
||||
|
||||
@router.get("/users")
|
||||
async def users() -> dict[str, Any]:
|
||||
return ok(get_platform_store().users())
|
||||
|
||||
|
||||
@router.post("/users")
|
||||
async def create_user(payload: dict[str, Any] = Body(...), request: FastAPIRequest = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
user = store.create_user(payload)
|
||||
store.record_audit(
|
||||
action="user.create",
|
||||
actor_id=_actor(request),
|
||||
target_type="user",
|
||||
target_id=user["id"],
|
||||
detail=f"username={user.get('username')}",
|
||||
)
|
||||
return ok(user)
|
||||
|
||||
|
||||
@router.put("/users/{user_id}")
|
||||
async def update_user(user_id: str, payload: dict[str, Any] = Body(...), request: FastAPIRequest = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
user = store.update_user(user_id, payload)
|
||||
except KeyError:
|
||||
raise fail(404, "user not found")
|
||||
store.record_audit(
|
||||
action="user.update",
|
||||
actor_id=_actor(request),
|
||||
target_type="user",
|
||||
target_id=user_id,
|
||||
detail=f"fields={','.join(payload.keys())}",
|
||||
)
|
||||
return ok(user)
|
||||
|
||||
|
||||
@router.delete("/users/{user_id}")
|
||||
async def delete_user(user_id: str, current_username: str | None = Query(default=None)) -> dict[str, Any]:
|
||||
try:
|
||||
get_platform_store().delete_user(user_id)
|
||||
return ok({"deleted": user_id, "current_username": current_username})
|
||||
except KeyError:
|
||||
raise fail(404, "user not found")
|
||||
except ValueError as exc:
|
||||
raise fail(400, str(exc))
|
||||
|
||||
|
||||
@router.post("/users/{user_id}/reset-password")
|
||||
async def reset_password(
|
||||
user_id: str,
|
||||
payload: dict[str, Any] = Body(default={}),
|
||||
request: FastAPIRequest = None,
|
||||
) -> dict[str, Any]:
|
||||
try:
|
||||
user = get_platform_store().reset_password(user_id, payload.get("password") or "")
|
||||
except KeyError:
|
||||
raise fail(404, "user not found")
|
||||
except ValueError as exc:
|
||||
raise fail(400, str(exc))
|
||||
get_platform_store().record_audit(
|
||||
action="user.reset_password",
|
||||
actor_id=_actor(request),
|
||||
target_type="user",
|
||||
target_id=user_id,
|
||||
detail=f"username={user.get('username')}",
|
||||
)
|
||||
return ok({"id": user_id})
|
||||
|
||||
|
||||
@router.get("/model-manage/local-models")
|
||||
async def local_models() -> dict[str, Any]:
|
||||
models = [{"path": item.get("path") or "", "name": item["name"]} for item in get_platform_store().models()]
|
||||
return ok({"models": models})
|
||||
|
||||
|
||||
@router.get("/model-manage/trained-models")
|
||||
async def trained_models() -> dict[str, Any]:
|
||||
return ok({"models": get_platform_store().trained_models()})
|
||||
|
||||
|
||||
@router.delete("/model-manage/trained-models/{model_id}")
|
||||
async def delete_trained_model(model_id: str, type: str = Query(default="merged")) -> dict[str, Any]:
|
||||
return ok({"deleted": model_id, "type": type})
|
||||
|
||||
|
||||
@router.get("/model-manage/name/{name}")
|
||||
async def model_by_name(name: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().model_by_name(name))
|
||||
except KeyError:
|
||||
raise fail(404, "model not found")
|
||||
|
||||
|
||||
@router.get("/model-manage")
|
||||
async def model_list() -> dict[str, Any]:
|
||||
return ok(get_platform_store().models())
|
||||
|
||||
|
||||
@router.post("/model-manage")
|
||||
async def create_model(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
return ok(get_platform_store().create_model(payload))
|
||||
|
||||
|
||||
@router.get("/model-manage/{model_id}")
|
||||
async def model_detail(model_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().model(model_id))
|
||||
except KeyError:
|
||||
raise fail(404, "model not found")
|
||||
|
||||
|
||||
@router.put("/model-manage/{model_id}")
|
||||
async def update_model(model_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_model(model_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "model not found")
|
||||
|
||||
|
||||
@router.put("/model-manage/{model_id}/purpose")
|
||||
async def update_model_purpose(model_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_model(model_id, {"purpose": payload.get("purpose", "training")}))
|
||||
except KeyError:
|
||||
raise fail(404, "model not found")
|
||||
|
||||
|
||||
@router.delete("/model-manage/{model_id}")
|
||||
async def delete_model(model_id: str) -> dict[str, Any]:
|
||||
get_platform_store().delete_model(model_id)
|
||||
return ok({"deleted": model_id})
|
||||
|
||||
|
||||
@router.post("/model-manage/merge")
|
||||
async def merge_model(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
return ok({"job_id": f"merge_{uuid.uuid4().hex[:12]}", "status": "queued", **payload})
|
||||
|
||||
|
||||
@router.get("/dataset-manage/preview/{file_id}")
|
||||
async def dataset_preview(file_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
row = get_platform_store().dataset_file(file_id)
|
||||
return ok({"content": row["content"]})
|
||||
except KeyError:
|
||||
raise fail(404, "dataset file not found")
|
||||
|
||||
|
||||
@router.get("/dataset-manage/versions/{file_id}")
|
||||
async def dataset_versions(file_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().file_versions(file_id))
|
||||
except KeyError:
|
||||
raise fail(404, "dataset file not found")
|
||||
|
||||
|
||||
@router.get("/dataset-manage/versions/{file_id}/{version_id}")
|
||||
async def dataset_version_content(file_id: str, version_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
row = get_platform_store().dataset_file(file_id)
|
||||
versions = get_platform_store().file_versions(file_id)["versions"]
|
||||
version = next((item for item in versions if item["id"] == version_id), None)
|
||||
if not version:
|
||||
raise KeyError(version_id)
|
||||
return ok({"version": version, "content": row["content"]})
|
||||
except KeyError:
|
||||
raise fail(404, "dataset version not found")
|
||||
|
||||
|
||||
@router.post("/dataset-manage/versions/{file_id}")
|
||||
async def create_dataset_version(file_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().create_file_version(file_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "dataset file not found")
|
||||
|
||||
|
||||
@router.put("/dataset-manage/versions/{file_id}/active")
|
||||
async def activate_dataset_version(file_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().activate_file_version(file_id, payload["version_id"]))
|
||||
except KeyError:
|
||||
raise fail(404, "dataset version not found")
|
||||
|
||||
|
||||
@router.delete("/dataset-manage/versions/{file_id}/{version_id}")
|
||||
async def delete_dataset_version(file_id: str, version_id: str) -> dict[str, Any]:
|
||||
return ok(get_platform_store().file_versions(file_id))
|
||||
|
||||
|
||||
@router.post("/dataset-manage/upload/{dataset_id}")
|
||||
async def upload_dataset_files(dataset_id: str, files: list[UploadFile] = File(default=[])) -> dict[str, Any]:
|
||||
created: list[dict[str, Any]] = []
|
||||
store = get_platform_store()
|
||||
try:
|
||||
store.dataset(dataset_id)
|
||||
except KeyError:
|
||||
raise fail(404, "dataset not found")
|
||||
with store.connect() as conn:
|
||||
for file in files:
|
||||
raw = await file.read()
|
||||
content = raw.decode("utf-8", errors="replace")
|
||||
created.append(store.add_dataset_file(conn, dataset_id, file.filename or "upload.jsonl", content))
|
||||
return ok({"files": created})
|
||||
|
||||
|
||||
@router.get("/dataset-manage/download/{dataset_id}")
|
||||
async def download_dataset(dataset_id: str) -> PlainTextResponse:
|
||||
dataset = get_platform_store().dataset(dataset_id)
|
||||
content = "\n".join([f"{file['name']}" for file in dataset.get("files", [])])
|
||||
return PlainTextResponse(content, media_type="text/plain")
|
||||
|
||||
|
||||
@router.get("/dataset-manage/download/{dataset_id}/{file_id}")
|
||||
async def download_dataset_file(dataset_id: str, file_id: str, version_id: str | None = Query(default=None)) -> PlainTextResponse:
|
||||
row = get_platform_store().dataset_file(file_id)
|
||||
return PlainTextResponse(row["content"], media_type="text/plain")
|
||||
|
||||
|
||||
@router.get("/dataset-manage")
|
||||
async def dataset_list() -> dict[str, Any]:
|
||||
return ok(get_platform_store().datasets())
|
||||
|
||||
|
||||
@router.post("/dataset-manage")
|
||||
async def create_dataset(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
dataset = get_platform_store().create_dataset(payload)
|
||||
return ok({"id": dataset["id"]})
|
||||
|
||||
|
||||
@router.get("/dataset-manage/{dataset_id}")
|
||||
async def dataset_detail(dataset_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().dataset(dataset_id))
|
||||
except KeyError:
|
||||
raise fail(404, "dataset not found")
|
||||
|
||||
|
||||
@router.put("/dataset-manage/{dataset_id}")
|
||||
async def update_dataset(dataset_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_dataset(dataset_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "dataset not found")
|
||||
|
||||
|
||||
@router.delete("/dataset-manage/{dataset_id}")
|
||||
async def delete_dataset(dataset_id: str) -> dict[str, Any]:
|
||||
get_platform_store().delete_dataset(dataset_id)
|
||||
return ok({"deleted": dataset_id})
|
||||
|
||||
|
||||
@router.get("/fine-tune/check-name")
|
||||
async def check_fine_tune_name(name: str = Query(...)) -> dict[str, Any]:
|
||||
exists = any(task["name"] == name for task in get_platform_store().tasks())
|
||||
return ok({"exists": exists})
|
||||
|
||||
|
||||
@router.get("/fine-tune/progress/{task_id}")
|
||||
async def fine_tune_progress(task_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().progress(task_id))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
|
||||
|
||||
@router.post("/fine-tune/tensorboard/start")
|
||||
async def tensorboard_start() -> dict[str, Any]:
|
||||
return ok({"status": "running", "url": "http://localhost:6006"})
|
||||
|
||||
|
||||
@router.get("/fine-tune")
|
||||
async def fine_tune_list() -> dict[str, Any]:
|
||||
return ok(get_platform_store().tasks())
|
||||
|
||||
|
||||
@router.post("/fine-tune")
|
||||
async def create_fine_tune(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
task = get_platform_store().create_task(apply_presets(payload))
|
||||
return ok({"id": task["id"]})
|
||||
except ValueError as exc:
|
||||
raise fail(400, str(exc))
|
||||
|
||||
|
||||
@router.post("/fine-tune/start")
|
||||
async def start_fine_tune(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().start_task(payload))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
except RuntimeError as exc:
|
||||
raise fail(409, str(exc))
|
||||
|
||||
|
||||
@router.get("/fine-tune/{task_id}")
|
||||
async def fine_tune_detail(task_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().task(task_id))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
|
||||
|
||||
@router.put("/fine-tune/{task_id}")
|
||||
async def update_fine_tune(task_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_task(task_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
|
||||
|
||||
@router.post("/fine-tune/stop/{task_id}")
|
||||
async def stop_fine_tune(task_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().stop_task(task_id))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
|
||||
|
||||
@router.post("/fine-tune/{task_id}/stop")
|
||||
async def stop_fine_tune_alt(task_id: str) -> dict[str, Any]:
|
||||
return await stop_fine_tune(task_id)
|
||||
|
||||
|
||||
@router.post("/fine-tune/pause/{task_id}")
|
||||
async def pause_fine_tune(task_id: str) -> dict[str, Any]:
|
||||
if not get_platform_store().pause_task(task_id):
|
||||
raise fail(409, "当前没有可暂停的训练进程")
|
||||
return ok({"paused": task_id})
|
||||
|
||||
|
||||
@router.post("/fine-tune/resume/{task_id}")
|
||||
async def resume_fine_tune(task_id: str) -> dict[str, Any]:
|
||||
if not get_platform_store().resume_task_engine(task_id):
|
||||
raise fail(409, "当前没有可恢复的训练进程")
|
||||
return ok({"resumed": task_id})
|
||||
|
||||
|
||||
@router.post("/fine-tune/cancel/{task_id}")
|
||||
async def cancel_fine_tune(task_id: str) -> dict[str, Any]:
|
||||
get_platform_store().cancel_task_engine(task_id)
|
||||
return ok({"canceled": task_id})
|
||||
|
||||
|
||||
@router.delete("/fine-tune/{task_id}")
|
||||
async def delete_fine_tune(task_id: str) -> dict[str, Any]:
|
||||
get_platform_store().delete_task(task_id)
|
||||
return ok({"deleted": task_id})
|
||||
|
||||
|
||||
@router.get("/fine-tune/{task_id}/overview")
|
||||
async def fine_tune_overview(task_id: str) -> dict[str, Any]:
|
||||
task = get_platform_store().task(task_id)
|
||||
return ok({"task": task, "progress": get_platform_store().progress(task_id)})
|
||||
|
||||
|
||||
@router.get("/fine-tune/{task_id}/checkpoints")
|
||||
async def fine_tune_checkpoints(task_id: str) -> dict[str, Any]:
|
||||
task = get_platform_store().task(task_id)
|
||||
checkpoints = []
|
||||
for step in [50, 100, 150]:
|
||||
if task.get("progress", 0) >= min(100, step // 2):
|
||||
checkpoints.append({"step": step, "path": f"/data/yg-ft/outputs/{task['name']}/checkpoint-{step}"})
|
||||
return ok(checkpoints)
|
||||
|
||||
|
||||
@router.get("/compute/nodes")
|
||||
async def compute_nodes() -> dict[str, Any]:
|
||||
return ok(get_platform_store().compute_nodes())
|
||||
|
||||
|
||||
@router.post("/compute/nodes")
|
||||
async def create_compute_node(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().create_compute_node(payload))
|
||||
except KeyError as exc:
|
||||
raise fail(400, f"missing field: {exc}")
|
||||
|
||||
|
||||
@router.put("/compute/nodes/{node_id}")
|
||||
async def update_compute_node(node_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_compute_node(node_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "compute node not found")
|
||||
|
||||
|
||||
@router.post("/compute/nodes/{node_id}/test-connection")
|
||||
async def test_compute_node(node_id: str) -> dict[str, Any]:
|
||||
return ok({"node_id": node_id, "success": True, "latency_ms": 12})
|
||||
|
||||
|
||||
@router.post("/compute/nodes/{node_id}/enable")
|
||||
async def enable_compute_node(node_id: str) -> dict[str, Any]:
|
||||
return ok(get_platform_store().update_compute_node(node_id, {"enabled": True, "scheduler_status": "online"}))
|
||||
|
||||
|
||||
@router.post("/compute/nodes/{node_id}/disable")
|
||||
async def disable_compute_node(node_id: str) -> dict[str, Any]:
|
||||
return ok(get_platform_store().update_compute_node(node_id, {"enabled": False, "scheduler_status": "offline"}))
|
||||
|
||||
|
||||
@router.post("/compute/nodes/{node_id}/drain")
|
||||
async def drain_compute_node(node_id: str) -> dict[str, Any]:
|
||||
return ok(get_platform_store().update_compute_node(node_id, {"scheduler_status": "draining"}))
|
||||
|
||||
|
||||
@router.get("/compute/nodes/{node_id}/replicas")
|
||||
async def compute_node_replicas(node_id: str) -> dict[str, Any]:
|
||||
return ok(get_platform_store().replicas(node_id))
|
||||
|
||||
|
||||
@router.get("/compute/gpus")
|
||||
async def compute_gpus() -> dict[str, Any]:
|
||||
return ok(get_platform_store().gpus())
|
||||
|
||||
|
||||
@router.get("/compute/queue")
|
||||
async def compute_queue() -> dict[str, Any]:
|
||||
return ok(get_platform_store().queue())
|
||||
|
||||
|
||||
@router.post("/internal/compute-sync/resources")
|
||||
async def create_compute_sync(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
sync_id = get_platform_store().create_sync_job(payload.get("target_node_id", "node_01"), payload)
|
||||
return ok(get_platform_store().sync_job(sync_id))
|
||||
|
||||
|
||||
@router.get("/internal/compute-sync/resources/{sync_id}")
|
||||
async def compute_sync_detail(sync_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().sync_job(sync_id))
|
||||
except KeyError:
|
||||
raise fail(404, "sync job not found")
|
||||
|
||||
|
||||
@router.get("/training-log-files")
|
||||
async def training_log_files() -> dict[str, Any]:
|
||||
return ok(get_platform_store().training_log_files())
|
||||
|
||||
|
||||
@router.get("/training-log-content")
|
||||
async def training_log_content(file: str = Query(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().training_log_content(file))
|
||||
except KeyError:
|
||||
raise fail(404, "training log not found")
|
||||
|
||||
|
||||
@router.get("/log-files")
|
||||
async def log_files(date: str | None = Query(default=None)) -> dict[str, Any]:
|
||||
return ok(get_platform_store().log_files(date))
|
||||
|
||||
|
||||
@router.get("/log-content")
|
||||
async def log_content(file: str = Query(...)) -> dict[str, Any]:
|
||||
return ok(get_platform_store().log_content(file))
|
||||
|
||||
|
||||
@router.post("/web-log")
|
||||
async def web_log(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
return ok({"received": True, **payload})
|
||||
|
||||
|
||||
# ===================== Project Management (§13.2) =====================
|
||||
|
||||
|
||||
@router.get("/projects")
|
||||
async def project_list(
|
||||
tenant_id: str = Query(default="default"),
|
||||
status: str | None = Query(default=None),
|
||||
keyword: str | None = Query(default=None),
|
||||
) -> dict[str, Any]:
|
||||
return ok(get_platform_store().projects(tenant_id, status, keyword))
|
||||
|
||||
|
||||
@router.post("/projects")
|
||||
async def create_project(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
project = get_platform_store().create_project(payload)
|
||||
return ok({"id": project["id"]})
|
||||
except KeyError as exc:
|
||||
raise fail(400, f"missing required field: {exc}")
|
||||
|
||||
|
||||
@router.get("/projects/{project_id}")
|
||||
async def project_detail(project_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().project(project_id))
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
|
||||
|
||||
@router.put("/projects/{project_id}")
|
||||
async def update_project(project_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_project(project_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
|
||||
|
||||
@router.post("/projects/{project_id}/activate")
|
||||
async def activate_project(project_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().activate_project(project_id))
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
|
||||
|
||||
@router.get("/projects/{project_id}/members")
|
||||
async def project_members(project_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().project_members(project_id))
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
|
||||
|
||||
@router.post("/projects/{project_id}/members")
|
||||
async def add_project_member(project_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
user_id = payload.get("user_id")
|
||||
if not user_id:
|
||||
raise fail(400, "user_id is required")
|
||||
try:
|
||||
return ok(get_platform_store().add_project_member(project_id, user_id, payload.get("role", "member")))
|
||||
except KeyError as exc:
|
||||
raise fail(404, str(exc))
|
||||
|
||||
|
||||
@router.put("/projects/{project_id}/members/{user_id}")
|
||||
async def update_project_member_role(project_id: str, user_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().update_project_member_role(project_id, user_id, payload.get("role", "member")))
|
||||
except KeyError as exc:
|
||||
raise fail(404, str(exc))
|
||||
|
||||
|
||||
@router.delete("/projects/{project_id}/members/{user_id}")
|
||||
async def remove_project_member(project_id: str, user_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
get_platform_store().remove_project_member(project_id, user_id)
|
||||
return ok({"deleted": user_id})
|
||||
except KeyError as exc:
|
||||
raise fail(404, str(exc))
|
||||
|
||||
|
||||
# ===================== Fine-tune Events (§7.1) =====================
|
||||
|
||||
@router.get("/fine-tune/{task_id}/events")
|
||||
async def fine_tune_events(task_id: str) -> StreamingResponse:
|
||||
import json as _json
|
||||
|
||||
async def event_stream():
|
||||
store = get_platform_store()
|
||||
try:
|
||||
events = store.task_events(task_id)
|
||||
for event in events:
|
||||
yield f"data: {_json.dumps(event, default=str)}\n\n"
|
||||
yield f"data: {_json.dumps({'type': 'done', 'data': {}}, default=str)}\n\n"
|
||||
except KeyError:
|
||||
yield f"data: {_json.dumps({'type': 'error', 'data': {'message': 'task not found'}}, default=str)}\n\n"
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
|
||||
# ===================== Fine-tune Retry & Resume (§13.9) =====================
|
||||
|
||||
|
||||
@router.post("/fine-tune/{task_id}/retry")
|
||||
async def retry_fine_tune(task_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().retry_task(task_id))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
except ValueError as exc:
|
||||
raise fail(400, str(exc))
|
||||
|
||||
|
||||
@router.post("/fine-tune/{task_id}/resume")
|
||||
async def resume_fine_tune(task_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
checkpoint_id = payload.get("checkpoint_id")
|
||||
if not checkpoint_id:
|
||||
raise fail(400, "checkpoint_id is required")
|
||||
try:
|
||||
return ok(get_platform_store().resume_task(task_id, checkpoint_id))
|
||||
except KeyError as exc:
|
||||
raise fail(404, str(exc))
|
||||
except ValueError as exc:
|
||||
raise fail(400, str(exc))
|
||||
|
||||
|
||||
# ===================== Checkpoint Management (§13.9) =====================
|
||||
|
||||
|
||||
@router.delete("/fine-tune/{task_id}/checkpoints/{checkpoint_id}")
|
||||
async def delete_checkpoint(task_id: str, checkpoint_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
get_platform_store().delete_checkpoint(checkpoint_id)
|
||||
return ok({"deleted": checkpoint_id})
|
||||
except KeyError:
|
||||
raise fail(404, "checkpoint not found")
|
||||
|
||||
|
||||
@router.put("/fine-tune/{task_id}/checkpoint-retention")
|
||||
async def set_checkpoint_retention(task_id: str, payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().set_checkpoint_retention(task_id, payload))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
|
||||
|
||||
@router.get("/fine-tune/{task_id}/checkpoint-retention")
|
||||
async def get_checkpoint_retention(task_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().get_checkpoint_retention(task_id))
|
||||
except KeyError:
|
||||
raise fail(404, "fine tune task not found")
|
||||
|
||||
|
||||
# ===================== Compute Jobs (§13.6) =====================
|
||||
|
||||
|
||||
@router.post("/compute/jobs")
|
||||
async def create_compute_job(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
try:
|
||||
job = get_platform_store().create_compute_job(payload)
|
||||
return ok({"id": job["id"]})
|
||||
except KeyError as exc:
|
||||
raise fail(400, f"missing required field: {exc}")
|
||||
|
||||
|
||||
@router.get("/compute/jobs/{job_id}")
|
||||
async def compute_job_detail(job_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().compute_job(job_id))
|
||||
except KeyError:
|
||||
raise fail(404, "compute job not found")
|
||||
|
||||
|
||||
@router.get("/compute/jobs")
|
||||
async def compute_jobs(task_id: str = Query(default=None)) -> dict[str, Any]:
|
||||
if task_id:
|
||||
return ok(get_platform_store().compute_jobs_by_task(task_id))
|
||||
return ok([])
|
||||
|
||||
|
||||
@router.post("/compute/jobs/{job_id}/stop")
|
||||
async def stop_compute_job(job_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().stop_compute_job(job_id))
|
||||
except KeyError:
|
||||
raise fail(404, "compute job not found")
|
||||
|
||||
|
||||
@router.get("/compute/jobs/{job_id}/logs")
|
||||
async def compute_job_logs(job_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().compute_job_logs(job_id))
|
||||
except KeyError as exc:
|
||||
raise fail(404, str(exc))
|
||||
|
||||
21
backend/app/api/v1/router.py
Normal file
21
backend/app/api/v1/router.py
Normal file
@@ -0,0 +1,21 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.v1.endpoints.platform import router as platform_router
|
||||
from app.api.v1.endpoints.health import router as health_router
|
||||
from app.modules.auth import router as auth_router
|
||||
from app.modules.system import router as system_router
|
||||
from app.modules.tenant import router as tenant_router
|
||||
from app.modules.project import router as project_router
|
||||
from app.modules.resource import router as resource_router
|
||||
from app.modules.approval import router as approval_router
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(health_router, tags=["health"])
|
||||
api_router.include_router(platform_router, tags=["platform"])
|
||||
api_router.include_router(auth_router, tags=["auth"])
|
||||
api_router.include_router(system_router, tags=["system"])
|
||||
api_router.include_router(tenant_router, tags=["tenant"])
|
||||
api_router.include_router(project_router, tags=["project"])
|
||||
api_router.include_router(resource_router, tags=["resource"])
|
||||
api_router.include_router(approval_router, tags=["approval"])
|
||||
|
||||
1
backend/app/core/__init__.py
Normal file
1
backend/app/core/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Core infrastructure modules."""
|
||||
60
backend/app/core/config.py
Normal file
60
backend/app/core/config.py
Normal file
@@ -0,0 +1,60 @@
|
||||
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)
|
||||
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)
|
||||
jwt_secret: str = os.getenv("JWT_SECRET", "dev-insecure-change-me")
|
||||
jwt_algorithm: str = os.getenv("JWT_ALGORITHM", "HS256")
|
||||
access_token_expire_minutes: int = _int_env("ACCESS_TOKEN_EXPIRE_MINUTES", 1440)
|
||||
|
||||
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
253
backend/app/core/logging.py
Normal 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)
|
||||
1
backend/app/db/__init__.py
Normal file
1
backend/app/db/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Database infrastructure package."""
|
||||
1946
backend/app/db/platform_store.py
Normal file
1946
backend/app/db/platform_store.py
Normal file
File diff suppressed because it is too large
Load Diff
40
backend/app/db/session.py
Normal file
40
backend/app/db/session.py
Normal file
@@ -0,0 +1,40 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
|
||||
DATABASE_URL = os.getenv("DATABASE_URL", "postgresql+psycopg://yg_ft:change_me@localhost:15432/yg_ft")
|
||||
|
||||
engine = create_engine(
|
||||
DATABASE_URL,
|
||||
pool_pre_ping=True,
|
||||
future=True,
|
||||
)
|
||||
SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False, expire_on_commit=False, future=True)
|
||||
|
||||
|
||||
def get_db() -> Generator[Session, None, None]:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def session_scope() -> Generator[Session, None, None]:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
yield db
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
214
backend/app/db/sql/001_platform_runtime.sql
Normal file
214
backend/app/db/sql/001_platform_runtime.sql
Normal file
@@ -0,0 +1,214 @@
|
||||
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 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,
|
||||
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
|
||||
);
|
||||
|
||||
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
|
||||
);
|
||||
|
||||
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,
|
||||
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
|
||||
);
|
||||
|
||||
-- ===================== Project / Tenant =====================
|
||||
|
||||
CREATE TABLE IF NOT EXISTS projects (
|
||||
id TEXT PRIMARY KEY,
|
||||
tenant_id TEXT NOT NULL DEFAULT 'default',
|
||||
name TEXT NOT NULL,
|
||||
code TEXT NOT NULL,
|
||||
description TEXT,
|
||||
quota TEXT,
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
create_time TEXT NOT NULL,
|
||||
create_by TEXT,
|
||||
updated_at TEXT
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS project_members (
|
||||
project_id TEXT NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
|
||||
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
role TEXT NOT NULL DEFAULT 'member',
|
||||
create_time TEXT NOT NULL,
|
||||
PRIMARY KEY (project_id, user_id)
|
||||
);
|
||||
|
||||
-- ===================== Fine-tune Checkpoints =====================
|
||||
|
||||
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,
|
||||
name TEXT NOT NULL,
|
||||
path TEXT NOT NULL,
|
||||
step INTEGER NOT NULL DEFAULT 0,
|
||||
loss DOUBLE PRECISION,
|
||||
is_best INTEGER NOT NULL DEFAULT 0,
|
||||
size_bytes BIGINT DEFAULT 0,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
-- ===================== Compute Jobs (internal) =====================
|
||||
|
||||
CREATE TABLE IF NOT EXISTS compute_jobs (
|
||||
id TEXT PRIMARY KEY,
|
||||
task_id TEXT NOT NULL REFERENCES fine_tune_tasks(id) ON DELETE CASCADE,
|
||||
node_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
type TEXT NOT NULL DEFAULT 'train',
|
||||
status TEXT NOT NULL DEFAULT 'pending',
|
||||
command TEXT,
|
||||
gpu_count INTEGER NOT NULL DEFAULT 1,
|
||||
priority INTEGER NOT NULL DEFAULT 0,
|
||||
timeout_seconds INTEGER,
|
||||
progress REAL DEFAULT 0,
|
||||
result TEXT,
|
||||
error_message TEXT,
|
||||
create_time TEXT NOT NULL,
|
||||
start_time TEXT,
|
||||
completed_at TEXT
|
||||
);
|
||||
|
||||
-- ===================== Indexes =====================
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_fine_tune_status ON fine_tune_tasks(status);
|
||||
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 INDEX IF NOT EXISTS idx_replicas_resource ON resource_replicas(resource_type, resource_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_projects_tenant ON projects(tenant_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_project_members_user ON project_members(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_checkpoints_task ON fine_tune_checkpoints(task_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_compute_jobs_task ON compute_jobs(task_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_compute_jobs_node ON compute_jobs(node_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_compute_jobs_status ON compute_jobs(status);
|
||||
|
||||
-- ===================== Migrations: extend fine_tune_tasks =====================
|
||||
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS project_id TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS tenant_id TEXT NOT NULL DEFAULT 'default';
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS dataset_name TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS model_name TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS trained_model_name TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS checkpoint_retention_policy TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS error_message TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS retry_count INTEGER NOT NULL DEFAULT 0;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS last_retry_at TEXT;
|
||||
ALTER TABLE fine_tune_tasks ADD COLUMN IF NOT EXISTS resumed_from_checkpoint_id TEXT;
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_fine_tune_project ON fine_tune_tasks(project_id);
|
||||
127
backend/app/db/sql/002_governance.sql
Normal file
127
backend/app/db/sql/002_governance.sql
Normal file
@@ -0,0 +1,127 @@
|
||||
-- A. 平台基础与企业治理:权限 / 角色 / 租户 / 审批 / 审计 / 留存
|
||||
-- 沿用 001 的约定:时间戳存 TEXT,布尔用 INTEGER,列表用 TEXT(JSON)
|
||||
|
||||
CREATE TABLE IF NOT EXISTS permissions (
|
||||
code TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
group_name TEXT NOT NULL,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS roles (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
display_name TEXT NOT NULL,
|
||||
permissions TEXT NOT NULL,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS role_permissions (
|
||||
role_id TEXT NOT NULL,
|
||||
permission TEXT NOT NULL,
|
||||
PRIMARY KEY (role_id, permission)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS user_permission_overrides (
|
||||
user_id TEXT NOT NULL,
|
||||
permission TEXT NOT NULL,
|
||||
granted INTEGER NOT NULL,
|
||||
PRIMARY KEY (user_id, permission)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tenants (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
code TEXT NOT NULL UNIQUE,
|
||||
status TEXT NOT NULL,
|
||||
owner_user_id TEXT,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tenant_users (
|
||||
tenant_id TEXT NOT NULL,
|
||||
user_id TEXT NOT NULL,
|
||||
role TEXT NOT NULL,
|
||||
PRIMARY KEY (tenant_id, user_id)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS resource_acl (
|
||||
resource_type TEXT NOT NULL,
|
||||
resource_id TEXT NOT NULL,
|
||||
principal_type TEXT NOT NULL,
|
||||
principal_id TEXT NOT NULL,
|
||||
permission TEXT NOT NULL,
|
||||
granted INTEGER NOT NULL,
|
||||
PRIMARY KEY (resource_type, resource_id, principal_type, principal_id, permission)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS approval_templates (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
steps TEXT NOT NULL,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS approval_instances (
|
||||
id TEXT PRIMARY KEY,
|
||||
template_id TEXT,
|
||||
resource_type TEXT NOT NULL,
|
||||
resource_id TEXT NOT NULL,
|
||||
applicant_id TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
current_step INTEGER NOT NULL DEFAULT 0,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS approval_steps (
|
||||
instance_id TEXT NOT NULL,
|
||||
step_index INTEGER NOT NULL,
|
||||
approver_id TEXT,
|
||||
status TEXT NOT NULL,
|
||||
comment TEXT,
|
||||
time TEXT,
|
||||
PRIMARY KEY (instance_id, step_index)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS audit_logs (
|
||||
id TEXT PRIMARY KEY,
|
||||
tenant_id TEXT,
|
||||
project_id TEXT,
|
||||
actor_id TEXT,
|
||||
action TEXT NOT NULL,
|
||||
target_type TEXT,
|
||||
target_id TEXT,
|
||||
detail TEXT,
|
||||
client_ip TEXT,
|
||||
time TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS retention_policies (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
resource_type TEXT NOT NULL,
|
||||
retention_days INTEGER NOT NULL,
|
||||
create_time TEXT NOT NULL
|
||||
);
|
||||
|
||||
-- ===================== 种子数据 =====================
|
||||
INSERT INTO permissions (code, name, group_name, create_time) VALUES
|
||||
('dashboard', '仪表盘', '概览', '2026-01-01T00:00:00Z'),
|
||||
('fine-tune', '模型微调', '模型', '2026-01-01T00:00:00Z'),
|
||||
('model-eval', '模型评估', '模型', '2026-01-01T00:00:00Z'),
|
||||
('model-inference', '模型推理', '模型', '2026-01-01T00:00:00Z'),
|
||||
('model-manage', '模型管理', '模型', '2026-01-01T00:00:00Z'),
|
||||
('dataset', '数据集', '数据', '2026-01-01T00:00:00Z'),
|
||||
('data-process', '数据处理', '数据', '2026-01-01T00:00:00Z'),
|
||||
('data-convert', '数据转换', '数据', '2026-01-01T00:00:00Z'),
|
||||
('compute', '算力管理', '算力', '2026-01-01T00:00:00Z'),
|
||||
('hardware', '硬件监控', '算力', '2026-01-01T00:00:00Z'),
|
||||
('logs', '日志查看', '运维', '2026-01-01T00:00:00Z'),
|
||||
('user-settings', '用户设置', '运维', '2026-01-01T00:00:00Z')
|
||||
ON CONFLICT (code) DO NOTHING;
|
||||
|
||||
INSERT INTO roles (id, name, display_name, permissions, create_time) VALUES
|
||||
('role_admin', 'admin', '管理员', '["dashboard","fine-tune","model-eval","model-inference","model-manage","dataset","data-process","data-convert","compute","hardware","logs","user-settings"]', '2026-01-01T00:00:00Z'),
|
||||
('role_operator','operator','操作员', '["dashboard","fine-tune","model-eval","model-inference","model-manage","dataset","data-process","data-convert","compute","hardware","logs"]', '2026-01-01T00:00:00Z'),
|
||||
('role_viewer', 'viewer', '访客', '["dashboard"]', '2026-01-01T00:00:00Z')
|
||||
ON CONFLICT (name) DO NOTHING;
|
||||
17
backend/app/db/sql/003_tenant_quota.sql
Normal file
17
backend/app/db/sql/003_tenant_quota.sql
Normal file
@@ -0,0 +1,17 @@
|
||||
-- 003: tenants 补充 quota / retention_policy 列(接口契约 13.1 需要)
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_name = 'tenants' AND column_name = 'quota'
|
||||
) THEN
|
||||
ALTER TABLE tenants ADD COLUMN quota TEXT;
|
||||
END IF;
|
||||
IF NOT EXISTS (
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_name = 'tenants' AND column_name = 'retention_policy_id'
|
||||
) THEN
|
||||
ALTER TABLE tenants ADD COLUMN retention_policy_id TEXT;
|
||||
END IF;
|
||||
END
|
||||
$$;
|
||||
26
backend/app/main.py
Normal file
26
backend/app/main.py
Normal file
@@ -0,0 +1,26 @@
|
||||
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
|
||||
|
||||
|
||||
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)
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
15
backend/app/modules/README.md
Normal file
15
backend/app/modules/README.md
Normal file
@@ -0,0 +1,15 @@
|
||||
# Backend Module Convention
|
||||
|
||||
每个业务模块建议保持一致结构:
|
||||
|
||||
```text
|
||||
module_name/
|
||||
__init__.py
|
||||
router.py # FastAPI router
|
||||
schemas.py # Pydantic request/response models
|
||||
service.py # Business orchestration
|
||||
repository.py # Database access
|
||||
permissions.py # Optional resource permission checks
|
||||
```
|
||||
|
||||
模块边界以 `docs/system-development-plan.md` 的页面模块开发工作包为准。
|
||||
5
backend/app/modules/approval/__init__.py
Normal file
5
backend/app/modules/approval/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
"""Approval workflow module."""
|
||||
|
||||
from app.modules.approval.router import router
|
||||
|
||||
__all__ = ["router"]
|
||||
67
backend/app/modules/approval/router.py
Normal file
67
backend/app/modules/approval/router.py
Normal file
@@ -0,0 +1,67 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Body
|
||||
from typing import Any
|
||||
|
||||
from app.api.v1.endpoints.platform import ok, fail
|
||||
from app.db.platform_store import get_platform_store
|
||||
|
||||
router = APIRouter(prefix="/approvals", tags=["approval"])
|
||||
|
||||
|
||||
@router.get("/templates")
|
||||
def list_templates() -> dict[str, Any]:
|
||||
return ok(get_platform_store().approval_templates())
|
||||
|
||||
|
||||
@router.post("/templates")
|
||||
def create_template(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
if not payload.get("name"):
|
||||
raise fail(400, "name 必填")
|
||||
return ok(get_platform_store().create_approval_template(payload))
|
||||
|
||||
|
||||
@router.get("")
|
||||
def list_instances(status: str | None = None) -> dict[str, Any]:
|
||||
return ok(get_platform_store().approval_instances(status=status))
|
||||
|
||||
|
||||
@router.post("")
|
||||
def create_instance(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
for field in ("resource_type", "resource_id", "applicant_id"):
|
||||
if not payload.get(field):
|
||||
raise fail(400, f"{field} 必填")
|
||||
try:
|
||||
return ok(get_platform_store().create_approval_instance(payload))
|
||||
except KeyError:
|
||||
raise fail(404, "template not found")
|
||||
|
||||
|
||||
@router.get("/{instance_id}")
|
||||
def get_instance(instance_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().approval_instance(instance_id))
|
||||
except KeyError:
|
||||
raise fail(404, "instance not found")
|
||||
|
||||
|
||||
@router.post("/{instance_id}/steps/{step_index}/decision")
|
||||
def decide(
|
||||
instance_id: str,
|
||||
step_index: int,
|
||||
payload: dict[str, Any] = Body(...),
|
||||
) -> dict[str, Any]:
|
||||
if not payload.get("approver_id"):
|
||||
raise fail(400, "approver_id 必填")
|
||||
try:
|
||||
return ok(
|
||||
get_platform_store().decide_approval_step(
|
||||
instance_id,
|
||||
step_index,
|
||||
approver_id=payload["approver_id"],
|
||||
approved=bool(payload.get("approved", False)),
|
||||
comment=payload.get("comment"),
|
||||
)
|
||||
)
|
||||
except (KeyError, ValueError) as e:
|
||||
raise fail(400, str(e))
|
||||
1
backend/app/modules/audit/__init__.py
Normal file
1
backend/app/modules/audit/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Audit log module."""
|
||||
3
backend/app/modules/auth/__init__.py
Normal file
3
backend/app/modules/auth/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
from app.modules.auth.router import router
|
||||
|
||||
__all__ = ["router"]
|
||||
20
backend/app/modules/auth/deps.py
Normal file
20
backend/app/modules/auth/deps.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import Depends, Header, HTTPException, status
|
||||
|
||||
from app.db.platform_store import get_platform_store
|
||||
from app.modules.auth.service import decode_access_token
|
||||
|
||||
|
||||
def get_current_user(authorization: str | None = Header(default=None)) -> dict:
|
||||
"""从 Bearer 令牌解析出当前登录用户,供受保护接口依赖使用。"""
|
||||
if not authorization or not authorization.lower().startswith("bearer "):
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="缺少认证令牌")
|
||||
token = authorization.split(" ", 1)[1].strip()
|
||||
user_id = decode_access_token(token)
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="令牌无效或已过期")
|
||||
user = get_platform_store().user_by_id(user_id)
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="用户不存在")
|
||||
return user
|
||||
30
backend/app/modules/auth/router.py
Normal file
30
backend/app/modules/auth/router.py
Normal file
@@ -0,0 +1,30 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.db.platform_store import get_platform_store
|
||||
from app.modules.auth.deps import get_current_user
|
||||
from app.modules.auth.service import create_access_token
|
||||
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
class LoginBody(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
@router.post("/login")
|
||||
def login(body: LoginBody) -> dict:
|
||||
user = get_platform_store().login(body.username, body.password)
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="用户名或密码错误")
|
||||
token = create_access_token(user["id"])
|
||||
return {"code": 0, "message": "ok", "data": {"token": token, "user": user}}
|
||||
|
||||
|
||||
@router.get("/me")
|
||||
def me(current_user: dict = Depends(get_current_user)) -> dict:
|
||||
return {"code": 0, "message": "ok", "data": current_user}
|
||||
33
backend/app/modules/auth/service.py
Normal file
33
backend/app/modules/auth/service.py
Normal file
@@ -0,0 +1,33 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import jwt
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from app.core.config import get_settings
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def create_access_token(user_id: str, expires_minutes: int | None = None) -> str:
|
||||
"""为指定用户签发 JWT 访问令牌。"""
|
||||
settings = get_settings()
|
||||
expire = _now() + timedelta(minutes=expires_minutes or settings.access_token_expire_minutes)
|
||||
payload = {
|
||||
"sub": user_id,
|
||||
"iat": int(_now().timestamp()),
|
||||
"exp": int(expire.timestamp()),
|
||||
}
|
||||
return jwt.encode(payload, settings.jwt_secret, algorithm=settings.jwt_algorithm)
|
||||
|
||||
|
||||
def decode_access_token(token: str) -> str | None:
|
||||
"""校验并返回令牌中的用户 ID;无效/过期返回 None。"""
|
||||
settings = get_settings()
|
||||
try:
|
||||
payload = jwt.decode(token, settings.jwt_secret, algorithms=[settings.jwt_algorithm])
|
||||
except jwt.PyJWTError:
|
||||
return None
|
||||
sub = payload.get("sub")
|
||||
return sub if isinstance(sub, str) else None
|
||||
1
backend/app/modules/compute_gateway/__init__.py
Normal file
1
backend/app/modules/compute_gateway/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Application-side compute platform gateway module."""
|
||||
1
backend/app/modules/data_process/__init__.py
Normal file
1
backend/app/modules/data_process/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Data processing module."""
|
||||
1
backend/app/modules/dataset/__init__.py
Normal file
1
backend/app/modules/dataset/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Dataset management module."""
|
||||
1
backend/app/modules/engine_registry/__init__.py
Normal file
1
backend/app/modules/engine_registry/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Training engine registry module."""
|
||||
1
backend/app/modules/eval/__init__.py
Normal file
1
backend/app/modules/eval/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Evaluation module."""
|
||||
1
backend/app/modules/file_gateway/__init__.py
Normal file
1
backend/app/modules/file_gateway/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Application-side file gateway module."""
|
||||
7
backend/app/modules/fine_tune/__init__.py
Normal file
7
backend/app/modules/fine_tune/__init__.py
Normal file
@@ -0,0 +1,7 @@
|
||||
"""模型训练模块(移植自模型服务 projects/backend)。
|
||||
|
||||
- service.py: 业务编排(preset 参数预设、train_type→stage 映射、启动/暂停/恢复/取消)。
|
||||
- runner.py: 真实训练执行器(基于 LLaMA-Factory 的 llamafactory-cli 子进程 + 实时 loss 监控)。
|
||||
|
||||
当前后端默认 COMPUTE_MODE=real 时由本模块真正驱动训练;本地无 GPU 用 simulator 时不触发,行为不变。
|
||||
"""
|
||||
423
backend/app/modules/fine_tune/runner.py
Normal file
423
backend/app/modules/fine_tune/runner.py
Normal file
@@ -0,0 +1,423 @@
|
||||
"""
|
||||
真实训练执行器(基于项目内 train.py 调用 LLaMA-Factory 库)
|
||||
|
||||
- 根据任务配置构建 `python train.py` 命令(非 llamafactory-cli 子命令)
|
||||
- 实时捕获训练日志与 trainer_log.jsonl 的 loss,回写到 PlatformStore
|
||||
- 支持 SIGSTOP / SIGCONT / SIGKILL 实现暂停 / 恢复 / 取消
|
||||
- 训练完成后自动绘制 loss 曲线(matplotlib 可选)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import signal
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from app.db.platform_store import get_platform_store
|
||||
|
||||
BACKEND_ROOT = Path(__file__).resolve().parent.parent.parent.parent
|
||||
DATA_DIR = BACKEND_ROOT / "data"
|
||||
DATASET_INFO_PATH = DATA_DIR / "dataset_info.json"
|
||||
DATASET_STORE_DIR = DATA_DIR / "fine_tune_datasets"
|
||||
OUTPUT_ROOT = DATA_DIR / "fine_tune_outputs"
|
||||
TRAIN_SCRIPT = BACKEND_ROOT / "train.py"
|
||||
|
||||
_running_processes: dict[str, "subprocess.Popen[Any]"] = {}
|
||||
|
||||
_stage_map = {"sft": "sft", "dpo": "dpo", "cpt": "pt", "cot": "sft"}
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# 工具函数
|
||||
# ──────────────────────────────────────────────
|
||||
def _log(store: Any, task_id: str, msg: str, log_type: str = "info") -> None:
|
||||
try:
|
||||
task = store.task(task_id)
|
||||
logs = list(task.get("logs") or [])
|
||||
logs.append({"time": datetime.now().strftime("%H:%M:%S"), "msg": msg, "type": log_type})
|
||||
store.update_task_runtime(task_id, extra={"logs": logs})
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _resolve_model_path(store: Any, name_or_path: str) -> str:
|
||||
"""把模型短名解析为本地绝对路径;已是合法路径则原样返回。"""
|
||||
if not name_or_path:
|
||||
return name_or_path
|
||||
p = Path(name_or_path)
|
||||
if p.exists() and (p / "config.json").exists():
|
||||
return str(p.resolve())
|
||||
for m in store.models() or []:
|
||||
if m.get("name") == name_or_path or m.get("path") == name_or_path:
|
||||
resolved = Path(m.get("path", ""))
|
||||
if resolved.exists():
|
||||
return str(resolved.resolve())
|
||||
return name_or_path
|
||||
|
||||
|
||||
def _materialize_dataset(store: Any, dataset_id: str) -> str:
|
||||
"""把数据集管理系统的 UUID 数据集落盘并注册进 dataset_info.json,返回 --dataset key。"""
|
||||
if not dataset_id or dataset_id == "identity":
|
||||
return dataset_id
|
||||
try:
|
||||
existing = json.loads(DATASET_INFO_PATH.read_text(encoding="utf-8")) if DATASET_INFO_PATH.exists() else {}
|
||||
except Exception:
|
||||
existing = {}
|
||||
if dataset_id in existing:
|
||||
return dataset_id
|
||||
try:
|
||||
ds = store.dataset(dataset_id)
|
||||
except Exception:
|
||||
return dataset_id
|
||||
files = ds.get("files") or []
|
||||
if not files:
|
||||
return dataset_id
|
||||
file_id = files[0].get("id")
|
||||
ext = files[0].get("ext", ".json")
|
||||
if ext not in (".json", ".jsonl"):
|
||||
ext = ".json"
|
||||
if not file_id:
|
||||
return dataset_id
|
||||
try:
|
||||
row = store.dataset_file(file_id)
|
||||
except Exception:
|
||||
return dataset_id
|
||||
content = row.get("content", "")
|
||||
DATASET_STORE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
actual = DATASET_STORE_DIR / f"{dataset_id}{ext}"
|
||||
actual.write_text(content, encoding="utf-8")
|
||||
rel = actual.relative_to(DATA_DIR)
|
||||
entry: dict[str, Any] = {"file_name": str(rel)}
|
||||
try:
|
||||
text = content.strip()
|
||||
if text.startswith("["):
|
||||
text = text[text.find("{") : text.find("}") + 1]
|
||||
sample = json.loads(text)
|
||||
if isinstance(sample, dict) and ("messages" in sample or "conversations" in sample):
|
||||
key = "messages" if "messages" in sample else "conversations"
|
||||
entry["formatting"] = "sharegpt"
|
||||
entry["columns"] = {"messages": key}
|
||||
except Exception:
|
||||
pass
|
||||
existing[dataset_id] = entry
|
||||
DATASET_INFO_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
DATASET_INFO_PATH.write_text(json.dumps(existing, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
return dataset_id
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# 训练主流程
|
||||
# ──────────────────────────────────────────────
|
||||
def run_training(task_id: str) -> None:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
task = store.task(task_id)
|
||||
except Exception:
|
||||
return
|
||||
|
||||
cfg = dict(task)
|
||||
mode = (cfg.get("train_type") or "sft").lower()
|
||||
stage = _stage_map.get(mode, "sft")
|
||||
|
||||
base_model = _resolve_model_path(store, cfg.get("base_model", ""))
|
||||
train_dataset = _materialize_dataset(store, cfg.get("train_dataset_id", ""))
|
||||
eval_dataset = _materialize_dataset(store, cfg.get("eval_dataset_id") or cfg.get("eval_dataset", ""))
|
||||
|
||||
finetuning_type = (cfg.get("train_method") or "lora").lower()
|
||||
OUTPUT_ROOT.mkdir(parents=True, exist_ok=True)
|
||||
output_dir = str(OUTPUT_ROOT / task["name"])
|
||||
|
||||
env = os.environ.copy()
|
||||
env.setdefault("HF_HUB_OFFLINE", "1")
|
||||
env.setdefault("TRANSFORMERS_OFFLINE", "1")
|
||||
env.setdefault("HF_ENDPOINT", "https://hf-mirror.com")
|
||||
|
||||
gpus = cfg.get("gpus") or [0]
|
||||
num_gpus = int(cfg.get("num_gpus", len(gpus)) or 1) or 1
|
||||
# 仅当显式指定非默认 GPU 时限制可见设备;多卡统一走 torchrun
|
||||
if gpus and str(gpus[0]) not in ("0", "gpu-0"):
|
||||
env["CUDA_VISIBLE_DEVICES"] = ",".join(str(g).replace("gpu-", "") for g in gpus)
|
||||
num_gpus = len(gpus)
|
||||
|
||||
if num_gpus > 1:
|
||||
_log(store, task_id, f"[INFO] 启用分布式训练: {num_gpus} 个 GPU", "info")
|
||||
cmd = [sys.executable, "-m", "torch.distributed.run", "--nproc_per_node", str(num_gpus), str(TRAIN_SCRIPT)]
|
||||
else:
|
||||
cmd = [sys.executable, str(TRAIN_SCRIPT)]
|
||||
cmd += [
|
||||
"--stage", stage,
|
||||
"--do_train", "True",
|
||||
"--model_name_or_path", base_model,
|
||||
"--dataset", train_dataset,
|
||||
"--dataset_dir", str(DATA_DIR),
|
||||
"--template", cfg.get("template", "default"),
|
||||
"--finetuning_type", finetuning_type,
|
||||
"--output_dir", output_dir,
|
||||
"--trust_remote_code", "True",
|
||||
"--overwrite_output_dir", "True",
|
||||
"--report_to", "none",
|
||||
"--learning_rate", str(cfg.get("learning_rate", "1e-5")),
|
||||
"--num_train_epochs", str(cfg.get("n_epochs", 3)),
|
||||
"--per_device_train_batch_size", str(cfg.get("batch_size", 4)),
|
||||
"--cutoff_len", str(cfg.get("max_length", 1024)),
|
||||
"--gradient_accumulation_steps", str(cfg.get("gradient_accumulation_steps", 4)),
|
||||
"--max_samples", str(cfg.get("max_samples", 100000)),
|
||||
"--lr_scheduler_type", cfg.get("lr_scheduler_type", "cosine"),
|
||||
"--warmup_ratio", str(cfg.get("warmup_ratio", 0.03)),
|
||||
"--max_grad_norm", str(cfg.get("max_grad_norm", "1.0")),
|
||||
"--optim", cfg.get("optim", "adamw_torch"),
|
||||
"--logging_steps", str(cfg.get("logging_steps", 10)),
|
||||
"--save_steps", str(cfg.get("save_steps", 100)),
|
||||
"--save_total_limit", str(cfg.get("save_total_limit", 5)),
|
||||
"--flash_attn", cfg.get("flash_attn", "auto"),
|
||||
]
|
||||
dtype = cfg.get("dtype", "bf16")
|
||||
if dtype == "bf16":
|
||||
cmd += ["--bf16", "True"]
|
||||
elif dtype == "fp16":
|
||||
cmd += ["--fp16", "True"]
|
||||
if finetuning_type == "lora":
|
||||
cmd += [
|
||||
"--lora_rank", str(cfg.get("lora_rank", 8)),
|
||||
"--lora_alpha", str(cfg.get("lora_alpha", 16)),
|
||||
"--lora_dropout", str(cfg.get("lora_dropout", "0.05")),
|
||||
"--lora_target", cfg.get("lora_target", "all"),
|
||||
]
|
||||
if cfg.get("do_eval", False):
|
||||
cmd += ["--do_eval", "True", "--eval_strategy", "steps",
|
||||
"--eval_steps", str(cfg.get("eval_steps", 50)),
|
||||
"--per_device_eval_batch_size", "4"]
|
||||
if eval_dataset:
|
||||
cmd += ["--eval_dataset", eval_dataset]
|
||||
else:
|
||||
cmd += ["--val_size", str(cfg.get("val_size", "0.1"))]
|
||||
if cfg.get("resume_from"):
|
||||
cmd += ["--resume_from_checkpoint", str(cfg["resume_from"])]
|
||||
|
||||
_log(store, task_id, f"[INFO] 训练任务启动: {mode.upper()} 微调")
|
||||
_log(store, task_id, f"[INFO] 基座模型: {base_model}")
|
||||
_log(store, task_id, f"[INFO] 数据集: {train_dataset}")
|
||||
_log(store, task_id, f"[INFO] 输出目录: {output_dir}")
|
||||
store.update_task_runtime(task_id, status="running", progress=10, process_id=None)
|
||||
|
||||
try:
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
cwd=str(BACKEND_ROOT),
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
env=env,
|
||||
start_new_session=True,
|
||||
)
|
||||
_running_processes[task_id] = process
|
||||
store.update_task_runtime(task_id, process_id=process.pid)
|
||||
_log(store, task_id, f"[INFO] 训练进程已启动 (PID: {process.pid})", "info")
|
||||
|
||||
log_file = Path(output_dir) / "trainer_log.jsonl"
|
||||
last_pos = 0
|
||||
stop_flag = threading.Event()
|
||||
finished = threading.Event()
|
||||
|
||||
def watch_logs() -> None:
|
||||
nonlocal last_pos
|
||||
while not stop_flag.is_set():
|
||||
time.sleep(1)
|
||||
if not log_file.exists():
|
||||
continue
|
||||
try:
|
||||
with open(log_file, "r", encoding="utf-8") as lf:
|
||||
lf.seek(last_pos)
|
||||
for line in lf:
|
||||
try:
|
||||
data = json.loads(line)
|
||||
if "loss" not in data:
|
||||
continue
|
||||
loss = float(data["loss"])
|
||||
history = list(store.task(task_id).get("loss_history") or [])
|
||||
history.append(loss)
|
||||
extra: dict[str, Any] = {"current_loss": loss, "loss_history": history}
|
||||
if data.get("lr") is not None:
|
||||
extra["learning_rate"] = float(str(data["lr"]).replace("'", ""))
|
||||
if data.get("epoch") is not None:
|
||||
extra["current_epoch"] = float(data["epoch"])
|
||||
if data.get("percentage") is not None:
|
||||
extra["progress"] = float(data["percentage"])
|
||||
if data.get("remaining_time"):
|
||||
extra["eta"] = str(data["remaining_time"])
|
||||
store.update_task_runtime(task_id, extra=extra)
|
||||
step = data.get("current_steps")
|
||||
total = data.get("total_steps")
|
||||
_log(store, task_id, f"Step {step}/{total} loss={loss:.4f}")
|
||||
try:
|
||||
if int(step) >= int(total):
|
||||
finished.set()
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
last_pos = lf.tell()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
watcher = threading.Thread(target=watch_logs, daemon=True)
|
||||
watcher.start()
|
||||
|
||||
q: "queue.Queue[str]" = queue.Queue()
|
||||
|
||||
def reader() -> None:
|
||||
try:
|
||||
for line in process.stdout:
|
||||
q.put(line.rstrip("\r\n"))
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
q.put("")
|
||||
|
||||
r = threading.Thread(target=reader, daemon=True)
|
||||
r.start()
|
||||
|
||||
while True:
|
||||
try:
|
||||
raw = q.get(timeout=1)
|
||||
except queue.Empty:
|
||||
if finished.is_set() or (process.poll() is not None and q.empty()):
|
||||
try:
|
||||
raw = q.get(timeout=1)
|
||||
except queue.Empty:
|
||||
break
|
||||
else:
|
||||
continue
|
||||
if not raw:
|
||||
break
|
||||
low = raw.lower()
|
||||
log_type = "error" if "error" in low else ("warn" if "warn" in low else "info")
|
||||
_log(store, task_id, raw, log_type)
|
||||
|
||||
stop_flag.set()
|
||||
try:
|
||||
process.stdout.close()
|
||||
except Exception:
|
||||
pass
|
||||
watcher.join(timeout=3)
|
||||
try:
|
||||
process.wait(timeout=60)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
|
||||
ret = process.returncode
|
||||
if ret == 0:
|
||||
store.update_task_runtime(task_id, status="completed", progress=100)
|
||||
_log(store, task_id, "[INFO] 训练完成!", "info")
|
||||
try:
|
||||
store.ensure_trained_model_for_task(task_id, cfg, output_dir)
|
||||
except Exception as exc:
|
||||
_log(store, task_id, f"[WARN] 登记训练产物失败: {exc}", "warn")
|
||||
try:
|
||||
_plot_loss_curve(task_id, output_dir, bool(cfg.get("do_eval", False)))
|
||||
except Exception as exc:
|
||||
_log(store, task_id, f"[WARN] Loss 曲线异常: {exc}", "warn")
|
||||
else:
|
||||
store.update_task_runtime(task_id, status="failed", extra={"error_message": f"训练退出码: {ret}"})
|
||||
_log(store, task_id, f"[ERROR] 训练退出码: {ret}", "error")
|
||||
except FileNotFoundError:
|
||||
store.update_task_runtime(task_id, status="failed", extra={"error_message": "未找到 train.py 或 Python,请确认后端根目录存在 train.py 且 LLaMA-Factory 已安装"})
|
||||
_log(store, task_id, "[ERROR] 未找到 train.py / Python,无法启动真实训练", "error")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
store.update_task_runtime(task_id, status="failed", extra={"error_message": str(exc)})
|
||||
_log(store, task_id, f"[ERROR] 训练异常: {exc}", "error")
|
||||
finally:
|
||||
_running_processes.pop(task_id, None)
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# 进程信号控制
|
||||
# ──────────────────────────────────────────────
|
||||
def _signal(task_id: str, sig: int) -> bool:
|
||||
process = _running_processes.get(task_id)
|
||||
if not process:
|
||||
return False
|
||||
try:
|
||||
try:
|
||||
pgid = os.getpgid(process.pid)
|
||||
os.killpg(pgid, sig)
|
||||
except (ProcessLookupError, PermissionError):
|
||||
process.send_signal(sig)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def pause(task_id: str) -> bool:
|
||||
ok = _signal(task_id, signal.SIGSTOP)
|
||||
if ok:
|
||||
get_platform_store().update_task_runtime(task_id, status="paused")
|
||||
_log(get_platform_store(), task_id, "[INFO] 训练已暂停", "info")
|
||||
return ok
|
||||
|
||||
|
||||
def resume(task_id: str) -> bool:
|
||||
ok = _signal(task_id, signal.SIGCONT)
|
||||
if ok:
|
||||
get_platform_store().update_task_runtime(task_id, status="running")
|
||||
_log(get_platform_store(), task_id, "[INFO] 训练已继续", "info")
|
||||
return ok
|
||||
|
||||
|
||||
def cancel(task_id: str) -> bool:
|
||||
ok = _signal(task_id, signal.SIGKILL)
|
||||
if ok:
|
||||
get_platform_store().update_task_runtime(task_id, status="failed", extra={"error_message": "用户中断训练"})
|
||||
_log(get_platform_store(), task_id, "[WARN] 用户中断了训练", "warn")
|
||||
return ok
|
||||
|
||||
|
||||
def _plot_loss_curve(task_id: str, output_dir: str, has_eval: bool = False) -> None:
|
||||
try:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
state_file = Path(output_dir) / "trainer_state.json"
|
||||
if not state_file.exists():
|
||||
return
|
||||
state = json.loads(state_file.read_text(encoding="utf-8"))
|
||||
log_history = state.get("log_history", [])
|
||||
if not log_history:
|
||||
return
|
||||
train_steps, train_losses, eval_steps, eval_losses = [], [], [], []
|
||||
for entry in log_history:
|
||||
if "loss" in entry and "step" in entry:
|
||||
train_steps.append(entry["step"])
|
||||
train_losses.append(entry["loss"])
|
||||
if has_eval and "eval_loss" in entry and "step" in entry:
|
||||
eval_steps.append(entry["step"])
|
||||
eval_losses.append(entry["eval_loss"])
|
||||
if not train_losses:
|
||||
return
|
||||
fig, ax = plt.subplots(figsize=(10, 5))
|
||||
ax.plot(train_steps, train_losses, label="Training Loss", color="#409eff", linewidth=1.5)
|
||||
ax.axhline(y=min(train_losses), color="#67c23a", linestyle="--", alpha=0.5,
|
||||
label=f"Min: {min(train_losses):.4f}")
|
||||
if eval_losses:
|
||||
ax.plot(eval_steps, eval_losses, label="Validation Loss", color="#f56c6c",
|
||||
linewidth=1.5, marker="o", markersize=3)
|
||||
ax.set_xlabel("Step")
|
||||
ax.set_ylabel("Loss")
|
||||
ax.set_title("Training Loss Curve")
|
||||
ax.legend()
|
||||
ax.grid(True, alpha=0.3)
|
||||
save_path = Path(output_dir) / "loss_curve.png"
|
||||
fig.savefig(str(save_path), dpi=150, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
_log(get_platform_store(), task_id, f"Loss 曲线已保存: {save_path}")
|
||||
except Exception:
|
||||
pass
|
||||
79
backend/app/modules/fine_tune/service.py
Normal file
79
backend/app/modules/fine_tune/service.py
Normal file
@@ -0,0 +1,79 @@
|
||||
"""
|
||||
模型训练业务编排(移植自模型服务 projects/backend 的 training_service)
|
||||
|
||||
- preset 参数预设(quick / standard / high)
|
||||
- train_type → stage 映射(sft/dpo/cpt/cot)
|
||||
- 训练任务的启动 / 暂停 / 恢复 / 取消(委托 runner 真实执行)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from app.db.platform_store import get_platform_store
|
||||
from app.modules.fine_tune import runner
|
||||
|
||||
PRESETS: dict[str, dict[str, Any]] = {
|
||||
"quick": {"learning_rate": "5e-5", "n_epochs": 1, "batch_size": 4, "lora_rank": 8},
|
||||
"standard": {"learning_rate": "2e-5", "n_epochs": 3, "batch_size": 8, "lora_rank": 16},
|
||||
"high": {"learning_rate": "1e-5", "n_epochs": 5, "batch_size": 4, "lora_rank": 32},
|
||||
}
|
||||
|
||||
|
||||
def apply_presets(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""根据 preset 字段补全缺失的超参;preset=custom 时不覆盖。"""
|
||||
payload = dict(payload)
|
||||
preset = payload.get("preset", "standard")
|
||||
if preset in PRESETS and payload.get("preset") != "custom":
|
||||
for key, value in PRESETS[preset].items():
|
||||
payload.setdefault(key, value)
|
||||
return payload
|
||||
|
||||
|
||||
def build_training_config(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""把前端创建/启动载荷标准化为执行器可消费的 config。"""
|
||||
payload = apply_presets(dict(payload))
|
||||
gpus = payload.get("gpus") or [0]
|
||||
return {
|
||||
"name": payload.get("name", ""),
|
||||
"description": payload.get("description", ""),
|
||||
"train_type": payload.get("train_type", "SFT"),
|
||||
"train_method": payload.get("train_method", "lora"),
|
||||
"template": payload.get("template", "qwen"),
|
||||
"base_model": payload.get("base_model", "") or payload.get("base_model_id", ""),
|
||||
"train_dataset_id": payload.get("train_dataset_id", ""),
|
||||
"eval_dataset_id": payload.get("eval_dataset_id", ""),
|
||||
"auto_merge": bool(payload.get("auto_merge", False)),
|
||||
"output_model_name": payload.get("output_model_name", ""),
|
||||
"gpus": gpus,
|
||||
"num_gpus": payload.get("num_gpus", len(gpus)),
|
||||
"batch_size": payload.get("batch_size", 2),
|
||||
"learning_rate": payload.get("learning_rate", 0.0002),
|
||||
"n_epochs": payload.get("n_epochs", 3),
|
||||
"save_steps": payload.get("save_steps", 50),
|
||||
"lr_scheduler_type": payload.get("lr_scheduler_type", "cosine"),
|
||||
"max_length": payload.get("max_length", 2048),
|
||||
"warmup_ratio": payload.get("warmup_ratio", 0.03),
|
||||
"weight_decay": payload.get("weight_decay", 0.01),
|
||||
"lora_rank": payload.get("lora_rank", 8),
|
||||
"lora_alpha": payload.get("lora_alpha", 16),
|
||||
"lora_dropout": payload.get("lora_dropout", 0.05),
|
||||
"resume_from": payload.get("resume_from"),
|
||||
}
|
||||
|
||||
|
||||
def launch_training(task_id: str) -> None:
|
||||
"""在后台线程启动真实训练。"""
|
||||
threading.Thread(target=runner.run_training, args=(task_id,), daemon=True).start()
|
||||
|
||||
|
||||
def pause(task_id: str) -> bool:
|
||||
return runner.pause(task_id)
|
||||
|
||||
|
||||
def resume(task_id: str) -> bool:
|
||||
return runner.resume(task_id)
|
||||
|
||||
|
||||
def cancel(task_id: str) -> bool:
|
||||
return runner.cancel(task_id)
|
||||
1
backend/app/modules/inference/__init__.py
Normal file
1
backend/app/modules/inference/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Inference and compare module."""
|
||||
1
backend/app/modules/model/__init__.py
Normal file
1
backend/app/modules/model/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Model registry module."""
|
||||
1
backend/app/modules/project/__init__.py
Normal file
1
backend/app/modules/project/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from app.modules.project.router import router
|
||||
170
backend/app/modules/project/router.py
Normal file
170
backend/app/modules/project/router.py
Normal file
@@ -0,0 +1,170 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Body, Request
|
||||
from typing import Any
|
||||
|
||||
from app.api.v1.endpoints.platform import ok, fail
|
||||
from app.db.platform_store import get_platform_store
|
||||
|
||||
router = APIRouter(prefix="/projects", tags=["project"])
|
||||
|
||||
|
||||
def _actor(request: Request) -> str | None:
|
||||
auth = request.headers.get("Authorization", "")
|
||||
token = auth.replace("Bearer ", "").strip()
|
||||
return token or None
|
||||
|
||||
|
||||
def _require_no_pending_approval(resource_type: str, resource_id: str) -> None:
|
||||
"""第 4 周:写操作审批拦截——存在待审批实例时拒绝执行。"""
|
||||
store = get_platform_store()
|
||||
pending = [
|
||||
i for i in store.approval_instances(status="pending")
|
||||
if i["resource_type"] == resource_type and i["resource_id"] == resource_id
|
||||
]
|
||||
if pending:
|
||||
raise fail(409, "存在待审批的变更,请先完成审批")
|
||||
|
||||
|
||||
@router.get("")
|
||||
def list_projects(
|
||||
tenant_id: str = "default",
|
||||
status: str | None = None,
|
||||
keyword: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return ok(
|
||||
get_platform_store().projects(
|
||||
tenant_id=tenant_id, status=status, keyword=keyword
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@router.post("")
|
||||
def create_project(payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
proj = store.create_project(payload)
|
||||
store.record_audit(
|
||||
action="project.create",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project",
|
||||
target_id=proj["id"],
|
||||
tenant_id=proj.get("tenant_id"),
|
||||
detail=f"name={proj.get('name')}",
|
||||
)
|
||||
return ok(proj)
|
||||
|
||||
|
||||
@router.get("/{project_id}")
|
||||
def get_project(project_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().project(project_id))
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
|
||||
|
||||
@router.put("/{project_id}")
|
||||
def update_project(project_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
proj = store.update_project(project_id, payload)
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
store.record_audit(
|
||||
action="project.update",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project",
|
||||
target_id=project_id,
|
||||
tenant_id=proj.get("tenant_id"),
|
||||
detail=f"fields={','.join(payload.keys())}",
|
||||
)
|
||||
return ok(proj)
|
||||
|
||||
|
||||
@router.post("/{project_id}/archive")
|
||||
def archive_project(project_id: str, request: Request = None) -> dict[str, Any]:
|
||||
_require_no_pending_approval("project", project_id)
|
||||
store = get_platform_store()
|
||||
try:
|
||||
proj = store.archive_project(project_id)
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
store.record_audit(
|
||||
action="project.archive",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project",
|
||||
target_id=project_id,
|
||||
tenant_id=proj.get("tenant_id"),
|
||||
)
|
||||
return ok(proj)
|
||||
|
||||
|
||||
@router.delete("/{project_id}")
|
||||
def delete_project(project_id: str, request: Request = None) -> dict[str, Any]:
|
||||
_require_no_pending_approval("project", project_id)
|
||||
store = get_platform_store()
|
||||
store.delete_project(project_id)
|
||||
store.record_audit(
|
||||
action="project.delete",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project",
|
||||
target_id=project_id,
|
||||
)
|
||||
return ok(None)
|
||||
|
||||
|
||||
@router.get("/{project_id}/members")
|
||||
def list_members(project_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().project_members(project_id))
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
|
||||
|
||||
@router.post("/{project_id}/members")
|
||||
def add_member(project_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
member = store.add_project_member(project_id, payload)
|
||||
except KeyError:
|
||||
raise fail(404, "project not found")
|
||||
store.record_audit(
|
||||
action="project.member.add",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project.member",
|
||||
target_id=project_id,
|
||||
detail=f"user_id={payload.get('user_id')},role={payload.get('role')}",
|
||||
)
|
||||
return ok(member)
|
||||
|
||||
|
||||
@router.put("/{project_id}/members/{user_id}")
|
||||
def update_member(
|
||||
project_id: str, user_id: str, payload: dict[str, Any] = Body(...), request: Request = None
|
||||
) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
member = store.update_project_member_role(project_id, user_id, payload)
|
||||
except KeyError:
|
||||
raise fail(404, "project or member not found")
|
||||
store.record_audit(
|
||||
action="project.member.update",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project.member",
|
||||
target_id=project_id,
|
||||
detail=f"user_id={user_id},role={payload.get('role')}",
|
||||
)
|
||||
return ok(member)
|
||||
|
||||
|
||||
@router.delete("/{project_id}/members/{user_id}")
|
||||
def remove_member(project_id: str, user_id: str, request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
store.remove_project_member(project_id, user_id)
|
||||
store.record_audit(
|
||||
action="project.member.remove",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="project.member",
|
||||
target_id=project_id,
|
||||
detail=f"user_id={user_id}",
|
||||
)
|
||||
return ok(None)
|
||||
1
backend/app/modules/resource/__init__.py
Normal file
1
backend/app/modules/resource/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from app.modules.resource.router import router
|
||||
25
backend/app/modules/resource/router.py
Normal file
25
backend/app/modules/resource/router.py
Normal file
@@ -0,0 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Body
|
||||
from typing import Any
|
||||
|
||||
from app.api.v1.endpoints.platform import ok, fail
|
||||
from app.db.platform_store import get_platform_store
|
||||
|
||||
router = APIRouter(prefix="/resources", tags=["resource"])
|
||||
|
||||
|
||||
@router.get("/{resource_type}/{resource_id}/acl")
|
||||
def get_acl(resource_type: str, resource_id: str) -> dict[str, Any]:
|
||||
return ok(get_platform_store().get_acl(resource_type, resource_id))
|
||||
|
||||
|
||||
@router.put("/{resource_type}/{resource_id}/acl")
|
||||
def set_acl(
|
||||
resource_type: str, resource_id: str, payload: dict[str, Any] = Body(...)
|
||||
) -> dict[str, Any]:
|
||||
return ok(
|
||||
get_platform_store().set_acl(
|
||||
resource_type, resource_id, payload.get("entries", [])
|
||||
)
|
||||
)
|
||||
1
backend/app/modules/retention/__init__.py
Normal file
1
backend/app/modules/retention/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Retention policy and cleanup module."""
|
||||
3
backend/app/modules/system/__init__.py
Normal file
3
backend/app/modules/system/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
from app.modules.system.router import router
|
||||
|
||||
__all__ = ["router"]
|
||||
93
backend/app/modules/system/router.py
Normal file
93
backend/app/modules/system/router.py
Normal file
@@ -0,0 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.db.platform_store import ALL_PERMISSIONS, get_platform_store
|
||||
|
||||
|
||||
router = APIRouter(prefix="/system", tags=["system"])
|
||||
|
||||
|
||||
@router.get("/permissions/codes")
|
||||
def permission_codes() -> dict:
|
||||
"""返回平台权限码清单(权限码接口)。"""
|
||||
return {"code": 0, "message": "ok", "data": {"codes": ALL_PERMISSIONS}}
|
||||
|
||||
|
||||
@router.get("/permissions")
|
||||
def permissions_overview() -> dict:
|
||||
"""返回权限码清单与角色定义。"""
|
||||
store = get_platform_store()
|
||||
return {
|
||||
"code": 0,
|
||||
"message": "ok",
|
||||
"data": {"codes": ALL_PERMISSIONS, "roles": store.roles()},
|
||||
}
|
||||
|
||||
|
||||
@router.get("/audit-logs")
|
||||
def audit_logs(
|
||||
tenant_id: str | None = Query(default=None, description="租户 ID"),
|
||||
project_id: str | None = Query(default=None, description="项目 ID"),
|
||||
actor_id: str | None = Query(default=None, description="操作人 ID"),
|
||||
action: str | None = Query(default=None, description="动作类型"),
|
||||
target_type: str | None = Query(default=None, description="目标类型"),
|
||||
start_time: str | None = Query(default=None, description="ISO8601 起始时间"),
|
||||
end_time: str | None = Query(default=None, description="ISO8601 结束时间"),
|
||||
limit: int = Query(default=50, ge=1, le=200),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
) -> dict:
|
||||
"""审计日志查询:按租户/项目/操作人/动作/目标类型/时间范围分页过滤。"""
|
||||
store = get_platform_store()
|
||||
result = store.audit_logs(
|
||||
tenant_id=tenant_id,
|
||||
project_id=project_id,
|
||||
actor_id=actor_id,
|
||||
action=action,
|
||||
target_type=target_type,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
return {"code": 0, "message": "ok", "data": result}
|
||||
|
||||
|
||||
@router.get("/audit-logs/export")
|
||||
def audit_logs_export(
|
||||
tenant_id: str | None = Query(default=None, description="租户 ID"),
|
||||
project_id: str | None = Query(default=None, description="项目 ID"),
|
||||
actor_id: str | None = Query(default=None, description="操作人 ID"),
|
||||
action: str | None = Query(default=None, description="动作类型"),
|
||||
target_type: str | None = Query(default=None, description="目标类型"),
|
||||
start_time: str | None = Query(default=None, description="ISO8601 起始时间"),
|
||||
end_time: str | None = Query(default=None, description="ISO8601 结束时间"),
|
||||
) -> StreamingResponse:
|
||||
"""审计日志导出:返回 CSV 流,与应用查询相同的过滤条件。"""
|
||||
store = get_platform_store()
|
||||
result = store.audit_logs(
|
||||
tenant_id=tenant_id,
|
||||
project_id=project_id,
|
||||
actor_id=actor_id,
|
||||
action=action,
|
||||
target_type=target_type,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
limit=10000,
|
||||
offset=0,
|
||||
)
|
||||
items = result["items"]
|
||||
columns = ["time", "tenant_id", "project_id", "actor_id", "action", "target_type", "target_id", "detail", "client_ip"]
|
||||
header = ",".join(columns) + "\n"
|
||||
|
||||
def iter_rows():
|
||||
yield header
|
||||
for row in items:
|
||||
yield ",".join(f'"{str(row.get(c, "") or "")}"' for c in columns) + "\n"
|
||||
|
||||
return StreamingResponse(
|
||||
iter_rows(),
|
||||
media_type="text/csv",
|
||||
headers={"Content-Disposition": "attachment; filename=audit_logs.csv"},
|
||||
)
|
||||
1
backend/app/modules/tenant/__init__.py
Normal file
1
backend/app/modules/tenant/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from app.modules.tenant.router import router
|
||||
98
backend/app/modules/tenant/router.py
Normal file
98
backend/app/modules/tenant/router.py
Normal file
@@ -0,0 +1,98 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, Body, Request
|
||||
from typing import Any
|
||||
|
||||
from app.api.v1.endpoints.platform import ok, fail
|
||||
from app.db.platform_store import get_platform_store
|
||||
|
||||
router = APIRouter(prefix="/tenants", tags=["tenant"])
|
||||
|
||||
|
||||
def _actor(request: Request) -> str | None:
|
||||
auth = request.headers.get("Authorization", "")
|
||||
token = auth.replace("Bearer ", "").strip()
|
||||
return token or None
|
||||
|
||||
|
||||
@router.get("")
|
||||
def list_tenants() -> dict[str, Any]:
|
||||
return ok(get_platform_store().tenants())
|
||||
|
||||
|
||||
@router.post("")
|
||||
def create_tenant(payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
tenant = store.create_tenant(payload)
|
||||
except KeyError as e:
|
||||
raise fail(400, f"missing field: {e}")
|
||||
store.record_audit(
|
||||
action="tenant.create",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="tenant",
|
||||
target_id=tenant["id"],
|
||||
tenant_id=tenant["id"],
|
||||
detail=f"name={tenant.get('name')}",
|
||||
)
|
||||
return ok(tenant)
|
||||
|
||||
|
||||
@router.get("/{tenant_id}")
|
||||
def get_tenant(tenant_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
return ok(get_platform_store().tenant(tenant_id))
|
||||
except KeyError:
|
||||
raise fail(404, "tenant not found")
|
||||
|
||||
|
||||
@router.put("/{tenant_id}")
|
||||
def update_tenant(tenant_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
tenant = store.update_tenant(tenant_id, payload)
|
||||
except KeyError:
|
||||
raise fail(404, "tenant not found")
|
||||
store.record_audit(
|
||||
action="tenant.update",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="tenant",
|
||||
target_id=tenant_id,
|
||||
tenant_id=tenant_id,
|
||||
detail=f"fields={','.join(payload.keys())}",
|
||||
)
|
||||
return ok(tenant)
|
||||
|
||||
|
||||
@router.put("/{tenant_id}/quota")
|
||||
def set_quota(tenant_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
tenant = store.set_tenant_quota(tenant_id, payload.get("quota", {}))
|
||||
except KeyError:
|
||||
raise fail(404, "tenant not found")
|
||||
store.record_audit(
|
||||
action="tenant.quota.set",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="tenant",
|
||||
target_id=tenant_id,
|
||||
tenant_id=tenant_id,
|
||||
)
|
||||
return ok(tenant)
|
||||
|
||||
|
||||
@router.put("/{tenant_id}/retention-policy")
|
||||
def set_retention(tenant_id: str, payload: dict[str, Any] = Body(...), request: Request = None) -> dict[str, Any]:
|
||||
store = get_platform_store()
|
||||
try:
|
||||
tenant = store.set_tenant_retention(tenant_id, payload.get("retention_policy_id"))
|
||||
except KeyError:
|
||||
raise fail(404, "tenant not found")
|
||||
store.record_audit(
|
||||
action="tenant.retention.set",
|
||||
actor_id=_actor(request) if request else None,
|
||||
target_type="tenant",
|
||||
target_id=tenant_id,
|
||||
tenant_id=tenant_id,
|
||||
)
|
||||
return ok(tenant)
|
||||
1
backend/app/schemas/__init__.py
Normal file
1
backend/app/schemas/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Shared schemas package."""
|
||||
1
backend/app/services/__init__.py
Normal file
1
backend/app/services/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Cross-module services package."""
|
||||
1
backend/app/workers/__init__.py
Normal file
1
backend/app/workers/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Background workers package."""
|
||||
32
backend/pyproject.toml
Normal file
32
backend/pyproject.toml
Normal file
@@ -0,0 +1,32 @@
|
||||
[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",
|
||||
]
|
||||
|
||||
[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"]
|
||||
12
backend/requirements.txt
Normal file
12
backend/requirements.txt
Normal file
@@ -0,0 +1,12 @@
|
||||
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
|
||||
22
backend/tests/test_auth_service.py
Normal file
22
backend/tests/test_auth_service.py
Normal file
@@ -0,0 +1,22 @@
|
||||
"""A 模块(平台基础与企业治理)第 1 周:鉴权与权限码基础测试。
|
||||
|
||||
不需要数据库连接,可直接运行:
|
||||
cd backend && python -m pytest tests/test_auth_service.py -q
|
||||
"""
|
||||
from app.db.platform_store import ALL_PERMISSIONS
|
||||
from app.modules.auth.service import create_access_token, decode_access_token
|
||||
|
||||
|
||||
def test_token_roundtrip():
|
||||
token = create_access_token("u_abc")
|
||||
assert decode_access_token(token) == "u_abc"
|
||||
|
||||
|
||||
def test_decode_invalid_token_returns_none():
|
||||
assert decode_access_token("not.a.valid.token") is None
|
||||
|
||||
|
||||
def test_permission_codes_present():
|
||||
for code in ("dashboard", "fine-tune", "model-manage", "user-settings", "logs"):
|
||||
assert code in ALL_PERMISSIONS
|
||||
assert len(ALL_PERMISSIONS) >= 12
|
||||
27
backend/train.py
Normal file
27
backend/train.py
Normal file
@@ -0,0 +1,27 @@
|
||||
# Copyright 2025 the LlamaFactory team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from llamafactory.train.tuner import run_exp
|
||||
|
||||
def main():
|
||||
run_exp()
|
||||
|
||||
|
||||
def _mp_fn(index):
|
||||
# For xla_spawn (TPUs)
|
||||
run_exp()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user