138 lines
4.3 KiB
Go
138 lines
4.3 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 ParseTask = &parseTaskDao{}
|
|
|
|
type parseTaskDao struct{}
|
|
|
|
func init() {
|
|
ctx := context.Background()
|
|
_, err := g.DB(consts.DbGroupDefault).Exec(ctx, `CREATE TABLE IF NOT EXISTS `+consts.TableNameParseTask+` (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
document_id INTEGER NOT NULL DEFAULT 0,
|
|
dataset_id INTEGER NOT NULL DEFAULT 0,
|
|
task_type TEXT NOT NULL DEFAULT 'parse',
|
|
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_parse_task table failed: %v", err)
|
|
}
|
|
if _, err := g.DB(consts.DbGroupDefault).Exec(ctx, "CREATE INDEX IF NOT EXISTS idx_kb_parse_task_status ON "+consts.TableNameParseTask+"(status)"); err != nil {
|
|
g.Log().Warningf(ctx, "create index idx_kb_parse_task_status failed: %v", err)
|
|
}
|
|
}
|
|
|
|
func (d *parseTaskDao) GetOne(ctx context.Context, id int64) (*entity.ParseTask, error) {
|
|
var m entity.ParseTask
|
|
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).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 *parseTaskDao) List(ctx context.Context, page, pageSize int) ([]*entity.ParseTask, int, error) {
|
|
if page < 1 {
|
|
page = 1
|
|
}
|
|
if pageSize < 1 {
|
|
pageSize = 20
|
|
}
|
|
total, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).Count()
|
|
if err != nil {
|
|
return nil, 0, err
|
|
}
|
|
var list []*entity.ParseTask
|
|
err = g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).
|
|
Page(page, pageSize).OrderDesc("id").Scan(&list)
|
|
if list == nil {
|
|
list = make([]*entity.ParseTask, 0)
|
|
}
|
|
return list, total, err
|
|
}
|
|
|
|
func (d *parseTaskDao) Insert(ctx context.Context, documentId, datasetId int64, taskType string) (int64, error) {
|
|
now := gtime.Now().Format("Y-m-d H:i:s")
|
|
r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).Data(g.Map{
|
|
"document_id": documentId,
|
|
"dataset_id": datasetId,
|
|
"task_type": taskType,
|
|
"status": consts.TaskStatusPending,
|
|
"created_at": now,
|
|
"updated_at": now,
|
|
}).Insert()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return r.LastInsertId()
|
|
}
|
|
|
|
func (d *parseTaskDao) UpdateStatus(ctx context.Context, id int64, status int, errorMsg string) error {
|
|
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).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 *parseTaskDao) DeleteByDocument(ctx context.Context, documentId int64) error {
|
|
_, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).
|
|
Where("document_id", documentId).Delete()
|
|
return err
|
|
}
|
|
|
|
// ResetRunning 启动恢复:进程异常退出会孤儿化 running 任务(轮询器只消费 pending),重置回 pending 让轮询器重新捡起
|
|
func (d *parseTaskDao) ResetRunning(ctx context.Context) (int64, error) {
|
|
r, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).
|
|
Data(g.Map{
|
|
"status": consts.TaskStatusPending,
|
|
"updated_at": gtime.Now().Format("Y-m-d H:i:s"),
|
|
}).Where("status", consts.TaskStatusRunning).Update()
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return r.RowsAffected()
|
|
}
|
|
|
|
// HasPending 该文档是否存在待处理/运行中任务(重复入队防护)
|
|
func (d *parseTaskDao) HasPending(ctx context.Context, documentId int64) (bool, error) {
|
|
n, err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).
|
|
Where("document_id", documentId).
|
|
WhereIn("status", []int{consts.TaskStatusPending, consts.TaskStatusRunning}).Count()
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
return n > 0, nil
|
|
}
|
|
|
|
func (d *parseTaskDao) NextPending(ctx context.Context) (*entity.ParseTask, error) {
|
|
var m entity.ParseTask
|
|
err := g.DB(consts.DbGroupDefault).Model(consts.TableNameParseTask).Ctx(ctx).
|
|
Where("status", consts.TaskStatusPending).OrderAsc("id").Scan(&m)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, nil
|
|
}
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &m, nil
|
|
}
|