Files
oci-portal/internal/service/task.go
T

858 lines
28 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"context"
"encoding/json"
"fmt"
"strings"
"sync"
"time"
"github.com/robfig/cron/v3"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"oci-portal/internal/model"
"oci-portal/internal/oci"
)
// taskRunTimeout 是单次任务执行的超时时间。
const taskRunTimeout = 10 * time.Minute
// taskLogKeep 是每个任务保留的执行日志条数。
const taskLogKeep = 100
// TaskService 管理后台任务的存储、cron 调度与执行。
type TaskService struct {
db *gorm.DB
configs *OciConfigService
notifier *Notifier
settings *SettingService
cron *cron.Cron
// aiGateway 供 AI 探测任务执行渠道探测,由 main 装配(可为 nil)
aiGateway *AiGatewayService
mu sync.Mutex
entries map[uint]cron.EntryID
}
// NewTaskService 组装依赖;notifier 传 nil 表示整体关闭通知,
// settings 供发送前按事件类型过滤(nil 视为全开)。调用 Start 后开始调度。
func NewTaskService(db *gorm.DB, configs *OciConfigService, notifier *Notifier, settings *SettingService) *TaskService {
return &TaskService{
db: db,
configs: configs,
notifier: notifier,
settings: settings,
cron: cron.New(),
entries: map[uint]cron.EntryID{},
}
}
// AttachAiGateway 注入 AI 网关服务,启用 AI 探测任务的执行与自动同步。
func (s *TaskService) AttachAiGateway(gw *AiGatewayService) { s.aiGateway = gw }
// Start 加载全部 active 任务注册调度并启动 cron。
func (s *TaskService) Start() error {
var tasks []model.Task
if err := s.db.Where("status = ?", model.TaskStatusActive).Find(&tasks).Error; err != nil {
return fmt.Errorf("load active tasks: %w", err)
}
for i := range tasks {
if err := s.schedule(&tasks[i]); err != nil {
return err
}
}
s.cron.Start()
return nil
}
// Stop 停止调度并等待执行中的任务收尾,
// 保证任务产生的异步通知都已进入 Notifier 的等待队列。
func (s *TaskService) Stop() {
<-s.cron.Stop().Done()
}
// healthCheckPayload 是测活任务参数;ociConfigIds 为空表示全部配置。
type healthCheckPayload struct {
OciConfigIDs []uint `json:"ociConfigIds"`
}
// costPayload 是成本同步任务参数;ociConfigIds 为空表示全部配置,
// 免费类别的配置在执行时跳过,不发起 Usage API 请求。
type costPayload struct {
OciConfigIDs []uint `json:"ociConfigIds"`
}
// snatchPayload 是抢机任务参数;count 为剩余台数(创建时即目标台数,
// 每次执行把剩余写回,直到抢满),totalCount 固定为创建时的目标台数,
// 供前端计算进度(旧任务缺省时前端回退用 count)。
// authFailCount 为连续 NotAuthenticated 失败计数,达阈值任务熔断停止。
type snatchPayload struct {
OciConfigID uint `json:"ociConfigId"`
Count int `json:"count"`
TotalCount int `json:"totalCount,omitempty"`
AuthFailCount int `json:"authFailCount,omitempty"`
Instance oci.CreateInstanceInput `json:"instance"`
}
// CreateTaskInput 是创建任务的输入。
type CreateTaskInput struct {
Name string
Type string
CronExpr string
Payload json.RawMessage
}
// CreateTask 校验并保存任务,立即进入调度;AI 探测任务全局唯一(系统自动管理)。
func (s *TaskService) CreateTask(ctx context.Context, in CreateTaskInput) (*model.Task, error) {
if in.Name == "" {
return nil, fmt.Errorf("create task: name is required")
}
if _, err := cron.ParseStandard(in.CronExpr); err != nil {
return nil, fmt.Errorf("create task: invalid cron %q: %w", in.CronExpr, err)
}
if err := validateTaskPayload(in.Type, in.Payload); err != nil {
return nil, err
}
if err := s.ensureAiProbeUnique(ctx, in.Type); err != nil {
return nil, err
}
task := &model.Task{
Name: in.Name,
Type: in.Type,
CronExpr: in.CronExpr,
Payload: string(normalizeSnatchPayload(in.Type, in.Payload)),
Status: model.TaskStatusActive,
}
if err := s.db.WithContext(ctx).Create(task).Error; err != nil {
return nil, fmt.Errorf("create task: %w", err)
}
if err := s.schedule(task); err != nil {
return nil, err
}
return task, nil
}
// ensureAiProbeUnique 拒绝重复创建 AI 探测任务(该类型随渠道数量自动管理)。
func (s *TaskService) ensureAiProbeUnique(ctx context.Context, taskType string) error {
if taskType != model.TaskTypeAiProbe {
return nil
}
var n int64
if err := s.db.WithContext(ctx).Model(&model.Task{}).
Where("type = ?", model.TaskTypeAiProbe).Count(&n).Error; err != nil {
return err
}
if n > 0 {
return fmt.Errorf("create task: AI 探测任务已存在,由系统自动管理")
}
return nil
}
// normalizeSnatchPayload 给抢机 payload 补全目标台数:count 默认 1、
// totalCount 缺省时固定为创建时的 count,供进度展示;解析失败原样返回
// validateTaskPayload 已在前面拦截非法 JSON)。
func normalizeSnatchPayload(taskType string, payload json.RawMessage) json.RawMessage {
if taskType != model.TaskTypeSnatch {
return payload
}
var p snatchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return payload
}
if p.Count <= 0 {
p.Count = 1
}
if p.TotalCount <= 0 {
p.TotalCount = p.Count
}
out, err := json.Marshal(p)
if err != nil {
return payload
}
return out
}
// validateTaskPayload 按任务类型校验参数 JSON。
func validateTaskPayload(taskType string, payload json.RawMessage) error {
switch taskType {
case model.TaskTypeHealthCheck:
var p healthCheckPayload
if len(payload) > 0 {
if err := json.Unmarshal(payload, &p); err != nil {
return fmt.Errorf("create task: invalid payload: %w", err)
}
}
return nil
case model.TaskTypeCost:
var p costPayload
if len(payload) > 0 {
if err := json.Unmarshal(payload, &p); err != nil {
return fmt.Errorf("create task: invalid payload: %w", err)
}
}
return nil
case model.TaskTypeSnatch:
var p snatchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return fmt.Errorf("create task: invalid payload: %w", err)
}
if p.OciConfigID == 0 {
return fmt.Errorf("create task: snatch payload requires ociConfigId")
}
return validateCreateInstance(p.Instance)
case model.TaskTypeAiProbe:
return nil // 无参数:探测全部渠道
default:
return fmt.Errorf("create task: unsupported type %q", taskType)
}
}
// UpdateTaskInput 是更新任务的输入;nil 字段不修改。
type UpdateTaskInput struct {
Name *string
CronExpr *string
Payload json.RawMessage
Status *string
}
// UpdateTask 修改任务并重新调度。
func (s *TaskService) UpdateTask(ctx context.Context, id uint, in UpdateTaskInput) (*model.Task, error) {
task, err := s.GetTask(ctx, id)
if err != nil {
return nil, err
}
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)
}
s.unschedule(task.ID)
if task.Status == model.TaskStatusActive {
if err := s.schedule(task); err != nil {
return nil, err
}
}
return task, nil
}
func applyTaskUpdate(task *model.Task, in UpdateTaskInput) error {
if in.Name != nil {
task.Name = *in.Name
}
if in.CronExpr != nil {
if _, err := cron.ParseStandard(*in.CronExpr); err != nil {
return fmt.Errorf("update task: invalid cron %q: %w", *in.CronExpr, err)
}
task.CronExpr = *in.CronExpr
}
if len(in.Payload) > 0 {
if err := validateTaskPayload(task.Type, in.Payload); err != nil {
return err
}
task.Payload = string(in.Payload)
}
if in.Status != nil {
if *in.Status != model.TaskStatusActive && *in.Status != model.TaskStatusPaused {
return fmt.Errorf("update task: status must be active or paused")
}
if task.Status == model.TaskStatusFailed && *in.Status == model.TaskStatusActive {
resetSnatchAuthFail(task)
}
task.Status = *in.Status
}
return nil
}
// resetSnatchAuthFail 清零抢机 payload 的连续鉴权失败计数,
// 供 failed 任务重新启用时调用,避免一恢复调度就再次熔断;非抢机任务不处理。
func resetSnatchAuthFail(task *model.Task) {
if task.Type != model.TaskTypeSnatch {
return
}
var p snatchPayload
if err := json.Unmarshal([]byte(task.Payload), &p); err != nil {
return
}
p.AuthFailCount = 0
writeSnatchPayload(task, &p)
}
// ListTasks 返回全部任务。
func (s *TaskService) ListTasks(ctx context.Context) ([]model.Task, error) {
tasks := make([]model.Task, 0)
if err := s.db.WithContext(ctx).Order("id").Find(&tasks).Error; err != nil {
return nil, fmt.Errorf("list tasks: %w", err)
}
return tasks, nil
}
// GetTask 返回单个任务。
func (s *TaskService) GetTask(ctx context.Context, id uint) (*model.Task, error) {
var task model.Task
if err := s.db.WithContext(ctx).First(&task, id).Error; err != nil {
return nil, fmt.Errorf("find task %d: %w", id, err)
}
return &task, nil
}
// DeleteTask 注销调度并删除任务与其日志;AI 探测任务由系统按渠道数量
// 自动创建/删除,拒绝手动删除。
func (s *TaskService) DeleteTask(ctx context.Context, id uint) error {
task, err := s.GetTask(ctx, id)
if err != nil {
return err
}
if task.Type == model.TaskTypeAiProbe {
return fmt.Errorf("delete task: AI 探测任务由系统自动管理,删除最后一个渠道时自动移除")
}
return s.removeTask(ctx, id)
}
// removeTask 注销调度并删除任务与其日志(内部路径,不做类型限制)。
func (s *TaskService) removeTask(ctx context.Context, id uint) error {
s.unschedule(id)
if err := s.db.WithContext(ctx).Delete(&model.Task{}, id).Error; err != nil {
return fmt.Errorf("delete task %d: %w", id, err)
}
if err := s.db.WithContext(ctx).Where("task_id = ?", id).Delete(&model.TaskLog{}).Error; err != nil {
return fmt.Errorf("delete task %d logs: %w", id, err)
}
return nil
}
// TaskLogs 返回任务最近的执行日志(时间倒序)。
func (s *TaskService) TaskLogs(ctx context.Context, id uint, limit int) ([]model.TaskLog, error) {
if limit <= 0 || limit > taskLogKeep {
limit = 50
}
logs := make([]model.TaskLog, 0)
err := s.db.WithContext(ctx).Where("task_id = ?", id).
Order("id desc").Limit(limit).Find(&logs).Error
if err != nil {
return nil, fmt.Errorf("list task %d logs: %w", id, err)
}
return logs, nil
}
// RunTaskNow 立即执行一次任务并返回本次日志。
func (s *TaskService) RunTaskNow(ctx context.Context, id uint) (*model.TaskLog, error) {
if _, err := s.GetTask(ctx, id); err != nil {
return nil, err
}
return s.execute(id), nil
}
// schedule 把任务注册进 cron 调度。
func (s *TaskService) schedule(task *model.Task) error {
s.mu.Lock()
defer s.mu.Unlock()
if _, ok := s.entries[task.ID]; ok {
return nil
}
taskID := task.ID
entry, err := s.cron.AddFunc(task.CronExpr, func() { s.execute(taskID) })
if err != nil {
return fmt.Errorf("schedule task %d: %w", task.ID, err)
}
s.entries[task.ID] = entry
return nil
}
// unschedule 把任务移出 cron 调度。
func (s *TaskService) unschedule(taskID uint) {
s.mu.Lock()
defer s.mu.Unlock()
if entry, ok := s.entries[taskID]; ok {
s.cron.Remove(entry)
delete(s.entries, taskID)
}
}
// execute 执行一次任务:加载 → 分派 → 更新任务状态并写日志,
// 前后状态交给通知判定,只在状态变化时推送。
func (s *TaskService) execute(taskID uint) *model.TaskLog {
ctx, cancel := context.WithTimeout(context.Background(), taskRunTimeout)
defer cancel()
task, err := s.GetTask(ctx, taskID)
if err != nil {
return nil
}
prev := taskSnapshot{Name: task.Name, Status: task.Status, LastError: task.LastError}
start := time.Now()
message, runErr := s.run(ctx, task)
now := time.Now()
task.LastRunAt = &now
task.RunCount++
task.LastError = ""
if runErr != nil {
task.LastError = oci.CompactError(runErr)
message = task.LastError
}
s.db.Save(task)
cur := taskSnapshot{Name: task.Name, Status: task.Status, LastError: task.LastError, Message: message}
s.notify(notifyEvents(prev, cur))
return s.appendLog(task.ID, runErr == nil, message, time.Since(start))
}
// run 按类型分派任务执行。
func (s *TaskService) run(ctx context.Context, task *model.Task) (string, error) {
switch task.Type {
case model.TaskTypeHealthCheck:
return s.runHealthCheck(ctx, task)
case model.TaskTypeCost:
return s.runCost(ctx, task)
case model.TaskTypeSnatch:
return s.runSnatch(ctx, task)
case model.TaskTypeAiProbe:
return s.runAiProbe(ctx)
default:
return "", fmt.Errorf("unsupported task type %q", task.Type)
}
}
// runAiProbe 逐渠道探测 AI 网关号池;网关未装配时报错。
func (s *TaskService) runAiProbe(ctx context.Context) (string, error) {
if s.aiGateway == nil {
return "", fmt.Errorf("ai gateway not attached")
}
msg, err := s.aiGateway.ProbeAll(ctx)
if err == nil {
s.warnDeprecatingModels(ctx)
}
return msg, err
}
// warnDeprecatingModels 对 30 天内即将退役或弃用的在池模型发 Telegram 提醒;
// 随每日探测执行,模型退役被同步剔除后自动停止,受通知管理 model_deprecated 开关控制。
func (s *TaskService) warnDeprecatingModels(ctx context.Context) {
if s.notifier == nil {
return
}
if s.settings != nil && !s.settings.NotifyEventEnabled(ctx, "model_deprecated") {
return
}
names, err := s.aiGateway.DeprecatingModels(ctx, 30*24*time.Hour)
if err != nil || len(names) == 0 {
return
}
s.notifier.SendTemplateAsync("model_deprecated", map[string]string{"models": strings.Join(names, "\n")})
}
// runHealthCheck 对范围内的配置逐个测活,汇总结果并触发失联通知。
func (s *TaskService) runHealthCheck(ctx context.Context, task *model.Task) (string, error) {
var p healthCheckPayload
if task.Payload != "" {
if err := json.Unmarshal([]byte(task.Payload), &p); err != nil {
return "", fmt.Errorf("parse payload: %w", err)
}
}
ids, err := s.targetConfigIDs(ctx, p.OciConfigIDs)
if err != nil {
return "", err
}
alive := 0
var deadAliases, failures []string
for _, id := range ids {
cfg, _, err := s.configs.Verify(ctx, id)
ok := err == nil && cfg.AliveStatus == model.AliveStatusAlive
s.saveCheckSnapshot(ctx, id, ok)
if !ok {
deadAliases = append(deadAliases, configAlias(cfg, id))
failures = append(failures, fmt.Sprintf("#%d %s", id, verifyFailReason(cfg, err)))
continue
}
alive++
}
s.notifyDeadAliases(ctx, task.ID, deadAliases)
msg := fmt.Sprintf("checked %d: %d alive, %d dead", len(ids), alive, len(deadAliases))
if len(failures) > 0 {
msg += "; " + strings.Join(failures, "; ")
}
return msg, nil
}
// configAlias 返回配置别名,配置加载失败时退回 #ID 表示。
func configAlias(cfg *model.OciConfig, id uint) string {
if cfg != nil && cfg.Alias != "" {
return cfg.Alias
}
return fmt.Sprintf("#%d", id)
}
func verifyFailReason(cfg *model.OciConfig, err error) string {
if err != nil {
return oci.CompactError(err)
}
return cfg.LastError
}
// saveCheckSnapshot 覆盖写入测活快照;仅存活时刷新实例数(默认区域口径),
// 失联时保留上次实例数,避免总览 KPI 因 key 失效而抖动。
func (s *TaskService) saveCheckSnapshot(ctx context.Context, cfgID uint, alive bool) {
status := model.AliveStatusDead
if alive {
status = model.AliveStatusAlive
}
snap := model.CheckSnapshot{OciConfigID: cfgID, AliveStatus: status, CheckedAt: time.Now()}
cols := []string{"alive_status", "checked_at"}
if alive {
if instances, err := s.configs.Instances(ctx, cfgID, "", ""); err == nil {
snap.InstanceCount = len(instances)
cols = append(cols, "instance_count")
}
}
s.db.WithContext(ctx).Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "oci_config_id"}},
DoUpdates: clause.AssignmentColumns(cols),
}).Create(&snap)
}
// runCost 对范围内配置同步近 7 天每日成本快照,免费类别跳过。
func (s *TaskService) runCost(ctx context.Context, task *model.Task) (string, error) {
var p costPayload
if task.Payload != "" {
if err := json.Unmarshal([]byte(task.Payload), &p); err != nil {
return "", fmt.Errorf("parse payload: %w", err)
}
}
ids, err := s.targetConfigIDs(ctx, p.OciConfigIDs)
if err != nil {
return "", err
}
synced, skipped := 0, 0
var failures []string
for _, id := range ids {
switch err := s.syncCostSnapshot(ctx, id); {
case err == errFreeAccountSkipped:
skipped++
case err != nil:
failures = append(failures, fmt.Sprintf("#%d %v", id, err))
default:
synced++
}
}
msg := fmt.Sprintf("synced usage for %d tenants, skipped %d free", synced, skipped)
if len(failures) > 0 {
msg += "; " + strings.Join(failures, "; ")
}
return msg, nil
}
// errFreeAccountSkipped 标记成本同步因免费类别被跳过。
var errFreeAccountSkipped = fmt.Errorf("free account skipped")
// syncCostSnapshot 拉取单配置近 7 天每日成本并按天覆盖写入快照。
func (s *TaskService) syncCostSnapshot(ctx context.Context, cfgID uint) error {
cfg, err := s.configs.Get(ctx, cfgID)
if err != nil {
return err
}
if cfg.AccountType == model.AccountTypeFree {
return errFreeAccountSkipped
}
end := time.Now().UTC()
items, err := s.configs.Costs(ctx, cfgID, oci.CostQuery{
StartTime: end.AddDate(0, 0, -7),
EndTime: end,
})
if err != nil {
return err
}
return s.saveCostSnapshots(ctx, cfgID, items)
}
// saveCostSnapshots 把成本条目按 UTC 日聚合后逐日 upsert。
func (s *TaskService) saveCostSnapshots(ctx context.Context, cfgID uint, items []oci.CostItem) error {
type bucket struct {
amount float64
currency string
}
byDay := map[string]*bucket{}
for _, item := range items {
if item.TimeStart == nil {
continue
}
day := item.TimeStart.UTC().Format("2006-01-02")
b, ok := byDay[day]
if !ok {
b = &bucket{currency: item.Currency}
byDay[day] = b
}
b.amount += float64(item.ComputedAmount)
}
now := time.Now()
for day, b := range byDay {
snap := model.CostSnapshot{
OciConfigID: cfgID, Day: day,
Amount: b.amount, Currency: b.currency, SyncedAt: now,
}
err := s.db.WithContext(ctx).Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "oci_config_id"}, {Name: "day"}},
DoUpdates: clause.AssignmentColumns([]string{"amount", "currency", "synced_at"}),
}).Create(&snap).Error
if err != nil {
return fmt.Errorf("save cost snapshot %s: %w", day, err)
}
}
return nil
}
// targetConfigIDs 解析任务作用范围;未指定时返回全部配置 ID。
func (s *TaskService) targetConfigIDs(ctx context.Context, ids []uint) ([]uint, error) {
if len(ids) > 0 {
return ids, nil
}
configs, err := s.configs.List(ctx)
if err != nil {
return nil, err
}
all := make([]uint, 0, len(configs))
for _, cfg := range configs {
all = append(all, cfg.ID)
}
return all, nil
}
// runSnatch 尝试创建实例;抢到目标台数后任务标记 succeeded 并停止调度,
// 部分成功把剩余台数写回 payload 下次继续;成功路径一并清零并写回连续
// 鉴权失败计数,失败路径交给 snatchFailure 做连续 NotAuthenticated 熔断
// 判定。字段落库由 execute 统一 Save。
func (s *TaskService) runSnatch(ctx context.Context, task *model.Task) (string, error) {
var p snatchPayload
if err := json.Unmarshal([]byte(task.Payload), &p); err != nil {
return "", fmt.Errorf("parse payload: %w", err)
}
if p.Count < 1 {
p.Count = 1
}
in, adNote, err := s.snatchInstanceInput(ctx, task, &p)
if err != nil {
return "", s.snatchFailure(ctx, task, &p, err)
}
instances, failures, err := s.configs.CreateInstances(ctx, p.OciConfigID, in, p.Count)
if err == nil && len(instances) == 0 {
err = fmt.Errorf("no instance created%s: %s", adNote, strings.Join(failures, "; "))
}
if err != nil {
return "", s.snatchFailure(ctx, task, &p, err)
}
p.AuthFailCount = 0
ids := make([]string, 0, len(instances))
for _, in := range instances {
ids = append(ids, in.ID)
}
remaining := p.Count - len(instances)
if remaining > 0 {
p.Count = remaining
writeSnatchPayload(task, &p)
return fmt.Sprintf("created %d (%s)%s, %d remaining", len(instances), strings.Join(ids, ","), adNote, remaining), nil
}
task.Status = model.TaskStatusSucceeded
writeSnatchPayload(task, &p)
s.unschedule(task.ID)
return fmt.Sprintf("created %d: %s%s", len(instances), strings.Join(ids, ","), adNote), nil
}
// snatchInstanceInput 组装本次创建参数:可用域显式指定时原样使用;
// 留空(自动)时按执行序号轮询区域全部可用域——ad-1、ad-2、ad-3 依次循环,
// 分摊单可用域容量不足。附加说明串供执行日志展示本次所用可用域。
func (s *TaskService) snatchInstanceInput(ctx context.Context, task *model.Task, p *snatchPayload) (oci.CreateInstanceInput, string, error) {
in := p.Instance
if in.AvailabilityDomain != "" {
return in, "", nil
}
ads, err := s.configs.AvailabilityDomains(ctx, p.OciConfigID, in.Region)
if err != nil {
return in, "", fmt.Errorf("list availability domains: %w", err)
}
if len(ads) == 0 {
return in, "", fmt.Errorf("region has no availability domain")
}
// execute 在 run 之后才递增 RunCount,此处即 0 起的本次执行序号
in.AvailabilityDomain = ads[task.RunCount%len(ads)]
return in, " @ " + in.AvailabilityDomain, nil
}
// snatchFailure 处理抢机单次失败:错误含 NotAuthenticated 时累计连续计数,
// 达阈值把任务置 failed 并移出调度(熔断);其他错误清零计数。计数写回 payload。
func (s *TaskService) snatchFailure(ctx context.Context, task *model.Task, p *snatchPayload, cause error) error {
if !strings.Contains(cause.Error(), "NotAuthenticated") {
p.AuthFailCount = 0
writeSnatchPayload(task, p)
return cause
}
p.AuthFailCount++
writeSnatchPayload(task, p)
if p.AuthFailCount < s.snatchAuthFailLimit(ctx) {
return cause
}
task.Status = model.TaskStatusFailed
s.unschedule(task.ID)
return fmt.Errorf("连续 %d 次 NotAuthenticated,任务已熔断停止: %w", p.AuthFailCount, cause)
}
// snatchAuthFailLimit 读取抢机熔断阈值;settings 未注入或读取失败按默认值。
func (s *TaskService) snatchAuthFailLimit(ctx context.Context) int {
if s.settings == nil {
return defaultSnatchAuthFailLimit
}
view, err := s.settings.TaskSettings(ctx)
if err != nil {
return defaultSnatchAuthFailLimit
}
return view.SnatchAuthFailLimit
}
// writeSnatchPayload 把最新抢机参数序列化写回任务(execute 统一落库)。
func writeSnatchPayload(task *model.Task, p *snatchPayload) {
if raw, err := json.Marshal(p); err == nil {
task.Payload = string(raw)
}
}
// appendLog 写入执行日志并裁剪超出保留数量的旧日志。
func (s *TaskService) appendLog(taskID uint, success bool, message string, elapsed time.Duration) *model.TaskLog {
entry := &model.TaskLog{
TaskID: taskID,
Success: success,
Message: message,
DurationMs: elapsed.Milliseconds(),
}
s.db.Create(entry)
s.db.Where("task_id = ? AND id NOT IN (?)", taskID,
s.db.Model(&model.TaskLog{}).Select("id").Where("task_id = ?", taskID).
Order("id desc").Limit(taskLogKeep),
).Delete(&model.TaskLog{})
return entry
}
// taskSnapshot 是通知判定所需的任务状态切片(执行前后各取一份)。
type taskSnapshot struct {
Name string
Status string
LastError string
Message string // 本次执行结果摘要,仅执行后快照填写
}
// notifyKind 是通知事件类型,与设置页「通知管理」开关一一对应。
type notifyKind string
const (
notifyTaskFail notifyKind = "task_fail"
notifyTaskRecover notifyKind = "task_recover"
notifySnatchSuccess notifyKind = "snatch_success"
notifyTenantDead notifyKind = "tenant_dead"
notifyTaskStop notifyKind = "task_stop" // 任务熔断停止(抢机连续鉴权失败达阈值)
)
// notifyEvent 是一条待发送的通知:类型供开关过滤与模板选择,Vars 为模板变量。
type notifyEvent struct {
Kind notifyKind
Vars map[string]string
}
// notifyEvents 比较执行前后的任务状态,返回需要推送的通知事件。
// 只在状态发生变化时产出:连续失败或持续正常都不重复发,防轰炸。
// 熔断翻转(置 failed)优先判定,该次只发任务停止、不叠加任务失败。
func notifyEvents(prev, cur taskSnapshot) []notifyEvent {
var events []notifyEvent
if prev.Status != model.TaskStatusSucceeded && cur.Status == model.TaskStatusSucceeded {
events = append(events, notifyEvent{notifySnatchSuccess, map[string]string{"task_name": cur.Name, "message": cur.Message}})
}
switch {
case prev.Status != model.TaskStatusFailed && cur.Status == model.TaskStatusFailed:
events = append(events, notifyEvent{notifyTaskStop, map[string]string{"task_name": cur.Name, "error": cur.LastError}})
case prev.LastError == "" && cur.LastError != "":
events = append(events, notifyEvent{notifyTaskFail, map[string]string{"task_name": cur.Name, "error": cur.LastError}})
case prev.LastError != "" && cur.LastError == "" && cur.Status != model.TaskStatusSucceeded:
// 恢复即成功收尾(抢机达成目标)时已有抢机成功通知,不再叠加恢复通知
events = append(events, notifyEvent{notifyTaskRecover, map[string]string{"task_name": cur.Name}})
}
return events
}
// notifyFilterTimeout 是发送前查询事件开关的超时时间(本地 SQLite,查询极快)。
const notifyFilterTimeout = 5 * time.Second
// notify 逐条按事件开关过滤后异步发送;notifier 未注入(nil)时整体关闭。
// 开关读取失败按开启降级(NotifyEventEnabled 内部兜底),不因设置异常漏发。
func (s *TaskService) notify(events []notifyEvent) {
if s.notifier == nil || len(events) == 0 {
return
}
ctx, cancel := context.WithTimeout(context.Background(), notifyFilterTimeout)
defer cancel()
for _, ev := range events {
if s.settings != nil && !s.settings.NotifyEventEnabled(ctx, string(ev.Kind)) {
continue
}
s.notifier.SendTemplateAsync(string(ev.Kind), ev.Vars)
}
}
// deadAliasKey 是测活任务留存上次失联别名集合的 Setting 键。
func deadAliasKey(taskID uint) string {
return fmt.Sprintf("health_dead_alias:%d", taskID)
}
// notifyDeadAliases 只在失联集合发生变化时推送失联通知,并留存本次集合;
// 集合不变(含持续失联)不重复发。notifier 未注入时整体跳过。
func (s *TaskService) notifyDeadAliases(ctx context.Context, taskID uint, aliases []string) {
if s.notifier == nil {
return
}
if sameStringSet(s.loadDeadAliases(ctx, taskID), aliases) {
return
}
s.saveDeadAliases(ctx, taskID, aliases)
if len(aliases) > 0 {
s.notify([]notifyEvent{{notifyTenantDead, map[string]string{"tenants": strings.Join(aliases, "、")}}})
}
}
// loadDeadAliases 读取任务上次记录的失联别名集合;无记录视为空集。
func (s *TaskService) loadDeadAliases(ctx context.Context, taskID uint) []string {
var st model.Setting
err := s.db.WithContext(ctx).First(&st, "key = ?", deadAliasKey(taskID)).Error
if err != nil || st.Value == "" {
return nil
}
var aliases []string
if err := json.Unmarshal([]byte(st.Value), &aliases); err != nil {
return nil
}
return aliases
}
// saveDeadAliases 覆盖保存本次失联别名集合;留存失败不影响任务执行。
func (s *TaskService) saveDeadAliases(ctx context.Context, taskID uint, aliases []string) {
raw, err := json.Marshal(aliases)
if err != nil {
return
}
s.db.WithContext(ctx).Save(&model.Setting{
Key: deadAliasKey(taskID), Value: string(raw), UpdatedAt: time.Now(),
})
}
// sameStringSet 判断两个字符串切片内容是否相同(忽略顺序,重复元素按次数计)。
func sameStringSet(a, b []string) bool {
if len(a) != len(b) {
return false
}
count := make(map[string]int, len(a))
for _, v := range a {
count[v]++
}
for _, v := range b {
count[v]--
if count[v] < 0 {
return false
}
}
return true
}