Files
observer/server/biz/dao/model_version.go
T
2026-08-28 14:31:13 +08:00

130 lines
4.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}