Files
rag-local/kb/dao/kg_entity_dao.go
T
2026-08-05 10:28:44 +08:00

118 lines
3.7 KiB
Go

package dao
import (
"context"
"database/sql"
"errors"
"strings"
"rag-local/kb/consts"
"rag-local/kb/model/entity"
"github.com/gogf/gf/v2/frame/g"
"github.com/gogf/gf/v2/os/gtime"
)
var KgEntity = &kgEntityDao{}
type kgEntityDao struct{}
func init() {
ctx := context.Background()
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameKgEntity+` (
id INTEGER PRIMARY KEY AUTOINCREMENT,
dataset_id INTEGER NOT NULL DEFAULT 0,
name TEXT NOT NULL DEFAULT '',
entity_type TEXT NOT NULL DEFAULT '',
chunk_id INTEGER NOT NULL DEFAULT 0,
created_at DATETIME DEFAULT (datetime('now','localtime')),
updated_at DATETIME DEFAULT (datetime('now','localtime')),
UNIQUE(dataset_id, name)
)`)
if err != nil {
g.Log().Warningf(ctx, "create kg_entity table failed: %v", err)
}
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_kg_entity_dataset ON "+consts.TableNameKgEntity+"(dataset_id)"); err != nil {
g.Log().Warningf(ctx, "create index idx_kg_entity_dataset failed: %v", err)
}
}
func (d *kgEntityDao) GetOne(ctx context.Context, id int64) (*entity.KgEntity, error) {
var m entity.KgEntity
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgEntity).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 *kgEntityDao) List(ctx context.Context, datasetId int64, page, pageSize int) ([]*entity.KgEntity, int, error) {
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
m := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgEntity).Ctx(ctx)
if datasetId > 0 {
m = m.Where("dataset_id", datasetId)
}
total, err := m.Count()
if err != nil {
return nil, 0, err
}
var list []*entity.KgEntity
err = m.Page(page, pageSize).OrderDesc("id").Scan(&list)
return list, total, err
}
// ListNames 数据集全部实体名(实体链接在内存中打分,本地库规模可控)
func (d *kgEntityDao) ListNames(ctx context.Context, datasetId int64) ([]string, error) {
var list []*entity.KgEntity
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgEntity).Ctx(ctx).
Fields("name").Where("dataset_id", datasetId).Scan(&list)
if err != nil {
return nil, err
}
names := make([]string, 0, len(list))
for _, e := range list {
names = append(names, e.Name)
}
return names, nil
}
// Upsert 按 (dataset_id, name) 去重,已存在则更新类型与来源
func (d *kgEntityDao) Upsert(ctx context.Context, datasetId, chunkId int64, name, entityType string) error {
now := gtime.Now().Format("Y-m-d H:i:s")
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `INSERT INTO `+consts.TableNameKgEntity+`
(dataset_id, name, entity_type, chunk_id, created_at, updated_at)
VALUES (?,?,?,?,?,?)
ON CONFLICT(dataset_id, name) DO UPDATE SET
entity_type=excluded.entity_type, chunk_id=excluded.chunk_id, updated_at=excluded.updated_at`,
datasetId, name, entityType, chunkId, now, now)
return err
}
func (d *kgEntityDao) DeleteByChunkIds(ctx context.Context, chunkIds []int64) error {
if len(chunkIds) == 0 {
return nil
}
placeholders := strings.TrimSuffix(strings.Repeat("?,", len(chunkIds)), ",")
args := make([]any, 0, len(chunkIds))
for _, id := range chunkIds {
args = append(args, id)
}
_, err := g.DB(consts.DbGroupDefault).Exec(ctx,
"DELETE FROM "+consts.TableNameKgEntity+" WHERE chunk_id IN ("+placeholders+")", args...)
return err
}
func (d *kgEntityDao) DeleteByDataset(ctx context.Context, datasetId int64) error {
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameKgEntity).Ctx(ctx).
Where("dataset_id", datasetId).Delete()
return err
}