194 lines
6.4 KiB
Go
194 lines
6.4 KiB
Go
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)
|
|
}
|
|
// 迁移:添加 case_id 列(关联我的案件)
|
|
for _, col := range []struct {
|
|
name string
|
|
ddl string
|
|
}{
|
|
{"case_id", "case_id INTEGER NOT NULL DEFAULT 0"},
|
|
} {
|
|
cnt, err := g.DB(consts.DbGroupDefault).Ctx(ctx).GetValue(ctx,
|
|
"SELECT COUNT(*) FROM pragma_table_info('"+consts.TableNameContractTask+"') WHERE name=?", col.name)
|
|
if err != nil || cnt.Int64() > 0 {
|
|
continue
|
|
}
|
|
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx,
|
|
"ALTER TABLE "+consts.TableNameContractTask+" ADD COLUMN "+col.ddl); err != nil {
|
|
g.Log().Warningf(ctx, "migrate kb_contract_task add column %s failed: %v", col.name, 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) ListByCaseId(ctx context.Context, caseId int64) ([]*entity.ContractTask, error) {
|
|
var list []*entity.ContractTask
|
|
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameContractTask).Ctx(ctx).
|
|
Where("case_id", caseId).OrderDesc("id").Scan(&list)
|
|
if list == nil {
|
|
list = make([]*entity.ContractTask, 0)
|
|
}
|
|
return list, err
|
|
}
|
|
|
|
func (d *contractTaskDao) Insert(ctx context.Context, filename, filePath, datasetIds string, caseId int64) (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,
|
|
"case_id": caseId,
|
|
"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()
|
|
}
|