package dao import ( "context" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/os/gtime" "observer-server/biz/consts" "observer-server/biz/model/entity" "observer-server/common" ) // ModelTraining 训练任务表 DAO:进度/日志尾部为高频更新(独立小事务,不走 Serial 串行, // 单行 UPDATE 天然原子,无并发写冲突);状态流转(发起/结束)走 service 单写者。 type modelTrainingDao struct{} var Training = &modelTrainingDao{} func init() { ctx := context.Background() _, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS model_training ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, status TEXT NOT NULL DEFAULT 'running', dataset_id INTEGER NOT NULL, imgsz INTEGER NOT NULL DEFAULT 704, epochs INTEGER NOT NULL DEFAULT 150, batch INTEGER NOT NULL DEFAULT 16, device TEXT NOT NULL DEFAULT '0', current_epoch INTEGER NOT NULL DEFAULT 0, total_epochs INTEGER NOT NULL DEFAULT 0, metrics TEXT, log_tail TEXT, pid INTEGER, error TEXT, started_at TEXT NOT NULL, finished_at TEXT, created_at TEXT NOT NULL )`) if err != nil { panic(err) } } // Insert 创建训练任务,返回自增 id func (d *modelTrainingDao) Insert(ctx context.Context, m *entity.ModelTraining) (int64, error) { res, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Data(g.Map{ "name": m.Name, "status": m.Status, "dataset_id": m.DatasetId, "imgsz": m.Imgsz, "epochs": m.Epochs, "batch": m.Batch, "device": m.Device, "current_epoch": m.CurrentEpoch, "total_epochs": m.TotalEpochs, "metrics": m.Metrics, "log_tail": m.LogTail, "pid": m.Pid, "error": m.Error, "started_at": m.StartedAt, "finished_at": m.FinishedAt, "created_at": m.CreatedAt, }).Insert() if err != nil { return 0, err } return res.LastInsertId() } // GetById 按主键查询,不存在返回 nil func (d *modelTrainingDao) GetById(ctx context.Context, id int64) (*entity.ModelTraining, error) { var e entity.ModelTraining err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Scan(&e) if err != nil { if common.IsNoRows(err) { return nil, nil } return nil, err } return &e, nil } // UpdateProgress 更新进度/指标/日志尾部(训练轮询高频调用) func (d *modelTrainingDao) UpdateProgress(ctx context.Context, id int64, currentEpoch, totalEpochs int, metrics, logTail string) error { data := g.Map{} if currentEpoch > 0 { data["current_epoch"] = currentEpoch } if totalEpochs > 0 { data["total_epochs"] = totalEpochs } if metrics != "" { data["metrics"] = metrics } if logTail != "" { data["log_tail"] = logTail } if len(data) == 0 { return nil } _, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Data(data).Update() return err } // UpdatePid 记录训练进程 pid(启动后写入,恢复扫描用) func (d *modelTrainingDao) UpdatePid(ctx context.Context, id int64, pid int) error { _, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id). Data(g.Map{"pid": pid}).Update() return err } // PageByStatus 训练任务分页:按状态筛选(创建时间倒序) func (d *modelTrainingDao) PageByStatus(ctx context.Context, status string, page, size int) ([]*entity.ModelTraining, int64, error) { base := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("status", status) total, err := base.Count() if err != nil { return nil, 0, err } var list []*entity.ModelTraining err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list) if err != nil { if common.IsNoRows(err) { return []*entity.ModelTraining{}, int64(total), nil } return nil, 0, err } return list, int64(total), nil } // Finish 结束任务(success/failed):状态 + 结束时间 + 指标 + 日志尾部 + 失败原因 func (d *modelTrainingDao) Finish(ctx context.Context, id int64, status, metrics, logTail, errMsg string) error { data := g.Map{"status": status, "finished_at": gtime.Now(), "log_tail": logTail} if metrics != "" { data["metrics"] = metrics } if errMsg != "" { data["error"] = errMsg } _, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Data(data).Update() return err } // Running 当前 running 任务(并发度 1 检查用;异常终态的失败任务同表记录,不算 running) func (d *modelTrainingDao) Running(ctx context.Context) (*entity.ModelTraining, error) { var e entity.ModelTraining err := g.DB().Model(consts.TableTraining).Ctx(ctx). Where("status", consts.TrainingStatusRunning).OrderAsc("id").Limit(1).Scan(&e) if err != nil { if common.IsNoRows(err) { return nil, nil } return nil, err } return &e, nil } // RunningByDataset 某数据集 running 任务(删数据集前检查用) func (d *modelTrainingDao) RunningByDataset(ctx context.Context, datasetId int64) (*entity.ModelTraining, error) { var e entity.ModelTraining err := g.DB().Model(consts.TableTraining).Ctx(ctx). Where("dataset_id", datasetId).Where("status", consts.TrainingStatusRunning).Limit(1).Scan(&e) if err != nil { if common.IsNoRows(err) { return nil, nil } return nil, err } return &e, nil } // ListRunning 全部 running 任务(Go 重启后恢复扫描用) func (d *modelTrainingDao) ListRunning(ctx context.Context) ([]*entity.ModelTraining, error) { var list []*entity.ModelTraining err := g.DB().Model(consts.TableTraining).Ctx(ctx). Where("status", consts.TrainingStatusRunning).OrderAsc("id").Scan(&list) if err != nil { if common.IsNoRows(err) { return []*entity.ModelTraining{}, nil } return nil, err } return list, nil } // FailUnstarted 重启恢复:running 且 pid 未落(发起准备阶段服务中断,进程已随服务消亡)的任务置 failed func (d *modelTrainingDao) FailUnstarted(ctx context.Context) error { _, err := g.DB().Model(consts.TableTraining).Ctx(ctx). Where("status", consts.TrainingStatusRunning).Where("pid", 0). Data(g.Map{ "status": consts.TrainingStatusFailed, "error": "服务器重启,训练发起未完成", "finished_at": gtime.Now(), }).Update() return err } // LatestByDatasets 批量取各数据集最新一条训练记录(列表卡片训练状态用; // IN 一次取回按 id 倒序,应用层按 dataset_id 去重;数据集表小、记录少,单次查询足够) func (d *modelTrainingDao) LatestByDatasets(ctx context.Context, datasetIds []int64) (map[int64]*entity.ModelTraining, error) { out := make(map[int64]*entity.ModelTraining) if len(datasetIds) == 0 { return out, nil } for start := 0; start < len(datasetIds); start += 100 { end := start + 100 if end > len(datasetIds) { end = len(datasetIds) } var list []*entity.ModelTraining err := g.DB().Model(consts.TableTraining).Ctx(ctx). WhereIn("dataset_id", datasetIds[start:end]).OrderDesc("id").Scan(&list) if err != nil { if common.IsNoRows(err) { continue } return nil, err } for _, t := range list { if _, ok := out[t.DatasetId]; !ok { out[t.DatasetId] = t } } } return out, nil } // Page 训练任务分页:创建时间倒序 func (d *modelTrainingDao) Page(ctx context.Context, page, size int) ([]*entity.ModelTraining, int64, error) { base := g.DB().Model(consts.TableTraining).Ctx(ctx) total, err := base.Count() if err != nil { return nil, 0, err } var list []*entity.ModelTraining err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list) if err != nil { if common.IsNoRows(err) { return []*entity.ModelTraining{}, int64(total), nil } return nil, 0, err } return list, int64(total), nil }