@@ -202,6 +202,38 @@ func normalizeSnatchPayload(taskType string, payload json.RawMessage) json.RawMe
|
||||
return out
|
||||
}
|
||||
|
||||
// mergeSnatchPayload 把编辑提交的 count 视为**新目标台数**:按旧 payload 的
|
||||
// 已完成数换算剩余、保留连续鉴权失败计数,目标不大于已完成数时拒绝。
|
||||
// 非抢机任务原样返回;旧 payload 不可解析时按新建任务归一化。
|
||||
func mergeSnatchPayload(task *model.Task, incoming json.RawMessage) (json.RawMessage, error) {
|
||||
if task.Type != model.TaskTypeSnatch {
|
||||
return incoming, nil
|
||||
}
|
||||
var old, next snatchPayload
|
||||
if err := json.Unmarshal([]byte(task.Payload), &old); err != nil {
|
||||
return normalizeSnatchPayload(task.Type, incoming), nil
|
||||
}
|
||||
if err := json.Unmarshal(incoming, &next); err != nil {
|
||||
return nil, fmt.Errorf("update task: invalid payload: %w", err)
|
||||
}
|
||||
oldTotal := old.TotalCount
|
||||
if oldTotal <= 0 {
|
||||
oldTotal = old.Count
|
||||
}
|
||||
done := oldTotal - old.Count
|
||||
if next.Count <= done {
|
||||
return nil, fmt.Errorf("%w(已完成 %d 台)", ErrSnatchTargetTooLow, done)
|
||||
}
|
||||
next.TotalCount = next.Count
|
||||
next.Count -= done
|
||||
next.AuthFailCount = old.AuthFailCount
|
||||
out, err := json.Marshal(next)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("update task: marshal payload: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// validateTaskPayload 按任务类型校验参数 JSON。
|
||||
func validateTaskPayload(taskType string, payload json.RawMessage) error {
|
||||
switch taskType {
|
||||
@@ -245,7 +277,8 @@ type UpdateTaskInput struct {
|
||||
Status *string
|
||||
}
|
||||
|
||||
// UpdateTask 修改任务并重新调度。
|
||||
// UpdateTask 修改任务并重新调度;写入用 updated_at 条件更新,
|
||||
// 防止陈旧快照整行覆盖并发执行刚落库的状态与进度。
|
||||
func (s *TaskService) UpdateTask(ctx context.Context, id uint, in UpdateTaskInput) (*model.Task, error) {
|
||||
task, err := s.GetTask(ctx, id)
|
||||
if err != nil {
|
||||
@@ -254,16 +287,39 @@ func (s *TaskService) UpdateTask(ctx context.Context, id uint, in UpdateTaskInpu
|
||||
if err := applyTaskUpdate(task, in); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Save(task).Error; err != nil {
|
||||
return nil, fmt.Errorf("update task %d: %w", id, err)
|
||||
fresh, err := s.persistTaskUpdate(ctx, task)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.unschedule(task.ID)
|
||||
if task.Status == model.TaskStatusActive {
|
||||
if err := s.schedule(task); err != nil {
|
||||
s.unschedule(fresh.ID)
|
||||
if fresh.Status == model.TaskStatusActive {
|
||||
if err := s.schedule(fresh); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return task, nil
|
||||
return fresh, nil
|
||||
}
|
||||
|
||||
// persistTaskUpdate 只写用户可编辑列;零命中说明执行侧已并发落库,返回冲突。
|
||||
func (s *TaskService) persistTaskUpdate(ctx context.Context, task *model.Task) (*model.Task, error) {
|
||||
res := s.db.WithContext(ctx).Model(&model.Task{}).
|
||||
Where("id = ? AND updated_at = ?", task.ID, task.UpdatedAt).
|
||||
Updates(map[string]any{
|
||||
"name": task.Name, "cron_expr": task.CronExpr,
|
||||
"payload": task.Payload, "status": task.Status,
|
||||
})
|
||||
if res.Error != nil {
|
||||
return nil, fmt.Errorf("update task %d: %w", task.ID, res.Error)
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return nil, ErrTaskConflict
|
||||
}
|
||||
// 重读用新变量:gorm 扫描 NULL 列到已有值的结构体时会保留旧值
|
||||
var fresh model.Task
|
||||
if err := s.db.WithContext(ctx).First(&fresh, task.ID).Error; err != nil {
|
||||
return nil, fmt.Errorf("reload task %d: %w", task.ID, err)
|
||||
}
|
||||
return &fresh, nil
|
||||
}
|
||||
|
||||
func applyTaskUpdate(task *model.Task, in UpdateTaskInput) error {
|
||||
@@ -280,7 +336,11 @@ func applyTaskUpdate(task *model.Task, in UpdateTaskInput) error {
|
||||
if err := validateTaskPayload(task.Type, in.Payload); err != nil {
|
||||
return err
|
||||
}
|
||||
task.Payload = string(in.Payload)
|
||||
merged, err := mergeSnatchPayload(task, in.Payload)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
task.Payload = string(merged)
|
||||
}
|
||||
if in.Status != nil {
|
||||
if *in.Status != model.TaskStatusActive && *in.Status != model.TaskStatusPaused {
|
||||
@@ -376,6 +436,12 @@ func (s *TaskService) RunTaskNow(ctx context.Context, id uint) (*model.TaskLog,
|
||||
// ErrTaskRunning 表示任务已有一次执行在途,拒绝重复触发。
|
||||
var ErrTaskRunning = errors.New("任务正在执行中,请稍候")
|
||||
|
||||
// ErrTaskConflict 表示任务在编辑期间被并发修改(如执行结果落库),须刷新重试。
|
||||
var ErrTaskConflict = errors.New("任务状态已变化,请刷新后重试")
|
||||
|
||||
// ErrSnatchTargetTooLow 表示编辑抢机任务时新目标台数不大于已完成数量。
|
||||
var ErrSnatchTargetTooLow = errors.New("目标台数须大于已完成数量")
|
||||
|
||||
// TriggerTask 异步触发一次任务执行并立即返回;执行结果照常落任务日志与通知,
|
||||
// 由前端轮询呈现。同一任务在途时返回 ErrTaskRunning。
|
||||
func (s *TaskService) TriggerTask(ctx context.Context, id uint) error {
|
||||
|
||||
Reference in New Issue
Block a user