Files
observer/server/biz/dao/model_version.go
T
2026-09-11 09:50:06 +08:00

136 lines
4.4 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 BIGSERIAL PRIMARY KEY,
dataset_id BIGINT NOT NULL,
variant VARCHAR(5) NOT NULL DEFAULT 's',
version VARCHAR(50) NOT NULL,
training_id BIGINT,
kind VARCHAR(20) NOT NULL DEFAULT 'species',
dataset_ids JSONB,
metrics JSONB,
labels JSONB NOT NULL,
sha256 VARCHAR(64) NOT NULL,
size_bytes BIGINT NOT NULL,
is_latest SMALLINT NOT NULL DEFAULT 0,
notes TEXT,
created_at TIMESTAMP 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,
"kind": m.Kind,
"dataset_ids": common.NilIfEmpty(m.DatasetIds),
"variant": m.Variant,
"version": m.Version,
"training_id": m.TrainingId,
"metrics": common.NilIfEmpty(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(发布前调用,s/n 两档互不影响)
func (d *modelVersionDao) ClearLatest(ctx context.Context, datasetId int64, variant string) error {
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).
Where("dataset_id", datasetId).Where("variant", variant).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
}