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

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
}