268 lines
9.1 KiB
Go
268 lines
9.1 KiB
Go
package dao
|
||
|
||
import (
|
||
"context"
|
||
"strconv"
|
||
|
||
"github.com/gogf/gf/v2/database/gdb"
|
||
"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"
|
||
)
|
||
|
||
// Dataset 数据集表 DAO:表小、读写低频,不做查询缓存;
|
||
// 图片计数由 service 在增删图片时维护,查询走普通读。
|
||
type datasetDao struct{}
|
||
|
||
var Dataset = &datasetDao{}
|
||
|
||
func init() {
|
||
ctx := context.Background()
|
||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS dataset (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
name TEXT NOT NULL UNIQUE,
|
||
source TEXT NOT NULL DEFAULT 'manual',
|
||
image_count INTEGER NOT NULL DEFAULT 0,
|
||
labeled_count INTEGER NOT NULL DEFAULT 0,
|
||
status TEXT NOT NULL DEFAULT 'building',
|
||
cover TEXT,
|
||
description TEXT,
|
||
name_prefix TEXT NOT NULL DEFAULT '',
|
||
sort_order INTEGER NOT NULL DEFAULT 0,
|
||
gen_species TEXT,
|
||
gen_tone TEXT,
|
||
gen_heights TEXT,
|
||
gen_scenes TEXT,
|
||
gen_actions TEXT,
|
||
gen_occlusions TEXT,
|
||
gen_classes TEXT,
|
||
created_at TEXT NOT NULL,
|
||
updated_at TEXT NOT NULL
|
||
)`)
|
||
if err != nil {
|
||
panic(err)
|
||
}
|
||
// 存量库迁移:生成图文件名前缀列(EnsureColumn 的 ddl 须自带列名)
|
||
common.EnsureColumn(ctx, consts.TableDataset, "name_prefix", "name_prefix TEXT NOT NULL DEFAULT ''")
|
||
// 存量库迁移:序号列(列表排序主键,升序;同号按创建时间倒序)
|
||
common.EnsureColumn(ctx, consts.TableDataset, "sort_order", "sort_order INTEGER NOT NULL DEFAULT 0")
|
||
// 存量库迁移:生成参数池列(VLM 自动生成,2026-08-28)
|
||
for _, c := range []struct{ name, ddl string }{
|
||
{"gen_species", "gen_species TEXT"},
|
||
{"gen_tone", "gen_tone TEXT"},
|
||
{"gen_heights", "gen_heights TEXT"},
|
||
{"gen_scenes", "gen_scenes TEXT"},
|
||
{"gen_actions", "gen_actions TEXT"},
|
||
{"gen_occlusions", "gen_occlusions TEXT"},
|
||
{"gen_classes", "gen_classes TEXT"},
|
||
} {
|
||
common.EnsureColumn(ctx, consts.TableDataset, c.name, c.ddl)
|
||
}
|
||
}
|
||
|
||
// Insert 新建数据集(name UNIQUE 由库兜底)
|
||
func (d *datasetDao) Insert(ctx context.Context, m *entity.Dataset) (int64, error) {
|
||
res, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Data(g.Map{
|
||
"name": m.Name,
|
||
"source": m.Source,
|
||
"name_prefix": m.NamePrefix,
|
||
"sort_order": m.SortOrder,
|
||
"image_count": m.ImageCount,
|
||
"labeled_count": m.LabeledCount,
|
||
"status": m.Status,
|
||
"created_at": m.CreatedAt,
|
||
"updated_at": m.UpdatedAt,
|
||
}).Insert()
|
||
if err != nil {
|
||
return 0, err
|
||
}
|
||
return res.LastInsertId()
|
||
}
|
||
|
||
// GetById 按主键查询,不存在返回 nil
|
||
func (d *datasetDao) GetById(ctx context.Context, id int64) (*entity.Dataset, error) {
|
||
var e entity.Dataset
|
||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Scan(&e)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return nil, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
return &e, nil
|
||
}
|
||
|
||
// GetByName 按名称查询(目录名即数据集名),不存在返回 nil
|
||
func (d *datasetDao) GetByName(ctx context.Context, name string) (*entity.Dataset, error) {
|
||
var e entity.Dataset
|
||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("name", name).Scan(&e)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return nil, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
return &e, nil
|
||
}
|
||
|
||
// GetByNamePrefix 按训练文件名前缀查询(模型发布命名,前缀唯一),不存在返回 nil
|
||
func (d *datasetDao) GetByNamePrefix(ctx context.Context, prefix string) (*entity.Dataset, error) {
|
||
var e entity.Dataset
|
||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("name_prefix", prefix).Scan(&e)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return nil, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
return &e, nil
|
||
}
|
||
|
||
// UpdateCounters 更新图片数/已标注数/状态(imgDelta 增量;labeledCount ≥0 时覆盖写)
|
||
func (d *datasetDao) UpdateCounters(ctx context.Context, id int64, imgDelta, labeledCount int64, status string) error {
|
||
data := g.Map{"updated_at": gtime.Now()}
|
||
if imgDelta != 0 {
|
||
data["image_count"] = gdb.Raw("image_count + " + strconv.FormatInt(imgDelta, 10))
|
||
}
|
||
if labeledCount >= 0 {
|
||
data["labeled_count"] = labeledCount
|
||
}
|
||
if status != "" {
|
||
data["status"] = status
|
||
}
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(data).Update()
|
||
return err
|
||
}
|
||
|
||
// UpdateConfigs 更新数据集展示配置(名称唯一性与目录迁移由 service 校验后传入;
|
||
// 仅更新非零字段,留空字段保留原值——前端提交整包配置,未填项不覆盖。
|
||
// AI 端点/训练机 SSH 为全局训练配置,不在此表维护)
|
||
func (d *datasetDao) UpdateConfigs(ctx context.Context, id int64, m *entity.Dataset) error {
|
||
data := g.Map{"updated_at": gtime.Now()}
|
||
if m.Name != "" {
|
||
data["name"] = m.Name
|
||
}
|
||
if m.NamePrefix != "" {
|
||
data["name_prefix"] = m.NamePrefix
|
||
}
|
||
if m.Cover != "" {
|
||
data["cover"] = m.Cover
|
||
}
|
||
if m.Description != "" {
|
||
data["description"] = m.Description
|
||
}
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(data).Update()
|
||
return err
|
||
}
|
||
|
||
// UpdatePools 更新生成参数池(gen_* 7 列,非空才覆盖;仅创建数据集时 VLM 生成调用。
|
||
// gen_species/gen_tone/gen_scenes/gen_actions/gen_occlusions/gen_classes 字符串空值不覆盖;
|
||
// gen_heights 数值 0 也不覆盖,避免生成时留空清库)
|
||
func (d *datasetDao) UpdatePools(ctx context.Context, id int64, m *entity.Dataset) error {
|
||
data := g.Map{"updated_at": gtime.Now()}
|
||
for col, v := range map[string]string{
|
||
"gen_species": m.GenSpecies, "gen_tone": m.GenTone,
|
||
"gen_scenes": m.GenScenes, "gen_actions": m.GenActions, "gen_occlusions": m.GenOcclusions,
|
||
"gen_classes": m.GenClasses,
|
||
} {
|
||
if v != "" {
|
||
data[col] = v
|
||
}
|
||
}
|
||
if m.GenHeights > 0 {
|
||
data["gen_heights"] = m.GenHeights
|
||
}
|
||
if len(data) == 1 {
|
||
return nil
|
||
}
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(data).Update()
|
||
return err
|
||
}
|
||
|
||
// UpdateDescription 更新描述(描述是自由文本,允许清空:空值也写入;
|
||
// UpdateConfigs 空值不覆盖,清空描述无法复用,2026-08-28)
|
||
func (d *datasetDao) UpdateDescription(ctx context.Context, id int64, description string) error {
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).
|
||
Data(g.Map{"description": description, "updated_at": gtime.Now()}).Update()
|
||
return err
|
||
}
|
||
|
||
// UpdateSortOrder 更新序号(列表排序主键,升序;同号按创建时间倒序)
|
||
func (d *datasetDao) UpdateSortOrder(ctx context.Context, id int64, sortOrder int64) error {
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).
|
||
Data(g.Map{"sort_order": sortOrder, "updated_at": gtime.Now()}).Update()
|
||
return err
|
||
}
|
||
|
||
// ClearCover 清空封面字段(删除封面用;UpdateConfigs 空值不覆盖,无法复用)
|
||
func (d *datasetDao) ClearCover(ctx context.Context, id int64) error {
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(g.Map{"cover": "", "updated_at": gtime.Now()}).Update()
|
||
return err
|
||
}
|
||
|
||
// UpdateStatus 更新状态(building|labeled|synced)
|
||
func (d *datasetDao) UpdateStatus(ctx context.Context, id int64, status string) error {
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).
|
||
Data(g.Map{"status": status, "updated_at": gtime.Now()}).Update()
|
||
return err
|
||
}
|
||
|
||
// DeleteById 删除数据集记录
|
||
func (d *datasetDao) DeleteById(ctx context.Context, id int64) error {
|
||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Delete()
|
||
return err
|
||
}
|
||
|
||
// PageByKeyword 数据集分页:名称模糊匹配(Like 命中目录名,数据库内仅元数据)
|
||
// 排序同 Page:序号升序,同号按创建时间倒序(id 降序)
|
||
func (d *datasetDao) PageByKeyword(ctx context.Context, keyword string, page, size int) ([]*entity.Dataset, int64, error) {
|
||
base := g.DB().Model(consts.TableDataset).Ctx(ctx).WhereLike("name", "%"+keyword+"%")
|
||
total, err := base.Count()
|
||
if err != nil {
|
||
return nil, 0, err
|
||
}
|
||
var list []*entity.Dataset
|
||
err = base.Order("sort_order ASC, id DESC").Limit((page-1)*size, size).Scan(&list)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return []*entity.Dataset{}, int64(total), nil
|
||
}
|
||
return nil, 0, err
|
||
}
|
||
return list, int64(total), nil
|
||
}
|
||
|
||
// ListAll 全量数据集(模型/训练列表组装数据集名用,数据集表小)
|
||
func (d *datasetDao) ListAll(ctx context.Context) ([]*entity.Dataset, error) {
|
||
var list []*entity.Dataset
|
||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).OrderAsc("id").Scan(&list)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return []*entity.Dataset{}, nil
|
||
}
|
||
return nil, err
|
||
}
|
||
return list, nil
|
||
}
|
||
|
||
// Page 数据集分页:序号升序,同号按创建时间倒序(id 降序)
|
||
func (d *datasetDao) Page(ctx context.Context, page, size int) ([]*entity.Dataset, int64, error) {
|
||
base := g.DB().Model(consts.TableDataset).Ctx(ctx)
|
||
total, err := base.Count()
|
||
if err != nil {
|
||
return nil, 0, err
|
||
}
|
||
var list []*entity.Dataset
|
||
err = base.Order("sort_order ASC, id DESC").Limit((page-1)*size, size).Scan(&list)
|
||
if err != nil {
|
||
if common.IsNoRows(err) {
|
||
return []*entity.Dataset{}, int64(total), nil
|
||
}
|
||
return nil, 0, err
|
||
}
|
||
return list, int64(total), nil
|
||
}
|