104 lines
3.3 KiB
Go
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
|
|
}
|