Files
rag-local/kb/service/dataset_service.go
T
2026-08-20 12:05:53 +08:00

119 lines
3.9 KiB
Go

package service
import (
"context"
"rag-local/kb/consts"
"rag-local/kb/dao"
"rag-local/kb/model/dto"
"rag-local/kb/model/entity"
"github.com/gogf/gf/v2/errors/gerror"
"github.com/gogf/gf/v2/frame/g"
)
var DatasetService = &datasetService{}
type datasetService struct{}
func (s *datasetService) List(ctx context.Context) ([]*entity.Dataset, error) {
return dao.Dataset.List(ctx)
}
func (s *datasetService) Save(ctx context.Context, req *dto.SaveDatasetReq) (int64, error) {
// 在 service 内部创建 entity,组装 dto 到 entity 的映射
m := &entity.Dataset{
Id: req.Id,
Name: req.Name,
Description: req.Description,
EmbeddingCfgId: req.EmbeddingCfgId,
ChunkSize: req.ChunkSize,
ChunkOverlap: req.ChunkOverlap,
ReactRounds: req.ReactRounds,
VecTopK: req.VecTopK,
FtsTopK: req.FtsTopK,
RerankTopK: req.RerankTopK,
RecallTopK: req.RecallTopK,
}
if m.Status == 0 {
m.Status = 1
}
if m.EmbeddingCfgId == 0 {
return 0, gerror.New("数据集必须绑定向量模型,请先选择向量模型")
}
if m.ReactRounds < 0 || m.ReactRounds > consts.MaxReactRounds {
return 0, gerror.Newf("智能体轮次需在 0~%d 之间(0=关闭)", consts.MaxReactRounds)
}
for _, v := range []int{m.VecTopK, m.FtsTopK, m.RerankTopK, m.RecallTopK} {
if v < -1 || v > consts.MaxRecallTopK {
return 0, gerror.Newf("召回数量需在 -1~%d 之间(-1=尽量多,0=全局默认)", consts.MaxRecallTopK)
}
}
// 0/-1 = 使用 config.yml 全局默认,落库前回填为实际值(Insert/Update 一致)
if m.ChunkSize <= 0 {
m.ChunkSize = g.Cfg().MustGet(ctx, "chunk.default_size", consts.DefaultChunkSize).Int()
}
if m.ChunkOverlap < 0 {
m.ChunkOverlap = g.Cfg().MustGet(ctx, "chunk.default_overlap", consts.DefaultChunkOverlap).Int()
}
if m.ReactRounds < 0 {
m.ReactRounds = g.Cfg().MustGet(ctx, "react.default_rounds", 0).Int()
}
if m.Id > 0 {
old, err := dao.Dataset.GetOne(ctx, m.Id)
if err != nil {
return 0, err
}
if err := dao.Dataset.Update(ctx, m); err != nil {
return 0, err
}
// 向量模型变更 → 入队重新向量化任务,保证 chunk 向量与提问检索模型一致
if old != nil && old.EmbeddingCfgId != m.EmbeddingCfgId {
if err := s.enqueueTask(ctx, m.Id, consts.TaskTypeReembed); err != nil {
g.Log().Warningf(ctx, "enqueue reembed for dataset %d failed: %v", m.Id, err)
}
}
// 分块大小/重叠变更 → 入队完整重新解析任务(重新分词+向量化)
if old != nil && (old.ChunkSize != m.ChunkSize || old.ChunkOverlap != m.ChunkOverlap) {
if err := s.enqueueTask(ctx, m.Id, consts.TaskTypeParse); err != nil {
g.Log().Warningf(ctx, "enqueue reparse for dataset %d failed: %v", m.Id, err)
}
}
return m.Id, nil
}
return dao.Dataset.Insert(ctx, m)
}
func (s *datasetService) enqueueTask(ctx context.Context, datasetId int64, taskType string) error {
docs, _, err := dao.Document.List(ctx, datasetId, 1, 100000)
if err != nil {
return err
}
for _, doc := range docs {
if _, err := dao.ParseTask.Insert(ctx, doc.Id, datasetId, taskType); err != nil {
return err
}
// 重新向量化入队即刻置为向量生成中,与单文档入队(ParseTaskService.Enqueue)行为一致
if taskType == consts.TaskTypeReembed {
if err := dao.Document.UpdateFields(ctx, doc.Id, g.Map{"status": consts.DocumentStatusEmbedding}); err != nil {
g.Log().Warningf(ctx, "mark doc %d embedding failed: %v", doc.Id, err)
}
}
}
return nil
}
func (s *datasetService) Delete(ctx context.Context, id int64) error {
// 有文档的数据集不允许删除
count, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameDocument).Ctx(ctx).
Where("dataset_id", id).Count()
if err != nil {
return err
}
if count > 0 {
return gerror.New("数据集下存在文档,无法删除")
}
return dao.Dataset.Delete(ctx, id)
}