第一次提交
This commit is contained in:
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)
|
||||
Reference in New Issue
Block a user