Files
rag-local/kb/service/dataset_service.go
T
2026-08-05 13:39:37 +08:00

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)
}