package dao import ( "context" "database/sql" "errors" "rag-local/kb/consts" "rag-local/kb/model/entity" "github.com/gogf/gf/v2/frame/g" "github.com/gogf/gf/v2/os/gtime" ) var ContractTask = &contractTaskDao{} type contractTaskDao struct{} func init() { ctx := context.Background() _, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameContractTask+` ( id INTEGER PRIMARY KEY AUTOINCREMENT, filename TEXT NOT NULL DEFAULT '', file_path TEXT NOT NULL DEFAULT '', dataset_ids TEXT NOT NULL DEFAULT '', status INTEGER NOT NULL DEFAULT 0, total_clauses INTEGER NOT NULL DEFAULT 0, done_clauses 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_task table failed: %v", err) } if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_kb_contract_task_status ON "+consts.TableNameContractTask+"(status)"); err != nil { g.Log().Warningf(ctx, "create index idx_kb_contract_task_status failed: %v", err) } } func (d *contractTaskDao) GetOne(ctx context.Context, id int64) (*entity.ContractTask, error) { var m entity.ContractTask err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).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 *contractTaskDao) List(ctx context.Context, page, pageSize int) ([]*entity.ContractTask, int, error) { if page < 1 { page = 1 } if pageSize < 1 { pageSize = 20 } total, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx).Count() if err != nil { return nil, 0, err } var list []*entity.ContractTask err = g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx). Page(page, pageSize).OrderDesc("id").Scan(&list) if list == nil { list = make([]*entity.ContractTask, 0) } return list, total, err } func (d *contractTaskDao) Insert(ctx context.Context, filename, filePath, datasetIds string) (int64, error) { now := gtime.Now().Format("Y-m-d H:i:s") r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx).Data(g.Map{ "filename": filename, "file_path": filePath, "dataset_ids": datasetIds, "status": consts.TaskStatusPending, "created_at": now, "updated_at": now, }).Insert() if err != nil { return 0, err } return r.LastInsertId() } func (d *contractTaskDao) UpdateStatus(ctx context.Context, id int64, status int, errorMsg string) error { _, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).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 } func (d *contractTaskDao) UpdateProgress(ctx context.Context, id int64, totalClauses, doneClauses int) error { _, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx).Data(g.Map{ "total_clauses": totalClauses, "done_clauses": doneClauses, "updated_at": gtime.Now().Format("Y-m-d H:i:s"), }).Where("id", id).Update() return err } // NextPending 取待处理任务;status IN (Pending, Running) 使进程重启后遗留的 running 任务重新进入轮询 func (d *contractTaskDao) NextPending(ctx context.Context) (*entity.ContractTask, error) { var m entity.ContractTask err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx). WhereIn("status", []int{consts.TaskStatusPending, consts.TaskStatusRunning}).OrderAsc("id").Scan(&m) if errors.Is(err, sql.ErrNoRows) { return nil, nil } if err != nil { return nil, err } return &m, nil } func (d *contractTaskDao) Delete(ctx context.Context, id int64) error { _, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx). Where("id", id).Delete() return err } // DeleteWithRelated 事务删除任务及其条款、标注、风险数据 func (d *contractTaskDao) DeleteWithRelated(ctx context.Context, id int64) error { tx, err := g.DB(consts.DbGroupDefault).Begin(ctx) if err != nil { return err } // Commit 成功后 IsClosed 为 true,跳过 Rollback,避免对已提交事务回滚产生报错日志 defer func() { if !tx.IsClosed() { _ = tx.Rollback() } }() // 单表约束:先取条款 id,再按 IN 分批删除(条款数可能超 SQLite 变量上限) clauseIds, err := tx.Model(consts.TableNameContractClause).Ctx(ctx).Fields("id").Where("task_id", id).Array() if err != nil { return err } for _, table := range []string{consts.TableNameContractMark, consts.TableNameContractRisk} { 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()) } if len(ids) > 0 { if _, err := tx.Model(table).Ctx(ctx).WhereIn("clause_id", ids).Delete(); err != nil { return err } } } } if _, err := tx.Model(consts.TableNameContractClause).Ctx(ctx).Where("task_id", id).Delete(); err != nil { return err } if _, err := tx.Model(consts.TableNameContractTask).Ctx(ctx).Where("id", id).Delete(); err != nil { return err } return tx.Commit() }