189 lines
6.2 KiB
Go
189 lines
6.2 KiB
Go
package dao
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
|
|
"rag-local/kb/consts"
|
|
"rag-local/kb/model/entity"
|
|
|
|
"github.com/gogf/gf/v2/frame/g"
|
|
"github.com/gogf/gf/v2/os/gtime"
|
|
)
|
|
|
|
var Dataset = &datasetDao{}
|
|
|
|
type datasetDao struct{}
|
|
|
|
func init() {
|
|
ctx := context.Background()
|
|
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameDataset+` (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
name TEXT NOT NULL DEFAULT '',
|
|
description TEXT NOT NULL DEFAULT '',
|
|
embedding_cfg_id INTEGER NOT NULL DEFAULT 0,
|
|
chunk_size INTEGER NOT NULL DEFAULT 800,
|
|
chunk_overlap INTEGER NOT NULL DEFAULT 150,
|
|
chunk_strategy TEXT NOT NULL DEFAULT 'title',
|
|
unit_pattern TEXT NOT NULL DEFAULT '',
|
|
context_pattern TEXT NOT NULL DEFAULT '',
|
|
status INTEGER NOT NULL DEFAULT 1,
|
|
created_at DATETIME DEFAULT (datetime('now','localtime')),
|
|
updated_at DATETIME DEFAULT (datetime('now','localtime'))
|
|
)`)
|
|
if err != nil {
|
|
g.Log().Warningf(ctx, "create kb_dataset table failed: %v", err)
|
|
}
|
|
// 旧表迁移:补充分块配置列(先查列是否已存在,避免 ALTER 报错刷日志)
|
|
for _, col := range []struct {
|
|
name string
|
|
ddl string
|
|
}{
|
|
{"chunk_size", "chunk_size INTEGER NOT NULL DEFAULT 800"},
|
|
{"chunk_overlap", "chunk_overlap INTEGER NOT NULL DEFAULT 150"},
|
|
{"chunk_strategy", "chunk_strategy TEXT NOT NULL DEFAULT 'title'"},
|
|
{"unit_pattern", "unit_pattern TEXT NOT NULL DEFAULT ''"},
|
|
{"context_pattern", "context_pattern TEXT NOT NULL DEFAULT ''"},
|
|
{"react_rounds", "react_rounds INTEGER NOT NULL DEFAULT 0"},
|
|
{"vec_top_k", "vec_top_k INTEGER NOT NULL DEFAULT 0"},
|
|
{"fts_top_k", "fts_top_k INTEGER NOT NULL DEFAULT 0"},
|
|
{"rerank_top_k", "rerank_top_k INTEGER NOT NULL DEFAULT 0"},
|
|
{"recall_top_k", "recall_top_k INTEGER NOT NULL DEFAULT 0"},
|
|
} {
|
|
cnt, err := g.DB(consts.DbGroupDefault).Ctx(ctx).GetValue(ctx,
|
|
"SELECT COUNT(*) FROM pragma_table_info('"+consts.TableNameDataset+"') WHERE name=?", col.name)
|
|
if err != nil || cnt.Int64() > 0 {
|
|
continue
|
|
}
|
|
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx,
|
|
"ALTER TABLE "+consts.TableNameDataset+" ADD COLUMN "+col.ddl); err != nil {
|
|
g.Log().Warningf(ctx, "migrate kb_dataset add column %s failed: %v", col.name, err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (d *datasetDao) GetOne(ctx context.Context, id int64) (*entity.Dataset, error) {
|
|
var m entity.Dataset
|
|
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).Where("id", id).Scan(&m)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &m, nil
|
|
}
|
|
|
|
func (d *datasetDao) List(ctx context.Context) ([]*entity.Dataset, error) {
|
|
var list []*entity.Dataset
|
|
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).OrderAsc("id").Scan(&list)
|
|
if list == nil {
|
|
list = make([]*entity.Dataset, 0)
|
|
}
|
|
return list, err
|
|
}
|
|
|
|
func (d *datasetDao) Insert(ctx context.Context, data *entity.Dataset) (int64, error) {
|
|
now := gtime.Now().Format("Y-m-d H:i:s")
|
|
r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).Data(g.Map{
|
|
"name": data.Name,
|
|
"description": data.Description,
|
|
"embedding_cfg_id": data.EmbeddingCfgId,
|
|
"chunk_size": data.ChunkSize,
|
|
"chunk_overlap": data.ChunkOverlap,
|
|
"react_rounds": data.ReactRounds,
|
|
"vec_top_k": data.VecTopK,
|
|
"fts_top_k": data.FtsTopK,
|
|
"rerank_top_k": data.RerankTopK,
|
|
"recall_top_k": data.RecallTopK,
|
|
"status": data.Status,
|
|
"created_at": now,
|
|
"updated_at": now,
|
|
}).Insert()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return r.LastInsertId()
|
|
}
|
|
|
|
func (d *datasetDao) Update(ctx context.Context, data *entity.Dataset) error {
|
|
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).Data(g.Map{
|
|
"name": data.Name,
|
|
"description": data.Description,
|
|
"embedding_cfg_id": data.EmbeddingCfgId,
|
|
"chunk_size": data.ChunkSize,
|
|
"chunk_overlap": data.ChunkOverlap,
|
|
"react_rounds": data.ReactRounds,
|
|
"vec_top_k": data.VecTopK,
|
|
"fts_top_k": data.FtsTopK,
|
|
"rerank_top_k": data.RerankTopK,
|
|
"recall_top_k": data.RecallTopK,
|
|
"status": data.Status,
|
|
"updated_at": gtime.Now().Format("Y-m-d H:i:s"),
|
|
}).Where("id", data.Id).Update()
|
|
return err
|
|
}
|
|
|
|
func (d *datasetDao) Delete(ctx context.Context, id int64) error {
|
|
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).Where("id", id).Delete()
|
|
return err
|
|
}
|
|
|
|
func (d *datasetDao) UpdateFields(ctx context.Context, id int64, data g.Map) error {
|
|
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).Data(data).
|
|
Where("id", id).Update()
|
|
return err
|
|
}
|
|
|
|
func (d *datasetDao) GetEmbeddingCfgId(ctx context.Context, id int64) (int64, error) {
|
|
r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).
|
|
Fields("embedding_cfg_id").Where("id", id).One()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
if r == nil {
|
|
return 0, nil
|
|
}
|
|
return r["embedding_cfg_id"].Int64(), nil
|
|
}
|
|
|
|
// GetReactRounds 读取数据集 ReAct 轮次配置(未配置/异常时返回 0 = 关闭智能体模式)
|
|
func (d *datasetDao) GetReactRounds(ctx context.Context, id int64) (int, error) {
|
|
r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).
|
|
Fields("react_rounds").Where("id", id).One()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
if r == nil {
|
|
return 0, nil
|
|
}
|
|
return r["react_rounds"].Int(), nil
|
|
}
|
|
|
|
// RecallParams 数据集召回数量配置(0=全局默认,-1=尽量多,>0=固定值)
|
|
type RecallParams struct {
|
|
VecTopK int
|
|
FtsTopK int
|
|
RerankTopK int
|
|
RecallTopK int
|
|
}
|
|
|
|
// GetRecallParams 读取数据集召回数量配置(未配置/异常时返回全 0 = 用全局默认)
|
|
func (d *datasetDao) GetRecallParams(ctx context.Context, id int64) (*RecallParams, error) {
|
|
r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDataset).Ctx(ctx).
|
|
Fields("vec_top_k,fts_top_k,rerank_top_k,recall_top_k").Where("id", id).One()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if r == nil {
|
|
return &RecallParams{}, nil
|
|
}
|
|
return &RecallParams{
|
|
VecTopK: r["vec_top_k"].Int(),
|
|
FtsTopK: r["fts_top_k"].Int(),
|
|
RerankTopK: r["rerank_top_k"].Int(),
|
|
RecallTopK: r["recall_top_k"].Int(),
|
|
}, nil
|
|
}
|