102 lines
2.8 KiB
Go
102 lines
2.8 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
|
|
"rag-local/kb/consts"
|
|
"rag-local/kb/dao"
|
|
"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, m *entity.Dataset) (int64, error) {
|
|
if m.Status == 0 {
|
|
m.Status = 1
|
|
}
|
|
if err := s.validateStrategy(m); err != nil {
|
|
return 0, err
|
|
}
|
|
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 || old.ChunkStrategy != m.ChunkStrategy) {
|
|
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
|
|
}
|
|
if m.ChunkSize <= 0 {
|
|
m.ChunkSize = consts.DefaultChunkSize
|
|
}
|
|
if m.ChunkOverlap < 0 {
|
|
m.ChunkOverlap = consts.DefaultChunkOverlap
|
|
}
|
|
if m.ChunkStrategy == "" {
|
|
m.ChunkStrategy = consts.ChunkStrategyTitle
|
|
}
|
|
return dao.Dataset.Insert(ctx, m)
|
|
}
|
|
|
|
func (s *datasetService) validateStrategy(m *entity.Dataset) error {
|
|
if m.ChunkStrategy == "" {
|
|
m.ChunkStrategy = consts.ChunkStrategyTitle
|
|
}
|
|
if m.EmbeddingCfgId == 0 {
|
|
return gerror.New("数据集必须绑定向量模型,请先选择向量模型")
|
|
}
|
|
// 语义分块无重叠参数,清零避免误导
|
|
if m.ChunkStrategy == consts.ChunkStrategySemantic {
|
|
m.ChunkOverlap = 0
|
|
}
|
|
return nil
|
|
}
|
|
|
|
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
|
|
}
|
|
}
|
|
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)
|
|
}
|