package dao import ( "context" "github.com/gogf/gf/v2/frame/g" "observer-server/biz/consts" "observer-server/biz/model/entity" "observer-server/common" ) // ModelVersion 模型版本表 DAO:每数据集独立版本序列;表小不做查询缓存。 type modelVersionDao struct{} var ModelVersion = &modelVersionDao{} func init() { ctx := context.Background() _, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS model_version ( id INTEGER PRIMARY KEY AUTOINCREMENT, dataset_id INTEGER NOT NULL, version TEXT NOT NULL, training_id INTEGER, metrics TEXT, labels TEXT NOT NULL, sha256 TEXT NOT NULL, size_bytes INTEGER NOT NULL, is_latest INTEGER NOT NULL DEFAULT 0, notes TEXT, created_at TEXT NOT NULL, UNIQUE (dataset_id, version) )`) if err != nil { panic(err) } } // Insert 插入模型版本(is_latest 由 service 先置 0 再插新行置 1) func (d *modelVersionDao) Insert(ctx context.Context, m *entity.ModelVersion) (int64, error) { res, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Data(g.Map{ "dataset_id": m.DatasetId, "version": m.Version, "training_id": m.TrainingId, "metrics": m.Metrics, "labels": m.Labels, "sha256": m.Sha256, "size_bytes": m.SizeBytes, "is_latest": m.IsLatest, "notes": m.Notes, "created_at": m.CreatedAt, }).Insert() if err != nil { return 0, err } return res.LastInsertId() } // ClearLatest 某数据集所有版本置 is_latest=0(发布前调用) func (d *modelVersionDao) ClearLatest(ctx context.Context, datasetId int64) error { _, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx). Where("dataset_id", datasetId).Data(g.Map{"is_latest": 0}).Update() return err } // DeleteByDataset 删除某数据集全部版本记录(删数据集级联) func (d *modelVersionDao) DeleteByDataset(ctx context.Context, datasetId int64) error { _, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("dataset_id", datasetId).Delete() return err } // ListAllLatest 全部数据集的当前生效版本(模型目录/热更新目录) func (d *modelVersionDao) ListAllLatest(ctx context.Context) ([]*entity.ModelVersion, error) { var list []*entity.ModelVersion err := g.DB().Model(consts.TableModelVersion).Ctx(ctx). Where("is_latest", 1).OrderAsc("dataset_id").Scan(&list) if err != nil { if common.IsNoRows(err) { return []*entity.ModelVersion{}, nil } return nil, err } return list, nil } // PublishedByTrainingIds 批量判断训练记录是否已发布过版本(IN 按 ≤100 分批,SQLite 变量数上限内) func (d *modelVersionDao) PublishedByTrainingIds(ctx context.Context, trainingIds []int64) (map[int64]bool, error) { published := make(map[int64]bool) for start := 0; start < len(trainingIds); start += 100 { end := start + 100 if end > len(trainingIds) { end = len(trainingIds) } var rows []*entity.ModelVersion err := g.DB().Model(consts.TableModelVersion).Ctx(ctx). WhereIn("training_id", trainingIds[start:end]).Fields("training_id").Scan(&rows) if err != nil && !common.IsNoRows(err) { return nil, err } for _, r := range rows { published[r.TrainingId] = true } } return published, nil } // PageByDataset 某数据集版本分页:发布时间倒序 func (d *modelVersionDao) PageByDataset(ctx context.Context, datasetId int64, page, size int) ([]*entity.ModelVersion, int64, error) { base := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("dataset_id", datasetId) total, err := base.Count() if err != nil { return nil, 0, err } var list []*entity.ModelVersion err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list) if err != nil { if common.IsNoRows(err) { return []*entity.ModelVersion{}, int64(total), nil } return nil, 0, err } return list, int64(total), nil } // DeleteById 删除版本记录 func (d *modelVersionDao) DeleteById(ctx context.Context, id int64) error { _, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("id", id).Delete() return err }