82 lines
2.7 KiB
Go
82 lines
2.7 KiB
Go
package dao
|
|
|
|
import (
|
|
"context"
|
|
"sort"
|
|
|
|
"rag-local/kb/consts"
|
|
"rag-local/kb/model/entity"
|
|
|
|
"github.com/gogf/gf/v2/frame/g"
|
|
)
|
|
|
|
var ContractMark = &contractMarkDao{}
|
|
|
|
type contractMarkDao struct{}
|
|
|
|
func init() {
|
|
ctx := context.Background()
|
|
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameContractMark+` (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
clause_id INTEGER NOT NULL DEFAULT 0,
|
|
chunk_id INTEGER NOT NULL DEFAULT 0,
|
|
dataset_id INTEGER NOT NULL DEFAULT 0,
|
|
law_title TEXT NOT NULL DEFAULT '',
|
|
law_item TEXT NOT NULL DEFAULT '',
|
|
content TEXT NOT NULL DEFAULT '',
|
|
reason TEXT NOT NULL DEFAULT '',
|
|
score REAL NOT NULL DEFAULT 0,
|
|
created_at DATETIME DEFAULT (datetime('now','localtime'))
|
|
)`)
|
|
if err != nil {
|
|
g.Log().Warningf(ctx, "create kb_contract_mark table failed: %v", err)
|
|
}
|
|
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_kb_contract_mark_clause ON "+consts.TableNameContractMark+"(clause_id)"); err != nil {
|
|
g.Log().Warningf(ctx, "create index idx_kb_contract_mark_clause failed: %v", err)
|
|
}
|
|
}
|
|
|
|
func (d *contractMarkDao) ListByClause(ctx context.Context, clauseId int64) ([]*entity.ContractMark, error) {
|
|
var list []*entity.ContractMark
|
|
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractMark).Ctx(ctx).
|
|
Where("clause_id", clauseId).OrderDesc("score").Scan(&list)
|
|
if list == nil {
|
|
list = make([]*entity.ContractMark, 0)
|
|
}
|
|
return list, err
|
|
}
|
|
|
|
// ListByTask 某任务的全部标注(单表约束:先取条款 id,再按 IN 分批查询,内存按分排序)
|
|
func (d *contractMarkDao) ListByTask(ctx context.Context, taskId int64) ([]*entity.ContractMark, error) {
|
|
clauseIds, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractClause).Ctx(ctx).
|
|
Fields("id").Where("task_id", taskId).Array()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var list []*entity.ContractMark
|
|
for start := 0; start < len(clauseIds); start += 100 {
|
|
end := min(start+100, len(clauseIds))
|
|
ids := make([]int64, 0, end-start)
|
|
for _, v := range clauseIds[start:end] {
|
|
ids = append(ids, v.Int64())
|
|
}
|
|
var part []*entity.ContractMark
|
|
if err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractMark).Ctx(ctx).
|
|
WhereIn("clause_id", ids).Scan(&part); err != nil {
|
|
return nil, err
|
|
}
|
|
list = append(list, part...)
|
|
}
|
|
if list == nil {
|
|
list = make([]*entity.ContractMark, 0)
|
|
}
|
|
sort.Slice(list, func(i, j int) bool { return list[i].Score > list[j].Score })
|
|
return list, nil
|
|
}
|
|
|
|
func (d *contractMarkDao) DeleteByClause(ctx context.Context, clauseId int64) error {
|
|
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractMark).Ctx(ctx).
|
|
Where("clause_id", clauseId).Delete()
|
|
return err
|
|
}
|