Files
oci-portal/internal/service/task_test.go
T
Wang Defa dbba1f4905
CI / test (push) Successful in 30s
Release / release (push) Successful in 49s
发布 0.1.0:通知渠道、告警规则、令牌版本与安全加固
2026-07-10 17:38:34 +08:00

804 lines
28 KiB
Go

package service
import (
"context"
"encoding/json"
"fmt"
"strings"
"sync"
"testing"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"oci-portal/internal/crypto"
"oci-portal/internal/model"
"oci-portal/internal/oci"
)
// snatchFake 可控 LaunchInstance 结果:前 failFirst 次失败模拟容量不足。
type snatchFake struct {
fakeClient
launches int
failFirst int
}
func (f *snatchFake) LaunchInstance(ctx context.Context, cred oci.Credentials, in oci.CreateInstanceInput) (oci.Instance, error) {
f.launches++
if f.launches <= f.failFirst {
return oci.Instance{}, fmt.Errorf("Out of host capacity")
}
return oci.Instance{ID: fmt.Sprintf("ocid1.instance.%d", f.launches), DisplayName: in.DisplayName}, nil
}
// scriptedSnatchFake 按脚本逐次控制 LaunchInstance 结果:nil 表示成功;
// 脚本耗尽后一律成功。
type scriptedSnatchFake struct {
fakeClient
mu sync.Mutex
script []error
launches int
}
func (f *scriptedSnatchFake) LaunchInstance(ctx context.Context, cred oci.Credentials, in oci.CreateInstanceInput) (oci.Instance, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.launches++
if len(f.script) > 0 {
err := f.script[0]
f.script = f.script[1:]
if err != nil {
return oci.Instance{}, err
}
}
return oci.Instance{ID: fmt.Sprintf("ocid1.instance.%d", f.launches), DisplayName: in.DisplayName}, nil
}
func newTaskEnv(t *testing.T, client oci.Client) (*TaskService, *OciConfigService, *gorm.DB) {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatalf("open in-memory sqlite: %v", err)
}
// :memory: 库每个连接彼此独立,锁单连接避免池内新连接拿到空库
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db handle: %v", err)
}
sqlDB.SetMaxOpenConns(1)
if err := db.AutoMigrate(
&model.OciConfig{}, &model.Task{}, &model.TaskLog{},
&model.CheckSnapshot{}, &model.CostSnapshot{}, &model.Setting{}, &model.Proxy{},
); err != nil {
t.Fatalf("auto migrate: %v", err)
}
cipher, err := crypto.NewCipher("test-data-key")
if err != nil {
t.Fatalf("new cipher: %v", err)
}
configs := NewOciConfigService(db, cipher, client)
tasks := NewTaskService(db, configs, nil, NewSettingService(db, cipher))
// 测试不等待公网 IP 分配:预算置 0,首查后立即回退
tasks.snatchIPWait = 0
return tasks, configs, db
}
func TestCreateTaskValidation(t *testing.T) {
tasks, _, _ := newTaskEnv(t, &fakeClient{})
tests := []struct {
name string
in CreateTaskInput
}{
{name: "非法 cron", in: CreateTaskInput{Name: "t", Type: model.TaskTypeHealthCheck, CronExpr: "not-cron"}},
{name: "未知类型", in: CreateTaskInput{Name: "t", Type: "nope", CronExpr: "* * * * *"}},
{name: "缺名称", in: CreateTaskInput{Type: model.TaskTypeHealthCheck, CronExpr: "* * * * *"}},
{
name: "抢机缺配置ID",
in: CreateTaskInput{Name: "t", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"instance":{"displayName":"a","shape":"s","imageId":"i"}}`)},
},
{
name: "抢机实例参数不全",
in: CreateTaskInput{Name: "t", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"instance":{"shape":"s"}}`)},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if _, err := tasks.CreateTask(context.Background(), tt.in); err == nil {
t.Error("CreateTask succeeded, want error")
}
})
}
}
func TestRunHealthCheckTask(t *testing.T) {
client := &fakeClient{tenancy: oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}}
tasks, configs, _ := newTaskEnv(t, client)
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
second := trialImportInput()
second.Alias = "第二个"
if _, err := configs.Import(ctx, second); err != nil {
t.Fatalf("import second: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "测活", Type: model.TaskTypeHealthCheck, CronExpr: "*/5 * * * *",
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
entry, err := tasks.RunTaskNow(ctx, task.ID)
if err != nil {
t.Fatalf("RunTaskNow: %v", err)
}
if !entry.Success || entry.Message != "checked 2: 2 alive, 0 dead" {
t.Errorf("log = %+v, want success checked 2", entry)
}
scoped, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "范围测活", Type: model.TaskTypeHealthCheck, CronExpr: "*/5 * * * *",
Payload: json.RawMessage(`{"ociConfigIds":[1]}`),
})
if err != nil {
t.Fatalf("CreateTask scoped: %v", err)
}
entry, err = tasks.RunTaskNow(ctx, scoped.ID)
if err != nil {
t.Fatalf("RunTaskNow scoped: %v", err)
}
if entry.Message != "checked 1: 1 alive, 0 dead" {
t.Errorf("scoped message = %q, want checked 1", entry.Message)
}
}
func TestRunSnatchTaskUntilSuccess(t *testing.T) {
client := &snatchFake{failFirst: 2}
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "抢机", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"instance":{"displayName":"snatch-vm","shape":"VM.Standard.A1.Flex","ocpus":4,"memoryInGBs":24,"imageId":"ocid1.image.test"}}`),
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
// 前两次容量不足:任务保持 active 并记录失败
for i := 1; i <= 2; i++ {
entry, err := tasks.RunTaskNow(ctx, task.ID)
if err != nil {
t.Fatalf("run %d: %v", i, err)
}
if entry.Success {
t.Fatalf("run %d succeeded, want capacity failure", i)
}
}
got, _ := tasks.GetTask(ctx, task.ID)
if got.Status != model.TaskStatusActive || got.RunCount != 2 || got.LastError == "" {
t.Errorf("after failures: status=%s runs=%d lastError=%q", got.Status, got.RunCount, got.LastError)
}
// 第三次成功:任务自动标记 succeeded
entry, err := tasks.RunTaskNow(ctx, task.ID)
if err != nil {
t.Fatalf("run 3: %v", err)
}
if !entry.Success {
t.Fatalf("run 3 = %+v, want success", entry)
}
got, _ = tasks.GetTask(ctx, task.ID)
if got.Status != model.TaskStatusSucceeded || got.LastError != "" {
t.Errorf("after success: status=%s lastError=%q, want succeeded", got.Status, got.LastError)
}
logs, err := tasks.TaskLogs(ctx, task.ID, 10)
if err != nil {
t.Fatalf("TaskLogs: %v", err)
}
if len(logs) != 3 || !logs[0].Success {
t.Errorf("logs = %d entries, newest success=%v; want 3 with newest success", len(logs), logs[0].Success)
}
}
// adRotationFake 记录每次 LaunchInstance 收到的可用域;创建一律按容量不足
// 失败,任务保持 active 以便连续执行验证轮询顺序。
type adRotationFake struct {
fakeClient
ads []string
got []string
}
func (f *adRotationFake) ListAvailabilityDomains(ctx context.Context, cred oci.Credentials, region string) ([]string, error) {
return f.ads, nil
}
func (f *adRotationFake) LaunchInstance(ctx context.Context, cred oci.Credentials, in oci.CreateInstanceInput) (oci.Instance, error) {
f.got = append(f.got, in.AvailabilityDomain)
return oci.Instance{}, fmt.Errorf("Out of host capacity")
}
// TestSnatchAdRotation 可用域留空(自动)时按执行序号轮询区域全部可用域,
// 显式指定时不轮询。
func TestSnatchAdRotation(t *testing.T) {
client := &adRotationFake{ads: []string{"ad-1", "ad-2", "ad-3"}}
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "抢机", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"instance":{"displayName":"vm","shape":"s","imageId":"i"}}`),
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
for i := 0; i < 4; i++ {
if _, err := tasks.RunTaskNow(ctx, task.ID); err != nil {
t.Fatalf("run %d: %v", i, err)
}
}
if got := strings.Join(client.got, ","); got != "ad-1,ad-2,ad-3,ad-1" {
t.Errorf("自动轮询序列 = %s, want ad-1,ad-2,ad-3,ad-1", got)
}
client.got = nil
fixed, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "定向", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"instance":{"displayName":"vm","shape":"s","imageId":"i","availabilityDomain":"ad-2"}}`),
})
if err != nil {
t.Fatalf("CreateTask fixed: %v", err)
}
for i := 0; i < 2; i++ {
if _, err := tasks.RunTaskNow(ctx, fixed.ID); err != nil {
t.Fatalf("run fixed %d: %v", i, err)
}
}
if got := strings.Join(client.got, ","); got != "ad-2,ad-2" {
t.Errorf("显式指定序列 = %s, want ad-2,ad-2", got)
}
}
func TestSnatchPartialSuccessKeepsRemaining(t *testing.T) {
client := &snatchFake{failFirst: 1} // 第 1 台失败、第 2 台成功
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "抢两台", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"count":2,"instance":{"displayName":"vm","shape":"s","imageId":"i"}}`),
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
entry, err := tasks.RunTaskNow(ctx, task.ID)
if err != nil {
t.Fatalf("RunTaskNow: %v", err)
}
if !entry.Success {
t.Fatalf("log = %+v, want partial success", entry)
}
got, _ := tasks.GetTask(ctx, task.ID)
if got.Status != model.TaskStatusActive {
t.Errorf("status = %s, want active(还差 1 台)", got.Status)
}
var p snatchPayload
if err := json.Unmarshal([]byte(got.Payload), &p); err != nil || p.Count != 1 {
t.Errorf("payload count = %d (%v), want 1", p.Count, err)
}
}
func TestUpdateTaskPauseResume(t *testing.T) {
tasks, _, _ := 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)
}
paused := model.TaskStatusPaused
updated, err := tasks.UpdateTask(ctx, task.ID, UpdateTaskInput{Status: &paused})
if err != nil {
t.Fatalf("pause: %v", err)
}
if updated.Status != model.TaskStatusPaused {
t.Errorf("status = %s, want paused", updated.Status)
}
if _, ok := tasks.entries[task.ID]; ok {
t.Error("paused task still scheduled")
}
active := model.TaskStatusActive
cronExpr := "*/30 * * * *"
updated, err = tasks.UpdateTask(ctx, task.ID, UpdateTaskInput{Status: &active, CronExpr: &cronExpr})
if err != nil {
t.Fatalf("resume: %v", err)
}
if updated.CronExpr != cronExpr || updated.Status != model.TaskStatusActive {
t.Errorf("updated = %+v, want active with new cron", updated)
}
if _, ok := tasks.entries[task.ID]; !ok {
t.Error("resumed task not scheduled")
}
}
// enabledTelegramInput 返回启用态的 Telegram 配置输入,供通知测试复用。
func enabledTelegramInput() *UpdateTelegramInput {
token := "123456:AAfake"
return &UpdateTelegramInput{Enabled: true, BotToken: &token, ChatID: "42"}
}
func TestSnatchRepeatedFailureNotifiesOnce(t *testing.T) {
client := &snatchFake{failFirst: 3}
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
srv, rec := newFakeTelegram(t, `{"ok":true}`)
tasks.notifier = newTestNotifier(t, srv.URL, enabledTelegramInput())
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "抢机", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"instance":{"displayName":"vm","shape":"s","imageId":"i"}}`),
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
// 连续失败 3 次:只发 1 条失败通知
for i := 1; i <= 3; i++ {
if _, err := tasks.RunTaskNow(ctx, task.ID); err != nil {
t.Fatalf("run %d: %v", i, err)
}
}
tasks.notifier.Wait()
if texts := rec.snapshot(); len(texts) != 1 || !strings.Contains(texts[0], "任务失败") {
t.Fatalf("失败阶段通知 = %v, want 仅 1 条任务失败", texts)
}
// 第 4 次抢满:追加 1 条抢机成功,不叠加恢复通知
if _, err := tasks.RunTaskNow(ctx, task.ID); err != nil {
t.Fatalf("run 4: %v", err)
}
tasks.notifier.Wait()
texts := rec.snapshot()
if len(texts) != 2 || !strings.Contains(texts[1], "抢机成功") {
t.Fatalf("成功后通知 = %v, want 追加 1 条抢机成功", texts)
}
}
// TestNotifyFiltersDisabledKinds 验证按事件类型过滤:关闭「任务失败」后
// 失败不再推送,任务恢复与抢机成功不受影响仍推送。
func TestNotifyFiltersDisabledKinds(t *testing.T) {
client := &snatchFake{failFirst: 3}
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
srv, rec := newFakeTelegram(t, `{"ok":true}`)
tasks.notifier = newTestNotifier(t, srv.URL, enabledTelegramInput())
ctx := context.Background()
off := NotifyEventsView{TaskRecover: true, SnatchSuccess: true, TenantDead: true, TaskStop: true}
if err := tasks.settings.UpdateNotifyEvents(ctx, off); err != nil {
t.Fatalf("UpdateNotifyEvents: %v", err)
}
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "抢两台", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(`{"ociConfigId":1,"count":2,"instance":{"displayName":"vm","shape":"s","imageId":"i"}}`),
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
// 三次执行依次触发:task_fail(已关闭)→ task_recover → snatch_success
steps := []struct {
name string
wantLen int // 执行后累计推送条数
wantKw string // 最新一条需包含的关键词;空表示本步无新增
}{
{name: "两台全失败:任务失败已关闭不发", wantLen: 0},
{name: "抢到 1 台剩 1:任务恢复仍发", wantLen: 1, wantKw: "任务恢复"},
{name: "抢满收尾:抢机成功仍发", wantLen: 2, wantKw: "抢机成功"},
}
for _, step := range steps {
if _, err := tasks.RunTaskNow(ctx, task.ID); err != nil {
t.Fatalf("%s: %v", step.name, err)
}
tasks.notifier.Wait()
texts := rec.snapshot()
if len(texts) != step.wantLen {
t.Fatalf("%s: 通知 = %v (%d 条), want %d 条", step.name, texts, len(texts), step.wantLen)
}
if step.wantKw != "" && !strings.Contains(texts[len(texts)-1], step.wantKw) {
t.Fatalf("%s: 最新通知 = %q, want contains %q", step.name, texts[len(texts)-1], step.wantKw)
}
}
}
// deadPickFake 按 UserOCID 控制测活结果,模拟部分租户失联。
type deadPickFake struct {
fakeClient
mu sync.Mutex
dead map[string]bool
}
func (f *deadPickFake) ValidateKey(ctx context.Context, cred oci.Credentials) (oci.TenancyInfo, error) {
f.mu.Lock()
defer f.mu.Unlock()
if f.dead[cred.UserOCID] {
return oci.TenancyInfo{}, fmt.Errorf("NotAuthenticated")
}
return f.tenancy, nil
}
func (f *deadPickFake) setDead(user string, dead bool) {
f.mu.Lock()
defer f.mu.Unlock()
f.dead[user] = dead
}
func TestHealthCheckNotifiesDeadAliasOnChange(t *testing.T) {
client := &deadPickFake{dead: map[string]bool{}}
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
srv, rec := newFakeTelegram(t, `{"ok":true}`)
tasks.notifier = newTestNotifier(t, srv.URL, enabledTelegramInput())
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import A: %v", err)
}
second := trialImportInput()
second.Alias = "second-tenant"
second.UserOCID = "ocid1.user.oc1..u2"
if _, err := configs.Import(ctx, second); err != nil {
t.Fatalf("import B: %v", err)
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "测活", Type: model.TaskTypeHealthCheck, CronExpr: "*/5 * * * *",
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
run := func(step string) {
t.Helper()
if _, err := tasks.RunTaskNow(ctx, task.ID); err != nil {
t.Fatalf("%s: %v", step, err)
}
tasks.notifier.Wait()
}
run("全部存活")
client.setDead("ocid1.user.oc1..u2", true)
run("B 失联") // 集合变化:发 1 条
run("B 持续失联") // 集合不变:不发
client.setDead("ocid1.user.oc1..u2", false)
run("B 恢复") // 集合清空:不发失联通知
client.setDead("ocid1.user.oc1..u2", true)
run("B 再失联") // 集合再次变化:再发 1 条
texts := rec.snapshot()
if len(texts) != 2 {
t.Fatalf("失联通知 %d 条 %v, want 2 条", len(texts), texts)
}
for i, text := range texts {
if !strings.Contains(text, "租户失联") || !strings.Contains(text, "second-tenant") {
t.Errorf("text[%d] = %q, want 含「租户失联」与失联别名", i, text)
}
}
}
// snatchAuthCount 解析任务 payload 中的连续鉴权失败计数。
func snatchAuthCount(t *testing.T, payload string) int {
t.Helper()
var p snatchPayload
if err := json.Unmarshal([]byte(payload), &p); err != nil {
t.Fatalf("parse payload %q: %v", payload, err)
}
return p.AuthFailCount
}
func TestSnatchAuthFuse(t *testing.T) {
notAuth := fmt.Errorf("NotAuthenticated: 鉴权失败")
capacity := fmt.Errorf("Out of host capacity")
tests := []struct {
name string
limit int // 0 表示不设置,用默认 3
count int // 目标台数;0 按 1 处理
script []error
runs int
wantStatus string
wantAuthCnt int
}{
{
name: "连续 2 次不熔断", script: []error{notAuth, notAuth}, runs: 2,
wantStatus: model.TaskStatusActive, wantAuthCnt: 2,
},
{
name: "默认阈值第 3 次熔断", script: []error{notAuth, notAuth, notAuth}, runs: 3,
wantStatus: model.TaskStatusFailed, wantAuthCnt: 3,
},
{
name: "中途成功清零", count: 2, script: []error{notAuth, notAuth, nil, capacity}, runs: 2,
wantStatus: model.TaskStatusActive, wantAuthCnt: 0,
},
{
name: "中途其他错误清零", script: []error{notAuth, capacity}, runs: 2,
wantStatus: model.TaskStatusActive, wantAuthCnt: 0,
},
{
name: "抢满收尾清零并写回", script: []error{notAuth, nil}, runs: 2,
wantStatus: model.TaskStatusSucceeded, wantAuthCnt: 0,
},
{
name: "limit=1 立即熔断", limit: 1, script: []error{notAuth}, runs: 1,
wantStatus: model.TaskStatusFailed, wantAuthCnt: 1,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
task, tasks := runAuthFuseCase(t, tt.limit, tt.count, tt.script, tt.runs)
got, err := tasks.GetTask(context.Background(), task.ID)
if err != nil {
t.Fatalf("GetTask: %v", err)
}
if got.Status != tt.wantStatus {
t.Errorf("status = %s, want %s", got.Status, tt.wantStatus)
}
if cnt := snatchAuthCount(t, got.Payload); cnt != tt.wantAuthCnt {
t.Errorf("authFailCount = %d, want %d", cnt, tt.wantAuthCnt)
}
assertAuthFuseSideEffects(t, tasks, got, tt.wantStatus, tt.wantAuthCnt)
})
}
}
// runAuthFuseCase 搭好环境并按用例参数执行 runs 次抢机任务。
func runAuthFuseCase(t *testing.T, limit, count int, script []error, runs int) (*model.Task, *TaskService) {
t.Helper()
client := &scriptedSnatchFake{script: script}
client.tenancy = oci.TenancyInfo{Name: "acme", HomeRegionKey: "FRA"}
tasks, configs, _ := newTaskEnv(t, client)
ctx := context.Background()
if _, err := configs.Import(ctx, trialImportInput()); err != nil {
t.Fatalf("import: %v", err)
}
if limit > 0 {
if err := tasks.settings.UpdateTaskSettings(ctx, TaskSettingsView{SnatchAuthFailLimit: limit}); err != nil {
t.Fatalf("UpdateTaskSettings: %v", err)
}
}
if count == 0 {
count = 1
}
task, err := tasks.CreateTask(ctx, CreateTaskInput{
Name: "抢机", Type: model.TaskTypeSnatch, CronExpr: "* * * * *",
Payload: json.RawMessage(fmt.Sprintf(
`{"ociConfigId":1,"count":%d,"instance":{"displayName":"vm","shape":"s","imageId":"i"}}`, count)),
})
if err != nil {
t.Fatalf("CreateTask: %v", err)
}
for i := 1; i <= runs; i++ {
if _, err := tasks.RunTaskNow(ctx, task.ID); err != nil {
t.Fatalf("run %d: %v", i, err)
}
}
return task, tasks
}
// assertAuthFuseSideEffects 断言熔断副作用:仅 active 保持调度(failed 熔断、
// succeeded 抢满均移出);failed 的 LastError 须带熔断说明,其余不得出现。
func assertAuthFuseSideEffects(t *testing.T, tasks *TaskService, got *model.Task, wantStatus string, wantCnt int) {
t.Helper()
if _, scheduled := tasks.entries[got.ID]; scheduled != (wantStatus == model.TaskStatusActive) {
t.Errorf("scheduled = %v, want %v", scheduled, wantStatus == model.TaskStatusActive)
}
fused := wantStatus == model.TaskStatusFailed
fuseMsg := fmt.Sprintf("连续 %d 次 NotAuthenticated,任务已熔断停止", wantCnt)
if fused && !strings.Contains(got.LastError, fuseMsg) {
t.Errorf("LastError = %q, want 含 %q", got.LastError, fuseMsg)
}
if !fused && strings.Contains(got.LastError, "任务已熔断停止") {
t.Errorf("LastError = %q, 不应出现熔断说明", got.LastError)
}
}
// TestNotifyEventsFailedTransitions 锁定熔断相关通知语义:置 failed 当次
// 只发任务停止不叠加任务失败;重新启用后恢复只发任务恢复。
func TestNotifyEventsFailedTransitions(t *testing.T) {
fuseErr := "连续 3 次 NotAuthenticated,任务已熔断停止: NotAuthenticated"
tests := []struct {
name string
prev, cur taskSnapshot
wantKinds []notifyKind
wantText string // 首条通知变量值需包含的关键词;空表示不校验
}{
{
name: "熔断翻转只发任务停止不叠加任务失败",
prev: taskSnapshot{Name: "抢机", Status: model.TaskStatusActive},
cur: taskSnapshot{Name: "抢机", Status: model.TaskStatusFailed, LastError: fuseErr},
wantKinds: []notifyKind{notifyTaskStop},
wantText: "任务已熔断停止",
},
{
name: "持续 failed 不重复发",
prev: taskSnapshot{Name: "抢机", Status: model.TaskStatusFailed, LastError: fuseErr},
cur: taskSnapshot{Name: "抢机", Status: model.TaskStatusFailed, LastError: fuseErr},
wantKinds: nil,
},
{
name: "failed 重新启用后成功只发任务恢复",
prev: taskSnapshot{Name: "抢机", Status: model.TaskStatusActive, LastError: fuseErr},
cur: taskSnapshot{Name: "抢机", Status: model.TaskStatusActive},
wantKinds: []notifyKind{notifyTaskRecover},
wantText: "抢机",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
events := notifyEvents(tt.prev, tt.cur)
if len(events) != len(tt.wantKinds) {
t.Fatalf("events = %+v (%d 条), want %v", events, len(events), tt.wantKinds)
}
for i, ev := range events {
if ev.Kind != tt.wantKinds[i] {
t.Errorf("kind[%d] = %s, want %s", i, ev.Kind, tt.wantKinds[i])
}
}
if tt.wantText != "" {
joined := ""
for _, v := range events[0].Vars {
joined += v + " "
}
if !strings.Contains(joined, tt.wantText) {
t.Errorf("vars = %q, want contains %q", joined, tt.wantText)
}
}
})
}
}
func TestUpdateTaskFailedToActiveResetsAuthFail(t *testing.T) {
notAuth := fmt.Errorf("NotAuthenticated: 鉴权失败")
task, tasks := runAuthFuseCase(t, 0, 1, []error{notAuth, notAuth, notAuth}, 3)
ctx := context.Background()
got, err := tasks.GetTask(ctx, task.ID)
if err != nil {
t.Fatalf("GetTask: %v", err)
}
if got.Status != model.TaskStatusFailed {
t.Fatalf("前置熔断未生效: status=%s", got.Status)
}
active := model.TaskStatusActive
updated, err := tasks.UpdateTask(ctx, task.ID, UpdateTaskInput{Status: &active})
if err != nil {
t.Fatalf("重新启用: %v", err)
}
if updated.Status != model.TaskStatusActive {
t.Errorf("status = %s, want active", updated.Status)
}
if cnt := snatchAuthCount(t, updated.Payload); cnt != 0 {
t.Errorf("authFailCount = %d, want 0(重新启用应清零)", cnt)
}
if _, ok := tasks.entries[task.ID]; !ok {
t.Error("重新启用后任务未回到调度")
}
}
// TestSnatchSuccessVars 锁定抢机成功通知变量的组装:租户/区域/类型/镜像,
// 以及 root 密码三种凭据模式与 IP 逐级回退;单项查询失败不阻塞其余字段。
func TestSnatchSuccessVars(t *testing.T) {
tests := []struct {
name string
client *fakeClient
in oci.CreateInstanceInput
insts []oci.Instance
want map[string]string
}{
{
name: "生成密码模式全字段",
client: &fakeClient{
tenancy: oci.TenancyInfo{Name: "mytenancy", HomeRegionKey: "FRA"},
instance: oci.Instance{PublicIP: "203.0.113.7"},
image: oci.Image{DisplayName: "Ubuntu-24.04"},
},
in: oci.CreateInstanceInput{Region: "ap-osaka-1", Shape: "VM.Standard.A1.Flex",
Ocpus: 4, MemoryInGBs: 24, ImageID: "ocid1.image.oc1..img"},
insts: []oci.Instance{{ID: "i1", FreeformTags: map[string]string{"RootPassword": "pw-from-tag"}}},
want: map[string]string{
"tenant": "试用期",
"region": "Japan Central (Osaka)",
"shape": "VM.Standard.A1.Flex",
"spec": "4C/24G",
"image": "Ubuntu-24.04",
"ip": "203.0.113.7",
"root_password": "pw-from-tag",
},
},
{
name: "显式密码_未知区域_IP回退内网_镜像失败回退OCID",
client: &fakeClient{
instance: oci.Instance{PrivateIP: "10.0.0.5"},
imageErr: fmt.Errorf("boom"),
},
in: oci.CreateInstanceInput{Region: "xx-nowhere-1", Shape: "VM.Standard.E2.1.Micro",
ImageID: "ocid1.image.oc1..img", RootPassword: "explicit-pw"},
insts: []oci.Instance{{ID: "i1"}},
want: map[string]string{
"tenant": "试用期",
"region": "xx-nowhere-1",
"shape": "VM.Standard.E2.1.Micro",
"spec": "-",
"image": "ocid1.image.oc1..img",
"ip": "10.0.0.5(内网)",
"root_password": "explicit-pw",
},
},
{
name: "SSH模式_引导卷_区域取配置默认_多台待分配",
client: &fakeClient{instanceErr: fmt.Errorf("boom")},
in: oci.CreateInstanceInput{Shape: "s",
BootVolumeID: "ocid1.bootvolume.oc1..bv", SSHPublicKey: "ssh-ed25519 AAAA"},
insts: []oci.Instance{{ID: "i1"}, {ID: "i2"}},
want: map[string]string{
"region": "Germany Central (Frankfurt)",
"image": "引导卷",
"ip": "待分配、待分配",
"root_password": "—(SSH 密钥登录)、—(SSH 密钥登录)",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tasks, configs, _ := newTaskEnv(t, tt.client)
cfg, err := configs.Import(context.Background(), trialImportInput())
if err != nil {
t.Fatalf("import: %v", err)
}
got := tasks.snatchSuccessVars(context.Background(),
&snatchPayload{OciConfigID: cfg.ID}, tt.in, tt.insts)
for k, want := range tt.want {
if got[k] != want {
t.Errorf("%s = %q, want %q", k, got[k], want)
}
}
})
}
}
// TestNotifyEventsSnatchVarsMerged 抢满成功事件应合并暂存的补充变量。
func TestNotifyEventsSnatchVarsMerged(t *testing.T) {
prev := taskSnapshot{Name: "抢机", Status: model.TaskStatusActive}
cur := taskSnapshot{Name: "抢机", Status: model.TaskStatusSucceeded, Message: "created 1",
SnatchVars: map[string]string{"tenant": "免费01·t", "ip": "203.0.113.7", "root_password": "pw"}}
events := notifyEvents(prev, cur)
if len(events) != 1 || events[0].Kind != notifySnatchSuccess {
t.Fatalf("events = %+v, want 1 条 snatch_success", events)
}
want := map[string]string{"task_name": "抢机", "message": "created 1",
"tenant": "免费01·t", "ip": "203.0.113.7", "root_password": "pw"}
for k, v := range want {
if events[0].Vars[k] != v {
t.Errorf("vars[%s] = %q, want %q", k, events[0].Vars[k], v)
}
}
}