118 lines
3.7 KiB
Go
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
|
|
}
|