Files
rag-local/kb/dao/chunk_dao.go
T
2026-08-11 11:19:04 +08:00

331 lines
11 KiB
Go

package dao
import (
"context"
"database/sql"
"errors"
"strconv"
"strings"
_ "modernc.org/sqlite/vec"
"rag-local/common"
"rag-local/kb/consts"
"rag-local/kb/model/domain"
"rag-local/kb/model/entity"
"github.com/gogf/gf/v2/errors/gerror"
"github.com/gogf/gf/v2/frame/g"
"github.com/gogf/gf/v2/os/gtime"
"github.com/gogf/gf/v2/text/gstr"
)
var Chunk = &chunkDao{}
type chunkDao struct{}
func init() {
ctx := context.Background()
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameChunk+` (
id INTEGER PRIMARY KEY AUTOINCREMENT,
dataset_id INTEGER NOT NULL DEFAULT 0,
document_id INTEGER NOT NULL DEFAULT 0,
seq INTEGER NOT NULL DEFAULT 0,
content TEXT NOT NULL DEFAULT '',
meta TEXT NOT NULL DEFAULT '',
created_at DATETIME DEFAULT (datetime('now','localtime'))
)`)
if err != nil {
g.Log().Warningf(ctx, "create kb_chunk table failed: %v", err)
}
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_kb_chunk_document ON "+consts.TableNameChunk+"(document_id)"); err != nil {
g.Log().Warningf(ctx, "create index idx_kb_chunk_document failed: %v", err)
}
// 向量虚拟表(维度取配置 vector.dim,切换维度需删表重建)
dim := g.Cfg().MustGet(ctx, "vector.dim", consts.DefaultEmbeddingDim).Int()
if dim < 1 {
dim = consts.DefaultEmbeddingDim
}
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE VIRTUAL TABLE IF NOT EXISTS `+consts.TableNameChunkVec+
` USING vec0(chunk_id INTEGER PRIMARY KEY, embedding float[`+strconv.Itoa(dim)+`])`); err != nil {
g.Log().Warningf(ctx, "create vec0 table failed: %v", err)
}
// 全文索引虚拟表(默认 unicode61 tokenizer;中文分词在应用层完成,content_tokens 存分词后空格连接文本)
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE VIRTUAL TABLE IF NOT EXISTS `+consts.TableNameChunkFts+
` USING fts5(chunk_id UNINDEXED, dataset_id UNINDEXED, title, content_tokens)`); err != nil {
g.Log().Warningf(ctx, "create fts5 table failed: %v", err)
}
}
func (d *chunkDao) GetOne(ctx context.Context, id int64) (*entity.Chunk, error) {
var m entity.Chunk
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).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 *chunkDao) ListByDocument(ctx context.Context, documentId int64, page, pageSize int) ([]*entity.Chunk, int, error) {
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
total, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx).Where("document_id", documentId).Count()
if err != nil {
return nil, 0, err
}
var list []*entity.Chunk
err = g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx).
Where("document_id", documentId).Page(page, pageSize).OrderAsc("seq").Scan(&list)
if list == nil {
list = make([]*entity.Chunk, 0)
}
return list, total, err
}
// InsertWithVec 事务内写入 chunk + 向量 + 全文索引,返回 chunk id
func (d *chunkDao) InsertWithVec(ctx context.Context, datasetId, documentId int64, seq int, content, meta, title string, vecJson string, dim int) (int64, error) {
tx, err := g.DB(consts.DbGroupDefault).Begin(ctx)
if err != nil {
return 0, err
}
// Commit 成功后 IsClosed 为 true,跳过 Rollback,避免对已提交事务回滚产生报错日志
defer func() {
if !tx.IsClosed() {
_ = tx.Rollback()
}
}()
r, err := tx.Model(consts.TableNameChunk).Ctx(ctx).Data(g.Map{
"dataset_id": datasetId,
"document_id": documentId,
"seq": seq,
"content": content,
"meta": meta,
"created_at": gtime.Now().Format("Y-m-d H:i:s"),
}).Insert()
if err != nil {
return 0, err
}
chunkId, _ := r.LastInsertId()
if chunkId == 0 {
return 0, gerror.New("chunk insert failed")
}
if vecJson != "" {
if _, err := tx.Exec("INSERT INTO "+consts.TableNameChunkVec+" (chunk_id, embedding) VALUES (?, vec_f32(?))", chunkId, vecJson); err != nil {
return 0, err
}
}
if _, err := tx.Exec("INSERT INTO "+consts.TableNameChunkFts+" (chunk_id, dataset_id, title, content_tokens) VALUES (?, ?, ?, ?)",
chunkId, datasetId, title, common.Tokenize(content)); err != nil {
return 0, err
}
if err := tx.Commit(); err != nil {
return 0, err
}
return chunkId, nil
}
// DeleteByDocument 事务内删除文档全部分块(chunk + 向量 + 全文索引)
func (d *chunkDao) DeleteByDocument(ctx context.Context, documentId int64) error {
tx, err := g.DB(consts.DbGroupDefault).Begin(ctx)
if err != nil {
return err
}
// Commit 成功后 IsClosed 为 true,跳过 Rollback,避免对已提交事务回滚产生报错日志
defer func() {
if !tx.IsClosed() {
_ = tx.Rollback()
}
}()
ids := tx.Model(consts.TableNameChunk).Ctx(ctx).Fields("id").Where("document_id", documentId)
r, err := ids.Array()
if err != nil {
return err
}
chunkIds := make([]int64, 0, len(r))
for _, v := range r {
chunkIds = append(chunkIds, v.Int64())
}
if len(chunkIds) > 0 {
placeholders := make([]string, 0, len(chunkIds))
args := make([]interface{}, 0, len(chunkIds))
for _, id := range chunkIds {
placeholders = append(placeholders, "?")
args = append(args, id)
}
in := gstr.Join(placeholders, ",")
if _, err := tx.Exec("DELETE FROM "+consts.TableNameChunkVec+" WHERE chunk_id IN ("+in+")", args...); err != nil {
return err
}
if _, err := tx.Exec("DELETE FROM "+consts.TableNameChunkFts+" WHERE chunk_id IN ("+in+")", args...); err != nil {
return err
}
// 知识图谱数据按 chunk 关联,重新解析时旧 chunk 消失,一并清理避免孤儿数据
if _, err := tx.Exec("DELETE FROM "+consts.TableNameKgRelation+" WHERE chunk_id IN ("+in+")", args...); err != nil {
return err
}
if _, err := tx.Exec("DELETE FROM "+consts.TableNameKgEntity+" WHERE chunk_id IN ("+in+")", args...); err != nil {
return err
}
}
if _, err := tx.Model(consts.TableNameChunk).Ctx(ctx).Where("document_id", documentId).Delete(); err != nil {
return err
}
return tx.Commit()
}
func (d *chunkDao) UpdateContent(ctx context.Context, id int64, content, vecJson, tokens string) error {
tx, err := g.DB(consts.DbGroupDefault).Begin(ctx)
if err != nil {
return err
}
// Commit 成功后 IsClosed 为 true,跳过 Rollback,避免对已提交事务回滚产生报错日志
defer func() {
if !tx.IsClosed() {
_ = tx.Rollback()
}
}()
if _, err := tx.Model(consts.TableNameChunk).Ctx(ctx).Data(g.Map{"content": content}).Where("id", id).Update(); err != nil {
return err
}
if vecJson != "" {
if _, err := tx.Exec("UPDATE "+consts.TableNameChunkVec+" SET embedding = vec_f32(?) WHERE chunk_id = ?", vecJson, id); err != nil {
return err
}
}
if _, err := tx.Exec("UPDATE "+consts.TableNameChunkFts+" SET content_tokens = ? WHERE chunk_id = ?", tokens, id); err != nil {
return err
}
return tx.Commit()
}
// UpdateVec 仅更新向量(重新向量化,不动文本与 FTS)
func (d *chunkDao) UpdateVec(ctx context.Context, id int64, vecJson string) error {
if vecJson == "" {
return nil
}
_, err := g.DB(consts.DbGroupDefault).Exec(ctx,
"UPDATE "+consts.TableNameChunkVec+" SET embedding = vec_f32(?) WHERE chunk_id = ?", vecJson, id)
return err
}
func (d *chunkDao) CountByDocument(ctx context.Context, documentId int64) (int, error) {
return g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx).Where("document_id", documentId).Count()
}
// ListByIds 批量按主键查分块(单表约束:IN ≤100 分批;调用方保证 ids 非空)
func (d *chunkDao) ListByIds(ctx context.Context, ids []int64) ([]*entity.Chunk, error) {
var list []*entity.Chunk
for start := 0; start < len(ids); start += 100 {
end := min(start+100, len(ids))
var part []*entity.Chunk
if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx).
WhereIn("id", ids[start:end]).Scan(&part); err != nil {
return nil, err
}
list = append(list, part...)
}
if list == nil {
list = make([]*entity.Chunk, 0)
}
return list, nil
}
// ListDocumentIdsByChunkIds 批量查 chunk 归属文档(单表约束:IN ≤100 分批;调用方保证 chunkIds 非空)
func (d *chunkDao) ListDocumentIdsByChunkIds(ctx context.Context, chunkIds []int64) (map[int64]int64, error) {
out := make(map[int64]int64, len(chunkIds))
for start := 0; start < len(chunkIds); start += 100 {
end := min(start+100, len(chunkIds))
var rows []*entity.Chunk
if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameChunk).Ctx(ctx).
Fields("id, document_id").WhereIn("id", chunkIds[start:end]).Scan(&rows); err != nil {
return nil, err
}
for _, r := range rows {
out[r.Id] = r.DocumentId
}
}
return out, nil
}
// VecSearch 向量 KNN 检索:vec0 取最近 topK*4 候选,再按数据集过滤(单表约束:vec0 无 dataset 列,
// 拆两条单表查询 + 内存过滤;候选已按距离升序,过滤后取前 topK 即等价原 JOIN 语义)
func (d *chunkDao) VecSearch(ctx context.Context, datasetId int64, vecJson string, topK int) ([]domain.VecHit, error) {
r, err := g.DB(consts.DbGroupDefault).Ctx(ctx).Raw(
`SELECT chunk_id, distance FROM `+consts.TableNameChunkVec+
` WHERE embedding MATCH ? ORDER BY distance LIMIT ?`,
vecJson, topK*4,
).All()
if err != nil {
return nil, err
}
if len(r) == 0 {
return nil, nil
}
cand := make([]domain.VecHit, 0, len(r))
ids := make([]int64, 0, len(r))
for _, row := range r {
id := row["chunk_id"].Int64()
cand = append(cand, domain.VecHit{ChunkId: id, Distance: row["distance"].Float64()})
ids = append(ids, id)
}
// 候选 id ≤ topK*4(默认 80),远低于 SQLite 变量上限,无需分批
placeholders := strings.TrimSuffix(strings.Repeat("?,", len(ids)), ",")
args := make([]any, 0, len(ids)+1)
for _, id := range ids {
args = append(args, id)
}
args = append(args, datasetId)
chunks, err := g.DB(consts.DbGroupDefault).Ctx(ctx).Raw(
"SELECT id FROM "+consts.TableNameChunk+" WHERE id IN ("+placeholders+") AND dataset_id = ?", args...).All()
if err != nil {
return nil, err
}
allowed := make(map[int64]bool, len(chunks))
for _, row := range chunks {
allowed[row["id"].Int64()] = true
}
hits := make([]domain.VecHit, 0, len(cand))
for _, h := range cand {
if allowed[h.ChunkId] {
hits = append(hits, h)
if len(hits) >= topK {
break
}
}
}
return hits, nil
}
// FtsSearch 全文检索(BM25):中文分词已在应用层完成,query 为空格连接的引号词串
func (d *chunkDao) FtsSearch(ctx context.Context, datasetId int64, query string, topK int) ([]domain.FtsHit, error) {
if strings.TrimSpace(query) == "" {
return nil, nil
}
r, err := g.DB(consts.DbGroupDefault).Ctx(ctx).Raw(
`SELECT chunk_id, bm25(`+consts.TableNameChunkFts+`) AS score FROM `+consts.TableNameChunkFts+
` WHERE `+consts.TableNameChunkFts+` MATCH ? AND dataset_id = ? ORDER BY score LIMIT ?`,
query, datasetId, topK,
).All()
if err != nil {
return nil, err
}
hits := make([]domain.FtsHit, 0, len(r))
for _, row := range r {
hits = append(hits, domain.FtsHit{
ChunkId: row["chunk_id"].Int64(),
Score: row["score"].Float64(),
})
}
return hits, nil
}