@@ -8,6 +8,7 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
@@ -835,3 +836,93 @@ func TestTriggerTaskAsyncAndDedup(t *testing.T) {
|
||||
t.Error("不存在的任务应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// snatchInstanceJSON 是编辑测试用的合法实例参数片段。
|
||||
const snatchInstanceJSON = `"instance":{"displayName":"vm","shape":"VM.Standard.A1.Flex","ocpus":4,"memoryInGBs":24,"imageId":"ocid1.image.test"}`
|
||||
|
||||
// TestUpdateSnatchTaskCountIsTarget 锁定编辑语义:提交的 count 是目标台数,
|
||||
// 按旧 payload 已完成数换算剩余、保留鉴权失败计数;目标不大于已完成数拒绝。
|
||||
func TestUpdateSnatchTaskCountIsTarget(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
oldPayload string
|
||||
newCount int
|
||||
wantErr error
|
||||
want snatchPayload
|
||||
}{
|
||||
{"部分成功后提高目标", `{"ociConfigId":1,"count":3,"totalCount":5,"authFailCount":2,` + snatchInstanceJSON + `}`,
|
||||
6, nil, snatchPayload{Count: 4, TotalCount: 6, AuthFailCount: 2}},
|
||||
{"原样保存不丢进度", `{"ociConfigId":1,"count":3,"totalCount":5,"authFailCount":2,` + snatchInstanceJSON + `}`,
|
||||
5, nil, snatchPayload{Count: 3, TotalCount: 5, AuthFailCount: 2}},
|
||||
{"目标不大于已完成拒绝", `{"ociConfigId":1,"count":3,"totalCount":5,` + snatchInstanceJSON + `}`,
|
||||
2, ErrSnatchTargetTooLow, snatchPayload{}},
|
||||
{"旧任务缺 totalCount 视为零完成", `{"ociConfigId":1,"count":3,` + snatchInstanceJSON + `}`,
|
||||
5, nil, snatchPayload{Count: 5, TotalCount: 5}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tasks, _, db := newTaskEnv(t, &fakeClient{})
|
||||
ctx := context.Background()
|
||||
task, err := tasks.CreateTask(ctx, CreateTaskInput{
|
||||
Name: "抢机", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
|
||||
Payload: json.RawMessage(tt.oldPayload),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTask: %v", err)
|
||||
}
|
||||
// 直写旧 payload,绕过创建期归一化,构造「执行过若干轮」的现场
|
||||
if err := db.Model(&model.Task{}).Where("id = ?", task.ID).
|
||||
Update("payload", tt.oldPayload).Error; err != nil {
|
||||
t.Fatalf("seed payload: %v", err)
|
||||
}
|
||||
incoming := fmt.Sprintf(`{"ociConfigId":1,"count":%d,`+snatchInstanceJSON+`}`, tt.newCount)
|
||||
updated, err := tasks.UpdateTask(ctx, task.ID, UpdateTaskInput{Payload: json.RawMessage(incoming)})
|
||||
if tt.wantErr != nil {
|
||||
if !errors.Is(err, tt.wantErr) {
|
||||
t.Fatalf("err = %v, want %v", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateTask: %v", err)
|
||||
}
|
||||
var got snatchPayload
|
||||
if err := json.Unmarshal([]byte(updated.Payload), &got); err != nil {
|
||||
t.Fatalf("payload: %v", err)
|
||||
}
|
||||
if got.Count != tt.want.Count || got.TotalCount != tt.want.TotalCount || got.AuthFailCount != tt.want.AuthFailCount {
|
||||
t.Errorf("payload = {count:%d total:%d authFail:%d}, want {count:%d total:%d authFail:%d}",
|
||||
got.Count, got.TotalCount, got.AuthFailCount, tt.want.Count, tt.want.TotalCount, tt.want.AuthFailCount)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestPersistTaskUpdateConflict 锁定编辑竞态防护:陈旧快照零命中返回冲突且不落库。
|
||||
func TestPersistTaskUpdateConflict(t *testing.T) {
|
||||
tasks, _, db := newTaskEnv(t, &fakeClient{})
|
||||
ctx := context.Background()
|
||||
task, err := tasks.CreateTask(ctx, CreateTaskInput{
|
||||
Name: "测活", Type: model.TaskTypeHealthCheck, CronExpr: "*/10 * * * *",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CreateTask: %v", err)
|
||||
}
|
||||
stale, err := tasks.GetTask(ctx, task.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetTask: %v", err)
|
||||
}
|
||||
// 模拟执行侧并发落库:行的 updated_at 前移,编辑快照过期
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
if err := db.Model(&model.Task{}).Where("id = ?", task.ID).Update("run_count", 1).Error; err != nil {
|
||||
t.Fatalf("simulate run persist: %v", err)
|
||||
}
|
||||
stale.Name = "renamed"
|
||||
if _, err := tasks.persistTaskUpdate(ctx, stale); !errors.Is(err, ErrTaskConflict) {
|
||||
t.Fatalf("err = %v, want ErrTaskConflict", err)
|
||||
}
|
||||
fresh, _ := tasks.GetTask(ctx, task.ID)
|
||||
if fresh.Name != "测活" || fresh.RunCount != 1 {
|
||||
t.Errorf("row = {name:%s runCount:%d}, want 执行结果保留且未被改名", fresh.Name, fresh.RunCount)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user