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

104 lines
3.3 KiB
Go

package dao
import (
"context"
"rag-local/kb/consts"
"rag-local/kb/model/entity"
"github.com/gogf/gf/v2/frame/g"
"github.com/gogf/gf/v2/os/gtime"
)
var ContractClause = &contractClauseDao{}
type contractClauseDao struct{}
func init() {
ctx := context.Background()
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameContractClause+` (
id INTEGER PRIMARY KEY AUTOINCREMENT,
task_id INTEGER NOT NULL DEFAULT 0,
seq INTEGER NOT NULL DEFAULT 0,
title TEXT NOT NULL DEFAULT '',
content TEXT NOT NULL DEFAULT '',
status INTEGER NOT NULL DEFAULT 0,
error_msg TEXT NOT NULL DEFAULT '',
created_at DATETIME DEFAULT (datetime('now','localtime')),
updated_at DATETIME DEFAULT (datetime('now','localtime'))
)`)
if err != nil {
g.Log().Warningf(ctx, "create kb_contract_clause table failed: %v", err)
}
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_kb_contract_clause_task ON "+consts.TableNameContractClause+"(task_id)"); err != nil {
g.Log().Warningf(ctx, "create index idx_kb_contract_clause_task failed: %v", err)
}
}
func (d *contractClauseDao) InsertAll(ctx context.Context, taskId int64, clauses []entity.ContractClause) error {
if len(clauses) == 0 {
return nil
}
now := gtime.Now().Format("Y-m-d H:i:s")
list := make(g.List, 0, len(clauses))
for i := range clauses {
list = append(list, g.Map{
"task_id": taskId,
"seq": clauses[i].Seq,
"title": clauses[i].Title,
"content": clauses[i].Content,
"status": consts.TaskStatusPending,
"error_msg": "",
"created_at": now,
"updated_at": now,
})
}
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.TableNameContractClause).Ctx(ctx).Data(list).Batch(100).Insert(); err != nil {
return err
}
return tx.Commit()
}
// UpdateStatuses 批量更新条款状态(单表约束:IN ≤100 分批)
func (d *contractClauseDao) UpdateStatuses(ctx context.Context, ids []int64, status int, errorMsg string) error {
now := gtime.Now().Format("Y-m-d H:i:s")
for start := 0; start < len(ids); start += 100 {
end := min(start+100, len(ids))
if _, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractClause).Ctx(ctx).
Data(g.Map{"status": status, "error_msg": errorMsg, "updated_at": now}).
WhereIn("id", ids[start:end]).Update(); err != nil {
return err
}
}
return nil
}
func (d *contractClauseDao) ListByTask(ctx context.Context, taskId int64) ([]*entity.ContractClause, error) {
var list []*entity.ContractClause
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractClause).Ctx(ctx).
Where("task_id", taskId).OrderAsc("seq").Scan(&list)
if list == nil {
list = make([]*entity.ContractClause, 0)
}
return list, err
}
func (d *contractClauseDao) UpdateStatus(ctx context.Context, id int64, status int, errorMsg string) error {
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractClause).Ctx(ctx).Data(g.Map{
"status": status,
"error_msg": errorMsg,
"updated_at": gtime.Now().Format("Y-m-d H:i:s"),
}).Where("id", id).Update()
return err
}