diff --git a/go.mod b/go.mod index 091894dc..0bc872bf 100644 --- a/go.mod +++ b/go.mod @@ -60,6 +60,7 @@ require ( github.com/tufanbarisyildirim/gonginx v0.0.0-20260220081509-8e17ce617db3 github.com/urfave/cli/v3 v3.10.1 github.com/valyala/fastjson v1.6.10 + github.com/wneessen/go-mail v0.8.1 github.com/xuri/excelize/v2 v2.11.0 go.yaml.in/yaml/v4 v4.0.0-rc.6 golang.org/x/crypto v0.54.0 diff --git a/go.sum b/go.sum index 0fb5968e..14aa1171 100644 --- a/go.sum +++ b/go.sum @@ -421,6 +421,8 @@ github.com/urfave/cli/v3 v3.10.1 h1:7Kx9H50hrHbRbyxgO1KP6/BcbiGRz0uYh5YyQ30JEEY= github.com/urfave/cli/v3 v3.10.1/go.mod h1:ysVLtOEmg2tOy6PknnYVhDoouyC/6N42TMeoMzskhso= github.com/valyala/fastjson v1.6.10 h1:/yjJg8jaVQdYR3arGxPE2X5z89xrlhS0eGXdv+ADTh4= github.com/valyala/fastjson v1.6.10/go.mod h1:e6FubmQouUNP73jtMLmcbxS6ydWIpOfhz34TSfO3JaE= +github.com/wneessen/go-mail v0.8.1 h1:tVcncj02/QySVFw3zr/kXOzZcuFQqBNT6K+Rbgm/pcM= +github.com/wneessen/go-mail v0.8.1/go.mod h1:dWZ61zadzCIyvB4y1/YzC5O7MrbbzBfPkARmbosdf8w= github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM= github.com/x448/float16 v0.8.4/go.mod h1:14CWIYCyZA/cWjXOioeEpHeN/83MdbZDRQHoFcYsOfg= github.com/xiang90/probing v0.0.0-20190116061207-43a291ad63a2/go.mod h1:UETIi67q53MR2AWcXfiuqkDkRtnGDLqkBTpCHuJHxtU= diff --git a/internal/apps/mysql/app.go b/internal/apps/mysql/app.go index cfceabbd..c8aab530 100644 --- a/internal/apps/mysql/app.go +++ b/internal/apps/mysql/app.go @@ -223,7 +223,7 @@ func (s *App) SetRootPassword(w http.ResponseWriter, r *http.Request) { } oldRootPassword, _ := s.settingRepo.Get(biz.SettingKeyMySQLRootPassword) - mysql, err := db.NewMySQL("root", oldRootPassword, s.getSock(), "unix") + mysql, err := db.NewMySQL(r.Context(), "root", oldRootPassword, s.getSock(), "unix") if err != nil { // 尝试安全模式直接改密 if err = db.MySQLResetRootPassword(req.Password); err != nil { diff --git a/internal/apps/pgadmin/app.go b/internal/apps/pgadmin/app.go index 8a7bf48a..fc84fa50 100644 --- a/internal/apps/pgadmin/app.go +++ b/internal/apps/pgadmin/app.go @@ -1,6 +1,7 @@ package pgadmin import ( + "context" "database/sql" "encoding/json" "errors" @@ -247,8 +248,8 @@ func (s *App) dumpExistingServers(email string) map[string]struct{} { // syncServers 将面板中全部 PostgreSQL 服务器合并注册到 pgAdmin,凭据写入 pgpass 实现免密 // 仅追加 pgAdmin 中缺失的服务器,不影响用户在 pgAdmin 中手动添加的内容 -func (s *App) syncServers(email string) error { - servers, _, err := s.databaseServerRepo.List(1, 10000, string(biz.DatabaseTypePostgresql)) +func (s *App) syncServers(ctx context.Context, email string) error { + servers, _, err := s.databaseServerRepo.List(ctx, 1, 10000, string(biz.DatabaseTypePostgresql)) if err != nil { return err } @@ -358,7 +359,7 @@ func (s *App) Login(w http.ResponseWriter, r *http.Request) { return } - if err = s.syncServers(email); err != nil { + if err = s.syncServers(r.Context(), email); err != nil { service.Error(w, http.StatusInternalServerError, s.t.Get("failed to sync servers to pgAdmin: %v", err)) return } diff --git a/internal/apps/phpmyadmin/app.go b/internal/apps/phpmyadmin/app.go index 40b6f250..859e2b13 100644 --- a/internal/apps/phpmyadmin/app.go +++ b/internal/apps/phpmyadmin/app.go @@ -124,7 +124,7 @@ func (s *App) Login(w http.ResponseWriter, r *http.Request) { return } - server, err := s.databaseServerRepo.Get(req.ServerID) + server, err := s.databaseServerRepo.Get(r.Context(), req.ServerID) if err != nil { service.Error(w, http.StatusInternalServerError, "%v", err) return diff --git a/internal/apps/postgresql/app.go b/internal/apps/postgresql/app.go index 30aa2412..277b46dc 100644 --- a/internal/apps/postgresql/app.go +++ b/internal/apps/postgresql/app.go @@ -214,7 +214,7 @@ func (s *App) SetPostgresPassword(w http.ResponseWriter, r *http.Request) { oldPassword, _ := s.settingRepo.Get(biz.SettingKeyPostgresPassword) port := s.getPort() - postgres, err := db.NewPostgres("postgres", oldPassword, "127.0.0.1", port) + postgres, err := db.NewPostgres(r.Context(), "postgres", oldPassword, "127.0.0.1", port) if err != nil { // 直接修改密码 if _, err = shell.Execf(`su - postgres -c "psql -p %d -c \"ALTER USER postgres WITH PASSWORD '%s';\""`, port, req.Password); err != nil { diff --git a/internal/biz/alert.go b/internal/biz/alert.go new file mode 100644 index 00000000..998533ca --- /dev/null +++ b/internal/biz/alert.go @@ -0,0 +1,985 @@ +package biz + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "log/slog" + "slices" + "strings" + "sync" + "time" + + "github.com/leonelquinteros/gotext" + "github.com/samber/do/v2" + "github.com/samber/lo" + lop "github.com/samber/lo/parallel" + "github.com/spf13/cast" + + "github.com/acepanel/panel/v3/internal/app" + "github.com/acepanel/panel/v3/internal/request" + "github.com/acepanel/panel/v3/pkg/apploader" + "github.com/acepanel/panel/v3/pkg/shell" + "github.com/acepanel/panel/v3/pkg/sshlog" + "github.com/acepanel/panel/v3/pkg/systemctl" + "github.com/acepanel/panel/v3/pkg/tools" + "github.com/acepanel/panel/v3/pkg/types" +) + +// 告警指标类型 +const ( + AlertTypeCPU = "cpu" // CPU 使用率 % + AlertTypeMemory = "memory" // 内存使用率 % + AlertTypeSwap = "swap" // Swap 使用率 % + AlertTypeLoad1 = "load1" // 1 分钟平均负载 + AlertTypeLoad5 = "load5" // 5 分钟平均负载 + AlertTypeLoad15 = "load15" // 15 分钟平均负载 + AlertTypeDisk = "disk" // 磁盘使用率 %,目标为挂载点 + AlertTypeDiskInode = "disk_inode" // 磁盘 inode 使用率 %,目标为挂载点 + AlertTypeDiskRead = "disk_read" // 磁盘读取速率 MB/s,目标为设备名 + AlertTypeDiskWrite = "disk_write" // 磁盘写入速率 MB/s,目标为设备名 + AlertTypeNetIn = "net_in" // 网卡下行速率 MB/s,目标为网卡名 + AlertTypeNetOut = "net_out" // 网卡上行速率 MB/s,目标为网卡名 + AlertTypeWebsite5xx = "website_5xx" // 网站本小时 5xx 次数,目标为网站名 + AlertTypeWebsiteError = "website_error" // 网站本小时错误率 %,目标为网站名 + AlertTypeService = "service" // 服务未运行,目标为服务名 + AlertTypeProject = "project" // 项目未运行,目标为项目名 + AlertTypeContainer = "container" // 容器未运行,目标为容器名 + AlertTypeApp = "app" // 应用未运行,目标为应用标识 + AlertTypeDatabase = "database" // 数据库服务器不可达,目标为服务器名 + AlertTypeCertExpire = "cert_expire" // 证书剩余天数,目标为域名 + AlertTypeWebsiteExpire = "website_expire" // 网站剩余天数,目标为网站名 +) + +// statusAlertTypes 状态类指标,语义固定为「不在运行」,不需要运算符与阈值 +var statusAlertTypes = []string{AlertTypeService, AlertTypeProject, AlertTypeContainer, AlertTypeApp, AlertTypeDatabase} + +// IsStatusAlert 是否为状态类指标 +func IsStatusAlert(typ string) bool { + return slices.Contains(statusAlertTypes, typ) +} + +const ( + // alertRetryDelay 通知发送失败后的重试间隔 + alertRetryDelay = 5 * time.Minute + // sshFailThreshold 单次检查窗口内触发爆破告警的失败次数 + sshFailThreshold uint = 5 + // sshFailSilence 同一来源两次爆破告警的最小间隔 + sshFailSilence = 30 * time.Minute +) + +// 告警比较运算符 +const ( + AlertOperatorGT = "gt" + AlertOperatorGTE = "gte" + AlertOperatorLT = "lt" + AlertOperatorLTE = "lte" +) + +// AlertRule 告警规则 +type AlertRule struct { + ID uint `gorm:"primaryKey" json:"id"` + Name string `gorm:"not null;default:''" json:"name"` + Type string `gorm:"not null;default:''" json:"type"` + Target string `gorm:"not null;default:''" json:"target"` // 挂载点/网卡/磁盘/服务名/域名,空表示全部 + Operator string `gorm:"not null;default:'gt'" json:"operator"` // gt/gte/lt/lte + Threshold float64 `gorm:"not null;default:0" json:"threshold"` + Duration uint `gorm:"not null;default:1" json:"duration"` // 连续满足次数 + Silence uint `gorm:"not null;default:30" json:"silence"` // 静默期(分钟) + Channels []uint `gorm:"serializer:json" json:"channels"` // 通知渠道 ID + Enabled bool `gorm:"not null;default:true" json:"enabled"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// Alert 告警记录 +type Alert struct { + ID uint `gorm:"primaryKey" json:"id"` + RuleID uint `gorm:"not null;default:0;index" json:"rule_id"` + RuleName string `gorm:"not null;default:''" json:"rule_name"` + Type string `gorm:"not null;default:''" json:"type"` + Target string `gorm:"not null;default:''" json:"target"` + Value float64 `gorm:"not null;default:0" json:"value"` + Message string `gorm:"not null;default:''" json:"message"` + Notified bool `gorm:"not null;default:false" json:"notified"` + CreatedAt time.Time `gorm:"index" json:"created_at"` +} + +// AlertMetric 单个目标的取值 +type AlertMetric struct { + Target string + Value float64 +} + +type AlertRepo interface { + ListRules(page, limit uint) ([]*AlertRule, int64, error) + AllRules() ([]*AlertRule, error) + GetRule(id uint) (*AlertRule, error) + CreateRule(rule *AlertRule) error + UpdateRule(rule *AlertRule) error + DeleteRule(id uint) error + AddAlert(alert *Alert) error + ListAlerts(page, limit uint) ([]*Alert, int64, error) + ClearAlerts() error + ClearAlertsBefore(t time.Time) error + CertExpiry() ([]*AlertMetric, error) + WebsiteExpiry() ([]*AlertMetric, error) + ProjectNames() ([]string, error) + DatabaseServers() ([]*DatabaseServer, error) + WebsiteHourStats() ([]*WebsiteHourStat, error) +} + +// WebsiteHourStat 网站当前小时的请求统计 +type WebsiteHourStat struct { + Site string + Requests uint64 + Errors uint64 + Status5xx uint64 +} + +// ioSnapshot 网卡/磁盘累计计数快照及据此算出的速率(字节/秒) +type ioSnapshot struct { + in uint64 + out uint64 + inRate uint64 + outRate uint64 + at time.Time +} + +type AlertUsecase struct { + repo AlertRepo + notify *NotifyUsecase + setting SettingRepo + container ContainerRepo + app AppRepo + database DatabaseServerRepo + loader *apploader.Loader + log *slog.Logger + t *gotext.Locale + + mu sync.Mutex + hits map[string]uint // 连续命中次数 + fired map[string]time.Time // 上次通知时间 + netSnaps map[string]ioSnapshot // 网卡累计流量快照 + diskSnaps map[string]ioSnapshot // 磁盘累计 IO 快照 + healthKeys map[string]struct{} // 已通知的健康问题 + sshFired map[string]time.Time // SSH 爆破上次通知时间 + sshAt time.Time // SSH 日志上次检查时间 + dbProbe map[string]bool // 本轮数据库连通性探测结果 + cleanedAt time.Time +} + +func NewAlertUsecase(i do.Injector) (*AlertUsecase, error) { + return &AlertUsecase{ + repo: do.MustInvoke[AlertRepo](i), + notify: do.MustInvoke[*NotifyUsecase](i), + setting: do.MustInvoke[SettingRepo](i), + container: do.MustInvoke[ContainerRepo](i), + app: do.MustInvoke[AppRepo](i), + database: do.MustInvoke[DatabaseServerRepo](i), + loader: do.MustInvoke[*apploader.Loader](i), + log: do.MustInvoke[*slog.Logger](i), + t: do.MustInvoke[*gotext.Locale](i), + hits: make(map[string]uint), + fired: make(map[string]time.Time), + netSnaps: make(map[string]ioSnapshot), + diskSnaps: make(map[string]ioSnapshot), + healthKeys: make(map[string]struct{}), + sshFired: make(map[string]time.Time), + }, nil +} + +func (uc *AlertUsecase) ListRules(page, limit uint) ([]*AlertRule, int64, error) { + return uc.repo.ListRules(page, limit) +} + +func (uc *AlertUsecase) GetRule(id uint) (*AlertRule, error) { + return uc.repo.GetRule(id) +} + +func (uc *AlertUsecase) CreateRule(ctx context.Context, req *request.AlertRuleCreate) (*AlertRule, error) { + rule := &AlertRule{ + Name: req.Name, + Type: req.Type, + Target: req.Target, + Operator: req.Operator, + Threshold: req.Threshold, + Duration: req.Duration, + Silence: req.Silence, + Channels: req.Channels, + Enabled: req.Enabled, + } + normalizeRule(rule) + + if err := uc.repo.CreateRule(rule); err != nil { + return nil, err + } + + uc.log.Info("alert rule created", slog.String("type", OperationTypeMonitor), slog.Uint64("operator_id", operatorID(ctx)), slog.String("name", req.Name)) + + return rule, nil +} + +func (uc *AlertUsecase) UpdateRule(ctx context.Context, req *request.AlertRuleUpdate) error { + rule, err := uc.repo.GetRule(req.ID) + if err != nil { + return err + } + + rule.Name = req.Name + rule.Type = req.Type + rule.Target = req.Target + rule.Operator = req.Operator + rule.Threshold = req.Threshold + rule.Duration = req.Duration + rule.Silence = req.Silence + rule.Channels = req.Channels + rule.Enabled = req.Enabled + normalizeRule(rule) + + if err = uc.repo.UpdateRule(rule); err != nil { + return err + } + + // 规则变更后重置命中状态,避免沿用旧阈值的计数 + uc.mu.Lock() + uc.clearState(rule.ID) + uc.mu.Unlock() + + uc.log.Info("alert rule updated", slog.String("type", OperationTypeMonitor), slog.Uint64("operator_id", operatorID(ctx)), slog.Uint64("id", uint64(req.ID)), slog.String("name", req.Name)) + + return nil +} + +func (uc *AlertUsecase) DeleteRule(ctx context.Context, id uint) error { + rule, err := uc.repo.GetRule(id) + if err != nil { + return err + } + if err = uc.repo.DeleteRule(id); err != nil { + return err + } + + uc.mu.Lock() + uc.clearState(id) + uc.mu.Unlock() + + uc.log.Info("alert rule deleted", slog.String("type", OperationTypeMonitor), slog.Uint64("operator_id", operatorID(ctx)), slog.Uint64("id", uint64(id)), slog.String("name", rule.Name)) + + return nil +} + +func (uc *AlertUsecase) ListAlerts(page, limit uint) ([]*Alert, int64, error) { + return uc.repo.ListAlerts(page, limit) +} + +func (uc *AlertUsecase) ClearAlerts() error { + return uc.repo.ClearAlerts() +} + +// Evaluate 评估全部启用的规则,命中则记录并通知 +func (uc *AlertUsecase) Evaluate(ctx context.Context) error { + rules, err := uc.repo.AllRules() + if err != nil { + return err + } + + // 探测结果按轮缓存,进入新一轮先失效 + uc.mu.Lock() + uc.dbProbe = nil + uc.mu.Unlock() + + uc.cleanup() + uc.checkHealth(ctx) + uc.checkSSH(ctx) + + var enabled []*AlertRule + for _, rule := range rules { + if rule.Enabled { + enabled = append(enabled, rule) + } + } + if len(enabled) == 0 { + uc.mu.Lock() + clear(uc.hits) + clear(uc.fired) + uc.mu.Unlock() + return nil + } + + info := tools.CurrentInfo(nil, nil) + now := time.Now() + uc.updateSnapshots(info, now) + + alive := make(map[string]struct{}) + for _, rule := range enabled { + metrics, err := uc.collect(ctx, rule, info) + if err != nil { + uc.log.Warn("failed to collect alert metric", slog.String("rule", rule.Name), slog.Any("err", err)) + continue + } + + for _, metric := range metrics { + key := stateKey(rule.ID, metric.Target) + alive[key] = struct{}{} + uc.evaluateMetric(ctx, rule, metric, key, now) + } + } + + // 清理已消失的目标状态 + uc.mu.Lock() + for key := range uc.hits { + if _, ok := alive[key]; !ok { + delete(uc.hits, key) + } + } + uc.mu.Unlock() + + return nil +} + +func (uc *AlertUsecase) evaluateMetric(ctx context.Context, rule *AlertRule, metric *AlertMetric, key string, now time.Time) { + uc.mu.Lock() + if !matchThreshold(rule, metric.Value) { + delete(uc.hits, key) + uc.mu.Unlock() + return + } + + uc.hits[key]++ + hits := uc.hits[key] + duration := max(rule.Duration, 1) + if hits < duration { + uc.mu.Unlock() + return + } + + // 静默期内不重复告警 + if last, ok := uc.fired[key]; ok && now.Sub(last) < time.Duration(rule.Silence)*time.Minute { + uc.mu.Unlock() + return + } + uc.fired[key] = now + uc.mu.Unlock() + + alert := &Alert{ + RuleID: rule.ID, + RuleName: rule.Name, + Type: rule.Type, + Target: metric.Target, + Value: metric.Value, + Message: uc.buildMessage(rule, metric), + } + + sent, err := uc.notify.Send(ctx, rule.Channels, uc.t.Get("[AcePanel] Alert: %s", rule.Name), NotifyBody(alert.Message, [][2]string{ + {uc.t.Get("Rule"), rule.Name}, + {uc.t.Get("Metric"), uc.metricLabel(rule.Type, metric.Target)}, + {uc.t.Get("Current Value"), uc.formatValue(rule.Type, metric.Value)}, + {uc.t.Get("Threshold"), uc.formatValue(rule.Type, rule.Threshold)}, + {uc.t.Get("Time"), now.Format(time.DateTime)}, + })) + if err != nil { + uc.log.Warn("failed to send alert notification", slog.String("rule", rule.Name), slog.Any("err", err)) + } + alert.Notified = sent > 0 + + // 一条都没送达时缩短静默期,让下一轮重试,避免临时故障吞掉整个静默窗口 + if len(rule.Channels) > 0 && sent == 0 { + uc.mu.Lock() + uc.fired[key] = now.Add(alertRetryDelay - time.Duration(rule.Silence)*time.Minute) + uc.mu.Unlock() + } + + if err := uc.repo.AddAlert(alert); err != nil { + uc.log.Warn("failed to save alert record", slog.String("rule", rule.Name), slog.Any("err", err)) + } +} + +// collect 采集规则对应的目标取值 +func (uc *AlertUsecase) collect(ctx context.Context, rule *AlertRule, info types.CurrentInfo) ([]*AlertMetric, error) { + switch rule.Type { + case AlertTypeCPU: + return []*AlertMetric{{Value: info.Percent}}, nil + + case AlertTypeMemory: + if info.Mem == nil { + return nil, nil + } + return []*AlertMetric{{Value: info.Mem.UsedPercent}}, nil + + case AlertTypeSwap: + if info.Swap == nil || info.Swap.Total == 0 { + return nil, nil + } + return []*AlertMetric{{Value: info.Swap.UsedPercent}}, nil + + case AlertTypeLoad1, AlertTypeLoad5, AlertTypeLoad15: + if info.Load == nil { + return nil, nil + } + switch rule.Type { + case AlertTypeLoad1: + return []*AlertMetric{{Value: info.Load.Load1}}, nil + case AlertTypeLoad5: + return []*AlertMetric{{Value: info.Load.Load5}}, nil + default: + return []*AlertMetric{{Value: info.Load.Load15}}, nil + } + + case AlertTypeDisk, AlertTypeDiskInode: + metrics := make([]*AlertMetric, 0) + for _, usage := range info.DiskUsage { + if rule.Target != "" && rule.Target != usage.Path { + continue + } + value := usage.UsedPercent + if rule.Type == AlertTypeDiskInode { + value = usage.InodesUsedPercent + } + metrics = append(metrics, &AlertMetric{Target: usage.Path, Value: value}) + } + return metrics, nil + + case AlertTypeNetIn, AlertTypeNetOut, AlertTypeDiskRead, AlertTypeDiskWrite: + return uc.rateMetrics(rule), nil + + case AlertTypeWebsite5xx, AlertTypeWebsiteError: + stats, err := uc.repo.WebsiteHourStats() + if err != nil { + return nil, err + } + metrics := make([]*AlertMetric, 0, len(stats)) + for _, item := range stats { + if rule.Target != "" && rule.Target != item.Site { + continue + } + if rule.Type == AlertTypeWebsite5xx { + metrics = append(metrics, &AlertMetric{Target: item.Site, Value: float64(item.Status5xx)}) + continue + } + // 无请求时错误率无意义 + if item.Requests == 0 { + continue + } + metrics = append(metrics, &AlertMetric{Target: item.Site, Value: float64(item.Errors) / float64(item.Requests) * 100}) + } + return metrics, nil + + case AlertTypeService: + if rule.Target == "" { + return nil, nil + } + return []*AlertMetric{{Target: rule.Target, Value: notRunning(systemctl.Status(rule.Target))}}, nil + + case AlertTypeProject: + names, err := uc.repo.ProjectNames() + if err != nil { + return nil, err + } + // 项目即 systemd 单元,单元名与项目名一致 + return lop.Map(filterNames(rule.Target, names), func(name string, _ int) *AlertMetric { + return &AlertMetric{Target: name, Value: notRunning(systemctl.Status(name))} + }), nil + + case AlertTypeContainer: + containers, err := uc.container.ListAll(containerSock(uc.setting)) + if err != nil { + return nil, err + } + metrics := make([]*AlertMetric, 0, len(containers)) + for _, item := range containers { + if rule.Target != "" && rule.Target != item.Name { + continue + } + value := float64(0) + if item.State != "running" { + value = 1 + } + metrics = append(metrics, &AlertMetric{Target: item.Name, Value: value}) + } + return metrics, nil + + case AlertTypeApp: + installed, err := uc.app.Installed() + if err != nil { + return nil, err + } + targets := lo.Filter(installed, func(item *App, _ int) bool { + if rule.Target != "" && rule.Target != item.Slug { + return false + } + _, ok := uc.loader.Get(item.Slug) + return ok + }) + // 状态查询逐个访问 systemd,并发采集 + return lo.Compact(lop.Map(targets, func(item *App, _ int) *AlertMetric { + a, _ := uc.loader.Get(item.Slug) + status := a.Status() + // 无 systemd 服务的应用没有运行状态 + if status == types.AppStatusNA { + return nil + } + value := float64(0) + if status != types.AppStatusRunning { + value = 1 + } + return &AlertMetric{Target: item.Slug, Value: value} + })), nil + + case AlertTypeDatabase: + servers, err := uc.repo.DatabaseServers() + if err != nil { + return nil, err + } + reachable := uc.probeDatabases(ctx, servers) + + metrics := make([]*AlertMetric, 0, len(servers)) + for _, item := range servers { + if rule.Target != "" && rule.Target != item.Name { + continue + } + value := float64(0) + if !reachable[item.Name] { + value = 1 + } + metrics = append(metrics, &AlertMetric{Target: item.Name, Value: value}) + } + return metrics, nil + + case AlertTypeCertExpire: + metrics, err := uc.repo.CertExpiry() + if err != nil { + return nil, err + } + return filterTarget(rule.Target, metrics), nil + + case AlertTypeWebsiteExpire: + metrics, err := uc.repo.WebsiteExpiry() + if err != nil { + return nil, err + } + return filterTarget(rule.Target, metrics), nil + } + + return nil, fmt.Errorf("unsupported alert type: %s", rule.Type) +} + +// notRunning 将运行状态转为告警取值,未运行记为 1,配合 normalizeRule 的 >=1 判定 +func notRunning(running bool, _ error) float64 { + if running { + return 0 + } + + return 1 +} + +// filterNames 按目标名筛选,目标为空表示全部 +func filterNames(target string, names []string) []string { + if target == "" { + return names + } + if slices.Contains(names, target) { + return []string{target} + } + + return nil +} + +// filterTarget 按目标名筛选,目标为空表示全部 +// 证书目标是逗号分隔的多域名,允许用其中任一域名命中 +func filterTarget(target string, metrics []*AlertMetric) []*AlertMetric { + if target == "" { + return metrics + } + + filtered := make([]*AlertMetric, 0) + for _, metric := range metrics { + if metric.Target == target || slices.Contains(strings.Split(metric.Target, ","), target) { + filtered = append(filtered, metric) + } + } + + return filtered +} + +// probeDatabases 并发探测数据库连通性,结果在本轮评估内复用,避免多条规则重复整批探测 +// 单台耗时上限由 pkg/db 各驱动的连接超时保证(5~10 秒),并发后整批不会拖长分钟级评估 +func (uc *AlertUsecase) probeDatabases(ctx context.Context, servers []*DatabaseServer) map[string]bool { + uc.mu.Lock() + cached := uc.dbProbe + uc.mu.Unlock() + if cached != nil { + return cached + } + + probe := lo.FromEntries(lop.Map(servers, func(item *DatabaseServer, _ int) lo.Entry[string, bool] { + return lo.Entry[string, bool]{Key: item.Name, Value: uc.database.CheckServer(ctx, item)} + })) + + uc.mu.Lock() + uc.dbProbe = probe + uc.mu.Unlock() + + return probe +} + +// rateMetrics 从快照读取速率指标(MB/s) +func (uc *AlertUsecase) rateMetrics(rule *AlertRule) []*AlertMetric { + uc.mu.Lock() + defer uc.mu.Unlock() + + snaps := uc.netSnaps + if rule.Type == AlertTypeDiskRead || rule.Type == AlertTypeDiskWrite { + snaps = uc.diskSnaps + } + + metrics := make([]*AlertMetric, 0, len(snaps)) + for name, snap := range snaps { + if rule.Target != "" && rule.Target != name { + continue + } + value := snap.inRate + if rule.Type == AlertTypeNetOut || rule.Type == AlertTypeDiskWrite { + value = snap.outRate + } + metrics = append(metrics, &AlertMetric{Target: name, Value: float64(value) / 1024 / 1024}) + } + + return metrics +} + +// updateSnapshots 用两次采集的累计值计算速率,结果暂存于快照 +func (uc *AlertUsecase) updateSnapshots(info types.CurrentInfo, now time.Time) { + uc.mu.Lock() + defer uc.mu.Unlock() + + nets := make(map[string]ioSnapshot, len(info.Net)) + for _, item := range info.Net { + if item.Name == "lo" { + continue + } + nets[item.Name] = rateOf(uc.netSnaps[item.Name], item.BytesRecv, item.BytesSent, now) + } + uc.netSnaps = nets + + disks := make(map[string]ioSnapshot, len(info.DiskIO)) + for _, item := range info.DiskIO { + disks[item.Name] = rateOf(uc.diskSnaps[item.Name], item.ReadBytes, item.WriteBytes, now) + } + uc.diskSnaps = disks +} + +// checkHealth 上报新出现的面板健康问题,问题恢复后重新出现会再次通知 +// 同步发送并只在送达后记入去重集合,否则一次发送失败就会让持续存在的问题再也不告警 +func (uc *AlertUsecase) checkHealth(ctx context.Context) { + issues := app.Health.Snapshot() + + uc.mu.Lock() + fresh := make([]app.HealthIssue, 0, len(issues)) + notified := make(map[string]struct{}, len(issues)) + for _, issue := range issues { + if _, ok := uc.healthKeys[issue.Key]; ok { + notified[issue.Key] = struct{}{} + continue + } + fresh = append(fresh, issue) + } + // 已消失的问题在此被丢弃,恢复后再次出现会重新通知 + uc.healthKeys = notified + uc.mu.Unlock() + + for _, issue := range fresh { + if err := uc.notify.SendEventSync(ctx, NotifyEventHealth, uc.t.Get("[AcePanel] Panel Health Issue"), NotifyBody(uc.t.Get("panel reported a health issue"), [][2]string{ + {uc.t.Get("Item"), issue.Key}, + {uc.t.Get("Level"), issue.Level}, + {uc.t.Get("Detail"), issue.Message}, + {uc.t.Get("Time"), issue.Since.Format(time.DateTime)}, + })); err != nil { + uc.log.Warn("failed to send health notification", slog.String("item", issue.Key), slog.Any("err", err)) + continue + } + + uc.mu.Lock() + uc.healthKeys[issue.Key] = struct{}{} + uc.mu.Unlock() + } +} + +// checkSSH 增量检查 sshd 日志,上报登录成功与爆破尝试 +func (uc *AlertUsecase) checkSSH(ctx context.Context) { + now := time.Now() + + uc.mu.Lock() + since := uc.sshAt + uc.sshAt = now + uc.mu.Unlock() + + // 首次仅记录时间,避免面板启动时把历史日志全推一遍 + if since.IsZero() { + return + } + + raw, err := shell.ExecfWithContext(ctx, `journalctl -u sshd -u ssh --no-pager -o json --since "@%d" 2>/dev/null`, since.Unix()) + if err != nil { + // 读取失败则回退检查点,下一轮重扫该窗口,避免丢掉这段时间的登录记录 + uc.mu.Lock() + uc.sshAt = since + uc.mu.Unlock() + return + } + if raw == "" { + return + } + + failures := make(map[string]uint) + scanner := bufio.NewScanner(strings.NewReader(raw)) + scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) + for scanner.Scan() { + var entry struct { + Message string `json:"MESSAGE"` + } + if json.Unmarshal(scanner.Bytes(), &entry) != nil { + continue + } + + record := sshlog.ParseMessage(entry.Message) + if record == nil { + continue + } + + switch record.Status { + case sshlog.StatusAccepted: + uc.notify.SendEvent(NotifyEventSSHLogin, uc.t.Get("[AcePanel] SSH Login"), NotifyBody(uc.t.Get("SSH login succeeded"), [][2]string{ + {uc.t.Get("Username"), record.User}, + {uc.t.Get("IP"), record.IP}, + {uc.t.Get("Method"), record.Method}, + {uc.t.Get("Time"), now.Format(time.DateTime)}, + })) + case sshlog.StatusFailed, sshlog.StatusInvalidUser: + failures[record.IP]++ + } + } + + for ip, count := range failures { + if count < sshFailThreshold || !uc.sshShouldFire(ip, now) { + continue + } + uc.notify.SendEvent(NotifyEventSSHBruteforce, uc.t.Get("[AcePanel] SSH Brute-force Attempts"), NotifyBody(uc.t.Get("too many failed SSH login attempts"), [][2]string{ + {uc.t.Get("IP"), ip}, + {uc.t.Get("Failed Attempts"), cast.ToString(count)}, + {uc.t.Get("Time"), now.Format(time.DateTime)}, + })) + } +} + +// sshShouldFire 判断某来源是否已过静默期,并顺带清理过期记录 +func (uc *AlertUsecase) sshShouldFire(ip string, now time.Time) bool { + uc.mu.Lock() + defer uc.mu.Unlock() + + for key, at := range uc.sshFired { + if now.Sub(at) > sshFailSilence { + delete(uc.sshFired, key) + } + } + + if at, ok := uc.sshFired[ip]; ok && now.Sub(at) < sshFailSilence { + return false + } + uc.sshFired[ip] = now + + return true +} + +// cleanup 按保留天数清理历史告警记录 +func (uc *AlertUsecase) cleanup() { + uc.mu.Lock() + if time.Since(uc.cleanedAt) < 6*time.Hour { + uc.mu.Unlock() + return + } + uc.cleanedAt = time.Now() + uc.mu.Unlock() + + days, err := uc.setting.GetInt(SettingKeyAlertLogDays, 30) + if err != nil || days <= 0 { + return + } + if err = uc.repo.ClearAlertsBefore(time.Now().AddDate(0, 0, -days)); err != nil { + uc.log.Warn("failed to clear expired alerts", slog.Any("err", err)) + } +} + +// clearState 清除某条规则的运行时状态,调用前需持有锁 +func (uc *AlertUsecase) clearState(ruleID uint) { + prefix := fmt.Sprintf("%d:", ruleID) + for key := range uc.hits { + if strings.HasPrefix(key, prefix) { + delete(uc.hits, key) + } + } + for key := range uc.fired { + if strings.HasPrefix(key, prefix) { + delete(uc.fired, key) + } + } +} + +func (uc *AlertUsecase) buildMessage(rule *AlertRule, metric *AlertMetric) string { + label := uc.metricLabel(rule.Type, metric.Target) + + switch rule.Type { + case AlertTypeService: + return uc.t.Get("service %s is not running", metric.Target) + case AlertTypeProject: + return uc.t.Get("project %s is not running", metric.Target) + case AlertTypeContainer: + return uc.t.Get("container %s is not running", metric.Target) + case AlertTypeApp: + return uc.t.Get("app %s is not running", metric.Target) + case AlertTypeDatabase: + return uc.t.Get("database server %s is unreachable", metric.Target) + case AlertTypeCertExpire: + return uc.t.Get("certificate %s expires in %s days", metric.Target, uc.formatValue(rule.Type, metric.Value)) + case AlertTypeWebsiteExpire: + return uc.t.Get("website %s expires in %s days", metric.Target, uc.formatValue(rule.Type, metric.Value)) + } + + return uc.t.Get("%s is %s, %s threshold %s", label, uc.formatValue(rule.Type, metric.Value), uc.operatorLabel(rule.Operator), uc.formatValue(rule.Type, rule.Threshold)) +} + +func (uc *AlertUsecase) metricLabel(typ, target string) string { + var label string + switch typ { + case AlertTypeCPU: + label = uc.t.Get("CPU usage") + case AlertTypeMemory: + label = uc.t.Get("memory usage") + case AlertTypeSwap: + label = uc.t.Get("swap usage") + case AlertTypeLoad1: + label = uc.t.Get("1 minute load") + case AlertTypeLoad5: + label = uc.t.Get("5 minutes load") + case AlertTypeLoad15: + label = uc.t.Get("15 minutes load") + case AlertTypeDisk: + label = uc.t.Get("disk usage") + case AlertTypeDiskInode: + label = uc.t.Get("disk inode usage") + case AlertTypeDiskRead: + label = uc.t.Get("disk read speed") + case AlertTypeDiskWrite: + label = uc.t.Get("disk write speed") + case AlertTypeNetIn: + label = uc.t.Get("network download speed") + case AlertTypeNetOut: + label = uc.t.Get("network upload speed") + case AlertTypeWebsite5xx: + label = uc.t.Get("website 5xx responses this hour") + case AlertTypeWebsiteError: + label = uc.t.Get("website error rate this hour") + case AlertTypeService: + label = uc.t.Get("service status") + case AlertTypeProject: + label = uc.t.Get("project status") + case AlertTypeContainer: + label = uc.t.Get("container status") + case AlertTypeApp: + label = uc.t.Get("app status") + case AlertTypeDatabase: + label = uc.t.Get("database server status") + case AlertTypeCertExpire: + label = uc.t.Get("certificate expiry") + case AlertTypeWebsiteExpire: + label = uc.t.Get("website expiry") + default: + label = typ + } + + if target != "" { + return fmt.Sprintf("%s (%s)", label, target) + } + + return label +} + +func (uc *AlertUsecase) operatorLabel(operator string) string { + switch operator { + case AlertOperatorGTE: + return uc.t.Get("greater than or equal to") + case AlertOperatorLT: + return uc.t.Get("less than") + case AlertOperatorLTE: + return uc.t.Get("less than or equal to") + default: + return uc.t.Get("greater than") + } +} + +func (uc *AlertUsecase) formatValue(typ string, value float64) string { + switch typ { + case AlertTypeCPU, AlertTypeMemory, AlertTypeSwap, AlertTypeDisk, AlertTypeDiskInode, AlertTypeWebsiteError: + return fmt.Sprintf("%.2f%%", value) + case AlertTypeNetIn, AlertTypeNetOut, AlertTypeDiskRead, AlertTypeDiskWrite: + return fmt.Sprintf("%.2f MB/s", value) + case AlertTypeCertExpire, AlertTypeWebsiteExpire, AlertTypeWebsite5xx: + return fmt.Sprintf("%.0f", value) + } + + if IsStatusAlert(typ) { + return uc.t.Get("not running") + } + + return fmt.Sprintf("%.2f", value) +} + +// normalizeRule 补齐规则的默认值,状态类规则固定为「不在运行」 +func normalizeRule(rule *AlertRule) { + if IsStatusAlert(rule.Type) { + rule.Operator = AlertOperatorGTE + rule.Threshold = 1 + } + if rule.Duration < 1 { + rule.Duration = 1 + } +} + +func matchThreshold(rule *AlertRule, value float64) bool { + switch rule.Operator { + case AlertOperatorGTE: + return value >= rule.Threshold + case AlertOperatorLT: + return value < rule.Threshold + case AlertOperatorLTE: + return value <= rule.Threshold + default: + return value > rule.Threshold + } +} + +// rateOf 依据上次累计值计算每秒增量,首次采集速率为 0 +func rateOf(prev ioSnapshot, in, out uint64, now time.Time) ioSnapshot { + snap := ioSnapshot{in: in, out: out, at: now} + if prev.at.IsZero() { + return snap + } + + elapsed := now.Sub(prev.at).Seconds() + if elapsed < 1 { + elapsed = 1 + } + if in >= prev.in { + snap.inRate = uint64(float64(in-prev.in) / elapsed) + } + if out >= prev.out { + snap.outRate = uint64(float64(out-prev.out) / elapsed) + } + + return snap +} + +func stateKey(ruleID uint, target string) string { + return fmt.Sprintf("%d:%s", ruleID, target) +} diff --git a/internal/biz/alert_test.go b/internal/biz/alert_test.go new file mode 100644 index 00000000..760e2c6b --- /dev/null +++ b/internal/biz/alert_test.go @@ -0,0 +1,29 @@ +package biz + +import "testing" + +// 多域名证书的目标是逗号分隔的域名列表,用其中任一域名都应能命中规则 +func TestFilterTargetMatchesAnyCertDomain(t *testing.T) { + metrics := []*AlertMetric{ + {Target: "a.example.com,b.example.com", Value: 5}, + {Target: "c.example.com", Value: 30}, + } + + cases := []struct { + target string + want int + }{ + {"", 2}, + {"a.example.com", 1}, + {"b.example.com", 1}, + {"c.example.com", 1}, + {"a.example.com,b.example.com", 1}, + {"missing.example.com", 0}, + } + + for _, c := range cases { + if got := filterTarget(c.target, metrics); len(got) != c.want { + t.Fatalf("target %q matched %d metrics, want %d", c.target, len(got), c.want) + } + } +} diff --git a/internal/biz/backup.go b/internal/biz/backup.go index 4b2e0b53..bb7bd11e 100644 --- a/internal/biz/backup.go +++ b/internal/biz/backup.go @@ -3,6 +3,10 @@ package biz import ( "context" "log/slog" + "time" + + "github.com/leonelquinteros/gotext" + "github.com/samber/do/v2" "github.com/acepanel/panel/v3/pkg/types" ) @@ -37,12 +41,19 @@ type BackupRepo interface { } type BackupUsecase struct { - repo BackupRepo - log *slog.Logger + repo BackupRepo + log *slog.Logger + notify *NotifyUsecase + t *gotext.Locale } -func NewBackupUsecase(repo BackupRepo, log *slog.Logger) *BackupUsecase { - return &BackupUsecase{repo: repo, log: log} +func NewBackupUsecase(i do.Injector) (*BackupUsecase, error) { + return &BackupUsecase{ + repo: do.MustInvoke[BackupRepo](i), + log: do.MustInvoke[*slog.Logger](i), + notify: do.MustInvoke[*NotifyUsecase](i), + t: do.MustInvoke[*gotext.Locale](i), + }, nil } func (uc *BackupUsecase) List(typ BackupType) ([]*types.BackupFile, error) { @@ -50,8 +61,22 @@ func (uc *BackupUsecase) List(typ BackupType) ([]*types.BackupFile, error) { } func (uc *BackupUsecase) Create(ctx context.Context, typ BackupType, target string, account uint) error { - // 审计留 repo:需区分预检早返回(不记)与备份执行失败(记 Warn),无法在此上移 - return uc.repo.Create(ctx, typ, target, account) + err := uc.repo.Create(ctx, typ, target, account) + if err == nil { + return nil + } + + // 定时备份由 CLI 执行,命令返回即退出,异步通知来不及发出,必须同步发送 + if sendErr := uc.notify.SendEventSync(ctx, NotifyEventBackup, uc.t.Get("[AcePanel] Backup Failed"), NotifyBody(uc.t.Get("backup task failed"), [][2]string{ + {uc.t.Get("Type"), string(typ)}, + {uc.t.Get("Target"), target}, + {uc.t.Get("Error"), err.Error()}, + {uc.t.Get("Time"), time.Now().Format(time.DateTime)}, + })); sendErr != nil { + uc.log.Warn("failed to send backup failure notification", slog.Any("err", sendErr)) + } + + return err } func (uc *BackupUsecase) CreatePanel() error { diff --git a/internal/biz/biz.go b/internal/biz/biz.go index 7d70a543..c6af1f44 100644 --- a/internal/biz/biz.go +++ b/internal/biz/biz.go @@ -7,14 +7,14 @@ import ( ) var Package = do.Package( - do.Lazy(NewAppUsecase), registry.Lazy2(NewBackupUsecase), do.Lazy(NewBackupAccountUsecase), + do.Lazy(NewAlertUsecase), do.Lazy(NewAppUsecase), do.Lazy(NewBackupUsecase), do.Lazy(NewBackupAccountUsecase), registry.Lazy(NewCacheUsecase), do.Lazy(NewCertUsecase), do.Lazy(NewCertAccountUsecase), registry.Lazy2(NewCertDNSUsecase), registry.Lazy2(NewContainerUsecase), registry.Lazy(NewContainerComposeUsecase), registry.Lazy2(NewContainerImageUsecase), registry.Lazy2(NewContainerNetworkUsecase), registry.Lazy2(NewContainerVolumeUsecase), registry.Lazy2(NewCronUsecase), do.Lazy(NewDatabaseUsecase), registry.Lazy(NewDatabaseRedisUsecase), registry.Lazy(NewDatabaseElasticsearchUsecase), do.Lazy(NewDatabaseServerUsecase), do.Lazy(NewDatabaseUserUsecase), do.Lazy(NewEnvironmentUsecase), do.Lazy(NewFileShareUsecase), registry.Lazy(NewLogUsecase), registry.Lazy2(NewMonitorUsecase), - do.Lazy(NewProjectUsecase), registry.Lazy2(NewSafeUsecase), registry.Lazy2(NewScanEventUsecase), + do.Lazy(NewNotifyUsecase), do.Lazy(NewProjectUsecase), registry.Lazy2(NewSafeUsecase), registry.Lazy2(NewScanEventUsecase), do.Lazy(NewSettingUsecase), registry.Lazy2(NewSSHUsecase), do.Lazy(NewTamperUsecase), registry.Lazy(NewTaskUsecase), do.Lazy(NewTemplateUsecase), do.Lazy(NewUserUsecase), registry.Lazy(NewUserPasskeyUsecase), registry.Lazy(NewUserTokenUsecase), do.Lazy(NewWebHookUsecase), do.Lazy(NewWebsiteUsecase), diff --git a/internal/biz/database.go b/internal/biz/database.go index 81bdfcc5..35224876 100644 --- a/internal/biz/database.go +++ b/internal/biz/database.go @@ -35,9 +35,9 @@ type Database struct { type DatabaseRepo interface { ListServers(typ string) ([]*DatabaseServer, error) - DatabasesOf(server *DatabaseServer) ([]*Database, error) - Operator(server *DatabaseServer) (db.Operator, error) - Mongo(server *DatabaseServer) (*db.MongoDB, error) + DatabasesOf(ctx context.Context, server *DatabaseServer) ([]*Database, error) + Operator(ctx context.Context, server *DatabaseServer) (db.Operator, error) + Mongo(ctx context.Context, server *DatabaseServer) (*db.MongoDB, error) } // DatabaseUsecase 数据库业务用例 @@ -59,7 +59,7 @@ func NewDatabaseUsecase(i do.Injector) (*DatabaseUsecase, error) { }, nil } -func (uc *DatabaseUsecase) List(page, limit uint, typ string) ([]*Database, int64, error) { +func (uc *DatabaseUsecase) List(ctx context.Context, page, limit uint, typ string) ([]*Database, int64, error) { servers, err := uc.repo.ListServers(typ) if err != nil { return nil, 0, err @@ -67,7 +67,7 @@ func (uc *DatabaseUsecase) List(page, limit uint, typ string) ([]*Database, int6 database := make([]*Database, 0) for _, server := range servers { - databases, err := uc.repo.DatabasesOf(server) + databases, err := uc.repo.DatabasesOf(ctx, server) if err != nil { continue } @@ -82,14 +82,14 @@ func (uc *DatabaseUsecase) List(page, limit uint, typ string) ([]*Database, int6 } func (uc *DatabaseUsecase) Create(ctx context.Context, req *request.DatabaseCreate) error { - server, err := uc.server.Get(req.ServerID) + server, err := uc.server.Get(ctx, req.ServerID) if err != nil { return err } // MongoDB 独立处理,不走 Operator 接口 if server.Type == DatabaseTypeMongoDB { - mongo, mongoErr := uc.repo.Mongo(server) + mongo, mongoErr := uc.repo.Mongo(ctx, server) if mongoErr != nil { return mongoErr } @@ -101,7 +101,7 @@ func (uc *DatabaseUsecase) Create(ctx context.Context, req *request.DatabaseCrea return nil } - operator, err := uc.repo.Operator(server) + operator, err := uc.repo.Operator(ctx, server) if err != nil { return err } @@ -203,14 +203,14 @@ func (uc *DatabaseUsecase) mysqlUserHosts(operator db.Operator, user string) []s } func (uc *DatabaseUsecase) Delete(ctx context.Context, serverID uint, name string) error { - server, err := uc.server.Get(serverID) + server, err := uc.server.Get(ctx, serverID) if err != nil { return err } switch server.Type { case DatabaseTypeMongoDB: - mongo, mongoErr := uc.repo.Mongo(server) + mongo, mongoErr := uc.repo.Mongo(ctx, server) if mongoErr != nil { return mongoErr } @@ -221,7 +221,7 @@ func (uc *DatabaseUsecase) Delete(ctx context.Context, serverID uint, name strin case DatabaseTypeSQLite: return errors.New(uc.t.Get("sqlite does not support dropping tables from here")) default: - operator, opErr := uc.repo.Operator(server) + operator, opErr := uc.repo.Operator(ctx, server) if opErr != nil { return opErr } @@ -237,15 +237,15 @@ func (uc *DatabaseUsecase) Delete(ctx context.Context, serverID uint, name strin return nil } -func (uc *DatabaseUsecase) Comment(req *request.DatabaseComment) error { - server, err := uc.server.Get(req.ServerID) +func (uc *DatabaseUsecase) Comment(ctx context.Context, req *request.DatabaseComment) error { + server, err := uc.server.Get(ctx, req.ServerID) if err != nil { return err } switch server.Type { case DatabaseTypePostgresql: - operator, opErr := uc.repo.Operator(server) + operator, opErr := uc.repo.Operator(ctx, server) if opErr != nil { return opErr } diff --git a/internal/biz/database_elasticsearch.go b/internal/biz/database_elasticsearch.go index f04943dd..152e2d73 100644 --- a/internal/biz/database_elasticsearch.go +++ b/internal/biz/database_elasticsearch.go @@ -1,18 +1,19 @@ package biz import ( + "context" "github.com/acepanel/panel/v3/internal/request" "github.com/acepanel/panel/v3/pkg/db" ) type DatabaseElasticsearchRepo interface { - Indices(req *request.DatabaseESIndices) ([]db.ESIndex, error) - IndexCreate(req *request.DatabaseESIndexCreate) error - IndexDelete(req *request.DatabaseESIndexDelete) error - Data(req *request.DatabaseESData) ([]db.ESDocument, int64, error) - DocumentGet(req *request.DatabaseESDocumentGet) (*db.ESDocument, error) - DocumentSet(req *request.DatabaseESDocumentSet) error - DocumentDelete(req *request.DatabaseESDocumentDelete) error + Indices(ctx context.Context, req *request.DatabaseESIndices) ([]db.ESIndex, error) + IndexCreate(ctx context.Context, req *request.DatabaseESIndexCreate) error + IndexDelete(ctx context.Context, req *request.DatabaseESIndexDelete) error + Data(ctx context.Context, req *request.DatabaseESData) ([]db.ESDocument, int64, error) + DocumentGet(ctx context.Context, req *request.DatabaseESDocumentGet) (*db.ESDocument, error) + DocumentSet(ctx context.Context, req *request.DatabaseESDocumentSet) error + DocumentDelete(ctx context.Context, req *request.DatabaseESDocumentDelete) error } // DatabaseElasticsearchUsecase Elasticsearch 业务用例 @@ -24,30 +25,30 @@ func NewDatabaseElasticsearchUsecase(repo DatabaseElasticsearchRepo) *DatabaseEl return &DatabaseElasticsearchUsecase{repo: repo} } -func (uc *DatabaseElasticsearchUsecase) Indices(req *request.DatabaseESIndices) ([]db.ESIndex, error) { - return uc.repo.Indices(req) +func (uc *DatabaseElasticsearchUsecase) Indices(ctx context.Context, req *request.DatabaseESIndices) ([]db.ESIndex, error) { + return uc.repo.Indices(ctx, req) } -func (uc *DatabaseElasticsearchUsecase) IndexCreate(req *request.DatabaseESIndexCreate) error { - return uc.repo.IndexCreate(req) +func (uc *DatabaseElasticsearchUsecase) IndexCreate(ctx context.Context, req *request.DatabaseESIndexCreate) error { + return uc.repo.IndexCreate(ctx, req) } -func (uc *DatabaseElasticsearchUsecase) IndexDelete(req *request.DatabaseESIndexDelete) error { - return uc.repo.IndexDelete(req) +func (uc *DatabaseElasticsearchUsecase) IndexDelete(ctx context.Context, req *request.DatabaseESIndexDelete) error { + return uc.repo.IndexDelete(ctx, req) } -func (uc *DatabaseElasticsearchUsecase) Data(req *request.DatabaseESData) ([]db.ESDocument, int64, error) { - return uc.repo.Data(req) +func (uc *DatabaseElasticsearchUsecase) Data(ctx context.Context, req *request.DatabaseESData) ([]db.ESDocument, int64, error) { + return uc.repo.Data(ctx, req) } -func (uc *DatabaseElasticsearchUsecase) DocumentGet(req *request.DatabaseESDocumentGet) (*db.ESDocument, error) { - return uc.repo.DocumentGet(req) +func (uc *DatabaseElasticsearchUsecase) DocumentGet(ctx context.Context, req *request.DatabaseESDocumentGet) (*db.ESDocument, error) { + return uc.repo.DocumentGet(ctx, req) } -func (uc *DatabaseElasticsearchUsecase) DocumentSet(req *request.DatabaseESDocumentSet) error { - return uc.repo.DocumentSet(req) +func (uc *DatabaseElasticsearchUsecase) DocumentSet(ctx context.Context, req *request.DatabaseESDocumentSet) error { + return uc.repo.DocumentSet(ctx, req) } -func (uc *DatabaseElasticsearchUsecase) DocumentDelete(req *request.DatabaseESDocumentDelete) error { - return uc.repo.DocumentDelete(req) +func (uc *DatabaseElasticsearchUsecase) DocumentDelete(ctx context.Context, req *request.DatabaseESDocumentDelete) error { + return uc.repo.DocumentDelete(ctx, req) } diff --git a/internal/biz/database_redis.go b/internal/biz/database_redis.go index 7320b09d..59a977d1 100644 --- a/internal/biz/database_redis.go +++ b/internal/biz/database_redis.go @@ -1,19 +1,20 @@ package biz import ( + "context" "github.com/acepanel/panel/v3/internal/request" "github.com/acepanel/panel/v3/pkg/db" ) type DatabaseRedisRepo interface { - Databases(req *request.DatabaseRedisDatabases) (int, error) - Data(req *request.DatabaseRedisData) ([]db.RedisKV, int, error) - KeyGet(req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) - KeySet(req *request.DatabaseRedisKeySet) error - KeyDelete(req *request.DatabaseRedisKeyDelete) error - KeyTTL(req *request.DatabaseRedisKeyTTL) error - KeyRename(req *request.DatabaseRedisKeyRename) error - Clear(req *request.DatabaseRedisClear) error + Databases(ctx context.Context, req *request.DatabaseRedisDatabases) (int, error) + Data(ctx context.Context, req *request.DatabaseRedisData) ([]db.RedisKV, int, error) + KeyGet(ctx context.Context, req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) + KeySet(ctx context.Context, req *request.DatabaseRedisKeySet) error + KeyDelete(ctx context.Context, req *request.DatabaseRedisKeyDelete) error + KeyTTL(ctx context.Context, req *request.DatabaseRedisKeyTTL) error + KeyRename(ctx context.Context, req *request.DatabaseRedisKeyRename) error + Clear(ctx context.Context, req *request.DatabaseRedisClear) error } // DatabaseRedisUsecase Redis 业务用例 @@ -25,34 +26,34 @@ func NewDatabaseRedisUsecase(repo DatabaseRedisRepo) *DatabaseRedisUsecase { return &DatabaseRedisUsecase{repo: repo} } -func (uc *DatabaseRedisUsecase) Databases(req *request.DatabaseRedisDatabases) (int, error) { - return uc.repo.Databases(req) +func (uc *DatabaseRedisUsecase) Databases(ctx context.Context, req *request.DatabaseRedisDatabases) (int, error) { + return uc.repo.Databases(ctx, req) } -func (uc *DatabaseRedisUsecase) Data(req *request.DatabaseRedisData) ([]db.RedisKV, int, error) { - return uc.repo.Data(req) +func (uc *DatabaseRedisUsecase) Data(ctx context.Context, req *request.DatabaseRedisData) ([]db.RedisKV, int, error) { + return uc.repo.Data(ctx, req) } -func (uc *DatabaseRedisUsecase) KeyGet(req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) { - return uc.repo.KeyGet(req) +func (uc *DatabaseRedisUsecase) KeyGet(ctx context.Context, req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) { + return uc.repo.KeyGet(ctx, req) } -func (uc *DatabaseRedisUsecase) KeySet(req *request.DatabaseRedisKeySet) error { - return uc.repo.KeySet(req) +func (uc *DatabaseRedisUsecase) KeySet(ctx context.Context, req *request.DatabaseRedisKeySet) error { + return uc.repo.KeySet(ctx, req) } -func (uc *DatabaseRedisUsecase) KeyDelete(req *request.DatabaseRedisKeyDelete) error { - return uc.repo.KeyDelete(req) +func (uc *DatabaseRedisUsecase) KeyDelete(ctx context.Context, req *request.DatabaseRedisKeyDelete) error { + return uc.repo.KeyDelete(ctx, req) } -func (uc *DatabaseRedisUsecase) KeyTTL(req *request.DatabaseRedisKeyTTL) error { - return uc.repo.KeyTTL(req) +func (uc *DatabaseRedisUsecase) KeyTTL(ctx context.Context, req *request.DatabaseRedisKeyTTL) error { + return uc.repo.KeyTTL(ctx, req) } -func (uc *DatabaseRedisUsecase) KeyRename(req *request.DatabaseRedisKeyRename) error { - return uc.repo.KeyRename(req) +func (uc *DatabaseRedisUsecase) KeyRename(ctx context.Context, req *request.DatabaseRedisKeyRename) error { + return uc.repo.KeyRename(ctx, req) } -func (uc *DatabaseRedisUsecase) Clear(req *request.DatabaseRedisClear) error { - return uc.repo.Clear(req) +func (uc *DatabaseRedisUsecase) Clear(ctx context.Context, req *request.DatabaseRedisClear) error { + return uc.repo.Clear(ctx, req) } diff --git a/internal/biz/database_server.go b/internal/biz/database_server.go index e0940a86..c90ebadc 100644 --- a/internal/biz/database_server.go +++ b/internal/biz/database_server.go @@ -1,6 +1,7 @@ package biz import ( + "context" "errors" "fmt" "log/slog" @@ -69,9 +70,9 @@ func (r *DatabaseServer) AfterFind(tx *gorm.DB) error { type DatabaseServerRepo interface { Count() (int64, error) - List(page, limit uint, typ string) ([]*DatabaseServer, int64, error) - Get(id uint) (*DatabaseServer, error) - GetByName(name string) (*DatabaseServer, error) + List(ctx context.Context, page, limit uint, typ string) ([]*DatabaseServer, int64, error) + Get(ctx context.Context, id uint) (*DatabaseServer, error) + GetByName(ctx context.Context, name string) (*DatabaseServer, error) Create(server *DatabaseServer) error Save(server *DatabaseServer) error UpdateRemark(req *request.DatabaseServerUpdateRemark) error @@ -81,8 +82,8 @@ type DatabaseServerRepo interface { ClearUsers(id uint) error ListUsers(serverID uint) ([]*DatabaseUser, error) CreateUser(user *DatabaseUser) error - Operator(server *DatabaseServer) (db.Operator, error) - CheckServer(server *DatabaseServer) bool + Operator(ctx context.Context, server *DatabaseServer) (db.Operator, error) + CheckServer(ctx context.Context, server *DatabaseServer) bool } // DatabaseServerUsecase 数据库服务器业务用例 @@ -104,19 +105,19 @@ func (uc *DatabaseServerUsecase) Count() (int64, error) { return uc.repo.Count() } -func (uc *DatabaseServerUsecase) List(page, limit uint, typ string) ([]*DatabaseServer, int64, error) { - return uc.repo.List(page, limit, typ) +func (uc *DatabaseServerUsecase) List(ctx context.Context, page, limit uint, typ string) ([]*DatabaseServer, int64, error) { + return uc.repo.List(ctx, page, limit, typ) } -func (uc *DatabaseServerUsecase) Get(id uint) (*DatabaseServer, error) { - return uc.repo.Get(id) +func (uc *DatabaseServerUsecase) Get(ctx context.Context, id uint) (*DatabaseServer, error) { + return uc.repo.Get(ctx, id) } -func (uc *DatabaseServerUsecase) GetByName(name string) (*DatabaseServer, error) { - return uc.repo.GetByName(name) +func (uc *DatabaseServerUsecase) GetByName(ctx context.Context, name string) (*DatabaseServer, error) { + return uc.repo.GetByName(ctx, name) } -func (uc *DatabaseServerUsecase) Create(req *request.DatabaseServerCreate) error { +func (uc *DatabaseServerUsecase) Create(ctx context.Context, req *request.DatabaseServerCreate) error { databaseServer := &DatabaseServer{ Name: req.Name, Type: DatabaseType(req.Type), @@ -127,15 +128,15 @@ func (uc *DatabaseServerUsecase) Create(req *request.DatabaseServerCreate) error Remark: req.Remark, } - if !uc.repo.CheckServer(databaseServer) { + if !uc.repo.CheckServer(ctx, databaseServer) { return errors.New(uc.t.Get("check server connection failed")) } return uc.repo.Create(databaseServer) } -func (uc *DatabaseServerUsecase) Update(req *request.DatabaseServerUpdate) error { - server, err := uc.repo.Get(req.ID) +func (uc *DatabaseServerUsecase) Update(ctx context.Context, req *request.DatabaseServerUpdate) error { + server, err := uc.repo.Get(ctx, req.ID) if err != nil { return err } @@ -147,7 +148,7 @@ func (uc *DatabaseServerUsecase) Update(req *request.DatabaseServerUpdate) error server.Password = req.Password server.Remark = req.Remark - if !uc.repo.CheckServer(server) { + if !uc.repo.CheckServer(ctx, server) { return errors.New(uc.t.Get("check server connection failed")) } @@ -174,8 +175,8 @@ func (uc *DatabaseServerUsecase) ClearUsers(id uint) error { return uc.repo.ClearUsers(id) } -func (uc *DatabaseServerUsecase) Sync(id uint) error { - server, err := uc.repo.Get(id) +func (uc *DatabaseServerUsecase) Sync(ctx context.Context, id uint) error { + server, err := uc.repo.Get(ctx, id) if err != nil { return err } @@ -191,7 +192,7 @@ func (uc *DatabaseServerUsecase) Sync(id uint) error { return err } - operator, err := uc.repo.Operator(server) + operator, err := uc.repo.Operator(ctx, server) if err != nil { return err } diff --git a/internal/biz/database_user.go b/internal/biz/database_user.go index ad53de74..ae91fb3f 100644 --- a/internal/biz/database_user.go +++ b/internal/biz/database_user.go @@ -68,10 +68,10 @@ func (r *DatabaseUser) AfterFind(tx *gorm.DB) error { type DatabaseUserRepo interface { Count() (int64, error) - List(page, limit uint, typ string) ([]*DatabaseUser, int64, error) - Get(id uint) (*DatabaseUser, error) - UpdateRemark(req *request.DatabaseUserUpdateRemark) error - Operator(server *DatabaseServer) (db.Operator, error) + List(ctx context.Context, page, limit uint, typ string) ([]*DatabaseUser, int64, error) + Get(ctx context.Context, id uint) (*DatabaseUser, error) + UpdateRemark(ctx context.Context, req *request.DatabaseUserUpdateRemark) error + Operator(ctx context.Context, server *DatabaseServer) (db.Operator, error) Upsert(user *DatabaseUser) error Save(user *DatabaseUser) error DeleteByID(id uint) error @@ -98,21 +98,21 @@ func (uc *DatabaseUserUsecase) Count() (int64, error) { return uc.repo.Count() } -func (uc *DatabaseUserUsecase) List(page, limit uint, typ string) ([]*DatabaseUser, int64, error) { - return uc.repo.List(page, limit, typ) +func (uc *DatabaseUserUsecase) List(ctx context.Context, page, limit uint, typ string) ([]*DatabaseUser, int64, error) { + return uc.repo.List(ctx, page, limit, typ) } -func (uc *DatabaseUserUsecase) Get(id uint) (*DatabaseUser, error) { - return uc.repo.Get(id) +func (uc *DatabaseUserUsecase) Get(ctx context.Context, id uint) (*DatabaseUser, error) { + return uc.repo.Get(ctx, id) } func (uc *DatabaseUserUsecase) Create(ctx context.Context, req *request.DatabaseUserCreate) error { - server, err := uc.server.Get(req.ServerID) + server, err := uc.server.Get(ctx, req.ServerID) if err != nil { return err } - operator, err := uc.repo.Operator(server) + operator, err := uc.repo.Operator(ctx, server) if err != nil { return err } @@ -151,18 +151,18 @@ func (uc *DatabaseUserUsecase) Create(ctx context.Context, req *request.Database return nil } -func (uc *DatabaseUserUsecase) Update(req *request.DatabaseUserUpdate) error { - user, err := uc.repo.Get(req.ID) +func (uc *DatabaseUserUsecase) Update(ctx context.Context, req *request.DatabaseUserUpdate) error { + user, err := uc.repo.Get(ctx, req.ID) if err != nil { return err } - server, err := uc.server.Get(user.ServerID) + server, err := uc.server.Get(ctx, user.ServerID) if err != nil { return err } - operator, err := uc.repo.Operator(server) + operator, err := uc.repo.Operator(ctx, server) if err != nil { return err } @@ -199,22 +199,22 @@ func (uc *DatabaseUserUsecase) Update(req *request.DatabaseUserUpdate) error { return uc.repo.Save(user) } -func (uc *DatabaseUserUsecase) UpdateRemark(req *request.DatabaseUserUpdateRemark) error { - return uc.repo.UpdateRemark(req) +func (uc *DatabaseUserUsecase) UpdateRemark(ctx context.Context, req *request.DatabaseUserUpdateRemark) error { + return uc.repo.UpdateRemark(ctx, req) } func (uc *DatabaseUserUsecase) Delete(ctx context.Context, id uint) error { - user, err := uc.repo.Get(id) + user, err := uc.repo.Get(ctx, id) if err != nil { return err } - server, err := uc.server.Get(user.ServerID) + server, err := uc.server.Get(ctx, user.ServerID) if err != nil { return err } - operator, err := uc.repo.Operator(server) + operator, err := uc.repo.Operator(ctx, server) if err != nil { return err } @@ -232,13 +232,13 @@ func (uc *DatabaseUserUsecase) Delete(ctx context.Context, id uint) error { return nil } -func (uc *DatabaseUserUsecase) DeleteByNames(serverID uint, names []string) error { - server, err := uc.server.Get(serverID) +func (uc *DatabaseUserUsecase) DeleteByNames(ctx context.Context, serverID uint, names []string) error { + server, err := uc.server.Get(ctx, serverID) if err != nil { return err } - operator, err := uc.repo.Operator(server) + operator, err := uc.repo.Operator(ctx, server) if err != nil { return err } diff --git a/internal/biz/monitor.go b/internal/biz/monitor.go index ed1f77e2..47c06526 100644 --- a/internal/biz/monitor.go +++ b/internal/biz/monitor.go @@ -43,11 +43,16 @@ func (uc *MonitorUsecase) GetSetting() (*request.MonitorSetting, error) { if err != nil { return nil, err } + alertDays, err := uc.setting.GetInt(SettingKeyAlertLogDays, 30) + if err != nil { + return nil, err + } setting := new(request.MonitorSetting) setting.Enabled = cast.ToBool(monitor) setting.Days = cast.ToUint(monitorDays) setting.Interval = uint(monitorInterval) + setting.AlertDays = uint(alertDays) return setting, nil } @@ -63,7 +68,7 @@ func (uc *MonitorUsecase) UpdateSetting(setting *request.MonitorSetting) error { return err } - return nil + return uc.setting.Set(SettingKeyAlertLogDays, cast.ToString(setting.AlertDays)) } func (uc *MonitorUsecase) Clear() error { diff --git a/internal/biz/notify.go b/internal/biz/notify.go new file mode 100644 index 00000000..ca154f73 --- /dev/null +++ b/internal/biz/notify.go @@ -0,0 +1,319 @@ +package biz + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "html" + "log/slog" + "slices" + "strings" + "time" + + "github.com/leonelquinteros/gotext" + "github.com/libtnb/utils/crypt" + "github.com/samber/do/v2" + "gorm.io/gorm" + + "github.com/acepanel/panel/v3/internal/app" + "github.com/acepanel/panel/v3/internal/request" + "github.com/acepanel/panel/v3/pkg/notify" +) + +// NotifyChannel 通知渠道 +type NotifyChannel struct { + ID uint `gorm:"primaryKey" json:"id"` + Name string `gorm:"not null;default:''" json:"name"` + Type string `gorm:"not null;default:''" json:"type"` // smtp + Config json.RawMessage `gorm:"not null;default:''" json:"config"` // 渠道配置,含凭据,落库前整体加密 + Enabled bool `gorm:"not null;default:true" json:"enabled"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +func (r *NotifyChannel) BeforeSave(tx *gorm.DB) error { + crypter, err := crypt.NewXChacha20Poly1305([]byte(app.Key)) + if err != nil { + return err + } + + encrypted, err := crypter.Encrypt(r.Config) + if err != nil { + return err + } + r.Config = json.RawMessage(encrypted) + + return nil +} + +func (r *NotifyChannel) AfterFind(tx *gorm.DB) error { + crypter, err := crypt.NewXChacha20Poly1305([]byte(app.Key)) + if err != nil { + return err + } + + if config, err := crypter.Decrypt(string(r.Config)); err == nil { + r.Config = config + } + + return nil +} + +// NotifyEvent 系统事件类型 +type NotifyEvent string + +const ( + NotifyEventCertRenew NotifyEvent = "cert_renew" // 证书续签失败 + NotifyEventBackup NotifyEvent = "backup" // 备份失败 + NotifyEventTaskFailed NotifyEvent = "task_failed" // 后台任务失败 + NotifyEventCronFailed NotifyEvent = "cron_failed" // 计划任务执行失败 + NotifyEventWebsiteExpire NotifyEvent = "website_expire" // 网站到期关停 + NotifyEventTamper NotifyEvent = "tamper" // 防篡改拦截 + NotifyEventHealth NotifyEvent = "health" // 面板健康问题 + NotifyEventLogin NotifyEvent = "login" // 面板登录 + NotifyEventLoginFailed NotifyEvent = "login_failed" // 面板登录失败过多 + NotifyEventSSHLogin NotifyEvent = "ssh_login" // SSH 登录 + NotifyEventSSHBruteforce NotifyEvent = "ssh_bruteforce" // SSH 爆破 +) + +type NotifyChannelRepo interface { + List(page, limit uint) ([]*NotifyChannel, int64, error) + All() ([]*NotifyChannel, error) + Get(id uint) (*NotifyChannel, error) + GetByIDs(ids []uint) ([]*NotifyChannel, error) + Create(channel *NotifyChannel) error + Update(channel *NotifyChannel) error + Delete(id uint) error +} + +// notifyMaxPending 异步事件通知的并发上限,防止高频事件堆积 goroutine +const notifyMaxPending = 32 + +type NotifyUsecase struct { + repo NotifyChannelRepo + setting SettingRepo + log *slog.Logger + t *gotext.Locale + pending chan struct{} +} + +func NewNotifyUsecase(i do.Injector) (*NotifyUsecase, error) { + return &NotifyUsecase{ + repo: do.MustInvoke[NotifyChannelRepo](i), + setting: do.MustInvoke[SettingRepo](i), + log: do.MustInvoke[*slog.Logger](i), + t: do.MustInvoke[*gotext.Locale](i), + pending: make(chan struct{}, notifyMaxPending), + }, nil +} + +func (uc *NotifyUsecase) List(page, limit uint) ([]*NotifyChannel, int64, error) { + return uc.repo.List(page, limit) +} + +func (uc *NotifyUsecase) All() ([]*NotifyChannel, error) { + return uc.repo.All() +} + +func (uc *NotifyUsecase) Get(id uint) (*NotifyChannel, error) { + return uc.repo.Get(id) +} + +func (uc *NotifyUsecase) Create(ctx context.Context, req *request.NotifyChannelCreate) (*NotifyChannel, error) { + // 提前构造一次,配置不合法直接拒绝入库 + if _, err := notify.New(req.Type, req.Config); err != nil { + return nil, errors.New(uc.t.Get("invalid channel config: %v", err)) + } + + channel := &NotifyChannel{ + Name: req.Name, + Type: req.Type, + Config: req.Config, + Enabled: req.Enabled, + } + if err := uc.repo.Create(channel); err != nil { + return nil, err + } + + uc.log.Info("notify channel created", slog.String("type", OperationTypeSetting), slog.Uint64("operator_id", operatorID(ctx)), slog.String("name", req.Name)) + + // 落库时 Config 已被加密,重新读取以返回解密后的实体 + return uc.repo.Get(channel.ID) +} + +func (uc *NotifyUsecase) Update(ctx context.Context, req *request.NotifyChannelUpdate) error { + channel, err := uc.repo.Get(req.ID) + if err != nil { + return err + } + if _, err = notify.New(req.Type, req.Config); err != nil { + return errors.New(uc.t.Get("invalid channel config: %v", err)) + } + + channel.Name = req.Name + channel.Type = req.Type + channel.Config = req.Config + channel.Enabled = req.Enabled + if err = uc.repo.Update(channel); err != nil { + return err + } + + uc.log.Info("notify channel updated", slog.String("type", OperationTypeSetting), slog.Uint64("operator_id", operatorID(ctx)), slog.Uint64("id", uint64(req.ID)), slog.String("name", req.Name)) + + return nil +} + +func (uc *NotifyUsecase) Delete(ctx context.Context, id uint) error { + channel, err := uc.repo.Get(id) + if err != nil { + return err + } + if err = uc.repo.Delete(id); err != nil { + return err + } + + uc.log.Info("notify channel deleted", slog.String("type", OperationTypeSetting), slog.Uint64("operator_id", operatorID(ctx)), slog.Uint64("id", uint64(id)), slog.String("name", channel.Name)) + + return nil +} + +// Test 向指定渠道发送一条测试消息 +func (uc *NotifyUsecase) Test(ctx context.Context, id uint) error { + channel, err := uc.repo.Get(id) + if err != nil { + return err + } + + return uc.dispatch(ctx, channel, uc.t.Get("[AcePanel] Test Notification"), + NotifyBody(uc.t.Get("This is a test notification from AcePanel, receiving it means the channel is configured correctly."), [][2]string{ + {uc.t.Get("Channel"), channel.Name}, + {uc.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) +} + +// Send 向指定渠道列表发送通知,返回成功发送的渠道数,部分渠道失败不影响其他渠道 +func (uc *NotifyUsecase) Send(ctx context.Context, channelIDs []uint, subject, body string) (int, error) { + if len(channelIDs) == 0 { + return 0, nil + } + + channels, err := uc.repo.GetByIDs(channelIDs) + if err != nil { + return 0, err + } + + var sent int + var errs []error + for _, channel := range channels { + if !channel.Enabled { + continue + } + if err = uc.dispatch(ctx, channel, subject, body); err != nil { + errs = append(errs, fmt.Errorf("%s: %w", channel.Name, err)) + continue + } + sent++ + } + + return sent, errors.Join(errs...) +} + +// SendEvent 发送系统事件通知,不阻塞业务流程 +// 待发送数超过上限时丢弃并告知,避免慢渠道拖垮调用方 +func (uc *NotifyUsecase) SendEvent(event NotifyEvent, subject, body string) { + select { + case uc.pending <- struct{}{}: + default: + uc.log.Warn("event notification dropped, too many pending sends", slog.String("event", string(event))) + return + } + + go func() { + defer func() { <-uc.pending }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() + if err := uc.SendEventSync(ctx, event, subject, body); err != nil { + uc.log.Warn("failed to send event notification", slog.String("event", string(event)), slog.Any("err", err)) + } + }() +} + +// SendEventSync 同步发送系统事件通知,未订阅该事件或未配置渠道时静默跳过 +// 供 CLI 等短生命周期进程使用,异步发送会随进程退出丢失 +func (uc *NotifyUsecase) SendEventSync(ctx context.Context, event NotifyEvent, subject, body string) error { + setting, err := uc.GetSetting() + if err != nil || len(setting.Channels) == 0 || !slices.Contains(setting.Events, string(event)) { + return err + } + + _, err = uc.Send(ctx, setting.Channels, subject, body) + + return err +} + +func (uc *NotifyUsecase) GetSetting() (*request.NotifySetting, error) { + events, err := uc.setting.GetSlice(SettingKeyNotifyEvents) + if err != nil { + return nil, err + } + channelsStr, err := uc.setting.Get(SettingKeyNotifyEventChannels) + if err != nil { + return nil, err + } + + channels := make([]uint, 0) + if channelsStr != "" { + _ = json.Unmarshal([]byte(channelsStr), &channels) + } + + return &request.NotifySetting{ + Events: events, + Channels: channels, + }, nil +} + +func (uc *NotifyUsecase) UpdateSetting(setting *request.NotifySetting) error { + if err := uc.setting.SetSlice(SettingKeyNotifyEvents, setting.Events); err != nil { + return err + } + + channels, err := json.Marshal(setting.Channels) + if err != nil { + return err + } + + return uc.setting.Set(SettingKeyNotifyEventChannels, string(channels)) +} + +func (uc *NotifyUsecase) dispatch(ctx context.Context, channel *NotifyChannel, subject, body string) error { + notifier, err := notify.New(channel.Type, channel.Config) + if err != nil { + return err + } + + return notifier.Send(ctx, ¬ify.Message{Subject: subject, Body: body}) +} + +// NotifyBody 构建通知正文,rows 为「名称,值」明细 +func NotifyBody(summary string, rows [][2]string) string { + var sb strings.Builder + sb.WriteString(`

`) + sb.WriteString(html.EscapeString(summary)) + sb.WriteString(`

`) + + if len(rows) > 0 { + sb.WriteString(``) + for _, row := range rows { + sb.WriteString(``) + } + sb.WriteString(`
`) + sb.WriteString(html.EscapeString(row[0])) + sb.WriteString(``) + sb.WriteString(html.EscapeString(row[1])) + sb.WriteString(`
`) + } + + return sb.String() +} diff --git a/internal/biz/setting.go b/internal/biz/setting.go index de75ef6d..15eb9e83 100644 --- a/internal/biz/setting.go +++ b/internal/biz/setting.go @@ -66,9 +66,12 @@ const ( SettingKeyIPDBPath SettingKey = "ipdb_path" SettingKeyInfoRan SettingKey = "info_ran" // info 命令是否已运行过 SettingKeyTamperEnabled SettingKey = "tamper_enabled" - SettingKeyTamperMode SettingKey = "tamper_mode" // chattr / ebpf - SettingKeyTamperBlockNew SettingKey = "tamper_block_new" // 新建受保护类型文件时删除拦截 - SettingKeyTamperLogDays SettingKey = "tamper_log_days" // 拦截日志保留天数 + SettingKeyTamperMode SettingKey = "tamper_mode" // chattr / ebpf + SettingKeyTamperBlockNew SettingKey = "tamper_block_new" // 新建受保护类型文件时删除拦截 + SettingKeyTamperLogDays SettingKey = "tamper_log_days" // 拦截日志保留天数 + SettingKeyNotifyEvents SettingKey = "notify_event_types" // 订阅的系统事件类型,JSON 数组 + SettingKeyNotifyEventChannels SettingKey = "notify_event_channels" // 接收系统事件的渠道 ID,JSON 数组 + SettingKeyAlertLogDays SettingKey = "alert_log_days" // 告警记录保留天数 ) type Setting struct { diff --git a/internal/biz/tamper.go b/internal/biz/tamper.go index ac52834a..b75cccb9 100644 --- a/internal/biz/tamper.go +++ b/internal/biz/tamper.go @@ -2,6 +2,7 @@ package biz import ( "errors" + "fmt" "log/slog" "os" "path/filepath" @@ -9,6 +10,7 @@ import ( "sync" "time" + "github.com/leonelquinteros/gotext" "github.com/samber/do/v2" "github.com/spf13/cast" @@ -64,20 +66,25 @@ type TamperRepo interface { type TamperUsecase struct { repo TamperRepo setting *SettingUsecase + notify *NotifyUsecase log *slog.Logger + t *gotext.Locale - mu sync.Mutex - mgr *tamper.Manager - buf []*TamperLog - bufMu sync.Mutex - drainC chan struct{} + mu sync.Mutex + mgr *tamper.Manager + buf []*TamperLog + bufMu sync.Mutex + notifyAt time.Time + drainC chan struct{} } func NewTamperUsecase(i do.Injector) (*TamperUsecase, error) { return &TamperUsecase{ repo: do.MustInvoke[TamperRepo](i), setting: do.MustInvoke[*SettingUsecase](i), + notify: do.MustInvoke[*NotifyUsecase](i), log: do.MustInvoke[*slog.Logger](i), + t: do.MustInvoke[*gotext.Locale](i), }, nil } @@ -277,7 +284,30 @@ func (uc *TamperUsecase) FlushLogs() { if err := uc.repo.AddLogs(logs); err != nil { uc.log.Warn("failed to persist tamper logs", slog.Any("err", err)) + return } + + uc.notifyBlocked(logs) +} + +// notifyBlocked 汇总上报拦截事件,5 分钟内不重复通知 +func (uc *TamperUsecase) notifyBlocked(logs []*TamperLog) { + uc.bufMu.Lock() + if time.Since(uc.notifyAt) < 5*time.Minute { + uc.bufMu.Unlock() + return + } + uc.notifyAt = time.Now() + uc.bufMu.Unlock() + + latest := logs[len(logs)-1] + uc.notify.SendEvent(NotifyEventTamper, uc.t.Get("[AcePanel] Tamper Protection Alert"), NotifyBody(uc.t.Get("tamper protection blocked file operations"), [][2]string{ + {uc.t.Get("Count"), cast.ToString(len(logs))}, + {uc.t.Get("Path"), latest.Path}, + {uc.t.Get("Operation"), latest.Op}, + {uc.t.Get("Process"), fmt.Sprintf("%s (%d)", latest.Comm, latest.PID)}, + {uc.t.Get("Time"), latest.CreatedAt.Format(time.DateTime)}, + })) } // CleanupLogs 清理过期日志 diff --git a/internal/biz/website.go b/internal/biz/website.go index 5ff0491a..c705b72b 100644 --- a/internal/biz/website.go +++ b/internal/biz/website.go @@ -135,7 +135,7 @@ func (uc *WebsiteUsecase) Create(ctx context.Context, req *request.WebsiteCreate // 创建数据库 name := "local_" + req.DBType if req.DB { - server, err := uc.databaseServer.GetByName(name) + server, err := uc.databaseServer.GetByName(ctx, name) if err != nil { return nil, errors.New(uc.t.Get("can't find %s database server, please add it first", name)) } @@ -190,12 +190,12 @@ func (uc *WebsiteUsecase) Delete(ctx context.Context, req *request.WebsiteDelete _ = uc.repo.RemoveFiles(website.Name, req.Path) if req.DB { - if mysql, err := uc.databaseServer.GetByName("local_mysql"); err == nil { - _ = uc.databaseUser.DeleteByNames(mysql.ID, []string{website.Name}) + if mysql, err := uc.databaseServer.GetByName(ctx, "local_mysql"); err == nil { + _ = uc.databaseUser.DeleteByNames(ctx, mysql.ID, []string{website.Name}) _ = uc.database.Delete(ctx, mysql.ID, website.Name) } - if postgres, err := uc.databaseServer.GetByName("local_postgresql"); err == nil { - _ = uc.databaseUser.DeleteByNames(postgres.ID, []string{website.Name}) + if postgres, err := uc.databaseServer.GetByName(ctx, "local_postgresql"); err == nil { + _ = uc.databaseUser.DeleteByNames(ctx, postgres.ID, []string{website.Name}) _ = uc.database.Delete(ctx, postgres.ID, website.Name) } } diff --git a/internal/bootstrap/queue.go b/internal/bootstrap/queue.go index bb88bdca..3ec1d70b 100644 --- a/internal/bootstrap/queue.go +++ b/internal/bootstrap/queue.go @@ -3,14 +3,21 @@ package bootstrap import ( "log/slog" + "github.com/leonelquinteros/gotext" "github.com/samber/do/v2" "gorm.io/gorm" + "github.com/acepanel/panel/v3/internal/biz" "github.com/acepanel/panel/v3/internal/taskqueue" "github.com/acepanel/panel/v3/pkg/types" ) // NewRunner 创建任务运行器 func NewRunner(i do.Injector) (types.TaskRunner, error) { - return taskqueue.NewRunner(do.MustInvoke[*gorm.DB](i), do.MustInvoke[*slog.Logger](i)), nil + return taskqueue.NewRunner( + do.MustInvoke[*gorm.DB](i), + do.MustInvoke[*slog.Logger](i), + do.MustInvoke[*biz.NotifyUsecase](i), + do.MustInvoke[*gotext.Locale](i), + ), nil } diff --git a/internal/command/command.go b/internal/command/command.go index 7935347c..cca28b39 100644 --- a/internal/command/command.go +++ b/internal/command/command.go @@ -21,7 +21,7 @@ var Package = do.Package( do.LazyNamed(Prefix+"bind-ip", BindIPCommand), do.LazyNamed(Prefix+"bind-ua", BindUACommand), do.LazyNamed(Prefix+"website", WebsiteCommand), do.LazyNamed(Prefix+"database", DatabaseCommand), do.LazyNamed(Prefix+"backup", BackupCommand), do.LazyNamed(Prefix+"restore", RestoreCommand), - do.LazyNamed(Prefix+"cutoff", CutoffCommand), + do.LazyNamed(Prefix+"cutoff", CutoffCommand), do.LazyNamed(Prefix+"cron", CronCommand), do.LazyNamed(Prefix+"app", AppCommand), do.LazyNamed(Prefix+"setting", SettingCommand), ) diff --git a/internal/command/cron.go b/internal/command/cron.go new file mode 100644 index 00000000..c0da3bcc --- /dev/null +++ b/internal/command/cron.go @@ -0,0 +1,42 @@ +package command + +import ( + "context" + + "github.com/leonelquinteros/gotext" + "github.com/samber/do/v2" + "github.com/urfave/cli/v3" + + "github.com/acepanel/panel/v3/internal/service" +) + +// CronCommand 计划任务命令组 +func CronCommand(i do.Injector) (*cli.Command, error) { + t := do.MustInvoke[*gotext.Locale](i) + return &cli.Command{ + Name: "cron", + Usage: t.Get("Cron task"), + Commands: []*cli.Command{ + { + Name: "failed", + Usage: t.Get("Report a failed cron task, called by the task wrapper script"), + Flags: []cli.Flag{ + &cli.UintFlag{ + Name: "id", + Aliases: []string{"i"}, + Usage: t.Get("Cron task ID"), + Required: true, + }, + &cli.IntFlag{ + Name: "code", + Aliases: []string{"c"}, + Usage: t.Get("Exit code"), + }, + }, + Action: func(ctx context.Context, cmd *cli.Command) error { + return do.MustInvoke[*service.CliService](i).CronFailed(ctx, cmd) + }, + }, + }, + }, nil +} diff --git a/internal/data/alert.go b/internal/data/alert.go new file mode 100644 index 00000000..a68b562c --- /dev/null +++ b/internal/data/alert.go @@ -0,0 +1,160 @@ +package data + +import ( + "strings" + "time" + + "github.com/samber/do/v2" + "github.com/samber/lo" + "gorm.io/gorm" + + "github.com/acepanel/panel/v3/internal/biz" + "github.com/acepanel/panel/v3/pkg/cert" +) + +type alertRepo struct { + db *gorm.DB + statDB *gorm.DB // 网站统计独立库 +} + +func NewAlertRepo(i do.Injector) (biz.AlertRepo, error) { + statDB, err := openDB("stat") + if err != nil { + return nil, err + } + + return &alertRepo{ + db: do.MustInvoke[*gorm.DB](i), + statDB: statDB, + }, nil +} + +func (r *alertRepo) ListRules(page, limit uint) ([]*biz.AlertRule, int64, error) { + rules := make([]*biz.AlertRule, 0) + var total int64 + err := r.db.Model(&biz.AlertRule{}).Order("id desc").Count(&total).Offset(int((page - 1) * limit)).Limit(int(limit)).Find(&rules).Error + return rules, total, err +} + +func (r *alertRepo) AllRules() ([]*biz.AlertRule, error) { + rules := make([]*biz.AlertRule, 0) + err := r.db.Order("id desc").Find(&rules).Error + return rules, err +} + +func (r *alertRepo) GetRule(id uint) (*biz.AlertRule, error) { + rule := new(biz.AlertRule) + if err := r.db.Where("id = ?", id).First(rule).Error; err != nil { + return nil, err + } + return rule, nil +} + +func (r *alertRepo) CreateRule(rule *biz.AlertRule) error { + return r.db.Create(rule).Error +} + +func (r *alertRepo) UpdateRule(rule *biz.AlertRule) error { + return r.db.Save(rule).Error +} + +func (r *alertRepo) DeleteRule(id uint) error { + return r.db.Where("id = ?", id).Delete(&biz.AlertRule{}).Error +} + +func (r *alertRepo) AddAlert(alert *biz.Alert) error { + return r.db.Create(alert).Error +} + +func (r *alertRepo) ListAlerts(page, limit uint) ([]*biz.Alert, int64, error) { + alerts := make([]*biz.Alert, 0) + var total int64 + err := r.db.Model(&biz.Alert{}).Order("id desc").Count(&total).Offset(int((page - 1) * limit)).Limit(int(limit)).Find(&alerts).Error + return alerts, total, err +} + +func (r *alertRepo) ClearAlerts() error { + return r.db.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&biz.Alert{}).Error +} + +func (r *alertRepo) ClearAlertsBefore(t time.Time) error { + return r.db.Where("created_at < ?", t).Delete(&biz.Alert{}).Error +} + +func (r *alertRepo) CertExpiry() ([]*biz.AlertMetric, error) { + certs := make([]*biz.Cert, 0) + if err := r.db.Find(&certs).Error; err != nil { + return nil, err + } + + metrics := make([]*biz.AlertMetric, 0, len(certs)) + for _, item := range certs { + if item.Cert == "" { + continue + } + decode, err := cert.ParseCert([]byte(item.Cert)) + if err != nil { + continue + } + // 一张证书一条指标,目标是逗号分隔的全部域名,规则可用其中任一域名匹配 + metrics = append(metrics, &biz.AlertMetric{ + Target: strings.Join(item.Domains, ","), + Value: time.Until(decode.NotAfter).Hours() / 24, + }) + } + + return metrics, nil +} + +// WebsiteHourStats 取各网站当前自然小时的请求统计 +// 统计按自然小时聚合,整点归零,因此规则语义是「本小时累计」而非滚动一小时 +func (r *alertRepo) WebsiteHourStats() ([]*biz.WebsiteHourStat, error) { + now := time.Now() + stats := make([]*biz.WebsiteStat, 0) + if err := r.statDB.Where("date = ? AND hour = ?", now.Format(time.DateOnly), now.Hour()).Find(&stats).Error; err != nil { + return nil, err + } + + return lo.Map(stats, func(item *biz.WebsiteStat, _ int) *biz.WebsiteHourStat { + return &biz.WebsiteHourStat{ + Site: item.Site, + Requests: item.Requests, + Errors: item.Errors, + Status5xx: item.Status5xx, + } + }), nil +} + +func (r *alertRepo) ProjectNames() ([]string, error) { + names := make([]string, 0) + err := r.db.Model(&biz.Project{}).Pluck("name", &names).Error + return names, err +} + +// DatabaseServers 取数据库服务器列表 +// 不复用 DatabaseServerRepo.List:它会串行探测每台服务器,探测由调用方按需并发进行 +func (r *alertRepo) DatabaseServers() ([]*biz.DatabaseServer, error) { + servers := make([]*biz.DatabaseServer, 0) + err := r.db.Order("id desc").Find(&servers).Error + return servers, err +} + +func (r *alertRepo) WebsiteExpiry() ([]*biz.AlertMetric, error) { + websites := make([]*biz.Website, 0) + if err := r.db.Where("expire_at IS NOT NULL").Find(&websites).Error; err != nil { + return nil, err + } + + metrics := make([]*biz.AlertMetric, 0, len(websites)) + for _, item := range websites { + if item.ExpireAt == nil { + continue + } + metrics = append(metrics, &biz.AlertMetric{ + Target: item.Name, + Value: time.Until(*item.ExpireAt).Hours() / 24, + }) + } + + return metrics, nil +} diff --git a/internal/data/backup.go b/internal/data/backup.go index bab355ac..14dfa51b 100644 --- a/internal/data/backup.go +++ b/internal/data/backup.go @@ -571,7 +571,7 @@ func (r *backupRepo) createMySQL(name string, storage storage.Storage, target st if err != nil { return err } - mysql, err := db.NewMySQL("root", rootPassword, "/tmp/mysql.sock", "unix") + mysql, err := db.NewMySQL(context.Background(), "root", rootPassword, "/tmp/mysql.sock", "unix") if err != nil { return err } @@ -629,7 +629,7 @@ func (r *backupRepo) createPostgres(name string, storage storage.Storage, target if err != nil { return err } - postgres, err := db.NewPostgres("postgres", postgresPassword, "127.0.0.1", 5432) + postgres, err := db.NewPostgres(context.Background(), "postgres", postgresPassword, "127.0.0.1", 5432) if err != nil { return err } @@ -937,7 +937,7 @@ func (r *backupRepo) restoreMySQL(backup, target string) error { if err != nil { return err } - mysql, err := db.NewMySQL("root", rootPassword, "/tmp/mysql.sock", "unix") + mysql, err := db.NewMySQL(context.Background(), "root", rootPassword, "/tmp/mysql.sock", "unix") if err != nil { return err } @@ -980,7 +980,7 @@ func (r *backupRepo) restorePostgres(backup, target string) error { if err != nil { return err } - postgres, err := db.NewPostgres("postgres", postgresPassword, "127.0.0.1", 5432) + postgres, err := db.NewPostgres(context.Background(), "postgres", postgresPassword, "127.0.0.1", 5432) if err != nil { return err } diff --git a/internal/data/cron.go b/internal/data/cron.go index 48816a35..0be78d30 100644 --- a/internal/data/cron.go +++ b/internal/data/cron.go @@ -113,27 +113,27 @@ func (r *cronRepo) RemoveScriptFiles(shellPath string) error { } // AddToSystem 添加到系统 +// 统一经 wrapper 脚本执行,以便捕获退出码并上报失败 func (r *cronRepo) AddToSystem(cron *biz.Cron) error { cmd := cron.Shell if cron.Config.Flock { lockFile := strings.TrimSuffix(cron.Shell, ".sh") + ".lock" - cmd = fmt.Sprintf("flock -xn %s %s", lockFile, cron.Shell) + // -E 指定未抢到锁时的退出码,与脚本自身失败区分,避免正常跳过被误报 + cmd = fmt.Sprintf("flock -xn -E %d %s %s", cronLockSkipCode, lockFile, cron.Shell) } - // 秒级任务:生成 wrapper 脚本,用每分钟触发模拟秒级执行 - if seconds := r.parseSeconds(cron.Time); seconds > 0 { - wrapperPath := strings.TrimSuffix(cron.Shell, ".sh") + "_wrapper.sh" - wrapperScript := r.generateWrapper(cmd, cron.Log, seconds) - if err := io.Write(wrapperPath, wrapperScript, 0700); err != nil { - return err - } - if _, err := shell.Execf(`( crontab -l; echo "* * * * * %s" ) | sort - | uniq - | crontab -`, wrapperPath); err != nil { - return err - } - return r.restartCron() + // 秒级任务由每分钟触发的 wrapper 内部循环模拟 + spec := cron.Time + seconds := r.parseSeconds(cron.Time) + if seconds > 0 { + spec = "* * * * *" } - if _, err := shell.Execf(`( crontab -l; echo "%s %s >> %s 2>&1" ) | sort - | uniq - | crontab -`, cron.Time, cmd, cron.Log); err != nil { + wrapperPath := strings.TrimSuffix(cron.Shell, ".sh") + "_wrapper.sh" + if err := io.Write(wrapperPath, r.generateWrapper(cron.ID, cmd, cron.Log, seconds), 0700); err != nil { + return err + } + if _, err := shell.Execf(`( crontab -l; echo "%s %s" ) | sort - | uniq - | crontab -`, spec, wrapperPath); err != nil { return err } @@ -277,17 +277,45 @@ func (r *cronRepo) parseSeconds(time string) int { return 0 } -// generateWrapper 生成秒级任务的 wrapper 脚本 -// 通过每分钟触发 + 循环 sleep 模拟秒级执行 -func (r *cronRepo) generateWrapper(cmd, logFile string, seconds int) string { +const ( + // wrapperPathEnv crontab 环境的 PATH 极简,需补全以便调用 acepanel + wrapperPathEnv = "export PATH=/bin:/sbin:/usr/bin:/usr/sbin:/usr/local/bin:/usr/local/sbin:$PATH" + // cronLockSkipCode flock 未抢到锁时的退出码,属正常跳过而非失败 + cronLockSkipCode = 200 +) + +// generateWrapper 生成任务的 wrapper 脚本,捕获退出码并上报失败 +// seconds 大于 0 时为秒级任务,用每分钟触发 + 循环 sleep 模拟 +func (r *cronRepo) generateWrapper(id uint, cmd, logFile string, seconds int) string { + if seconds <= 0 { + return fmt.Sprintf(`#!/bin/bash +%s + +%s >> %s 2>&1 +code=$? +if [ $code -ne 0 ] && [ $code -ne %d ]; then + acepanel cron failed -i %d -c $code >/dev/null 2>&1 +fi +exit $code +`, wrapperPathEnv, cmd, logFile, cronLockSkipCode, id) + } + + // 并发执行拿不到子进程退出码,用标记文件汇总本分钟是否出错 count := 60 / seconds return fmt.Sprintf(`#!/bin/bash +%s + INTERVAL=%d COUNT=%d +FLAG=$(mktemp) for i in $(seq 1 $COUNT); do - %s >> %s 2>&1 & + ( %s >> %s 2>&1; c=$?; [ $c -ne 0 ] && [ $c -ne %d ] && echo $c >> "$FLAG" ) & [ $i -lt $COUNT ] && sleep $INTERVAL done wait -`, seconds, count, cmd, logFile) +if [ -s "$FLAG" ]; then + acepanel cron failed -i %d -c "$(tail -n 1 "$FLAG")" >/dev/null 2>&1 +fi +rm -f "$FLAG" +`, wrapperPathEnv, seconds, count, cmd, logFile, cronLockSkipCode, id) } diff --git a/internal/data/data.go b/internal/data/data.go index 15b9675c..22fe269a 100644 --- a/internal/data/data.go +++ b/internal/data/data.go @@ -5,13 +5,14 @@ import ( ) var Package = do.Package( - do.Lazy(NewAppRepo), do.Lazy(NewBackupRepo), do.Lazy(NewBackupAccountRepo), + do.Lazy(NewAlertRepo), do.Lazy(NewAppRepo), do.Lazy(NewBackupRepo), do.Lazy(NewBackupAccountRepo), do.Lazy(NewCacheRepo), do.Lazy(NewCertRepo), do.Lazy(NewCertAccountRepo), do.Lazy(NewCertDNSRepo), do.Lazy(NewContainerRepo), do.Lazy(NewContainerComposeRepo), do.Lazy(NewContainerImageRepo), do.Lazy(NewContainerNetworkRepo), do.Lazy(NewContainerVolumeRepo), do.Lazy(NewCronRepo), do.Lazy(NewDatabaseRepo), do.Lazy(NewDatabaseRedisRepo), do.Lazy(NewDatabaseElasticsearchRepo), do.Lazy(NewDatabaseServerRepo), do.Lazy(NewDatabaseUserRepo), do.Lazy(NewEnvironmentRepo), do.Lazy(NewFileShareRepo), do.Lazy(NewLogRepo), do.Lazy(NewMonitorRepo), + do.Lazy(NewNotifyChannelRepo), do.Lazy(NewProjectRepo), do.Lazy(NewSafeRepo), do.Lazy(NewScanEventRepo), do.Lazy(NewSettingRepo), do.Lazy(NewSSHRepo), do.Lazy(NewTamperRepo), do.Lazy(NewTaskRepo), do.Lazy(NewTemplateRepo), do.Lazy(NewUserRepo), do.Lazy(NewUserPasskeyRepo), diff --git a/internal/data/database.go b/internal/data/database.go index 1c2cc03f..95d80d79 100644 --- a/internal/data/database.go +++ b/internal/data/database.go @@ -1,6 +1,7 @@ package data import ( + "context" "fmt" "slices" @@ -35,11 +36,11 @@ func (r *databaseRepo) ListServers(typ string) ([]*biz.DatabaseServer, error) { } // DatabasesOf 列出单个服务器上的数据库 -func (r *databaseRepo) DatabasesOf(server *biz.DatabaseServer) ([]*biz.Database, error) { +func (r *databaseRepo) DatabasesOf(ctx context.Context, server *biz.DatabaseServer) ([]*biz.Database, error) { database := make([]*biz.Database, 0) switch server.Type { case biz.DatabaseTypeMongoDB: - mongo, err := db.NewMongoDB(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + mongo, err := db.NewMongoDB(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err != nil { return nil, err } @@ -71,7 +72,7 @@ func (r *databaseRepo) DatabasesOf(server *biz.DatabaseServer) ([]*biz.Database, } sqlite.Close() default: - operator, err := r.Operator(server) + operator, err := r.Operator(ctx, server) if err != nil { return nil, err } @@ -93,22 +94,22 @@ func (r *databaseRepo) DatabasesOf(server *biz.DatabaseServer) ([]*biz.Database, } // Mongo 构建 MongoDB 客户端 -func (r *databaseRepo) Mongo(server *biz.DatabaseServer) (*db.MongoDB, error) { - return db.NewMongoDB(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) +func (r *databaseRepo) Mongo(ctx context.Context, server *biz.DatabaseServer) (*db.MongoDB, error) { + return db.NewMongoDB(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) } -func (r *databaseRepo) Operator(server *biz.DatabaseServer) (db.Operator, error) { +func (r *databaseRepo) Operator(ctx context.Context, server *biz.DatabaseServer) (db.Operator, error) { switch server.Type { case biz.DatabaseTypeMysql: - return newMySQLOperator(server.Username, server.Password, server.Host, server.Port) + return newMySQLOperator(ctx, server.Username, server.Password, server.Host, server.Port) case biz.DatabaseTypePostgresql: - postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) + postgres, err := db.NewPostgres(ctx, server.Username, server.Password, server.Host, server.Port) if err != nil { return nil, err } return postgres, nil case biz.DatabaseTypeClickHouse: - clickhouse, err := db.NewClickHouse(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + clickhouse, err := db.NewClickHouse(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err != nil { return nil, err } diff --git a/internal/data/database_elasticsearch.go b/internal/data/database_elasticsearch.go index b4623949..2469d886 100644 --- a/internal/data/database_elasticsearch.go +++ b/internal/data/database_elasticsearch.go @@ -1,6 +1,7 @@ package data import ( + "context" "errors" "fmt" "log/slog" @@ -28,8 +29,8 @@ func NewDatabaseElasticsearchRepo(i do.Injector) (biz.DatabaseElasticsearchRepo, }, nil } -func (r *databaseElasticsearchRepo) Indices(req *request.DatabaseESIndices) ([]db.ESIndex, error) { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) Indices(ctx context.Context, req *request.DatabaseESIndices) ([]db.ESIndex, error) { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return nil, err } @@ -38,8 +39,8 @@ func (r *databaseElasticsearchRepo) Indices(req *request.DatabaseESIndices) ([]d return client.Indices() } -func (r *databaseElasticsearchRepo) IndexCreate(req *request.DatabaseESIndexCreate) error { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) IndexCreate(ctx context.Context, req *request.DatabaseESIndexCreate) error { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return err } @@ -48,8 +49,8 @@ func (r *databaseElasticsearchRepo) IndexCreate(req *request.DatabaseESIndexCrea return client.IndexCreate(req.Name) } -func (r *databaseElasticsearchRepo) IndexDelete(req *request.DatabaseESIndexDelete) error { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) IndexDelete(ctx context.Context, req *request.DatabaseESIndexDelete) error { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return err } @@ -58,8 +59,8 @@ func (r *databaseElasticsearchRepo) IndexDelete(req *request.DatabaseESIndexDele return client.IndexDelete(req.Name) } -func (r *databaseElasticsearchRepo) Data(req *request.DatabaseESData) ([]db.ESDocument, int64, error) { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) Data(ctx context.Context, req *request.DatabaseESData) ([]db.ESDocument, int64, error) { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return nil, 0, err } @@ -68,8 +69,8 @@ func (r *databaseElasticsearchRepo) Data(req *request.DatabaseESData) ([]db.ESDo return client.Search(req.Index, req.Search, int(req.Page), int(req.Limit)) } -func (r *databaseElasticsearchRepo) DocumentGet(req *request.DatabaseESDocumentGet) (*db.ESDocument, error) { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) DocumentGet(ctx context.Context, req *request.DatabaseESDocumentGet) (*db.ESDocument, error) { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return nil, err } @@ -78,8 +79,8 @@ func (r *databaseElasticsearchRepo) DocumentGet(req *request.DatabaseESDocumentG return client.DocumentGet(req.Index, req.ID) } -func (r *databaseElasticsearchRepo) DocumentSet(req *request.DatabaseESDocumentSet) error { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) DocumentSet(ctx context.Context, req *request.DatabaseESDocumentSet) error { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return err } @@ -91,8 +92,8 @@ func (r *databaseElasticsearchRepo) DocumentSet(req *request.DatabaseESDocumentS return client.DocumentUpdate(req.Index, req.ID, req.Body) } -func (r *databaseElasticsearchRepo) DocumentDelete(req *request.DatabaseESDocumentDelete) error { - client, err := r.getClient(req.ServerID) +func (r *databaseElasticsearchRepo) DocumentDelete(ctx context.Context, req *request.DatabaseESDocumentDelete) error { + client, err := r.getClient(ctx, req.ServerID) if err != nil { return err } @@ -102,7 +103,7 @@ func (r *databaseElasticsearchRepo) DocumentDelete(req *request.DatabaseESDocume } // getClient 根据服务器 ID 创建 Elasticsearch 客户端 -func (r *databaseElasticsearchRepo) getClient(serverID uint) (*db.Elasticsearch, error) { +func (r *databaseElasticsearchRepo) getClient(ctx context.Context, serverID uint) (*db.Elasticsearch, error) { server := new(biz.DatabaseServer) if err := r.orm.Where("id = ?", serverID).First(server).Error; err != nil { return nil, errors.New(r.t.Get("server not found")) @@ -111,7 +112,7 @@ func (r *databaseElasticsearchRepo) getClient(serverID uint) (*db.Elasticsearch, return nil, errors.New(r.t.Get("server is not Elasticsearch type")) } - client, err := db.NewElasticsearch(fmt.Sprintf("%s:%d", server.Host, server.Port), server.Username, server.Password) + client, err := db.NewElasticsearch(ctx, fmt.Sprintf("%s:%d", server.Host, server.Port), server.Username, server.Password) if err != nil { return nil, errors.New(r.t.Get("failed to connect to Elasticsearch: %v", err)) } diff --git a/internal/data/database_redis.go b/internal/data/database_redis.go index 28ebb228..5629bfa2 100644 --- a/internal/data/database_redis.go +++ b/internal/data/database_redis.go @@ -1,6 +1,7 @@ package data import ( + "context" "errors" "fmt" "log/slog" @@ -28,8 +29,8 @@ func NewDatabaseRedisRepo(i do.Injector) (biz.DatabaseRedisRepo, error) { }, nil } -func (r *databaseRedisRepo) Databases(req *request.DatabaseRedisDatabases) (int, error) { - client, err := r.getClient(req.ServerID, 0) +func (r *databaseRedisRepo) Databases(ctx context.Context, req *request.DatabaseRedisDatabases) (int, error) { + client, err := r.getClient(ctx, req.ServerID, 0) if err != nil { return 0, err } @@ -38,8 +39,8 @@ func (r *databaseRedisRepo) Databases(req *request.DatabaseRedisDatabases) (int, return client.Database() } -func (r *databaseRedisRepo) Data(req *request.DatabaseRedisData) ([]db.RedisKV, int, error) { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) Data(ctx context.Context, req *request.DatabaseRedisData) ([]db.RedisKV, int, error) { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return nil, 0, err } @@ -53,8 +54,8 @@ func (r *databaseRedisRepo) Data(req *request.DatabaseRedisData) ([]db.RedisKV, return client.Search(pattern, int(req.Page), int(req.Limit)) } -func (r *databaseRedisRepo) KeyGet(req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) KeyGet(ctx context.Context, req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return nil, err } @@ -63,8 +64,8 @@ func (r *databaseRedisRepo) KeyGet(req *request.DatabaseRedisKeyGet) (*db.RedisK return client.Get(req.Key) } -func (r *databaseRedisRepo) KeySet(req *request.DatabaseRedisKeySet) error { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) KeySet(ctx context.Context, req *request.DatabaseRedisKeySet) error { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return err } @@ -73,8 +74,8 @@ func (r *databaseRedisRepo) KeySet(req *request.DatabaseRedisKeySet) error { return client.SetKey(req.Key, req.Value, req.Type, req.TTL) } -func (r *databaseRedisRepo) KeyDelete(req *request.DatabaseRedisKeyDelete) error { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) KeyDelete(ctx context.Context, req *request.DatabaseRedisKeyDelete) error { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return err } @@ -83,8 +84,8 @@ func (r *databaseRedisRepo) KeyDelete(req *request.DatabaseRedisKeyDelete) error return client.Del(req.Key) } -func (r *databaseRedisRepo) KeyTTL(req *request.DatabaseRedisKeyTTL) error { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) KeyTTL(ctx context.Context, req *request.DatabaseRedisKeyTTL) error { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return err } @@ -93,8 +94,8 @@ func (r *databaseRedisRepo) KeyTTL(req *request.DatabaseRedisKeyTTL) error { return client.Expire(req.Key, req.TTL) } -func (r *databaseRedisRepo) KeyRename(req *request.DatabaseRedisKeyRename) error { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) KeyRename(ctx context.Context, req *request.DatabaseRedisKeyRename) error { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return err } @@ -103,8 +104,8 @@ func (r *databaseRedisRepo) KeyRename(req *request.DatabaseRedisKeyRename) error return client.Rename(req.OldKey, req.NewKey) } -func (r *databaseRedisRepo) Clear(req *request.DatabaseRedisClear) error { - client, err := r.getClient(req.ServerID, req.DB) +func (r *databaseRedisRepo) Clear(ctx context.Context, req *request.DatabaseRedisClear) error { + client, err := r.getClient(ctx, req.ServerID, req.DB) if err != nil { return err } @@ -114,7 +115,7 @@ func (r *databaseRedisRepo) Clear(req *request.DatabaseRedisClear) error { } // getClient 根据服务器 ID 创建 Redis 客户端并选择指定数据库 -func (r *databaseRedisRepo) getClient(serverID uint, dbIndex int) (*db.Redis, error) { +func (r *databaseRedisRepo) getClient(ctx context.Context, serverID uint, dbIndex int) (*db.Redis, error) { server := new(biz.DatabaseServer) if err := r.orm.Where("id = ?", serverID).First(server).Error; err != nil { return nil, errors.New(r.t.Get("server not found")) @@ -123,7 +124,7 @@ func (r *databaseRedisRepo) getClient(serverID uint, dbIndex int) (*db.Redis, er return nil, errors.New(r.t.Get("server is not Redis type")) } - client, err := db.NewRedis(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + client, err := db.NewRedis(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err != nil { return nil, errors.New(r.t.Get("failed to connect to Redis: %v", err)) } diff --git a/internal/data/database_server.go b/internal/data/database_server.go index 3e827723..6d8090ee 100644 --- a/internal/data/database_server.go +++ b/internal/data/database_server.go @@ -1,6 +1,7 @@ package data import ( + "context" "fmt" "path/filepath" "slices" @@ -33,7 +34,7 @@ func (r *databaseServerRepo) Count() (int64, error) { return count, nil } -func (r *databaseServerRepo) List(page, limit uint, typ string) ([]*biz.DatabaseServer, int64, error) { +func (r *databaseServerRepo) List(ctx context.Context, page, limit uint, typ string) ([]*biz.DatabaseServer, int64, error) { databaseServer := make([]*biz.DatabaseServer, 0) var total int64 query := r.db.Model(&biz.DatabaseServer{}).Order("id desc") @@ -43,30 +44,30 @@ func (r *databaseServerRepo) List(page, limit uint, typ string) ([]*biz.Database err := query.Count(&total).Offset(int((page - 1) * limit)).Limit(int(limit)).Find(&databaseServer).Error for server := range slices.Values(databaseServer) { - r.CheckServer(server) + r.CheckServer(ctx, server) } return databaseServer, total, err } -func (r *databaseServerRepo) Get(id uint) (*biz.DatabaseServer, error) { +func (r *databaseServerRepo) Get(ctx context.Context, id uint) (*biz.DatabaseServer, error) { databaseServer := new(biz.DatabaseServer) if err := r.db.Where("id = ?", id).First(databaseServer).Error; err != nil { return nil, err } - r.CheckServer(databaseServer) + r.CheckServer(ctx, databaseServer) return databaseServer, nil } -func (r *databaseServerRepo) GetByName(name string) (*biz.DatabaseServer, error) { +func (r *databaseServerRepo) GetByName(ctx context.Context, name string) (*biz.DatabaseServer, error) { databaseServer := new(biz.DatabaseServer) if err := r.db.Where("name = ?", name).First(databaseServer).Error; err != nil { return nil, err } - r.CheckServer(databaseServer) + r.CheckServer(ctx, databaseServer) return databaseServer, nil } @@ -121,25 +122,25 @@ func (r *databaseServerRepo) CreateUser(user *biz.DatabaseUser) error { return r.db.Create(user).Error } -// CheckServer 检查服务器连接 -func (r *databaseServerRepo) CheckServer(server *biz.DatabaseServer) bool { +// CheckServer 检查服务器连接,ctx 取消时立即放弃探测 +func (r *databaseServerRepo) CheckServer(ctx context.Context, server *biz.DatabaseServer) bool { switch server.Type { case biz.DatabaseTypeMysql, biz.DatabaseTypePostgresql, biz.DatabaseTypeClickHouse: - operator, err := r.Operator(server) + operator, err := r.Operator(ctx, server) if err == nil { operator.Close() server.Status = biz.DatabaseServerStatusValid return true } case biz.DatabaseTypeRedis: - redis, err := db.NewRedis(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + redis, err := db.NewRedis(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err == nil { redis.Close() server.Status = biz.DatabaseServerStatusValid return true } case biz.DatabaseTypeMongoDB: - mongo, err := db.NewMongoDB(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + mongo, err := db.NewMongoDB(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err == nil { mongo.Close() server.Status = biz.DatabaseServerStatusValid @@ -153,7 +154,7 @@ func (r *databaseServerRepo) CheckServer(server *biz.DatabaseServer) bool { return true } case biz.DatabaseTypeElasticsearch: - es, err := db.NewElasticsearch(fmt.Sprintf("%s:%d", server.Host, server.Port), server.Username, server.Password) + es, err := db.NewElasticsearch(ctx, fmt.Sprintf("%s:%d", server.Host, server.Port), server.Username, server.Password) if err == nil { es.Close() server.Status = biz.DatabaseServerStatusValid @@ -166,18 +167,18 @@ func (r *databaseServerRepo) CheckServer(server *biz.DatabaseServer) bool { } // Operator 获取数据库操作句柄 -func (r *databaseServerRepo) Operator(server *biz.DatabaseServer) (db.Operator, error) { +func (r *databaseServerRepo) Operator(ctx context.Context, server *biz.DatabaseServer) (db.Operator, error) { switch server.Type { case biz.DatabaseTypeMysql: - return newMySQLOperator(server.Username, server.Password, server.Host, server.Port) + return newMySQLOperator(ctx, server.Username, server.Password, server.Host, server.Port) case biz.DatabaseTypePostgresql: - postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) + postgres, err := db.NewPostgres(ctx, server.Username, server.Password, server.Host, server.Port) if err != nil { return nil, err } return postgres, nil case biz.DatabaseTypeClickHouse: - clickhouse, err := db.NewClickHouse(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + clickhouse, err := db.NewClickHouse(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err != nil { return nil, err } @@ -189,13 +190,14 @@ func (r *databaseServerRepo) Operator(server *biz.DatabaseServer) (db.Operator, // newMySQLOperator 构建 MySQL 操作句柄 // 本地 MySQL 优先使用 unix socket 连接:开启 skip-name-resolve 后 TCP 127.0.0.1 无法反解为 localhost,默认的 root@localhost 账户会匹配失败 -func newMySQLOperator(username, password, host string, port uint) (db.Operator, error) { +func newMySQLOperator(ctx context.Context, username, password, host string, port uint) (db.Operator, error) { if sock := localMySQLSocket(host); sock != "" { - if mysql, err := db.NewMySQL(username, password, sock, "unix"); err == nil { + if mysql, err := db.NewMySQL(ctx, username, password, sock, "unix"); err == nil { return mysql, nil } } - return db.NewMySQL(username, password, fmt.Sprintf("%s:%d", host, port)) + + return db.NewMySQL(ctx, username, password, fmt.Sprintf("%s:%d", host, port)) } // localMySQLSocket 返回本地 MySQL 的 unix socket 路径,非本地或未探测到返回空 diff --git a/internal/data/database_user.go b/internal/data/database_user.go index 6d8b47b4..8845251b 100644 --- a/internal/data/database_user.go +++ b/internal/data/database_user.go @@ -1,6 +1,7 @@ package data import ( + "context" "fmt" "slices" @@ -31,7 +32,7 @@ func (r *databaseUserRepo) Count() (int64, error) { return count, nil } -func (r *databaseUserRepo) List(page, limit uint, typ string) ([]*biz.DatabaseUser, int64, error) { +func (r *databaseUserRepo) List(ctx context.Context, page, limit uint, typ string) ([]*biz.DatabaseUser, int64, error) { user := make([]*biz.DatabaseUser, 0) var total int64 query := r.db.Model(&biz.DatabaseUser{}).Preload("Server").Order("id desc") @@ -41,25 +42,25 @@ func (r *databaseUserRepo) List(page, limit uint, typ string) ([]*biz.DatabaseUs err := query.Count(&total).Offset(int((page - 1) * limit)).Limit(int(limit)).Find(&user).Error for u := range slices.Values(user) { - r.fillUser(u) + r.fillUser(ctx, u) } return user, total, err } -func (r *databaseUserRepo) Get(id uint) (*biz.DatabaseUser, error) { +func (r *databaseUserRepo) Get(ctx context.Context, id uint) (*biz.DatabaseUser, error) { user := new(biz.DatabaseUser) if err := r.db.Preload("Server").Where("id = ?", id).First(user).Error; err != nil { return nil, err } - r.fillUser(user) + r.fillUser(ctx, user) return user, nil } -func (r *databaseUserRepo) UpdateRemark(req *request.DatabaseUserUpdateRemark) error { - user, err := r.Get(req.ID) +func (r *databaseUserRepo) UpdateRemark(ctx context.Context, req *request.DatabaseUserUpdateRemark) error { + user, err := r.Get(ctx, req.ID) if err != nil { return err } @@ -70,18 +71,18 @@ func (r *databaseUserRepo) UpdateRemark(req *request.DatabaseUserUpdateRemark) e } // Operator 获取数据库操作句柄 -func (r *databaseUserRepo) Operator(server *biz.DatabaseServer) (db.Operator, error) { +func (r *databaseUserRepo) Operator(ctx context.Context, server *biz.DatabaseServer) (db.Operator, error) { switch server.Type { case biz.DatabaseTypeMysql: - return newMySQLOperator(server.Username, server.Password, server.Host, server.Port) + return newMySQLOperator(ctx, server.Username, server.Password, server.Host, server.Port) case biz.DatabaseTypePostgresql: - postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) + postgres, err := db.NewPostgres(ctx, server.Username, server.Password, server.Host, server.Port) if err != nil { return nil, err } return postgres, nil case biz.DatabaseTypeClickHouse: - clickhouse, err := db.NewClickHouse(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + clickhouse, err := db.NewClickHouse(ctx, server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err != nil { return nil, err } @@ -125,17 +126,17 @@ func (r *databaseUserRepo) DeleteByServerNames(serverID uint, names []string) er return r.db.Where("server_id = ? AND username IN ?", serverID, names).Delete(&biz.DatabaseUser{}).Error } -func (r *databaseUserRepo) fillUser(user *biz.DatabaseUser) { +func (r *databaseUserRepo) fillUser(ctx context.Context, user *biz.DatabaseUser) { server, err := r.loadServer(user.ServerID) if err == nil { - operator, err := r.Operator(server) + operator, err := r.Operator(ctx, server) if err == nil { defer operator.Close() switch server.Type { case biz.DatabaseTypeMysql: privileges, _ := operator.UserPrivileges(user.Username, user.Host) user.Privileges = privileges - if mysql2, err := newMySQLOperator(user.Username, user.Password, server.Host, server.Port); err == nil { + if mysql2, err := newMySQLOperator(ctx, user.Username, user.Password, server.Host, server.Port); err == nil { mysql2.Close() user.Status = biz.DatabaseUserStatusValid } else { @@ -144,7 +145,7 @@ func (r *databaseUserRepo) fillUser(user *biz.DatabaseUser) { case biz.DatabaseTypePostgresql: privileges, _ := operator.UserPrivileges(user.Username) user.Privileges = privileges - if postgres2, err := db.NewPostgres(user.Username, user.Password, server.Host, server.Port); err == nil { + if postgres2, err := db.NewPostgres(ctx, user.Username, user.Password, server.Host, server.Port); err == nil { postgres2.Close() user.Status = biz.DatabaseUserStatusValid } else { @@ -153,7 +154,7 @@ func (r *databaseUserRepo) fillUser(user *biz.DatabaseUser) { case biz.DatabaseTypeClickHouse: privileges, _ := operator.UserPrivileges(user.Username) user.Privileges = privileges - if ch2, err := db.NewClickHouse(user.Username, user.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)); err == nil { + if ch2, err := db.NewClickHouse(ctx, user.Username, user.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)); err == nil { ch2.Close() user.Status = biz.DatabaseUserStatusValid } else { diff --git a/internal/data/notify.go b/internal/data/notify.go new file mode 100644 index 00000000..23493cc7 --- /dev/null +++ b/internal/data/notify.go @@ -0,0 +1,113 @@ +package data + +import ( + "encoding/json" + + "github.com/samber/do/v2" + "github.com/samber/lo" + "gorm.io/gorm" + + "github.com/acepanel/panel/v3/internal/biz" +) + +type notifyChannelRepo struct { + db *gorm.DB +} + +func NewNotifyChannelRepo(i do.Injector) (biz.NotifyChannelRepo, error) { + return ¬ifyChannelRepo{ + db: do.MustInvoke[*gorm.DB](i), + }, nil +} + +func (r *notifyChannelRepo) List(page, limit uint) ([]*biz.NotifyChannel, int64, error) { + channels := make([]*biz.NotifyChannel, 0) + var total int64 + err := r.db.Model(&biz.NotifyChannel{}).Order("id desc").Count(&total).Offset(int((page - 1) * limit)).Limit(int(limit)).Find(&channels).Error + return channels, total, err +} + +func (r *notifyChannelRepo) All() ([]*biz.NotifyChannel, error) { + channels := make([]*biz.NotifyChannel, 0) + err := r.db.Order("id desc").Find(&channels).Error + return channels, err +} + +func (r *notifyChannelRepo) Get(id uint) (*biz.NotifyChannel, error) { + channel := new(biz.NotifyChannel) + if err := r.db.Where("id = ?", id).First(channel).Error; err != nil { + return nil, err + } + return channel, nil +} + +func (r *notifyChannelRepo) GetByIDs(ids []uint) ([]*biz.NotifyChannel, error) { + channels := make([]*biz.NotifyChannel, 0) + if len(ids) == 0 { + return channels, nil + } + err := r.db.Where("id IN ?", ids).Find(&channels).Error + return channels, err +} + +func (r *notifyChannelRepo) Create(channel *biz.NotifyChannel) error { + return r.db.Create(channel).Error +} + +func (r *notifyChannelRepo) Update(channel *biz.NotifyChannel) error { + return r.db.Save(channel).Error +} + +// Delete 删除渠道,同时清理告警规则与事件设置中的引用 +// 残留引用不会报错,只会让通知静默失效,因此必须一并清除 +func (r *notifyChannelRepo) Delete(id uint) error { + return r.db.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("id = ?", id).Delete(&biz.NotifyChannel{}).Error; err != nil { + return err + } + + rules := make([]*biz.AlertRule, 0) + if err := tx.Find(&rules).Error; err != nil { + return err + } + for _, rule := range rules { + channels := lo.Without(rule.Channels, id) + if len(channels) == len(rule.Channels) { + continue + } + rule.Channels = channels + if err := tx.Save(rule).Error; err != nil { + return err + } + } + + return r.removeEventChannel(tx, id) + }) +} + +// removeEventChannel 从系统事件通知设置中移除指定渠道 +func (r *notifyChannelRepo) removeEventChannel(tx *gorm.DB, id uint) error { + setting := new(biz.Setting) + if err := tx.Where("key = ?", biz.SettingKeyNotifyEventChannels).First(setting).Error; err != nil { + // 未配置过事件通知 + return nil + } + + channels := make([]uint, 0) + if json.Unmarshal([]byte(setting.Value), &channels) != nil { + return nil + } + + remain := lo.Without(channels, id) + if len(remain) == len(channels) { + return nil + } + + value, err := json.Marshal(remain) + if err != nil { + return err + } + setting.Value = string(value) + + return tx.Save(setting).Error +} diff --git a/internal/data/notify_test.go b/internal/data/notify_test.go new file mode 100644 index 00000000..a402968c --- /dev/null +++ b/internal/data/notify_test.go @@ -0,0 +1,124 @@ +package data + +import ( + "encoding/json" + "testing" + + "github.com/libtnb/sqlite" + "gorm.io/gorm" + + "github.com/acepanel/panel/v3/internal/app" + "github.com/acepanel/panel/v3/internal/biz" +) + +func newDBForTest(t *testing.T) *gorm.DB { + t.Helper() + // 渠道配置落库前会用 app.Key 加密 + app.Key = "0123456789abcdef0123456789abcdef" + db, err := gorm.Open(sqlite.Open("file::memory:"), &gorm.Config{SkipDefaultTransaction: true}) + if err != nil { + t.Fatal(err) + } + if err = db.AutoMigrate(&biz.NotifyChannel{}, &biz.AlertRule{}, &biz.Alert{}, &biz.Setting{}); err != nil { + t.Fatal(err) + } + return db +} + +// 渠道配置以原始 JSON 存取,读回后应与写入完全一致 +func TestNotifyChannelConfigRoundTrip(t *testing.T) { + repo := ¬ifyChannelRepo{db: newDBForTest(t)} + + config := json.RawMessage(`{"host":"smtp.example.com","port":465,"encryption":"ssl","to":["a@b.c"]}`) + channel := &biz.NotifyChannel{Name: "mail", Type: "smtp", Config: config, Enabled: true} + if err := repo.Create(channel); err != nil { + t.Fatalf("create: %v", err) + } + + got, err := repo.Get(channel.ID) + if err != nil { + t.Fatalf("get: %v", err) + } + + var want, actual map[string]any + if err = json.Unmarshal(config, &want); err != nil { + t.Fatal(err) + } + if err = json.Unmarshal(got.Config, &actual); err != nil { + t.Fatalf("unmarshal stored config %q: %v", string(got.Config), err) + } + if len(actual) != len(want) || actual["host"] != want["host"] { + t.Fatalf("config mismatch: %s", string(got.Config)) + } +} + +// 告警规则的渠道列表以 JSON 序列化存储 +func TestAlertRuleChannelsRoundTrip(t *testing.T) { + repo := &alertRepo{db: newDBForTest(t)} + + rule := &biz.AlertRule{Name: "cpu", Type: biz.AlertTypeCPU, Operator: biz.AlertOperatorGT, Threshold: 90, Duration: 3, Silence: 30, Channels: []uint{1, 2}, Enabled: true} + if err := repo.CreateRule(rule); err != nil { + t.Fatalf("create: %v", err) + } + + got, err := repo.GetRule(rule.ID) + if err != nil { + t.Fatalf("get: %v", err) + } + if len(got.Channels) != 2 || got.Channels[0] != 1 || got.Channels[1] != 2 { + t.Fatalf("channels mismatch: %v", got.Channels) + } +} + +// 删除渠道后,告警规则与事件设置中的引用应一并清除,否则通知会静默失效 +func TestNotifyChannelDeleteCleansReferences(t *testing.T) { + db := newDBForTest(t) + channels := ¬ifyChannelRepo{db: db} + alerts := &alertRepo{db: db} + + kept := &biz.NotifyChannel{Name: "kept", Type: "smtp", Config: json.RawMessage(`{}`), Enabled: true} + if err := channels.Create(kept); err != nil { + t.Fatalf("create kept: %v", err) + } + removed := &biz.NotifyChannel{Name: "removed", Type: "smtp", Config: json.RawMessage(`{}`), Enabled: true} + if err := channels.Create(removed); err != nil { + t.Fatalf("create removed: %v", err) + } + + rule := &biz.AlertRule{Name: "cpu", Type: biz.AlertTypeCPU, Operator: biz.AlertOperatorGT, Threshold: 90, Duration: 1, Channels: []uint{kept.ID, removed.ID}, Enabled: true} + if err := alerts.CreateRule(rule); err != nil { + t.Fatalf("create rule: %v", err) + } + + events, err := json.Marshal([]uint{kept.ID, removed.ID}) + if err != nil { + t.Fatal(err) + } + if err = db.Create(&biz.Setting{Key: biz.SettingKeyNotifyEventChannels, Value: string(events)}).Error; err != nil { + t.Fatal(err) + } + + if err = channels.Delete(removed.ID); err != nil { + t.Fatalf("delete: %v", err) + } + + got, err := alerts.GetRule(rule.ID) + if err != nil { + t.Fatalf("get rule: %v", err) + } + if len(got.Channels) != 1 || got.Channels[0] != kept.ID { + t.Fatalf("rule channels not cleaned: %v", got.Channels) + } + + setting := new(biz.Setting) + if err = db.Where("key = ?", biz.SettingKeyNotifyEventChannels).First(setting).Error; err != nil { + t.Fatalf("get setting: %v", err) + } + var remain []uint + if err = json.Unmarshal([]byte(setting.Value), &remain); err != nil { + t.Fatal(err) + } + if len(remain) != 1 || remain[0] != kept.ID { + t.Fatalf("event channels not cleaned: %v", remain) + } +} diff --git a/internal/job/alert.go b/internal/job/alert.go new file mode 100644 index 00000000..16189c83 --- /dev/null +++ b/internal/job/alert.go @@ -0,0 +1,40 @@ +package job + +import ( + "context" + "log/slog" + + "github.com/samber/do/v2" + + "github.com/acepanel/panel/v3/internal/app" + "github.com/acepanel/panel/v3/internal/biz" +) + +// Alert 告警规则评估任务 +type Alert struct { + log *slog.Logger + alertRepo *biz.AlertUsecase +} + +// NewAlert 构造告警评估任务 +func NewAlert(i do.Injector) (Job, error) { + return Job{ + Spec: "* * * * *", + Task: &Alert{ + log: do.MustInvoke[*slog.Logger](i), + alertRepo: do.MustInvoke[*biz.AlertUsecase](i), + }, + }, nil +} + +func (r *Alert) Run(ctx context.Context) error { + if app.Status != app.StatusNormal { + return nil + } + + if err := r.alertRepo.Evaluate(ctx); err != nil { + r.log.Warn("failed to evaluate alert rules", slog.Any("err", err)) + } + + return nil +} diff --git a/internal/job/cert_renew.go b/internal/job/cert_renew.go index e75acf8b..b6efeb1d 100644 --- a/internal/job/cert_renew.go +++ b/internal/job/cert_renew.go @@ -6,8 +6,10 @@ import ( "log/slog" "os" "path/filepath" + "strings" "time" + "github.com/leonelquinteros/gotext" "github.com/samber/do/v2" "gorm.io/gorm" @@ -27,6 +29,8 @@ type CertRenew struct { settingRepo *biz.SettingUsecase certRepo *biz.CertUsecase certAccountRepo *biz.CertAccountUsecase + notifyRepo *biz.NotifyUsecase + t *gotext.Locale } // NewCertRenew 构造证书续签任务 @@ -40,6 +44,8 @@ func NewCertRenew(i do.Injector) (Job, error) { settingRepo: do.MustInvoke[*biz.SettingUsecase](i), certRepo: do.MustInvoke[*biz.CertUsecase](i), certAccountRepo: do.MustInvoke[*biz.CertAccountUsecase](i), + notifyRepo: do.MustInvoke[*biz.NotifyUsecase](i), + t: do.MustInvoke[*gotext.Locale](i), }, }, nil } @@ -75,6 +81,7 @@ func (r *CertRenew) Run(_ context.Context) error { if time.Now().After(cert.RenewalInfo.SelectedTime) { if _, err := r.certRepo.Renew(cert.ID); err != nil { r.log.Warn("failed to renew certificate", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0), slog.Any("err", err)) + r.notifyFailed(strings.Join(cert.Domains, ", "), err) } } } @@ -96,6 +103,7 @@ func (r *CertRenew) Run(_ context.Context) error { newCrt, newKey, err := pkgcert.GenerateSelfSigned(tools.CollectLocalNames()) if err != nil { r.log.Warn("failed to generate self-signed certificate", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0), slog.Any("err", err)) + r.notifyFailed(r.t.Get("panel certificate"), err) return nil } if err = r.settingRepo.UpdateCert(&request.SettingCert{ @@ -103,6 +111,7 @@ func (r *CertRenew) Run(_ context.Context) error { Key: string(newKey), }); err != nil { r.log.Warn("failed to update panel certificate", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0), slog.Any("err", err)) + r.notifyFailed(r.t.Get("panel certificate"), err) return nil } r.log.Info("panel self-signed certificate renewed", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0)) @@ -143,11 +152,13 @@ func (r *CertRenew) Run(_ context.Context) error { account, err := r.certAccountRepo.GetDefault(user.ID) if err != nil { r.log.Warn("failed to get panel ACME account", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0), slog.Any("err", err)) + r.notifyFailed(r.t.Get("panel certificate"), err) return nil } crt, key, err := r.certRepo.ObtainPanel(account, ips) if err != nil { r.log.Warn("failed to obtain panel certificate via ACME", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0), slog.Any("err", err)) + r.notifyFailed(r.t.Get("panel certificate"), err) return nil } @@ -156,6 +167,7 @@ func (r *CertRenew) Run(_ context.Context) error { Key: string(key), }); err != nil { r.log.Warn("failed to update panel certificate", slog.String("type", biz.OperationTypeCert), slog.Uint64("operator_id", 0), slog.Any("err", err)) + r.notifyFailed(r.t.Get("panel certificate"), err) return nil } @@ -164,3 +176,12 @@ func (r *CertRenew) Run(_ context.Context) error { return nil } + +// notifyFailed 上报证书续签失败 +func (r *CertRenew) notifyFailed(target string, err error) { + r.notifyRepo.SendEvent(biz.NotifyEventCertRenew, r.t.Get("[AcePanel] Certificate Renewal Failed"), biz.NotifyBody(r.t.Get("certificate renewal failed"), [][2]string{ + {r.t.Get("Certificate"), target}, + {r.t.Get("Error"), err.Error()}, + {r.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) +} diff --git a/internal/job/job.go b/internal/job/job.go index 7a0a6a78..73f5250d 100644 --- a/internal/job/job.go +++ b/internal/job/job.go @@ -15,6 +15,7 @@ type Job struct { } var Package = do.Package( + do.LazyNamed(Prefix+"alert", NewAlert), do.LazyNamed(Prefix+"monitoring", NewMonitoring), do.LazyNamed(Prefix+"firewall_scan", NewFirewallScan), do.LazyNamed(Prefix+"cert_renew", NewCertRenew), diff --git a/internal/job/website_expire.go b/internal/job/website_expire.go index 27347fa1..c906280f 100644 --- a/internal/job/website_expire.go +++ b/internal/job/website_expire.go @@ -5,6 +5,7 @@ import ( "log/slog" "time" + "github.com/leonelquinteros/gotext" "github.com/samber/do/v2" "gorm.io/gorm" @@ -17,6 +18,8 @@ type WebsiteExpire struct { db *gorm.DB log *slog.Logger websiteRepo *biz.WebsiteUsecase + notifyRepo *biz.NotifyUsecase + t *gotext.Locale } // NewWebsiteExpire 构造网站到期检查任务 @@ -27,6 +30,8 @@ func NewWebsiteExpire(i do.Injector) (Job, error) { db: do.MustInvoke[*gorm.DB](i), log: do.MustInvoke[*slog.Logger](i), websiteRepo: do.MustInvoke[*biz.WebsiteUsecase](i), + notifyRepo: do.MustInvoke[*biz.NotifyUsecase](i), + t: do.MustInvoke[*gotext.Locale](i), }, }, nil } @@ -50,6 +55,10 @@ func (r *WebsiteExpire) Run(_ context.Context) error { continue } r.log.Info("website expired and disabled", slog.String("name", website.Name), slog.Time("expire_at", *website.ExpireAt)) + r.notifyRepo.SendEvent(biz.NotifyEventWebsiteExpire, r.t.Get("[AcePanel] Website Expired"), biz.NotifyBody(r.t.Get("website expired and has been disabled"), [][2]string{ + {r.t.Get("Website"), website.Name}, + {r.t.Get("Expire Time"), website.ExpireAt.Format(time.DateTime)}, + })) } return nil } diff --git a/internal/migration/v1.go b/internal/migration/v1.go index e4b5dc12..c0621de7 100644 --- a/internal/migration/v1.go +++ b/internal/migration/v1.go @@ -149,4 +149,13 @@ func init() { return nil }, }) + Migrations = append(Migrations, &gormigrate.Migration{ + ID: "20260725-add-notify-and-alert", + Migrate: func(tx *gorm.DB) error { + return tx.AutoMigrate(&biz.NotifyChannel{}, &biz.AlertRule{}, &biz.Alert{}) + }, + Rollback: func(tx *gorm.DB) error { + return tx.Migrator().DropTable(&biz.NotifyChannel{}, &biz.AlertRule{}, &biz.Alert{}) + }, + }) } diff --git a/internal/request/alert.go b/internal/request/alert.go new file mode 100644 index 00000000..c2b060da --- /dev/null +++ b/internal/request/alert.go @@ -0,0 +1,26 @@ +package request + +type AlertRuleCreate struct { + Name string `json:"name" form:"name" validate:"required"` + Type string `json:"type" form:"type" validate:"required && in:cpu,memory,swap,load1,load5,load15,disk,disk_inode,disk_read,disk_write,net_in,net_out,website_5xx,website_error,service,project,container,app,database,cert_expire,website_expire"` + Target string `json:"target" form:"target"` + Operator string `json:"operator" form:"operator" validate:"required && in:gt,gte,lt,lte"` + Threshold float64 `json:"threshold" form:"threshold"` + Duration uint `json:"duration" form:"duration" validate:"min:1 && max:60"` + Silence uint `json:"silence" form:"silence" validate:"min:0 && max:1440"` + Channels []uint `json:"channels" form:"channels"` + Enabled bool `json:"enabled" form:"enabled"` +} + +type AlertRuleUpdate struct { + ID uint `json:"id" form:"id" uri:"id" validate:"required && exists:alert_rules,id"` + Name string `json:"name" form:"name" validate:"required"` + Type string `json:"type" form:"type" validate:"required && in:cpu,memory,swap,load1,load5,load15,disk,disk_inode,disk_read,disk_write,net_in,net_out,website_5xx,website_error,service,project,container,app,database,cert_expire,website_expire"` + Target string `json:"target" form:"target"` + Operator string `json:"operator" form:"operator" validate:"required && in:gt,gte,lt,lte"` + Threshold float64 `json:"threshold" form:"threshold"` + Duration uint `json:"duration" form:"duration" validate:"min:1 && max:60"` + Silence uint `json:"silence" form:"silence" validate:"min:0 && max:1440"` + Channels []uint `json:"channels" form:"channels"` + Enabled bool `json:"enabled" form:"enabled"` +} diff --git a/internal/request/monitor.go b/internal/request/monitor.go index 36b2f533..856721dc 100644 --- a/internal/request/monitor.go +++ b/internal/request/monitor.go @@ -1,9 +1,10 @@ package request type MonitorSetting struct { - Enabled bool `json:"enabled"` - Days uint `json:"days"` - Interval uint `json:"interval" validate:"required && min:1 && max:120"` // 采集间隔(分钟),最小 1 + Enabled bool `json:"enabled"` + Days uint `json:"days"` + Interval uint `json:"interval" validate:"required && min:1 && max:120"` // 采集间隔(分钟),最小 1 + AlertDays uint `json:"alert_days" validate:"min:1 && max:365"` // 告警记录保留天数 } type MonitorList struct { diff --git a/internal/request/notify.go b/internal/request/notify.go new file mode 100644 index 00000000..3c27f426 --- /dev/null +++ b/internal/request/notify.go @@ -0,0 +1,23 @@ +package request + +import "encoding/json" + +type NotifyChannelCreate struct { + Name string `json:"name" form:"name" validate:"required"` + Type string `json:"type" form:"type" validate:"required && in:smtp"` + Config json.RawMessage `json:"config" form:"config"` + Enabled bool `json:"enabled" form:"enabled"` +} + +type NotifyChannelUpdate struct { + ID uint `json:"id" form:"id" uri:"id" validate:"required && exists:notify_channels,id"` + Name string `json:"name" form:"name" validate:"required"` + Type string `json:"type" form:"type" validate:"required && in:smtp"` + Config json.RawMessage `json:"config" form:"config"` + Enabled bool `json:"enabled" form:"enabled"` +} + +type NotifySetting struct { + Events []string `json:"events" form:"events"` + Channels []uint `json:"channels" form:"channels"` +} diff --git a/internal/route/alert.go b/internal/route/alert.go new file mode 100644 index 00000000..17281ce3 --- /dev/null +++ b/internal/route/alert.go @@ -0,0 +1,26 @@ +package route + +import ( + "net/http" + + "github.com/samber/do/v2" + + "github.com/acepanel/panel/v3/internal/biz" + "github.com/acepanel/panel/v3/internal/request" + "github.com/acepanel/panel/v3/internal/service" +) + +// AlertRoutes 告警路由 +func AlertRoutes(i do.Injector) (Endpoints, error) { + svc := do.MustInvoke[*service.AlertService](i) + + return Endpoints{ + {Method: http.MethodGet, Path: "/api/alert/rule", Handler: svc.ListRules, Summary: "告警规则列表", Tags: []string{"告警"}, Request: request.Paginate{}, Response: service.Envelope[service.Page[*biz.AlertRule]]{}}, + {Method: http.MethodPost, Path: "/api/alert/rule", Handler: svc.CreateRule, Summary: "创建告警规则", Tags: []string{"告警"}, Request: request.AlertRuleCreate{}, Response: service.Envelope[biz.AlertRule]{}}, + {Method: http.MethodGet, Path: "/api/alert/rule/{id}", Handler: svc.GetRule, Summary: "获取告警规则", Tags: []string{"告警"}, Request: request.ID{}, Response: service.Envelope[biz.AlertRule]{}}, + {Method: http.MethodPut, Path: "/api/alert/rule/{id}", Handler: svc.UpdateRule, Summary: "更新告警规则", Tags: []string{"告警"}, Request: request.AlertRuleUpdate{}}, + {Method: http.MethodDelete, Path: "/api/alert/rule/{id}", Handler: svc.DeleteRule, Summary: "删除告警规则", Tags: []string{"告警"}, Request: request.ID{}}, + {Method: http.MethodGet, Path: "/api/alert/record", Handler: svc.List, Summary: "告警记录列表", Tags: []string{"告警"}, Request: request.Paginate{}, Response: service.Envelope[service.Page[*biz.Alert]]{}}, + {Method: http.MethodPost, Path: "/api/alert/record/clear", Handler: svc.Clear, Summary: "清空告警记录", Tags: []string{"告警"}}, + }, nil +} diff --git a/internal/route/notify.go b/internal/route/notify.go new file mode 100644 index 00000000..df776a37 --- /dev/null +++ b/internal/route/notify.go @@ -0,0 +1,28 @@ +package route + +import ( + "net/http" + + "github.com/samber/do/v2" + + "github.com/acepanel/panel/v3/internal/biz" + "github.com/acepanel/panel/v3/internal/request" + "github.com/acepanel/panel/v3/internal/service" +) + +// NotifyRoutes 通知渠道路由 +func NotifyRoutes(i do.Injector) (Endpoints, error) { + svc := do.MustInvoke[*service.NotifyService](i) + + return Endpoints{ + {Method: http.MethodGet, Path: "/api/notify/channel", Handler: svc.List, Summary: "通知渠道列表", Tags: []string{"通知"}, Request: request.Paginate{}, Response: service.Envelope[service.Page[*biz.NotifyChannel]]{}}, + {Method: http.MethodGet, Path: "/api/notify/channel/all", Handler: svc.All, Summary: "全部通知渠道", Tags: []string{"通知"}, Response: service.Envelope[[]*biz.NotifyChannel]{}}, + {Method: http.MethodPost, Path: "/api/notify/channel", Handler: svc.Create, Summary: "创建通知渠道", Tags: []string{"通知"}, Request: request.NotifyChannelCreate{}, Response: service.Envelope[biz.NotifyChannel]{}}, + {Method: http.MethodGet, Path: "/api/notify/channel/{id}", Handler: svc.Get, Summary: "获取通知渠道", Tags: []string{"通知"}, Request: request.ID{}, Response: service.Envelope[biz.NotifyChannel]{}}, + {Method: http.MethodPut, Path: "/api/notify/channel/{id}", Handler: svc.Update, Summary: "更新通知渠道", Tags: []string{"通知"}, Request: request.NotifyChannelUpdate{}}, + {Method: http.MethodDelete, Path: "/api/notify/channel/{id}", Handler: svc.Delete, Summary: "删除通知渠道", Tags: []string{"通知"}, Request: request.ID{}}, + {Method: http.MethodPost, Path: "/api/notify/channel/{id}/test", Handler: svc.Test, Summary: "测试通知渠道", Tags: []string{"通知"}, Request: request.ID{}}, + {Method: http.MethodGet, Path: "/api/notify/setting", Handler: svc.GetSetting, Summary: "获取事件通知设置", Tags: []string{"通知"}, Response: service.Envelope[request.NotifySetting]{}}, + {Method: http.MethodPost, Path: "/api/notify/setting", Handler: svc.UpdateSetting, Summary: "更新事件通知设置", Tags: []string{"通知"}, Request: request.NotifySetting{}}, + }, nil +} diff --git a/internal/route/route.go b/internal/route/route.go index 0d9229da..31dcbf7a 100644 --- a/internal/route/route.go +++ b/internal/route/route.go @@ -33,6 +33,7 @@ var Package = do.Package( do.LazyNamed(RoutePrefix+"ssh", SSHRoutes), do.LazyNamed(RoutePrefix+"systemctl", SystemctlRoutes), do.LazyNamed(RoutePrefix+"setting", SettingRoutes), do.LazyNamed(RoutePrefix+"log", LogRoutes), do.LazyNamed(RoutePrefix+"monitor", MonitorRoutes), do.LazyNamed(RoutePrefix+"webhook", WebHookRoutes), + do.LazyNamed(RoutePrefix+"notify", NotifyRoutes), do.LazyNamed(RoutePrefix+"alert", AlertRoutes), do.LazyNamed(RoutePrefix+"template", TemplateRoutes), do.LazyNamed(RoutePrefix+"toolbox_network", ToolboxNetworkRoutes), do.LazyNamed(RoutePrefix+"toolbox_system", ToolboxSystemRoutes), do.LazyNamed(RoutePrefix+"toolbox_benchmark", ToolboxBenchmarkRoutes), do.LazyNamed(RoutePrefix+"toolbox_ssh", ToolboxSSHRoutes), do.LazyNamed(RoutePrefix+"toolbox_disk", ToolboxDiskRoutes), diff --git a/internal/service/alert.go b/internal/service/alert.go new file mode 100644 index 00000000..95ea7695 --- /dev/null +++ b/internal/service/alert.go @@ -0,0 +1,130 @@ +package service + +import ( + "net/http" + + "github.com/libtnb/chix/v2" + "github.com/samber/do/v2" + + "github.com/acepanel/panel/v3/internal/biz" + "github.com/acepanel/panel/v3/internal/request" +) + +type AlertService struct { + alertRepo *biz.AlertUsecase +} + +func NewAlertService(i do.Injector) (*AlertService, error) { + return &AlertService{ + alertRepo: do.MustInvoke[*biz.AlertUsecase](i), + }, nil +} + +func (s *AlertService) ListRules(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.Paginate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + rules, total, err := s.alertRepo.ListRules(req.Page, req.Limit) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, chix.M{ + "total": total, + "items": rules, + }) +} + +func (s *AlertService) GetRule(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.ID](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + rule, err := s.alertRepo.GetRule(req.ID) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, rule) +} + +func (s *AlertService) CreateRule(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.AlertRuleCreate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + rule, err := s.alertRepo.CreateRule(r.Context(), req) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, rule) +} + +func (s *AlertService) UpdateRule(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.AlertRuleUpdate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + if err = s.alertRepo.UpdateRule(r.Context(), req); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} + +func (s *AlertService) DeleteRule(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.ID](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + if err = s.alertRepo.DeleteRule(r.Context(), req.ID); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} + +func (s *AlertService) List(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.Paginate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + alerts, total, err := s.alertRepo.ListAlerts(req.Page, req.Limit) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, chix.M{ + "total": total, + "items": alerts, + }) +} + +func (s *AlertService) Clear(w http.ResponseWriter, r *http.Request) { + if err := s.alertRepo.ClearAlerts(); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} diff --git a/internal/service/cli.go b/internal/service/cli.go index 1f7d1743..db4d5543 100644 --- a/internal/service/cli.go +++ b/internal/service/cli.go @@ -52,6 +52,8 @@ type CliService struct { databaseServerRepo *biz.DatabaseServerUsecase certRepo *biz.CertUsecase certAccountRepo *biz.CertAccountUsecase + cronRepo *biz.CronUsecase + notifyRepo *biz.NotifyUsecase hash hash.Hasher } @@ -72,6 +74,8 @@ func NewCliService(i do.Injector) (*CliService, error) { databaseServerRepo: do.MustInvoke[*biz.DatabaseServerUsecase](i), certRepo: do.MustInvoke[*biz.CertUsecase](i), certAccountRepo: do.MustInvoke[*biz.CertAccountUsecase](i), + cronRepo: do.MustInvoke[*biz.CronUsecase](i), + notifyRepo: do.MustInvoke[*biz.NotifyUsecase](i), hash: hash.NewArgon2id(), }, nil } @@ -671,7 +675,7 @@ func (s *CliService) DatabaseAddServer(ctx context.Context, cmd *cli.Command) er Remark: cmd.String("remark"), } - if err := s.databaseServerRepo.Create(req); err != nil { + if err := s.databaseServerRepo.Create(ctx, req); err != nil { return err } @@ -680,7 +684,7 @@ func (s *CliService) DatabaseAddServer(ctx context.Context, cmd *cli.Command) er } func (s *CliService) DatabaseDeleteServer(ctx context.Context, cmd *cli.Command) error { - server, err := s.databaseServerRepo.GetByName(cmd.String("name")) + server, err := s.databaseServerRepo.GetByName(ctx, cmd.String("name")) if err != nil { return err } @@ -1144,3 +1148,24 @@ checkPort: return nil } + +// CronFailed 上报计划任务执行失败,由任务 wrapper 脚本调用 +func (s *CliService) CronFailed(ctx context.Context, cmd *cli.Command) error { + cron, err := s.cronRepo.Get(cmd.Uint("id")) + if err != nil { + return err + } + + // 附带日志尾部,便于直接定位问题 + tail, _ := shell.Execf("tail -n 20 %s", cron.Log) + + return s.notifyRepo.SendEventSync(ctx, biz.NotifyEventCronFailed, s.t.Get("[AcePanel] Cron Task Failed"), + biz.NotifyBody(s.t.Get("cron task exited abnormally"), [][2]string{ + {s.t.Get("Task"), cron.Name}, + {s.t.Get("Schedule"), cron.Time}, + {s.t.Get("Exit Code"), cast.ToString(cmd.Int("code"))}, + {s.t.Get("Log"), cron.Log}, + {s.t.Get("Output"), tail}, + {s.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) +} diff --git a/internal/service/database.go b/internal/service/database.go index 1e14e097..8f99996d 100644 --- a/internal/service/database.go +++ b/internal/service/database.go @@ -27,7 +27,7 @@ func (s *DatabaseService) List(w http.ResponseWriter, r *http.Request) { return } - databases, total, err := s.databaseRepo.List(req.Page, req.Limit, req.Type) + databases, total, err := s.databaseRepo.List(r.Context(), req.Page, req.Limit, req.Type) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -76,7 +76,7 @@ func (s *DatabaseService) Comment(w http.ResponseWriter, r *http.Request) { return } - if err = s.databaseRepo.Comment(req); err != nil { + if err = s.databaseRepo.Comment(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } diff --git a/internal/service/database_elasticsearch.go b/internal/service/database_elasticsearch.go index f260b880..f13c3882 100644 --- a/internal/service/database_elasticsearch.go +++ b/internal/service/database_elasticsearch.go @@ -25,7 +25,7 @@ func (s *DatabaseElasticsearchService) Indices(w http.ResponseWriter, r *http.Re return } - indices, err := s.repo.Indices(req) + indices, err := s.repo.Indices(r.Context(), req) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -41,7 +41,7 @@ func (s *DatabaseElasticsearchService) IndexCreate(w http.ResponseWriter, r *htt return } - if err = s.repo.IndexCreate(req); err != nil { + if err = s.repo.IndexCreate(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -56,7 +56,7 @@ func (s *DatabaseElasticsearchService) IndexDelete(w http.ResponseWriter, r *htt return } - if err = s.repo.IndexDelete(req); err != nil { + if err = s.repo.IndexDelete(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -71,7 +71,7 @@ func (s *DatabaseElasticsearchService) Data(w http.ResponseWriter, r *http.Reque return } - items, total, err := s.repo.Data(req) + items, total, err := s.repo.Data(r.Context(), req) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -90,7 +90,7 @@ func (s *DatabaseElasticsearchService) DocumentGet(w http.ResponseWriter, r *htt return } - doc, err := s.repo.DocumentGet(req) + doc, err := s.repo.DocumentGet(r.Context(), req) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -106,7 +106,7 @@ func (s *DatabaseElasticsearchService) DocumentSet(w http.ResponseWriter, r *htt return } - if err = s.repo.DocumentSet(req); err != nil { + if err = s.repo.DocumentSet(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -121,7 +121,7 @@ func (s *DatabaseElasticsearchService) DocumentDelete(w http.ResponseWriter, r * return } - if err = s.repo.DocumentDelete(req); err != nil { + if err = s.repo.DocumentDelete(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } diff --git a/internal/service/database_redis.go b/internal/service/database_redis.go index 67595e25..9b314cd7 100644 --- a/internal/service/database_redis.go +++ b/internal/service/database_redis.go @@ -25,7 +25,7 @@ func (s *DatabaseRedisService) Databases(w http.ResponseWriter, r *http.Request) return } - count, err := s.repo.Databases(req) + count, err := s.repo.Databases(r.Context(), req) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -41,7 +41,7 @@ func (s *DatabaseRedisService) Data(w http.ResponseWriter, r *http.Request) { return } - items, total, err := s.repo.Data(req) + items, total, err := s.repo.Data(r.Context(), req) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -60,7 +60,7 @@ func (s *DatabaseRedisService) KeyGet(w http.ResponseWriter, r *http.Request) { return } - kv, err := s.repo.KeyGet(req) + kv, err := s.repo.KeyGet(r.Context(), req) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -76,7 +76,7 @@ func (s *DatabaseRedisService) KeySet(w http.ResponseWriter, r *http.Request) { return } - if err = s.repo.KeySet(req); err != nil { + if err = s.repo.KeySet(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -91,7 +91,7 @@ func (s *DatabaseRedisService) KeyDelete(w http.ResponseWriter, r *http.Request) return } - if err = s.repo.KeyDelete(req); err != nil { + if err = s.repo.KeyDelete(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -106,7 +106,7 @@ func (s *DatabaseRedisService) KeyTTL(w http.ResponseWriter, r *http.Request) { return } - if err = s.repo.KeyTTL(req); err != nil { + if err = s.repo.KeyTTL(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -121,7 +121,7 @@ func (s *DatabaseRedisService) KeyRename(w http.ResponseWriter, r *http.Request) return } - if err = s.repo.KeyRename(req); err != nil { + if err = s.repo.KeyRename(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -136,7 +136,7 @@ func (s *DatabaseRedisService) Clear(w http.ResponseWriter, r *http.Request) { return } - if err = s.repo.Clear(req); err != nil { + if err = s.repo.Clear(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } diff --git a/internal/service/database_server.go b/internal/service/database_server.go index 7b9cdc37..abf63dd4 100644 --- a/internal/service/database_server.go +++ b/internal/service/database_server.go @@ -27,7 +27,7 @@ func (s *DatabaseServerService) List(w http.ResponseWriter, r *http.Request) { return } - servers, total, err := s.databaseServerRepo.List(req.Page, req.Limit, req.Type) + servers, total, err := s.databaseServerRepo.List(r.Context(), req.Page, req.Limit, req.Type) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -46,7 +46,7 @@ func (s *DatabaseServerService) Create(w http.ResponseWriter, r *http.Request) { return } - if err = s.databaseServerRepo.Create(req); err != nil { + if err = s.databaseServerRepo.Create(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -61,7 +61,7 @@ func (s *DatabaseServerService) Get(w http.ResponseWriter, r *http.Request) { return } - server, err := s.databaseServerRepo.Get(req.ID) + server, err := s.databaseServerRepo.Get(r.Context(), req.ID) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -77,7 +77,7 @@ func (s *DatabaseServerService) Update(w http.ResponseWriter, r *http.Request) { return } - if err = s.databaseServerRepo.Update(req); err != nil { + if err = s.databaseServerRepo.Update(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -122,7 +122,7 @@ func (s *DatabaseServerService) Sync(w http.ResponseWriter, r *http.Request) { return } - if err = s.databaseServerRepo.Sync(req.ID); err != nil { + if err = s.databaseServerRepo.Sync(r.Context(), req.ID); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } diff --git a/internal/service/database_user.go b/internal/service/database_user.go index 501bb9db..71a047a6 100644 --- a/internal/service/database_user.go +++ b/internal/service/database_user.go @@ -27,7 +27,7 @@ func (s *DatabaseUserService) List(w http.ResponseWriter, r *http.Request) { return } - users, total, err := s.databaseUserRepo.List(req.Page, req.Limit, req.Type) + users, total, err := s.databaseUserRepo.List(r.Context(), req.Page, req.Limit, req.Type) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -61,7 +61,7 @@ func (s *DatabaseUserService) Get(w http.ResponseWriter, r *http.Request) { return } - user, err := s.databaseUserRepo.Get(req.ID) + user, err := s.databaseUserRepo.Get(r.Context(), req.ID) if err != nil { Error(w, http.StatusInternalServerError, "%v", err) return @@ -77,7 +77,7 @@ func (s *DatabaseUserService) Update(w http.ResponseWriter, r *http.Request) { return } - if err = s.databaseUserRepo.Update(req); err != nil { + if err = s.databaseUserRepo.Update(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } @@ -92,7 +92,7 @@ func (s *DatabaseUserService) UpdateRemark(w http.ResponseWriter, r *http.Reques return } - if err = s.databaseUserRepo.UpdateRemark(req); err != nil { + if err = s.databaseUserRepo.UpdateRemark(r.Context(), req); err != nil { Error(w, http.StatusInternalServerError, "%v", err) return } diff --git a/internal/service/helper.go b/internal/service/helper.go index 4e274cda..77802744 100644 --- a/internal/service/helper.go +++ b/internal/service/helper.go @@ -3,7 +3,9 @@ package service import ( "errors" "fmt" + "net" "net/http" + "net/netip" "slices" "strings" @@ -13,6 +15,25 @@ import ( "github.com/acepanel/panel/v3/internal/request" ) +// clientIP 提取客户端 IP,优先取配置的真实 IP 头 +// 代理头通常给的是裸 IP,只有 RemoteAddr 带端口,两种形态都要能解析 +func clientIP(r *http.Request, ipHeader string) string { + ip := r.RemoteAddr + if ipHeader != "" && r.Header.Get(ipHeader) != "" { + ip = strings.Split(r.Header.Get(ipHeader), ",")[0] + } + ip = strings.TrimSpace(ip) + + if addr, err := netip.ParseAddr(ip); err == nil { + return addr.String() + } + if host, _, err := net.SplitHostPort(ip); err == nil { + return host + } + + return r.RemoteAddr +} + // SuccessResponse 通用成功响应 type SuccessResponse struct { Msg string `json:"msg"` diff --git a/internal/service/helper_test.go b/internal/service/helper_test.go new file mode 100644 index 00000000..c85a3d2f --- /dev/null +++ b/internal/service/helper_test.go @@ -0,0 +1,38 @@ +package service + +import ( + "net/http/httptest" + "testing" +) + +// 代理头给裸 IP、RemoteAddr 带端口,两种形态都要归一到同一个 IP +func TestClientIP(t *testing.T) { + cases := []struct { + name string + remoteAddr string + header string + value string + want string + }{ + {"remote addr", "1.2.3.4:5678", "", "", "1.2.3.4"}, + {"remote addr ipv6", "[2001:db8::1]:5678", "", "", "2001:db8::1"}, + {"bare ip header", "10.0.0.1:5678", "X-Real-IP", "1.2.3.4", "1.2.3.4"}, + {"bare ipv6 header", "10.0.0.1:5678", "X-Real-IP", "2001:db8::1", "2001:db8::1"}, + {"forwarded chain", "10.0.0.1:5678", "X-Forwarded-For", "1.2.3.4, 10.0.0.1", "1.2.3.4"}, + {"header with port", "10.0.0.1:5678", "X-Real-IP", "1.2.3.4:9999", "1.2.3.4"}, + {"header empty falls back", "1.2.3.4:5678", "X-Real-IP", "", "1.2.3.4"}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + r := httptest.NewRequest("POST", "/api/user/login", nil) + r.RemoteAddr = c.remoteAddr + if c.value != "" { + r.Header.Set(c.header, c.value) + } + if got := clientIP(r, c.header); got != c.want { + t.Fatalf("clientIP = %q, want %q", got, c.want) + } + }) + } +} diff --git a/internal/service/home.go b/internal/service/home.go index cf94083f..8de0d4cf 100644 --- a/internal/service/home.go +++ b/internal/service/home.go @@ -180,7 +180,7 @@ func (s *HomeService) CountInfo(w http.ResponseWriter, r *http.Request) { var databaseCount int if mysqlInstalled { rootPassword, _ := s.settingRepo.Get(biz.SettingKeyMySQLRootPassword) - mysql, err := db.NewMySQL("root", rootPassword, "/tmp/mysql.sock", "unix") + mysql, err := db.NewMySQL(r.Context(), "root", rootPassword, "/tmp/mysql.sock", "unix") if err == nil { defer mysql.Close() databases, err := mysql.Databases() @@ -190,8 +190,8 @@ func (s *HomeService) CountInfo(w http.ResponseWriter, r *http.Request) { } } if postgresqlInstalled { - if server, err := s.databaseServerRepo.GetByName("local_postgresql"); err == nil { - if postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port); err == nil { + if server, err := s.databaseServerRepo.GetByName(r.Context(), "local_postgresql"); err == nil { + if postgres, err := db.NewPostgres(r.Context(), server.Username, server.Password, server.Host, server.Port); err == nil { defer postgres.Close() if databases, err := postgres.Databases(); err == nil { databaseCount += len(databases) diff --git a/internal/service/log.go b/internal/service/log.go index a3f9b600..9233e956 100644 --- a/internal/service/log.go +++ b/internal/service/log.go @@ -8,7 +8,6 @@ import ( "io" "net/http" "os" - "regexp" "strconv" "strings" "time" @@ -19,6 +18,7 @@ import ( "github.com/acepanel/panel/v3/internal/biz" "github.com/acepanel/panel/v3/internal/request" "github.com/acepanel/panel/v3/pkg/shell" + "github.com/acepanel/panel/v3/pkg/sshlog" "github.com/acepanel/panel/v3/pkg/types" ) @@ -28,14 +28,6 @@ const ( sshLogChunkMax int64 = 64 * 1024 * 1024 ) -// SSH 日志正则 -var ( - sshAccepted = regexp.MustCompile(`Accepted\s+(\S+)\s+for\s+(\S+)\s+from\s+(\S+)\s+port\s+(\d+)`) - sshFailed = regexp.MustCompile(`Failed\s+(\S+)\s+for\s+(?:invalid user\s+)?(\S+)\s+from\s+(\S+)\s+port\s+(\d+)`) - sshInvalidUser = regexp.MustCompile(`Invalid user\s+(\S+)\s+from\s+(\S+)\s+port\s+(\d+)`) - sshDisconnect = regexp.MustCompile(`Disconnected from\s+(?:authenticating\s+)?user\s+(\S+)\s+(\S+)\s+port\s+(\d+)`) -) - type LogService struct { t *gotext.Locale logRepo *biz.LogUsecase @@ -126,7 +118,7 @@ func (s *LogService) sshFromJournalctl(limit int) ([]types.SSHLoginLog, error) { continue } - record := parseSSHMessage(entry.Message) + record := sshlog.ParseMessage(entry.Message) if record == nil { continue } @@ -198,7 +190,7 @@ func (s *LogService) sshFromLogFile(limit int) ([]types.SSHLoginLog, error) { } // 新读的块时间上更早,前置到已收集 logs 之前 - logs = append(parseSSHChunk(buf), logs...) + logs = append(sshlog.ParseChunk(buf), logs...) window *= 2 } @@ -208,87 +200,3 @@ func (s *LogService) sshFromLogFile(limit int) ([]types.SSHLoginLog, error) { return logs, nil } - -// parseSSHChunk 从连续日志字节中解析 SSH 登录记录 -func parseSSHChunk(data []byte) []types.SSHLoginLog { - if len(data) == 0 { - return nil - } - var logs []types.SSHLoginLog - scanner := bufio.NewScanner(bytes.NewReader(data)) - scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) - for scanner.Scan() { - line := scanner.Text() - if !strings.Contains(line, "sshd[") { - continue - } - record := parseSSHMessage(line) - if record == nil { - continue - } - // 从行首解析 syslog 时间戳(如 "Feb 11 08:30:01") - record.Time = parseSSHLogTime(line) - logs = append(logs, *record) - } - return logs -} - -// parseSSHMessage 从日志消息中提取 SSH 登录信息 -func parseSSHMessage(msg string) *types.SSHLoginLog { - if m := sshAccepted.FindStringSubmatch(msg); m != nil { - return &types.SSHLoginLog{ - Method: m[1], - User: m[2], - IP: m[3], - Port: m[4], - Status: "accepted", - } - } - if m := sshFailed.FindStringSubmatch(msg); m != nil { - return &types.SSHLoginLog{ - Method: m[1], - User: m[2], - IP: m[3], - Port: m[4], - Status: "failed", - } - } - if m := sshInvalidUser.FindStringSubmatch(msg); m != nil { - return &types.SSHLoginLog{ - User: m[1], - IP: m[2], - Port: m[3], - Method: "-", - Status: "invalid_user", - } - } - if m := sshDisconnect.FindStringSubmatch(msg); m != nil { - return &types.SSHLoginLog{ - User: m[1], - IP: m[2], - Port: m[3], - Method: "-", - Status: "disconnected", - } - } - return nil -} - -// parseSSHLogTime 从 syslog 格式行中解析时间 -func parseSSHLogTime(line string) string { - // syslog 格式:Mon DD HH:MM:SS(前 15 个字符) - if len(line) < 15 { - return "-" - } - ts := line[:15] - // 使用当前年份补全 - t, err := time.Parse("Jan 2 15:04:05", ts) - if err != nil { - t, err = time.Parse("Jan 2 15:04:05", ts) - if err != nil { - return "-" - } - } - t = t.AddDate(time.Now().Year(), 0, 0) - return t.Format("2006-01-02 15:04:05") -} diff --git a/internal/service/login_guard.go b/internal/service/login_guard.go new file mode 100644 index 00000000..73e28c69 --- /dev/null +++ b/internal/service/login_guard.go @@ -0,0 +1,85 @@ +package service + +import ( + "sync" + "time" +) + +const ( + // bruteforceWindow 失败计数的统计窗口 + bruteforceWindow = 10 * time.Minute + // bruteforceThreshold 窗口内触发告警的失败次数 + bruteforceThreshold = 5 + // bruteforceSilence 同一来源两次告警的最小间隔 + bruteforceSilence = 30 * time.Minute + // bruteforceMaxEntries 计数条目上限,防伪造来源耗尽内存 + bruteforceMaxEntries = 10000 +) + +type loginFailEntry struct { + count uint + since time.Time + notified time.Time +} + +// loginGuard 按来源 IP 统计登录失败,用于爆破告警 +type loginGuard struct { + mu sync.Mutex + entries map[string]*loginFailEntry +} + +func newLoginGuard() *loginGuard { + return &loginGuard{entries: make(map[string]*loginFailEntry)} +} + +// Fail 记录一次失败,返回窗口内失败次数与是否应当告警 +func (g *loginGuard) Fail(ip string) (uint, bool) { + now := time.Now() + + g.mu.Lock() + defer g.mu.Unlock() + + g.sweep(now) + + entry, ok := g.entries[ip] + if !ok { + if len(g.entries) >= bruteforceMaxEntries { + return 1, false + } + g.entries[ip] = &loginFailEntry{count: 1, since: now} + return 1, false + } + + // 超出窗口重新计数 + if now.Sub(entry.since) > bruteforceWindow { + entry.count, entry.since = 1, now + return 1, false + } + + entry.count++ + if entry.count < bruteforceThreshold { + return entry.count, false + } + if !entry.notified.IsZero() && now.Sub(entry.notified) < bruteforceSilence { + return entry.count, false + } + entry.notified = now + + return entry.count, true +} + +// Reset 登录成功后清除该来源的失败计数 +func (g *loginGuard) Reset(ip string) { + g.mu.Lock() + defer g.mu.Unlock() + delete(g.entries, ip) +} + +// sweep 清理超出窗口且未在静默期内的条目,调用方需持有锁 +func (g *loginGuard) sweep(now time.Time) { + for ip, entry := range g.entries { + if now.Sub(entry.since) > bruteforceWindow && now.Sub(entry.notified) > bruteforceSilence { + delete(g.entries, ip) + } + } +} diff --git a/internal/service/notify.go b/internal/service/notify.go new file mode 100644 index 00000000..bff514da --- /dev/null +++ b/internal/service/notify.go @@ -0,0 +1,152 @@ +package service + +import ( + "net/http" + + "github.com/libtnb/chix/v2" + "github.com/samber/do/v2" + + "github.com/acepanel/panel/v3/internal/biz" + "github.com/acepanel/panel/v3/internal/request" +) + +type NotifyService struct { + notifyRepo *biz.NotifyUsecase +} + +func NewNotifyService(i do.Injector) (*NotifyService, error) { + return &NotifyService{ + notifyRepo: do.MustInvoke[*biz.NotifyUsecase](i), + }, nil +} + +func (s *NotifyService) List(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.Paginate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + channels, total, err := s.notifyRepo.List(req.Page, req.Limit) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, chix.M{ + "total": total, + "items": channels, + }) +} + +func (s *NotifyService) All(w http.ResponseWriter, r *http.Request) { + channels, err := s.notifyRepo.All() + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, channels) +} + +func (s *NotifyService) Get(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.ID](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + channel, err := s.notifyRepo.Get(req.ID) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, channel) +} + +func (s *NotifyService) Create(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.NotifyChannelCreate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + channel, err := s.notifyRepo.Create(r.Context(), req) + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, channel) +} + +func (s *NotifyService) Update(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.NotifyChannelUpdate](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + if err = s.notifyRepo.Update(r.Context(), req); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} + +func (s *NotifyService) Delete(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.ID](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + if err = s.notifyRepo.Delete(r.Context(), req.ID); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} + +func (s *NotifyService) Test(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.ID](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + if err = s.notifyRepo.Test(r.Context(), req.ID); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} + +func (s *NotifyService) GetSetting(w http.ResponseWriter, r *http.Request) { + setting, err := s.notifyRepo.GetSetting() + if err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, setting) +} + +func (s *NotifyService) UpdateSetting(w http.ResponseWriter, r *http.Request) { + req, err := Bind[request.NotifySetting](r) + if err != nil { + Error(w, http.StatusUnprocessableEntity, "%v", err) + return + } + + if err = s.notifyRepo.UpdateSetting(req); err != nil { + Error(w, http.StatusInternalServerError, "%v", err) + return + } + + Success(w, nil) +} diff --git a/internal/service/service.go b/internal/service/service.go index a9e1f912..0cb6c2e7 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -5,7 +5,7 @@ import ( ) var Package = do.Package( - do.Lazy(NewAppService), do.Lazy(NewBackupService), do.Lazy(NewBackupStorageService), + do.Lazy(NewAlertService), do.Lazy(NewAppService), do.Lazy(NewBackupService), do.Lazy(NewBackupStorageService), do.Lazy(NewCertService), do.Lazy(NewCertAccountService), do.Lazy(NewCertDNSService), do.Lazy(NewCliService), do.Lazy(NewContainerService), do.Lazy(NewContainerComposeService), do.Lazy(NewContainerImageService), do.Lazy(NewContainerNetworkService), do.Lazy(NewContainerVolumeService), @@ -15,7 +15,7 @@ var Package = do.Package( do.Lazy(NewEnvironmentNodejsService), do.Lazy(NewEnvironmentPHPService), do.Lazy(NewEnvironmentPythonService), do.Lazy(NewEnvironmentDotnetService), do.Lazy(NewFileService), do.Lazy(NewFileShareService), do.Lazy(NewFirewallService), do.Lazy(NewFirewallScanService), do.Lazy(NewHomeService), do.Lazy(NewLogService), - do.Lazy(NewMonitorService), do.Lazy(NewProcessService), do.Lazy(NewProjectService), + do.Lazy(NewMonitorService), do.Lazy(NewNotifyService), do.Lazy(NewProcessService), do.Lazy(NewProjectService), do.Lazy(NewSafeService), do.Lazy(NewSettingService), do.Lazy(NewSSHService), do.Lazy(NewSystemctlService), do.Lazy(NewTamperService), do.Lazy(NewTaskService), do.Lazy(NewTemplateService), do.Lazy(NewUserService), do.Lazy(NewUserPasskeyService), do.Lazy(NewUserTokenService), diff --git a/internal/service/toolbox_migration.go b/internal/service/toolbox_migration.go index 07943d58..ca93efd8 100644 --- a/internal/service/toolbox_migration.go +++ b/internal/service/toolbox_migration.go @@ -199,7 +199,7 @@ func (s *ToolboxMigrationService) GetItems(w http.ResponseWriter, r *http.Reques } // 数据库列表 - databases, _, err := s.databaseRepo.List(1, 10000, "") + databases, _, err := s.databaseRepo.List(r.Context(), 1, 10000, "") if err != nil { Error(w, http.StatusInternalServerError, s.t.Get("failed to get database list: %v", err)) return @@ -213,7 +213,7 @@ func (s *ToolboxMigrationService) GetItems(w http.ResponseWriter, r *http.Reques } // 数据库用户列表 - databaseUsers, _, err := s.databaseUserRepo.List(1, 10000, "") + databaseUsers, _, err := s.databaseUserRepo.List(r.Context(), 1, 10000, "") if err != nil { Error(w, http.StatusInternalServerError, s.t.Get("failed to get database user list: %v", err)) return @@ -485,6 +485,8 @@ func (s *ToolboxMigrationService) migrateWebsite(conn *request.ToolboxMigrationC // migrateDatabase 迁移单个数据库 func (s *ToolboxMigrationService) migrateDatabase(conn *request.ToolboxMigrationConnection, db *request.ToolboxMigrationDatabase) { + // 迁移脱离请求生命周期,在独立 goroutine 中运行 + ctx := context.Background() displayName := fmt.Sprintf("%s (%s)", db.Name, db.Type) result := types.MigrationItemResult{ Type: "database", @@ -497,7 +499,7 @@ func (s *ToolboxMigrationService) migrateDatabase(conn *request.ToolboxMigration s.addLog(fmt.Sprintf("[%s] %s: %s", s.t.Get("Database"), s.t.Get("start migrating"), displayName)) // 取本地数据库服务器信息 - dbServer, err := s.databaseServerRepo.Get(db.ServerID) + dbServer, err := s.databaseServerRepo.Get(ctx, db.ServerID) if err != nil { s.failResult("database", displayName, s.t.Get("failed to get database server: %v", err)) return @@ -678,6 +680,8 @@ func (s *ToolboxMigrationService) clickHouseTables(conn, database string, onlyVi // migrateDatabaseUser 迁移单个数据库用户 func (s *ToolboxMigrationService) migrateDatabaseUser(conn *request.ToolboxMigrationConnection, user *request.ToolboxMigrationDatabaseUser) { + // 迁移脱离请求生命周期,在独立 goroutine 中运行 + ctx := context.Background() displayName := fmt.Sprintf("%s@%s (%s)", user.Username, user.Host, user.Type) result := types.MigrationItemResult{ Type: "database_user", @@ -690,14 +694,14 @@ func (s *ToolboxMigrationService) migrateDatabaseUser(conn *request.ToolboxMigra s.addLog(fmt.Sprintf("[%s] %s: %s", s.t.Get("Database User"), s.t.Get("start migrating"), displayName)) // 获取本地用户详情(含权限) - userDetail, err := s.databaseUserRepo.Get(user.ID) + userDetail, err := s.databaseUserRepo.Get(ctx, user.ID) if err != nil { s.failResult("database_user", displayName, s.t.Get("failed to get database user detail: %v", err)) return } // 获取本地数据库服务器信息 - dbServer, err := s.databaseServerRepo.Get(user.ServerID) + dbServer, err := s.databaseServerRepo.Get(ctx, user.ServerID) if err != nil { s.failResult("database_user", displayName, s.t.Get("failed to get database server: %v", err)) return diff --git a/internal/service/user.go b/internal/service/user.go index 20ff003d..f9109411 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -8,9 +8,7 @@ import ( "encoding/gob" "fmt" "image/png" - "net" "net/http" - "strings" "time" "github.com/dchest/captcha" @@ -31,19 +29,23 @@ import ( const loginFailThreshold = 3 type UserService struct { - t *gotext.Locale - conf *config.Config - session *sessions.Manager - userRepo *biz.UserUsecase + t *gotext.Locale + conf *config.Config + session *sessions.Manager + userRepo *biz.UserUsecase + notifyRepo *biz.NotifyUsecase + guard *loginGuard } func NewUserService(i do.Injector) (*UserService, error) { gob.Register(rsa.PrivateKey{}) // 必须注册 rsa.PrivateKey 类型否则无法反序列化 session 中的 key return &UserService{ - t: do.MustInvoke[*gotext.Locale](i), - conf: do.MustInvoke[*config.Config](i), - session: do.MustInvoke[*sessions.Manager](i), - userRepo: do.MustInvoke[*biz.UserUsecase](i), + t: do.MustInvoke[*gotext.Locale](i), + conf: do.MustInvoke[*config.Config](i), + session: do.MustInvoke[*sessions.Manager](i), + userRepo: do.MustInvoke[*biz.UserUsecase](i), + notifyRepo: do.MustInvoke[*biz.NotifyUsecase](i), + guard: newLoginGuard(), }, nil } @@ -130,22 +132,28 @@ func (s *UserService) Login(w http.ResponseWriter, r *http.Request) { return } + ip := clientIP(r, s.conf.HTTP.IPHeader) + decryptedUsername, _ := rsacrypto.DecryptData(&key, req.Username) decryptedPassword, _ := rsacrypto.DecryptData(&key, req.Password) user, err := s.userRepo.CheckPassword(string(decryptedUsername), string(decryptedPassword)) if err != nil { - sess.Put("login_fail_count", failCount+1) + s.loginFailed(r, sess, ip, string(decryptedUsername), failCount) Error(w, http.StatusForbidden, "%v", err) return } + // 2FA 失败同样计入,否则已知密码者可无限尝试验证码而不触发告警 if user.TwoFA != "" { if valid := totp.Validate(req.PassCode, user.TwoFA); !valid { + s.loginFailed(r, sess, ip, string(decryptedUsername), failCount) Error(w, http.StatusForbidden, s.t.Get("invalid 2FA code")) return } } + s.guard.Reset(ip) + // 重新生成会话 ID if err = sess.Regenerate(true); err != nil { Error(w, http.StatusInternalServerError, "%v", err) @@ -154,16 +162,6 @@ func (s *UserService) Login(w http.ResponseWriter, r *http.Request) { // 安全登录下,将当前客户端与会话绑定 // 安全登录只在未启用面板 HTTPS 时生效 - ip := r.RemoteAddr - ipHeader := s.conf.HTTP.IPHeader - if ipHeader != "" && r.Header.Get(ipHeader) != "" { - ip = strings.Split(r.Header.Get(ipHeader), ",")[0] - } - ip, _, err = net.SplitHostPort(strings.TrimSpace(ip)) - if err != nil { - ip = r.RemoteAddr - } - if req.SafeLogin && !s.conf.HTTP.IsHTTPS() { sess.Put("safe_login", true) sess.Put("safe_client", fmt.Sprintf("%x", sha256.Sum256([]byte(ip)))) @@ -177,9 +175,34 @@ func (s *UserService) Login(w http.ResponseWriter, r *http.Request) { sess.Forget("key") sess.Forget("login_fail_count") + s.notifyRepo.SendEvent(biz.NotifyEventLogin, s.t.Get("[AcePanel] Panel Login"), biz.NotifyBody(s.t.Get("panel login detected"), [][2]string{ + {s.t.Get("Username"), user.Username}, + {s.t.Get("IP"), ip}, + {s.t.Get("User Agent"), r.UserAgent()}, + {s.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) + Success(w, nil) } +// loginFailed 记录一次登录失败,同来源短时间内失败过多时告警 +func (s *UserService) loginFailed(r *http.Request, sess *sessions.Session, ip, username string, failCount int) { + sess.Put("login_fail_count", failCount+1) + + count, exceeded := s.guard.Fail(ip) + if !exceeded { + return + } + + s.notifyRepo.SendEvent(biz.NotifyEventLoginFailed, s.t.Get("[AcePanel] Suspicious Login Attempts"), biz.NotifyBody(s.t.Get("too many failed panel login attempts"), [][2]string{ + {s.t.Get("IP"), ip}, + {s.t.Get("Username"), username}, + {s.t.Get("Failed Attempts"), cast.ToString(count)}, + {s.t.Get("User Agent"), r.UserAgent()}, + {s.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) +} + func (s *UserService) Logout(w http.ResponseWriter, r *http.Request) { sess, err := s.session.GetSession(r) if err != nil { diff --git a/internal/service/user_passkey.go b/internal/service/user_passkey.go index f7d2f218..0a07845b 100644 --- a/internal/service/user_passkey.go +++ b/internal/service/user_passkey.go @@ -31,6 +31,7 @@ type UserPasskeyService struct { session *sessions.Manager userPasskeyRepo *biz.UserPasskeyUsecase userRepo *biz.UserUsecase + notifyRepo *biz.NotifyUsecase } func NewUserPasskeyService(i do.Injector) (*UserPasskeyService, error) { @@ -42,6 +43,7 @@ func NewUserPasskeyService(i do.Injector) (*UserPasskeyService, error) { session: do.MustInvoke[*sessions.Manager](i), userPasskeyRepo: do.MustInvoke[*biz.UserPasskeyUsecase](i), userRepo: do.MustInvoke[*biz.UserUsecase](i), + notifyRepo: do.MustInvoke[*biz.NotifyUsecase](i), }, nil } @@ -296,6 +298,14 @@ func (s *UserPasskeyService) FinishLogin(w http.ResponseWriter, r *http.Request) sess.Forget("safe_login") sess.Forget("safe_client") + s.notifyRepo.SendEvent(biz.NotifyEventLogin, s.t.Get("[AcePanel] Panel Login"), biz.NotifyBody(s.t.Get("panel login detected"), [][2]string{ + {s.t.Get("Username"), wUser.Inner.Username}, + {s.t.Get("Method"), s.t.Get("passkey")}, + {s.t.Get("IP"), clientIP(r, s.conf.HTTP.IPHeader)}, + {s.t.Get("User Agent"), r.UserAgent()}, + {s.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) + Success(w, nil) } diff --git a/internal/taskqueue/runner.go b/internal/taskqueue/runner.go index b060e796..e223fe27 100644 --- a/internal/taskqueue/runner.go +++ b/internal/taskqueue/runner.go @@ -10,6 +10,7 @@ import ( "sync" "time" + "github.com/leonelquinteros/gotext" "gorm.io/gorm" "github.com/acepanel/panel/v3/internal/app" @@ -17,10 +18,17 @@ import ( "github.com/acepanel/panel/v3/pkg/shell" ) +// Notifier 事件通知,由 biz.NotifyUsecase 实现 +type Notifier interface { + SendEvent(event biz.NotifyEvent, subject, body string) +} + type Runner struct { - db *gorm.DB - log *slog.Logger - notify chan struct{} + db *gorm.DB + log *slog.Logger + notifier Notifier + t *gotext.Locale + notify chan struct{} mu sync.Mutex currentID uint // 当前运行的任务 ID @@ -28,11 +36,13 @@ type Runner struct { } // NewRunner 创建任务运行器 -func NewRunner(db *gorm.DB, log *slog.Logger) *Runner { +func NewRunner(db *gorm.DB, log *slog.Logger, notifier Notifier, t *gotext.Locale) *Runner { return &Runner{ - db: db, - log: log, - notify: make(chan struct{}, 1), + db: db, + log: log, + notifier: notifier, + t: t, + notify: make(chan struct{}, 1), } } @@ -153,6 +163,16 @@ func (r *Runner) execute(ctx context.Context, task *biz.Task) { } r.log.Warn("background task did not finish", slog.Any("task_id", task.ID), slog.Any("status", status), slog.Any("err", err)) _ = r.db.Model(task).Update("status", status).Error + + // 用户主动取消不算故障 + if status == biz.TaskStatusFailed { + r.notifier.SendEvent(biz.NotifyEventTaskFailed, r.t.Get("[AcePanel] Background Task Failed"), biz.NotifyBody(r.t.Get("background task failed"), [][2]string{ + {r.t.Get("Task"), task.Name}, + {r.t.Get("Log"), logFile}, + {r.t.Get("Error"), err.Error()}, + {r.t.Get("Time"), time.Now().Format(time.DateTime)}, + })) + } return } diff --git a/internal/taskqueue/runner_test.go b/internal/taskqueue/runner_test.go index 5c699e09..924a1ccd 100644 --- a/internal/taskqueue/runner_test.go +++ b/internal/taskqueue/runner_test.go @@ -7,6 +7,7 @@ import ( "testing" "time" + "github.com/leonelquinteros/gotext" "github.com/libtnb/sqlite" "gorm.io/gorm" @@ -28,9 +29,13 @@ func newRunnerForTest(t *testing.T) *Runner { if err = db.AutoMigrate(&biz.Task{}); err != nil { t.Fatal(err) } - return NewRunner(db, slog.New(slog.NewTextHandler(os.Stderr, nil))) + return NewRunner(db, slog.New(slog.NewTextHandler(os.Stderr, nil)), stubNotifier{}, gotext.NewLocale("", "en")) } +type stubNotifier struct{} + +func (stubNotifier) SendEvent(biz.NotifyEvent, string, string) {} + // 等待任务进入指定状态 func waitStatus(t *testing.T, db *gorm.DB, id uint, status biz.TaskStatus, timeout time.Duration) *biz.Task { t.Helper() diff --git a/mocks/biz/AlertRepo.go b/mocks/biz/AlertRepo.go new file mode 100644 index 00000000..1b0703ad --- /dev/null +++ b/mocks/biz/AlertRepo.go @@ -0,0 +1,844 @@ +// Code generated by mockery. DO NOT EDIT. + +package biz + +import ( + biz "github.com/acepanel/panel/v3/internal/biz" + mock "github.com/stretchr/testify/mock" + + time "time" +) + +// AlertRepo is an autogenerated mock type for the AlertRepo type +type AlertRepo struct { + mock.Mock +} + +type AlertRepo_Expecter struct { + mock *mock.Mock +} + +func (_m *AlertRepo) EXPECT() *AlertRepo_Expecter { + return &AlertRepo_Expecter{mock: &_m.Mock} +} + +// AddAlert provides a mock function with given fields: alert +func (_m *AlertRepo) AddAlert(alert *biz.Alert) error { + ret := _m.Called(alert) + + if len(ret) == 0 { + panic("no return value specified for AddAlert") + } + + var r0 error + if rf, ok := ret.Get(0).(func(*biz.Alert) error); ok { + r0 = rf(alert) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// AlertRepo_AddAlert_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddAlert' +type AlertRepo_AddAlert_Call struct { + *mock.Call +} + +// AddAlert is a helper method to define mock.On call +// - alert *biz.Alert +func (_e *AlertRepo_Expecter) AddAlert(alert interface{}) *AlertRepo_AddAlert_Call { + return &AlertRepo_AddAlert_Call{Call: _e.mock.On("AddAlert", alert)} +} + +func (_c *AlertRepo_AddAlert_Call) Run(run func(alert *biz.Alert)) *AlertRepo_AddAlert_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(*biz.Alert)) + }) + return _c +} + +func (_c *AlertRepo_AddAlert_Call) Return(_a0 error) *AlertRepo_AddAlert_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AlertRepo_AddAlert_Call) RunAndReturn(run func(*biz.Alert) error) *AlertRepo_AddAlert_Call { + _c.Call.Return(run) + return _c +} + +// AllRules provides a mock function with no fields +func (_m *AlertRepo) AllRules() ([]*biz.AlertRule, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for AllRules") + } + + var r0 []*biz.AlertRule + var r1 error + if rf, ok := ret.Get(0).(func() ([]*biz.AlertRule, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []*biz.AlertRule); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.AlertRule) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_AllRules_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AllRules' +type AlertRepo_AllRules_Call struct { + *mock.Call +} + +// AllRules is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) AllRules() *AlertRepo_AllRules_Call { + return &AlertRepo_AllRules_Call{Call: _e.mock.On("AllRules")} +} + +func (_c *AlertRepo_AllRules_Call) Run(run func()) *AlertRepo_AllRules_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_AllRules_Call) Return(_a0 []*biz.AlertRule, _a1 error) *AlertRepo_AllRules_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_AllRules_Call) RunAndReturn(run func() ([]*biz.AlertRule, error)) *AlertRepo_AllRules_Call { + _c.Call.Return(run) + return _c +} + +// CertExpiry provides a mock function with no fields +func (_m *AlertRepo) CertExpiry() ([]*biz.AlertMetric, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for CertExpiry") + } + + var r0 []*biz.AlertMetric + var r1 error + if rf, ok := ret.Get(0).(func() ([]*biz.AlertMetric, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []*biz.AlertMetric); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.AlertMetric) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_CertExpiry_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CertExpiry' +type AlertRepo_CertExpiry_Call struct { + *mock.Call +} + +// CertExpiry is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) CertExpiry() *AlertRepo_CertExpiry_Call { + return &AlertRepo_CertExpiry_Call{Call: _e.mock.On("CertExpiry")} +} + +func (_c *AlertRepo_CertExpiry_Call) Run(run func()) *AlertRepo_CertExpiry_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_CertExpiry_Call) Return(_a0 []*biz.AlertMetric, _a1 error) *AlertRepo_CertExpiry_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_CertExpiry_Call) RunAndReturn(run func() ([]*biz.AlertMetric, error)) *AlertRepo_CertExpiry_Call { + _c.Call.Return(run) + return _c +} + +// ClearAlerts provides a mock function with no fields +func (_m *AlertRepo) ClearAlerts() error { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for ClearAlerts") + } + + var r0 error + if rf, ok := ret.Get(0).(func() error); ok { + r0 = rf() + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// AlertRepo_ClearAlerts_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ClearAlerts' +type AlertRepo_ClearAlerts_Call struct { + *mock.Call +} + +// ClearAlerts is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) ClearAlerts() *AlertRepo_ClearAlerts_Call { + return &AlertRepo_ClearAlerts_Call{Call: _e.mock.On("ClearAlerts")} +} + +func (_c *AlertRepo_ClearAlerts_Call) Run(run func()) *AlertRepo_ClearAlerts_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_ClearAlerts_Call) Return(_a0 error) *AlertRepo_ClearAlerts_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AlertRepo_ClearAlerts_Call) RunAndReturn(run func() error) *AlertRepo_ClearAlerts_Call { + _c.Call.Return(run) + return _c +} + +// ClearAlertsBefore provides a mock function with given fields: t +func (_m *AlertRepo) ClearAlertsBefore(t time.Time) error { + ret := _m.Called(t) + + if len(ret) == 0 { + panic("no return value specified for ClearAlertsBefore") + } + + var r0 error + if rf, ok := ret.Get(0).(func(time.Time) error); ok { + r0 = rf(t) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// AlertRepo_ClearAlertsBefore_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ClearAlertsBefore' +type AlertRepo_ClearAlertsBefore_Call struct { + *mock.Call +} + +// ClearAlertsBefore is a helper method to define mock.On call +// - t time.Time +func (_e *AlertRepo_Expecter) ClearAlertsBefore(t interface{}) *AlertRepo_ClearAlertsBefore_Call { + return &AlertRepo_ClearAlertsBefore_Call{Call: _e.mock.On("ClearAlertsBefore", t)} +} + +func (_c *AlertRepo_ClearAlertsBefore_Call) Run(run func(t time.Time)) *AlertRepo_ClearAlertsBefore_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(time.Time)) + }) + return _c +} + +func (_c *AlertRepo_ClearAlertsBefore_Call) Return(_a0 error) *AlertRepo_ClearAlertsBefore_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AlertRepo_ClearAlertsBefore_Call) RunAndReturn(run func(time.Time) error) *AlertRepo_ClearAlertsBefore_Call { + _c.Call.Return(run) + return _c +} + +// CreateRule provides a mock function with given fields: rule +func (_m *AlertRepo) CreateRule(rule *biz.AlertRule) error { + ret := _m.Called(rule) + + if len(ret) == 0 { + panic("no return value specified for CreateRule") + } + + var r0 error + if rf, ok := ret.Get(0).(func(*biz.AlertRule) error); ok { + r0 = rf(rule) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// AlertRepo_CreateRule_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateRule' +type AlertRepo_CreateRule_Call struct { + *mock.Call +} + +// CreateRule is a helper method to define mock.On call +// - rule *biz.AlertRule +func (_e *AlertRepo_Expecter) CreateRule(rule interface{}) *AlertRepo_CreateRule_Call { + return &AlertRepo_CreateRule_Call{Call: _e.mock.On("CreateRule", rule)} +} + +func (_c *AlertRepo_CreateRule_Call) Run(run func(rule *biz.AlertRule)) *AlertRepo_CreateRule_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(*biz.AlertRule)) + }) + return _c +} + +func (_c *AlertRepo_CreateRule_Call) Return(_a0 error) *AlertRepo_CreateRule_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AlertRepo_CreateRule_Call) RunAndReturn(run func(*biz.AlertRule) error) *AlertRepo_CreateRule_Call { + _c.Call.Return(run) + return _c +} + +// DatabaseServers provides a mock function with no fields +func (_m *AlertRepo) DatabaseServers() ([]*biz.DatabaseServer, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for DatabaseServers") + } + + var r0 []*biz.DatabaseServer + var r1 error + if rf, ok := ret.Get(0).(func() ([]*biz.DatabaseServer, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []*biz.DatabaseServer); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.DatabaseServer) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_DatabaseServers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DatabaseServers' +type AlertRepo_DatabaseServers_Call struct { + *mock.Call +} + +// DatabaseServers is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) DatabaseServers() *AlertRepo_DatabaseServers_Call { + return &AlertRepo_DatabaseServers_Call{Call: _e.mock.On("DatabaseServers")} +} + +func (_c *AlertRepo_DatabaseServers_Call) Run(run func()) *AlertRepo_DatabaseServers_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_DatabaseServers_Call) Return(_a0 []*biz.DatabaseServer, _a1 error) *AlertRepo_DatabaseServers_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_DatabaseServers_Call) RunAndReturn(run func() ([]*biz.DatabaseServer, error)) *AlertRepo_DatabaseServers_Call { + _c.Call.Return(run) + return _c +} + +// DeleteRule provides a mock function with given fields: id +func (_m *AlertRepo) DeleteRule(id uint) error { + ret := _m.Called(id) + + if len(ret) == 0 { + panic("no return value specified for DeleteRule") + } + + var r0 error + if rf, ok := ret.Get(0).(func(uint) error); ok { + r0 = rf(id) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// AlertRepo_DeleteRule_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteRule' +type AlertRepo_DeleteRule_Call struct { + *mock.Call +} + +// DeleteRule is a helper method to define mock.On call +// - id uint +func (_e *AlertRepo_Expecter) DeleteRule(id interface{}) *AlertRepo_DeleteRule_Call { + return &AlertRepo_DeleteRule_Call{Call: _e.mock.On("DeleteRule", id)} +} + +func (_c *AlertRepo_DeleteRule_Call) Run(run func(id uint)) *AlertRepo_DeleteRule_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint)) + }) + return _c +} + +func (_c *AlertRepo_DeleteRule_Call) Return(_a0 error) *AlertRepo_DeleteRule_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AlertRepo_DeleteRule_Call) RunAndReturn(run func(uint) error) *AlertRepo_DeleteRule_Call { + _c.Call.Return(run) + return _c +} + +// GetRule provides a mock function with given fields: id +func (_m *AlertRepo) GetRule(id uint) (*biz.AlertRule, error) { + ret := _m.Called(id) + + if len(ret) == 0 { + panic("no return value specified for GetRule") + } + + var r0 *biz.AlertRule + var r1 error + if rf, ok := ret.Get(0).(func(uint) (*biz.AlertRule, error)); ok { + return rf(id) + } + if rf, ok := ret.Get(0).(func(uint) *biz.AlertRule); ok { + r0 = rf(id) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*biz.AlertRule) + } + } + + if rf, ok := ret.Get(1).(func(uint) error); ok { + r1 = rf(id) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_GetRule_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'GetRule' +type AlertRepo_GetRule_Call struct { + *mock.Call +} + +// GetRule is a helper method to define mock.On call +// - id uint +func (_e *AlertRepo_Expecter) GetRule(id interface{}) *AlertRepo_GetRule_Call { + return &AlertRepo_GetRule_Call{Call: _e.mock.On("GetRule", id)} +} + +func (_c *AlertRepo_GetRule_Call) Run(run func(id uint)) *AlertRepo_GetRule_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint)) + }) + return _c +} + +func (_c *AlertRepo_GetRule_Call) Return(_a0 *biz.AlertRule, _a1 error) *AlertRepo_GetRule_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_GetRule_Call) RunAndReturn(run func(uint) (*biz.AlertRule, error)) *AlertRepo_GetRule_Call { + _c.Call.Return(run) + return _c +} + +// ListAlerts provides a mock function with given fields: page, limit +func (_m *AlertRepo) ListAlerts(page uint, limit uint) ([]*biz.Alert, int64, error) { + ret := _m.Called(page, limit) + + if len(ret) == 0 { + panic("no return value specified for ListAlerts") + } + + var r0 []*biz.Alert + var r1 int64 + var r2 error + if rf, ok := ret.Get(0).(func(uint, uint) ([]*biz.Alert, int64, error)); ok { + return rf(page, limit) + } + if rf, ok := ret.Get(0).(func(uint, uint) []*biz.Alert); ok { + r0 = rf(page, limit) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.Alert) + } + } + + if rf, ok := ret.Get(1).(func(uint, uint) int64); ok { + r1 = rf(page, limit) + } else { + r1 = ret.Get(1).(int64) + } + + if rf, ok := ret.Get(2).(func(uint, uint) error); ok { + r2 = rf(page, limit) + } else { + r2 = ret.Error(2) + } + + return r0, r1, r2 +} + +// AlertRepo_ListAlerts_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAlerts' +type AlertRepo_ListAlerts_Call struct { + *mock.Call +} + +// ListAlerts is a helper method to define mock.On call +// - page uint +// - limit uint +func (_e *AlertRepo_Expecter) ListAlerts(page interface{}, limit interface{}) *AlertRepo_ListAlerts_Call { + return &AlertRepo_ListAlerts_Call{Call: _e.mock.On("ListAlerts", page, limit)} +} + +func (_c *AlertRepo_ListAlerts_Call) Run(run func(page uint, limit uint)) *AlertRepo_ListAlerts_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint), args[1].(uint)) + }) + return _c +} + +func (_c *AlertRepo_ListAlerts_Call) Return(_a0 []*biz.Alert, _a1 int64, _a2 error) *AlertRepo_ListAlerts_Call { + _c.Call.Return(_a0, _a1, _a2) + return _c +} + +func (_c *AlertRepo_ListAlerts_Call) RunAndReturn(run func(uint, uint) ([]*biz.Alert, int64, error)) *AlertRepo_ListAlerts_Call { + _c.Call.Return(run) + return _c +} + +// ListRules provides a mock function with given fields: page, limit +func (_m *AlertRepo) ListRules(page uint, limit uint) ([]*biz.AlertRule, int64, error) { + ret := _m.Called(page, limit) + + if len(ret) == 0 { + panic("no return value specified for ListRules") + } + + var r0 []*biz.AlertRule + var r1 int64 + var r2 error + if rf, ok := ret.Get(0).(func(uint, uint) ([]*biz.AlertRule, int64, error)); ok { + return rf(page, limit) + } + if rf, ok := ret.Get(0).(func(uint, uint) []*biz.AlertRule); ok { + r0 = rf(page, limit) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.AlertRule) + } + } + + if rf, ok := ret.Get(1).(func(uint, uint) int64); ok { + r1 = rf(page, limit) + } else { + r1 = ret.Get(1).(int64) + } + + if rf, ok := ret.Get(2).(func(uint, uint) error); ok { + r2 = rf(page, limit) + } else { + r2 = ret.Error(2) + } + + return r0, r1, r2 +} + +// AlertRepo_ListRules_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListRules' +type AlertRepo_ListRules_Call struct { + *mock.Call +} + +// ListRules is a helper method to define mock.On call +// - page uint +// - limit uint +func (_e *AlertRepo_Expecter) ListRules(page interface{}, limit interface{}) *AlertRepo_ListRules_Call { + return &AlertRepo_ListRules_Call{Call: _e.mock.On("ListRules", page, limit)} +} + +func (_c *AlertRepo_ListRules_Call) Run(run func(page uint, limit uint)) *AlertRepo_ListRules_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint), args[1].(uint)) + }) + return _c +} + +func (_c *AlertRepo_ListRules_Call) Return(_a0 []*biz.AlertRule, _a1 int64, _a2 error) *AlertRepo_ListRules_Call { + _c.Call.Return(_a0, _a1, _a2) + return _c +} + +func (_c *AlertRepo_ListRules_Call) RunAndReturn(run func(uint, uint) ([]*biz.AlertRule, int64, error)) *AlertRepo_ListRules_Call { + _c.Call.Return(run) + return _c +} + +// ProjectNames provides a mock function with no fields +func (_m *AlertRepo) ProjectNames() ([]string, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for ProjectNames") + } + + var r0 []string + var r1 error + if rf, ok := ret.Get(0).(func() ([]string, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []string); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]string) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_ProjectNames_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ProjectNames' +type AlertRepo_ProjectNames_Call struct { + *mock.Call +} + +// ProjectNames is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) ProjectNames() *AlertRepo_ProjectNames_Call { + return &AlertRepo_ProjectNames_Call{Call: _e.mock.On("ProjectNames")} +} + +func (_c *AlertRepo_ProjectNames_Call) Run(run func()) *AlertRepo_ProjectNames_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_ProjectNames_Call) Return(_a0 []string, _a1 error) *AlertRepo_ProjectNames_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_ProjectNames_Call) RunAndReturn(run func() ([]string, error)) *AlertRepo_ProjectNames_Call { + _c.Call.Return(run) + return _c +} + +// UpdateRule provides a mock function with given fields: rule +func (_m *AlertRepo) UpdateRule(rule *biz.AlertRule) error { + ret := _m.Called(rule) + + if len(ret) == 0 { + panic("no return value specified for UpdateRule") + } + + var r0 error + if rf, ok := ret.Get(0).(func(*biz.AlertRule) error); ok { + r0 = rf(rule) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// AlertRepo_UpdateRule_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRule' +type AlertRepo_UpdateRule_Call struct { + *mock.Call +} + +// UpdateRule is a helper method to define mock.On call +// - rule *biz.AlertRule +func (_e *AlertRepo_Expecter) UpdateRule(rule interface{}) *AlertRepo_UpdateRule_Call { + return &AlertRepo_UpdateRule_Call{Call: _e.mock.On("UpdateRule", rule)} +} + +func (_c *AlertRepo_UpdateRule_Call) Run(run func(rule *biz.AlertRule)) *AlertRepo_UpdateRule_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(*biz.AlertRule)) + }) + return _c +} + +func (_c *AlertRepo_UpdateRule_Call) Return(_a0 error) *AlertRepo_UpdateRule_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *AlertRepo_UpdateRule_Call) RunAndReturn(run func(*biz.AlertRule) error) *AlertRepo_UpdateRule_Call { + _c.Call.Return(run) + return _c +} + +// WebsiteExpiry provides a mock function with no fields +func (_m *AlertRepo) WebsiteExpiry() ([]*biz.AlertMetric, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for WebsiteExpiry") + } + + var r0 []*biz.AlertMetric + var r1 error + if rf, ok := ret.Get(0).(func() ([]*biz.AlertMetric, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []*biz.AlertMetric); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.AlertMetric) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_WebsiteExpiry_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'WebsiteExpiry' +type AlertRepo_WebsiteExpiry_Call struct { + *mock.Call +} + +// WebsiteExpiry is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) WebsiteExpiry() *AlertRepo_WebsiteExpiry_Call { + return &AlertRepo_WebsiteExpiry_Call{Call: _e.mock.On("WebsiteExpiry")} +} + +func (_c *AlertRepo_WebsiteExpiry_Call) Run(run func()) *AlertRepo_WebsiteExpiry_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_WebsiteExpiry_Call) Return(_a0 []*biz.AlertMetric, _a1 error) *AlertRepo_WebsiteExpiry_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_WebsiteExpiry_Call) RunAndReturn(run func() ([]*biz.AlertMetric, error)) *AlertRepo_WebsiteExpiry_Call { + _c.Call.Return(run) + return _c +} + +// WebsiteHourStats provides a mock function with no fields +func (_m *AlertRepo) WebsiteHourStats() ([]*biz.WebsiteHourStat, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for WebsiteHourStats") + } + + var r0 []*biz.WebsiteHourStat + var r1 error + if rf, ok := ret.Get(0).(func() ([]*biz.WebsiteHourStat, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []*biz.WebsiteHourStat); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.WebsiteHourStat) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// AlertRepo_WebsiteHourStats_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'WebsiteHourStats' +type AlertRepo_WebsiteHourStats_Call struct { + *mock.Call +} + +// WebsiteHourStats is a helper method to define mock.On call +func (_e *AlertRepo_Expecter) WebsiteHourStats() *AlertRepo_WebsiteHourStats_Call { + return &AlertRepo_WebsiteHourStats_Call{Call: _e.mock.On("WebsiteHourStats")} +} + +func (_c *AlertRepo_WebsiteHourStats_Call) Run(run func()) *AlertRepo_WebsiteHourStats_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *AlertRepo_WebsiteHourStats_Call) Return(_a0 []*biz.WebsiteHourStat, _a1 error) *AlertRepo_WebsiteHourStats_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *AlertRepo_WebsiteHourStats_Call) RunAndReturn(run func() ([]*biz.WebsiteHourStat, error)) *AlertRepo_WebsiteHourStats_Call { + _c.Call.Return(run) + return _c +} + +// NewAlertRepo creates a new instance of AlertRepo. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewAlertRepo(t interface { + mock.TestingT + Cleanup(func()) +}) *AlertRepo { + mock := &AlertRepo{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/mocks/biz/DatabaseElasticsearchRepo.go b/mocks/biz/DatabaseElasticsearchRepo.go index 7c696753..83169db0 100644 --- a/mocks/biz/DatabaseElasticsearchRepo.go +++ b/mocks/biz/DatabaseElasticsearchRepo.go @@ -3,6 +3,8 @@ package biz import ( + context "context" + db "github.com/acepanel/panel/v3/pkg/db" mock "github.com/stretchr/testify/mock" @@ -22,9 +24,9 @@ func (_m *DatabaseElasticsearchRepo) EXPECT() *DatabaseElasticsearchRepo_Expecte return &DatabaseElasticsearchRepo_Expecter{mock: &_m.Mock} } -// Data provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) Data(req *request.DatabaseESData) ([]db.ESDocument, int64, error) { - ret := _m.Called(req) +// Data provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) Data(ctx context.Context, req *request.DatabaseESData) ([]db.ESDocument, int64, error) { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for Data") @@ -33,25 +35,25 @@ func (_m *DatabaseElasticsearchRepo) Data(req *request.DatabaseESData) ([]db.ESD var r0 []db.ESDocument var r1 int64 var r2 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESData) ([]db.ESDocument, int64, error)); ok { - return rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESData) ([]db.ESDocument, int64, error)); ok { + return rf(ctx, req) } - if rf, ok := ret.Get(0).(func(*request.DatabaseESData) []db.ESDocument); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESData) []db.ESDocument); ok { + r0 = rf(ctx, req) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]db.ESDocument) } } - if rf, ok := ret.Get(1).(func(*request.DatabaseESData) int64); ok { - r1 = rf(req) + if rf, ok := ret.Get(1).(func(context.Context, *request.DatabaseESData) int64); ok { + r1 = rf(ctx, req) } else { r1 = ret.Get(1).(int64) } - if rf, ok := ret.Get(2).(func(*request.DatabaseESData) error); ok { - r2 = rf(req) + if rf, ok := ret.Get(2).(func(context.Context, *request.DatabaseESData) error); ok { + r2 = rf(ctx, req) } else { r2 = ret.Error(2) } @@ -65,14 +67,15 @@ type DatabaseElasticsearchRepo_Data_Call struct { } // Data is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESData -func (_e *DatabaseElasticsearchRepo_Expecter) Data(req interface{}) *DatabaseElasticsearchRepo_Data_Call { - return &DatabaseElasticsearchRepo_Data_Call{Call: _e.mock.On("Data", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) Data(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_Data_Call { + return &DatabaseElasticsearchRepo_Data_Call{Call: _e.mock.On("Data", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_Data_Call) Run(run func(req *request.DatabaseESData)) *DatabaseElasticsearchRepo_Data_Call { +func (_c *DatabaseElasticsearchRepo_Data_Call) Run(run func(ctx context.Context, req *request.DatabaseESData)) *DatabaseElasticsearchRepo_Data_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESData)) + run(args[0].(context.Context), args[1].(*request.DatabaseESData)) }) return _c } @@ -82,22 +85,22 @@ func (_c *DatabaseElasticsearchRepo_Data_Call) Return(_a0 []db.ESDocument, _a1 i return _c } -func (_c *DatabaseElasticsearchRepo_Data_Call) RunAndReturn(run func(*request.DatabaseESData) ([]db.ESDocument, int64, error)) *DatabaseElasticsearchRepo_Data_Call { +func (_c *DatabaseElasticsearchRepo_Data_Call) RunAndReturn(run func(context.Context, *request.DatabaseESData) ([]db.ESDocument, int64, error)) *DatabaseElasticsearchRepo_Data_Call { _c.Call.Return(run) return _c } -// DocumentDelete provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) DocumentDelete(req *request.DatabaseESDocumentDelete) error { - ret := _m.Called(req) +// DocumentDelete provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) DocumentDelete(ctx context.Context, req *request.DatabaseESDocumentDelete) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for DocumentDelete") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESDocumentDelete) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESDocumentDelete) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -111,14 +114,15 @@ type DatabaseElasticsearchRepo_DocumentDelete_Call struct { } // DocumentDelete is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESDocumentDelete -func (_e *DatabaseElasticsearchRepo_Expecter) DocumentDelete(req interface{}) *DatabaseElasticsearchRepo_DocumentDelete_Call { - return &DatabaseElasticsearchRepo_DocumentDelete_Call{Call: _e.mock.On("DocumentDelete", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) DocumentDelete(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_DocumentDelete_Call { + return &DatabaseElasticsearchRepo_DocumentDelete_Call{Call: _e.mock.On("DocumentDelete", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_DocumentDelete_Call) Run(run func(req *request.DatabaseESDocumentDelete)) *DatabaseElasticsearchRepo_DocumentDelete_Call { +func (_c *DatabaseElasticsearchRepo_DocumentDelete_Call) Run(run func(ctx context.Context, req *request.DatabaseESDocumentDelete)) *DatabaseElasticsearchRepo_DocumentDelete_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESDocumentDelete)) + run(args[0].(context.Context), args[1].(*request.DatabaseESDocumentDelete)) }) return _c } @@ -128,14 +132,14 @@ func (_c *DatabaseElasticsearchRepo_DocumentDelete_Call) Return(_a0 error) *Data return _c } -func (_c *DatabaseElasticsearchRepo_DocumentDelete_Call) RunAndReturn(run func(*request.DatabaseESDocumentDelete) error) *DatabaseElasticsearchRepo_DocumentDelete_Call { +func (_c *DatabaseElasticsearchRepo_DocumentDelete_Call) RunAndReturn(run func(context.Context, *request.DatabaseESDocumentDelete) error) *DatabaseElasticsearchRepo_DocumentDelete_Call { _c.Call.Return(run) return _c } -// DocumentGet provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) DocumentGet(req *request.DatabaseESDocumentGet) (*db.ESDocument, error) { - ret := _m.Called(req) +// DocumentGet provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) DocumentGet(ctx context.Context, req *request.DatabaseESDocumentGet) (*db.ESDocument, error) { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for DocumentGet") @@ -143,19 +147,19 @@ func (_m *DatabaseElasticsearchRepo) DocumentGet(req *request.DatabaseESDocument var r0 *db.ESDocument var r1 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESDocumentGet) (*db.ESDocument, error)); ok { - return rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESDocumentGet) (*db.ESDocument, error)); ok { + return rf(ctx, req) } - if rf, ok := ret.Get(0).(func(*request.DatabaseESDocumentGet) *db.ESDocument); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESDocumentGet) *db.ESDocument); ok { + r0 = rf(ctx, req) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*db.ESDocument) } } - if rf, ok := ret.Get(1).(func(*request.DatabaseESDocumentGet) error); ok { - r1 = rf(req) + if rf, ok := ret.Get(1).(func(context.Context, *request.DatabaseESDocumentGet) error); ok { + r1 = rf(ctx, req) } else { r1 = ret.Error(1) } @@ -169,14 +173,15 @@ type DatabaseElasticsearchRepo_DocumentGet_Call struct { } // DocumentGet is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESDocumentGet -func (_e *DatabaseElasticsearchRepo_Expecter) DocumentGet(req interface{}) *DatabaseElasticsearchRepo_DocumentGet_Call { - return &DatabaseElasticsearchRepo_DocumentGet_Call{Call: _e.mock.On("DocumentGet", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) DocumentGet(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_DocumentGet_Call { + return &DatabaseElasticsearchRepo_DocumentGet_Call{Call: _e.mock.On("DocumentGet", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_DocumentGet_Call) Run(run func(req *request.DatabaseESDocumentGet)) *DatabaseElasticsearchRepo_DocumentGet_Call { +func (_c *DatabaseElasticsearchRepo_DocumentGet_Call) Run(run func(ctx context.Context, req *request.DatabaseESDocumentGet)) *DatabaseElasticsearchRepo_DocumentGet_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESDocumentGet)) + run(args[0].(context.Context), args[1].(*request.DatabaseESDocumentGet)) }) return _c } @@ -186,22 +191,22 @@ func (_c *DatabaseElasticsearchRepo_DocumentGet_Call) Return(_a0 *db.ESDocument, return _c } -func (_c *DatabaseElasticsearchRepo_DocumentGet_Call) RunAndReturn(run func(*request.DatabaseESDocumentGet) (*db.ESDocument, error)) *DatabaseElasticsearchRepo_DocumentGet_Call { +func (_c *DatabaseElasticsearchRepo_DocumentGet_Call) RunAndReturn(run func(context.Context, *request.DatabaseESDocumentGet) (*db.ESDocument, error)) *DatabaseElasticsearchRepo_DocumentGet_Call { _c.Call.Return(run) return _c } -// DocumentSet provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) DocumentSet(req *request.DatabaseESDocumentSet) error { - ret := _m.Called(req) +// DocumentSet provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) DocumentSet(ctx context.Context, req *request.DatabaseESDocumentSet) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for DocumentSet") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESDocumentSet) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESDocumentSet) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -215,14 +220,15 @@ type DatabaseElasticsearchRepo_DocumentSet_Call struct { } // DocumentSet is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESDocumentSet -func (_e *DatabaseElasticsearchRepo_Expecter) DocumentSet(req interface{}) *DatabaseElasticsearchRepo_DocumentSet_Call { - return &DatabaseElasticsearchRepo_DocumentSet_Call{Call: _e.mock.On("DocumentSet", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) DocumentSet(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_DocumentSet_Call { + return &DatabaseElasticsearchRepo_DocumentSet_Call{Call: _e.mock.On("DocumentSet", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_DocumentSet_Call) Run(run func(req *request.DatabaseESDocumentSet)) *DatabaseElasticsearchRepo_DocumentSet_Call { +func (_c *DatabaseElasticsearchRepo_DocumentSet_Call) Run(run func(ctx context.Context, req *request.DatabaseESDocumentSet)) *DatabaseElasticsearchRepo_DocumentSet_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESDocumentSet)) + run(args[0].(context.Context), args[1].(*request.DatabaseESDocumentSet)) }) return _c } @@ -232,22 +238,22 @@ func (_c *DatabaseElasticsearchRepo_DocumentSet_Call) Return(_a0 error) *Databas return _c } -func (_c *DatabaseElasticsearchRepo_DocumentSet_Call) RunAndReturn(run func(*request.DatabaseESDocumentSet) error) *DatabaseElasticsearchRepo_DocumentSet_Call { +func (_c *DatabaseElasticsearchRepo_DocumentSet_Call) RunAndReturn(run func(context.Context, *request.DatabaseESDocumentSet) error) *DatabaseElasticsearchRepo_DocumentSet_Call { _c.Call.Return(run) return _c } -// IndexCreate provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) IndexCreate(req *request.DatabaseESIndexCreate) error { - ret := _m.Called(req) +// IndexCreate provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) IndexCreate(ctx context.Context, req *request.DatabaseESIndexCreate) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for IndexCreate") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESIndexCreate) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESIndexCreate) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -261,14 +267,15 @@ type DatabaseElasticsearchRepo_IndexCreate_Call struct { } // IndexCreate is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESIndexCreate -func (_e *DatabaseElasticsearchRepo_Expecter) IndexCreate(req interface{}) *DatabaseElasticsearchRepo_IndexCreate_Call { - return &DatabaseElasticsearchRepo_IndexCreate_Call{Call: _e.mock.On("IndexCreate", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) IndexCreate(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_IndexCreate_Call { + return &DatabaseElasticsearchRepo_IndexCreate_Call{Call: _e.mock.On("IndexCreate", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_IndexCreate_Call) Run(run func(req *request.DatabaseESIndexCreate)) *DatabaseElasticsearchRepo_IndexCreate_Call { +func (_c *DatabaseElasticsearchRepo_IndexCreate_Call) Run(run func(ctx context.Context, req *request.DatabaseESIndexCreate)) *DatabaseElasticsearchRepo_IndexCreate_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESIndexCreate)) + run(args[0].(context.Context), args[1].(*request.DatabaseESIndexCreate)) }) return _c } @@ -278,22 +285,22 @@ func (_c *DatabaseElasticsearchRepo_IndexCreate_Call) Return(_a0 error) *Databas return _c } -func (_c *DatabaseElasticsearchRepo_IndexCreate_Call) RunAndReturn(run func(*request.DatabaseESIndexCreate) error) *DatabaseElasticsearchRepo_IndexCreate_Call { +func (_c *DatabaseElasticsearchRepo_IndexCreate_Call) RunAndReturn(run func(context.Context, *request.DatabaseESIndexCreate) error) *DatabaseElasticsearchRepo_IndexCreate_Call { _c.Call.Return(run) return _c } -// IndexDelete provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) IndexDelete(req *request.DatabaseESIndexDelete) error { - ret := _m.Called(req) +// IndexDelete provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) IndexDelete(ctx context.Context, req *request.DatabaseESIndexDelete) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for IndexDelete") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESIndexDelete) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESIndexDelete) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -307,14 +314,15 @@ type DatabaseElasticsearchRepo_IndexDelete_Call struct { } // IndexDelete is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESIndexDelete -func (_e *DatabaseElasticsearchRepo_Expecter) IndexDelete(req interface{}) *DatabaseElasticsearchRepo_IndexDelete_Call { - return &DatabaseElasticsearchRepo_IndexDelete_Call{Call: _e.mock.On("IndexDelete", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) IndexDelete(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_IndexDelete_Call { + return &DatabaseElasticsearchRepo_IndexDelete_Call{Call: _e.mock.On("IndexDelete", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_IndexDelete_Call) Run(run func(req *request.DatabaseESIndexDelete)) *DatabaseElasticsearchRepo_IndexDelete_Call { +func (_c *DatabaseElasticsearchRepo_IndexDelete_Call) Run(run func(ctx context.Context, req *request.DatabaseESIndexDelete)) *DatabaseElasticsearchRepo_IndexDelete_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESIndexDelete)) + run(args[0].(context.Context), args[1].(*request.DatabaseESIndexDelete)) }) return _c } @@ -324,14 +332,14 @@ func (_c *DatabaseElasticsearchRepo_IndexDelete_Call) Return(_a0 error) *Databas return _c } -func (_c *DatabaseElasticsearchRepo_IndexDelete_Call) RunAndReturn(run func(*request.DatabaseESIndexDelete) error) *DatabaseElasticsearchRepo_IndexDelete_Call { +func (_c *DatabaseElasticsearchRepo_IndexDelete_Call) RunAndReturn(run func(context.Context, *request.DatabaseESIndexDelete) error) *DatabaseElasticsearchRepo_IndexDelete_Call { _c.Call.Return(run) return _c } -// Indices provides a mock function with given fields: req -func (_m *DatabaseElasticsearchRepo) Indices(req *request.DatabaseESIndices) ([]db.ESIndex, error) { - ret := _m.Called(req) +// Indices provides a mock function with given fields: ctx, req +func (_m *DatabaseElasticsearchRepo) Indices(ctx context.Context, req *request.DatabaseESIndices) ([]db.ESIndex, error) { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for Indices") @@ -339,19 +347,19 @@ func (_m *DatabaseElasticsearchRepo) Indices(req *request.DatabaseESIndices) ([] var r0 []db.ESIndex var r1 error - if rf, ok := ret.Get(0).(func(*request.DatabaseESIndices) ([]db.ESIndex, error)); ok { - return rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESIndices) ([]db.ESIndex, error)); ok { + return rf(ctx, req) } - if rf, ok := ret.Get(0).(func(*request.DatabaseESIndices) []db.ESIndex); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseESIndices) []db.ESIndex); ok { + r0 = rf(ctx, req) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]db.ESIndex) } } - if rf, ok := ret.Get(1).(func(*request.DatabaseESIndices) error); ok { - r1 = rf(req) + if rf, ok := ret.Get(1).(func(context.Context, *request.DatabaseESIndices) error); ok { + r1 = rf(ctx, req) } else { r1 = ret.Error(1) } @@ -365,14 +373,15 @@ type DatabaseElasticsearchRepo_Indices_Call struct { } // Indices is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseESIndices -func (_e *DatabaseElasticsearchRepo_Expecter) Indices(req interface{}) *DatabaseElasticsearchRepo_Indices_Call { - return &DatabaseElasticsearchRepo_Indices_Call{Call: _e.mock.On("Indices", req)} +func (_e *DatabaseElasticsearchRepo_Expecter) Indices(ctx interface{}, req interface{}) *DatabaseElasticsearchRepo_Indices_Call { + return &DatabaseElasticsearchRepo_Indices_Call{Call: _e.mock.On("Indices", ctx, req)} } -func (_c *DatabaseElasticsearchRepo_Indices_Call) Run(run func(req *request.DatabaseESIndices)) *DatabaseElasticsearchRepo_Indices_Call { +func (_c *DatabaseElasticsearchRepo_Indices_Call) Run(run func(ctx context.Context, req *request.DatabaseESIndices)) *DatabaseElasticsearchRepo_Indices_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseESIndices)) + run(args[0].(context.Context), args[1].(*request.DatabaseESIndices)) }) return _c } @@ -382,7 +391,7 @@ func (_c *DatabaseElasticsearchRepo_Indices_Call) Return(_a0 []db.ESIndex, _a1 e return _c } -func (_c *DatabaseElasticsearchRepo_Indices_Call) RunAndReturn(run func(*request.DatabaseESIndices) ([]db.ESIndex, error)) *DatabaseElasticsearchRepo_Indices_Call { +func (_c *DatabaseElasticsearchRepo_Indices_Call) RunAndReturn(run func(context.Context, *request.DatabaseESIndices) ([]db.ESIndex, error)) *DatabaseElasticsearchRepo_Indices_Call { _c.Call.Return(run) return _c } diff --git a/mocks/biz/DatabaseRedisRepo.go b/mocks/biz/DatabaseRedisRepo.go index 03422d85..1fc4792d 100644 --- a/mocks/biz/DatabaseRedisRepo.go +++ b/mocks/biz/DatabaseRedisRepo.go @@ -3,6 +3,8 @@ package biz import ( + context "context" + db "github.com/acepanel/panel/v3/pkg/db" mock "github.com/stretchr/testify/mock" @@ -22,17 +24,17 @@ func (_m *DatabaseRedisRepo) EXPECT() *DatabaseRedisRepo_Expecter { return &DatabaseRedisRepo_Expecter{mock: &_m.Mock} } -// Clear provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) Clear(req *request.DatabaseRedisClear) error { - ret := _m.Called(req) +// Clear provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) Clear(ctx context.Context, req *request.DatabaseRedisClear) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for Clear") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisClear) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisClear) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -46,14 +48,15 @@ type DatabaseRedisRepo_Clear_Call struct { } // Clear is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisClear -func (_e *DatabaseRedisRepo_Expecter) Clear(req interface{}) *DatabaseRedisRepo_Clear_Call { - return &DatabaseRedisRepo_Clear_Call{Call: _e.mock.On("Clear", req)} +func (_e *DatabaseRedisRepo_Expecter) Clear(ctx interface{}, req interface{}) *DatabaseRedisRepo_Clear_Call { + return &DatabaseRedisRepo_Clear_Call{Call: _e.mock.On("Clear", ctx, req)} } -func (_c *DatabaseRedisRepo_Clear_Call) Run(run func(req *request.DatabaseRedisClear)) *DatabaseRedisRepo_Clear_Call { +func (_c *DatabaseRedisRepo_Clear_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisClear)) *DatabaseRedisRepo_Clear_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisClear)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisClear)) }) return _c } @@ -63,14 +66,14 @@ func (_c *DatabaseRedisRepo_Clear_Call) Return(_a0 error) *DatabaseRedisRepo_Cle return _c } -func (_c *DatabaseRedisRepo_Clear_Call) RunAndReturn(run func(*request.DatabaseRedisClear) error) *DatabaseRedisRepo_Clear_Call { +func (_c *DatabaseRedisRepo_Clear_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisClear) error) *DatabaseRedisRepo_Clear_Call { _c.Call.Return(run) return _c } -// Data provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) Data(req *request.DatabaseRedisData) ([]db.RedisKV, int, error) { - ret := _m.Called(req) +// Data provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) Data(ctx context.Context, req *request.DatabaseRedisData) ([]db.RedisKV, int, error) { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for Data") @@ -79,25 +82,25 @@ func (_m *DatabaseRedisRepo) Data(req *request.DatabaseRedisData) ([]db.RedisKV, var r0 []db.RedisKV var r1 int var r2 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisData) ([]db.RedisKV, int, error)); ok { - return rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisData) ([]db.RedisKV, int, error)); ok { + return rf(ctx, req) } - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisData) []db.RedisKV); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisData) []db.RedisKV); ok { + r0 = rf(ctx, req) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]db.RedisKV) } } - if rf, ok := ret.Get(1).(func(*request.DatabaseRedisData) int); ok { - r1 = rf(req) + if rf, ok := ret.Get(1).(func(context.Context, *request.DatabaseRedisData) int); ok { + r1 = rf(ctx, req) } else { r1 = ret.Get(1).(int) } - if rf, ok := ret.Get(2).(func(*request.DatabaseRedisData) error); ok { - r2 = rf(req) + if rf, ok := ret.Get(2).(func(context.Context, *request.DatabaseRedisData) error); ok { + r2 = rf(ctx, req) } else { r2 = ret.Error(2) } @@ -111,14 +114,15 @@ type DatabaseRedisRepo_Data_Call struct { } // Data is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisData -func (_e *DatabaseRedisRepo_Expecter) Data(req interface{}) *DatabaseRedisRepo_Data_Call { - return &DatabaseRedisRepo_Data_Call{Call: _e.mock.On("Data", req)} +func (_e *DatabaseRedisRepo_Expecter) Data(ctx interface{}, req interface{}) *DatabaseRedisRepo_Data_Call { + return &DatabaseRedisRepo_Data_Call{Call: _e.mock.On("Data", ctx, req)} } -func (_c *DatabaseRedisRepo_Data_Call) Run(run func(req *request.DatabaseRedisData)) *DatabaseRedisRepo_Data_Call { +func (_c *DatabaseRedisRepo_Data_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisData)) *DatabaseRedisRepo_Data_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisData)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisData)) }) return _c } @@ -128,14 +132,14 @@ func (_c *DatabaseRedisRepo_Data_Call) Return(_a0 []db.RedisKV, _a1 int, _a2 err return _c } -func (_c *DatabaseRedisRepo_Data_Call) RunAndReturn(run func(*request.DatabaseRedisData) ([]db.RedisKV, int, error)) *DatabaseRedisRepo_Data_Call { +func (_c *DatabaseRedisRepo_Data_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisData) ([]db.RedisKV, int, error)) *DatabaseRedisRepo_Data_Call { _c.Call.Return(run) return _c } -// Databases provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) Databases(req *request.DatabaseRedisDatabases) (int, error) { - ret := _m.Called(req) +// Databases provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) Databases(ctx context.Context, req *request.DatabaseRedisDatabases) (int, error) { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for Databases") @@ -143,17 +147,17 @@ func (_m *DatabaseRedisRepo) Databases(req *request.DatabaseRedisDatabases) (int var r0 int var r1 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisDatabases) (int, error)); ok { - return rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisDatabases) (int, error)); ok { + return rf(ctx, req) } - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisDatabases) int); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisDatabases) int); ok { + r0 = rf(ctx, req) } else { r0 = ret.Get(0).(int) } - if rf, ok := ret.Get(1).(func(*request.DatabaseRedisDatabases) error); ok { - r1 = rf(req) + if rf, ok := ret.Get(1).(func(context.Context, *request.DatabaseRedisDatabases) error); ok { + r1 = rf(ctx, req) } else { r1 = ret.Error(1) } @@ -167,14 +171,15 @@ type DatabaseRedisRepo_Databases_Call struct { } // Databases is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisDatabases -func (_e *DatabaseRedisRepo_Expecter) Databases(req interface{}) *DatabaseRedisRepo_Databases_Call { - return &DatabaseRedisRepo_Databases_Call{Call: _e.mock.On("Databases", req)} +func (_e *DatabaseRedisRepo_Expecter) Databases(ctx interface{}, req interface{}) *DatabaseRedisRepo_Databases_Call { + return &DatabaseRedisRepo_Databases_Call{Call: _e.mock.On("Databases", ctx, req)} } -func (_c *DatabaseRedisRepo_Databases_Call) Run(run func(req *request.DatabaseRedisDatabases)) *DatabaseRedisRepo_Databases_Call { +func (_c *DatabaseRedisRepo_Databases_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisDatabases)) *DatabaseRedisRepo_Databases_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisDatabases)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisDatabases)) }) return _c } @@ -184,22 +189,22 @@ func (_c *DatabaseRedisRepo_Databases_Call) Return(_a0 int, _a1 error) *Database return _c } -func (_c *DatabaseRedisRepo_Databases_Call) RunAndReturn(run func(*request.DatabaseRedisDatabases) (int, error)) *DatabaseRedisRepo_Databases_Call { +func (_c *DatabaseRedisRepo_Databases_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisDatabases) (int, error)) *DatabaseRedisRepo_Databases_Call { _c.Call.Return(run) return _c } -// KeyDelete provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) KeyDelete(req *request.DatabaseRedisKeyDelete) error { - ret := _m.Called(req) +// KeyDelete provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) KeyDelete(ctx context.Context, req *request.DatabaseRedisKeyDelete) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for KeyDelete") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisKeyDelete) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisKeyDelete) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -213,14 +218,15 @@ type DatabaseRedisRepo_KeyDelete_Call struct { } // KeyDelete is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisKeyDelete -func (_e *DatabaseRedisRepo_Expecter) KeyDelete(req interface{}) *DatabaseRedisRepo_KeyDelete_Call { - return &DatabaseRedisRepo_KeyDelete_Call{Call: _e.mock.On("KeyDelete", req)} +func (_e *DatabaseRedisRepo_Expecter) KeyDelete(ctx interface{}, req interface{}) *DatabaseRedisRepo_KeyDelete_Call { + return &DatabaseRedisRepo_KeyDelete_Call{Call: _e.mock.On("KeyDelete", ctx, req)} } -func (_c *DatabaseRedisRepo_KeyDelete_Call) Run(run func(req *request.DatabaseRedisKeyDelete)) *DatabaseRedisRepo_KeyDelete_Call { +func (_c *DatabaseRedisRepo_KeyDelete_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisKeyDelete)) *DatabaseRedisRepo_KeyDelete_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisKeyDelete)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisKeyDelete)) }) return _c } @@ -230,14 +236,14 @@ func (_c *DatabaseRedisRepo_KeyDelete_Call) Return(_a0 error) *DatabaseRedisRepo return _c } -func (_c *DatabaseRedisRepo_KeyDelete_Call) RunAndReturn(run func(*request.DatabaseRedisKeyDelete) error) *DatabaseRedisRepo_KeyDelete_Call { +func (_c *DatabaseRedisRepo_KeyDelete_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisKeyDelete) error) *DatabaseRedisRepo_KeyDelete_Call { _c.Call.Return(run) return _c } -// KeyGet provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) KeyGet(req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) { - ret := _m.Called(req) +// KeyGet provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) KeyGet(ctx context.Context, req *request.DatabaseRedisKeyGet) (*db.RedisKV, error) { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for KeyGet") @@ -245,19 +251,19 @@ func (_m *DatabaseRedisRepo) KeyGet(req *request.DatabaseRedisKeyGet) (*db.Redis var r0 *db.RedisKV var r1 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisKeyGet) (*db.RedisKV, error)); ok { - return rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisKeyGet) (*db.RedisKV, error)); ok { + return rf(ctx, req) } - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisKeyGet) *db.RedisKV); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisKeyGet) *db.RedisKV); ok { + r0 = rf(ctx, req) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*db.RedisKV) } } - if rf, ok := ret.Get(1).(func(*request.DatabaseRedisKeyGet) error); ok { - r1 = rf(req) + if rf, ok := ret.Get(1).(func(context.Context, *request.DatabaseRedisKeyGet) error); ok { + r1 = rf(ctx, req) } else { r1 = ret.Error(1) } @@ -271,14 +277,15 @@ type DatabaseRedisRepo_KeyGet_Call struct { } // KeyGet is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisKeyGet -func (_e *DatabaseRedisRepo_Expecter) KeyGet(req interface{}) *DatabaseRedisRepo_KeyGet_Call { - return &DatabaseRedisRepo_KeyGet_Call{Call: _e.mock.On("KeyGet", req)} +func (_e *DatabaseRedisRepo_Expecter) KeyGet(ctx interface{}, req interface{}) *DatabaseRedisRepo_KeyGet_Call { + return &DatabaseRedisRepo_KeyGet_Call{Call: _e.mock.On("KeyGet", ctx, req)} } -func (_c *DatabaseRedisRepo_KeyGet_Call) Run(run func(req *request.DatabaseRedisKeyGet)) *DatabaseRedisRepo_KeyGet_Call { +func (_c *DatabaseRedisRepo_KeyGet_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisKeyGet)) *DatabaseRedisRepo_KeyGet_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisKeyGet)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisKeyGet)) }) return _c } @@ -288,22 +295,22 @@ func (_c *DatabaseRedisRepo_KeyGet_Call) Return(_a0 *db.RedisKV, _a1 error) *Dat return _c } -func (_c *DatabaseRedisRepo_KeyGet_Call) RunAndReturn(run func(*request.DatabaseRedisKeyGet) (*db.RedisKV, error)) *DatabaseRedisRepo_KeyGet_Call { +func (_c *DatabaseRedisRepo_KeyGet_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisKeyGet) (*db.RedisKV, error)) *DatabaseRedisRepo_KeyGet_Call { _c.Call.Return(run) return _c } -// KeyRename provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) KeyRename(req *request.DatabaseRedisKeyRename) error { - ret := _m.Called(req) +// KeyRename provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) KeyRename(ctx context.Context, req *request.DatabaseRedisKeyRename) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for KeyRename") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisKeyRename) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisKeyRename) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -317,14 +324,15 @@ type DatabaseRedisRepo_KeyRename_Call struct { } // KeyRename is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisKeyRename -func (_e *DatabaseRedisRepo_Expecter) KeyRename(req interface{}) *DatabaseRedisRepo_KeyRename_Call { - return &DatabaseRedisRepo_KeyRename_Call{Call: _e.mock.On("KeyRename", req)} +func (_e *DatabaseRedisRepo_Expecter) KeyRename(ctx interface{}, req interface{}) *DatabaseRedisRepo_KeyRename_Call { + return &DatabaseRedisRepo_KeyRename_Call{Call: _e.mock.On("KeyRename", ctx, req)} } -func (_c *DatabaseRedisRepo_KeyRename_Call) Run(run func(req *request.DatabaseRedisKeyRename)) *DatabaseRedisRepo_KeyRename_Call { +func (_c *DatabaseRedisRepo_KeyRename_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisKeyRename)) *DatabaseRedisRepo_KeyRename_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisKeyRename)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisKeyRename)) }) return _c } @@ -334,22 +342,22 @@ func (_c *DatabaseRedisRepo_KeyRename_Call) Return(_a0 error) *DatabaseRedisRepo return _c } -func (_c *DatabaseRedisRepo_KeyRename_Call) RunAndReturn(run func(*request.DatabaseRedisKeyRename) error) *DatabaseRedisRepo_KeyRename_Call { +func (_c *DatabaseRedisRepo_KeyRename_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisKeyRename) error) *DatabaseRedisRepo_KeyRename_Call { _c.Call.Return(run) return _c } -// KeySet provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) KeySet(req *request.DatabaseRedisKeySet) error { - ret := _m.Called(req) +// KeySet provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) KeySet(ctx context.Context, req *request.DatabaseRedisKeySet) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for KeySet") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisKeySet) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisKeySet) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -363,14 +371,15 @@ type DatabaseRedisRepo_KeySet_Call struct { } // KeySet is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisKeySet -func (_e *DatabaseRedisRepo_Expecter) KeySet(req interface{}) *DatabaseRedisRepo_KeySet_Call { - return &DatabaseRedisRepo_KeySet_Call{Call: _e.mock.On("KeySet", req)} +func (_e *DatabaseRedisRepo_Expecter) KeySet(ctx interface{}, req interface{}) *DatabaseRedisRepo_KeySet_Call { + return &DatabaseRedisRepo_KeySet_Call{Call: _e.mock.On("KeySet", ctx, req)} } -func (_c *DatabaseRedisRepo_KeySet_Call) Run(run func(req *request.DatabaseRedisKeySet)) *DatabaseRedisRepo_KeySet_Call { +func (_c *DatabaseRedisRepo_KeySet_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisKeySet)) *DatabaseRedisRepo_KeySet_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisKeySet)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisKeySet)) }) return _c } @@ -380,22 +389,22 @@ func (_c *DatabaseRedisRepo_KeySet_Call) Return(_a0 error) *DatabaseRedisRepo_Ke return _c } -func (_c *DatabaseRedisRepo_KeySet_Call) RunAndReturn(run func(*request.DatabaseRedisKeySet) error) *DatabaseRedisRepo_KeySet_Call { +func (_c *DatabaseRedisRepo_KeySet_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisKeySet) error) *DatabaseRedisRepo_KeySet_Call { _c.Call.Return(run) return _c } -// KeyTTL provides a mock function with given fields: req -func (_m *DatabaseRedisRepo) KeyTTL(req *request.DatabaseRedisKeyTTL) error { - ret := _m.Called(req) +// KeyTTL provides a mock function with given fields: ctx, req +func (_m *DatabaseRedisRepo) KeyTTL(ctx context.Context, req *request.DatabaseRedisKeyTTL) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for KeyTTL") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseRedisKeyTTL) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseRedisKeyTTL) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -409,14 +418,15 @@ type DatabaseRedisRepo_KeyTTL_Call struct { } // KeyTTL is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseRedisKeyTTL -func (_e *DatabaseRedisRepo_Expecter) KeyTTL(req interface{}) *DatabaseRedisRepo_KeyTTL_Call { - return &DatabaseRedisRepo_KeyTTL_Call{Call: _e.mock.On("KeyTTL", req)} +func (_e *DatabaseRedisRepo_Expecter) KeyTTL(ctx interface{}, req interface{}) *DatabaseRedisRepo_KeyTTL_Call { + return &DatabaseRedisRepo_KeyTTL_Call{Call: _e.mock.On("KeyTTL", ctx, req)} } -func (_c *DatabaseRedisRepo_KeyTTL_Call) Run(run func(req *request.DatabaseRedisKeyTTL)) *DatabaseRedisRepo_KeyTTL_Call { +func (_c *DatabaseRedisRepo_KeyTTL_Call) Run(run func(ctx context.Context, req *request.DatabaseRedisKeyTTL)) *DatabaseRedisRepo_KeyTTL_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseRedisKeyTTL)) + run(args[0].(context.Context), args[1].(*request.DatabaseRedisKeyTTL)) }) return _c } @@ -426,7 +436,7 @@ func (_c *DatabaseRedisRepo_KeyTTL_Call) Return(_a0 error) *DatabaseRedisRepo_Ke return _c } -func (_c *DatabaseRedisRepo_KeyTTL_Call) RunAndReturn(run func(*request.DatabaseRedisKeyTTL) error) *DatabaseRedisRepo_KeyTTL_Call { +func (_c *DatabaseRedisRepo_KeyTTL_Call) RunAndReturn(run func(context.Context, *request.DatabaseRedisKeyTTL) error) *DatabaseRedisRepo_KeyTTL_Call { _c.Call.Return(run) return _c } diff --git a/mocks/biz/DatabaseRepo.go b/mocks/biz/DatabaseRepo.go index 96e44d6c..b7c81365 100644 --- a/mocks/biz/DatabaseRepo.go +++ b/mocks/biz/DatabaseRepo.go @@ -3,7 +3,10 @@ package biz import ( + context "context" + biz "github.com/acepanel/panel/v3/internal/biz" + db "github.com/acepanel/panel/v3/pkg/db" mock "github.com/stretchr/testify/mock" @@ -22,9 +25,9 @@ func (_m *DatabaseRepo) EXPECT() *DatabaseRepo_Expecter { return &DatabaseRepo_Expecter{mock: &_m.Mock} } -// DatabasesOf provides a mock function with given fields: server -func (_m *DatabaseRepo) DatabasesOf(server *biz.DatabaseServer) ([]*biz.Database, error) { - ret := _m.Called(server) +// DatabasesOf provides a mock function with given fields: ctx, server +func (_m *DatabaseRepo) DatabasesOf(ctx context.Context, server *biz.DatabaseServer) ([]*biz.Database, error) { + ret := _m.Called(ctx, server) if len(ret) == 0 { panic("no return value specified for DatabasesOf") @@ -32,19 +35,19 @@ func (_m *DatabaseRepo) DatabasesOf(server *biz.DatabaseServer) ([]*biz.Database var r0 []*biz.Database var r1 error - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) ([]*biz.Database, error)); ok { - return rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) ([]*biz.Database, error)); ok { + return rf(ctx, server) } - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) []*biz.Database); ok { - r0 = rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) []*biz.Database); ok { + r0 = rf(ctx, server) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*biz.Database) } } - if rf, ok := ret.Get(1).(func(*biz.DatabaseServer) error); ok { - r1 = rf(server) + if rf, ok := ret.Get(1).(func(context.Context, *biz.DatabaseServer) error); ok { + r1 = rf(ctx, server) } else { r1 = ret.Error(1) } @@ -58,14 +61,15 @@ type DatabaseRepo_DatabasesOf_Call struct { } // DatabasesOf is a helper method to define mock.On call +// - ctx context.Context // - server *biz.DatabaseServer -func (_e *DatabaseRepo_Expecter) DatabasesOf(server interface{}) *DatabaseRepo_DatabasesOf_Call { - return &DatabaseRepo_DatabasesOf_Call{Call: _e.mock.On("DatabasesOf", server)} +func (_e *DatabaseRepo_Expecter) DatabasesOf(ctx interface{}, server interface{}) *DatabaseRepo_DatabasesOf_Call { + return &DatabaseRepo_DatabasesOf_Call{Call: _e.mock.On("DatabasesOf", ctx, server)} } -func (_c *DatabaseRepo_DatabasesOf_Call) Run(run func(server *biz.DatabaseServer)) *DatabaseRepo_DatabasesOf_Call { +func (_c *DatabaseRepo_DatabasesOf_Call) Run(run func(ctx context.Context, server *biz.DatabaseServer)) *DatabaseRepo_DatabasesOf_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*biz.DatabaseServer)) + run(args[0].(context.Context), args[1].(*biz.DatabaseServer)) }) return _c } @@ -75,7 +79,7 @@ func (_c *DatabaseRepo_DatabasesOf_Call) Return(_a0 []*biz.Database, _a1 error) return _c } -func (_c *DatabaseRepo_DatabasesOf_Call) RunAndReturn(run func(*biz.DatabaseServer) ([]*biz.Database, error)) *DatabaseRepo_DatabasesOf_Call { +func (_c *DatabaseRepo_DatabasesOf_Call) RunAndReturn(run func(context.Context, *biz.DatabaseServer) ([]*biz.Database, error)) *DatabaseRepo_DatabasesOf_Call { _c.Call.Return(run) return _c } @@ -138,9 +142,9 @@ func (_c *DatabaseRepo_ListServers_Call) RunAndReturn(run func(string) ([]*biz.D return _c } -// Mongo provides a mock function with given fields: server -func (_m *DatabaseRepo) Mongo(server *biz.DatabaseServer) (*db.MongoDB, error) { - ret := _m.Called(server) +// Mongo provides a mock function with given fields: ctx, server +func (_m *DatabaseRepo) Mongo(ctx context.Context, server *biz.DatabaseServer) (*db.MongoDB, error) { + ret := _m.Called(ctx, server) if len(ret) == 0 { panic("no return value specified for Mongo") @@ -148,19 +152,19 @@ func (_m *DatabaseRepo) Mongo(server *biz.DatabaseServer) (*db.MongoDB, error) { var r0 *db.MongoDB var r1 error - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) (*db.MongoDB, error)); ok { - return rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) (*db.MongoDB, error)); ok { + return rf(ctx, server) } - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) *db.MongoDB); ok { - r0 = rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) *db.MongoDB); ok { + r0 = rf(ctx, server) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*db.MongoDB) } } - if rf, ok := ret.Get(1).(func(*biz.DatabaseServer) error); ok { - r1 = rf(server) + if rf, ok := ret.Get(1).(func(context.Context, *biz.DatabaseServer) error); ok { + r1 = rf(ctx, server) } else { r1 = ret.Error(1) } @@ -174,14 +178,15 @@ type DatabaseRepo_Mongo_Call struct { } // Mongo is a helper method to define mock.On call +// - ctx context.Context // - server *biz.DatabaseServer -func (_e *DatabaseRepo_Expecter) Mongo(server interface{}) *DatabaseRepo_Mongo_Call { - return &DatabaseRepo_Mongo_Call{Call: _e.mock.On("Mongo", server)} +func (_e *DatabaseRepo_Expecter) Mongo(ctx interface{}, server interface{}) *DatabaseRepo_Mongo_Call { + return &DatabaseRepo_Mongo_Call{Call: _e.mock.On("Mongo", ctx, server)} } -func (_c *DatabaseRepo_Mongo_Call) Run(run func(server *biz.DatabaseServer)) *DatabaseRepo_Mongo_Call { +func (_c *DatabaseRepo_Mongo_Call) Run(run func(ctx context.Context, server *biz.DatabaseServer)) *DatabaseRepo_Mongo_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*biz.DatabaseServer)) + run(args[0].(context.Context), args[1].(*biz.DatabaseServer)) }) return _c } @@ -191,14 +196,14 @@ func (_c *DatabaseRepo_Mongo_Call) Return(_a0 *db.MongoDB, _a1 error) *DatabaseR return _c } -func (_c *DatabaseRepo_Mongo_Call) RunAndReturn(run func(*biz.DatabaseServer) (*db.MongoDB, error)) *DatabaseRepo_Mongo_Call { +func (_c *DatabaseRepo_Mongo_Call) RunAndReturn(run func(context.Context, *biz.DatabaseServer) (*db.MongoDB, error)) *DatabaseRepo_Mongo_Call { _c.Call.Return(run) return _c } -// Operator provides a mock function with given fields: server -func (_m *DatabaseRepo) Operator(server *biz.DatabaseServer) (db.Operator, error) { - ret := _m.Called(server) +// Operator provides a mock function with given fields: ctx, server +func (_m *DatabaseRepo) Operator(ctx context.Context, server *biz.DatabaseServer) (db.Operator, error) { + ret := _m.Called(ctx, server) if len(ret) == 0 { panic("no return value specified for Operator") @@ -206,19 +211,19 @@ func (_m *DatabaseRepo) Operator(server *biz.DatabaseServer) (db.Operator, error var r0 db.Operator var r1 error - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) (db.Operator, error)); ok { - return rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) (db.Operator, error)); ok { + return rf(ctx, server) } - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) db.Operator); ok { - r0 = rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) db.Operator); ok { + r0 = rf(ctx, server) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(db.Operator) } } - if rf, ok := ret.Get(1).(func(*biz.DatabaseServer) error); ok { - r1 = rf(server) + if rf, ok := ret.Get(1).(func(context.Context, *biz.DatabaseServer) error); ok { + r1 = rf(ctx, server) } else { r1 = ret.Error(1) } @@ -232,14 +237,15 @@ type DatabaseRepo_Operator_Call struct { } // Operator is a helper method to define mock.On call +// - ctx context.Context // - server *biz.DatabaseServer -func (_e *DatabaseRepo_Expecter) Operator(server interface{}) *DatabaseRepo_Operator_Call { - return &DatabaseRepo_Operator_Call{Call: _e.mock.On("Operator", server)} +func (_e *DatabaseRepo_Expecter) Operator(ctx interface{}, server interface{}) *DatabaseRepo_Operator_Call { + return &DatabaseRepo_Operator_Call{Call: _e.mock.On("Operator", ctx, server)} } -func (_c *DatabaseRepo_Operator_Call) Run(run func(server *biz.DatabaseServer)) *DatabaseRepo_Operator_Call { +func (_c *DatabaseRepo_Operator_Call) Run(run func(ctx context.Context, server *biz.DatabaseServer)) *DatabaseRepo_Operator_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*biz.DatabaseServer)) + run(args[0].(context.Context), args[1].(*biz.DatabaseServer)) }) return _c } @@ -249,7 +255,7 @@ func (_c *DatabaseRepo_Operator_Call) Return(_a0 db.Operator, _a1 error) *Databa return _c } -func (_c *DatabaseRepo_Operator_Call) RunAndReturn(run func(*biz.DatabaseServer) (db.Operator, error)) *DatabaseRepo_Operator_Call { +func (_c *DatabaseRepo_Operator_Call) RunAndReturn(run func(context.Context, *biz.DatabaseServer) (db.Operator, error)) *DatabaseRepo_Operator_Call { _c.Call.Return(run) return _c } diff --git a/mocks/biz/DatabaseServerRepo.go b/mocks/biz/DatabaseServerRepo.go index 2806c810..7efae8b3 100644 --- a/mocks/biz/DatabaseServerRepo.go +++ b/mocks/biz/DatabaseServerRepo.go @@ -3,7 +3,10 @@ package biz import ( + context "context" + biz "github.com/acepanel/panel/v3/internal/biz" + db "github.com/acepanel/panel/v3/pkg/db" mock "github.com/stretchr/testify/mock" @@ -24,17 +27,17 @@ func (_m *DatabaseServerRepo) EXPECT() *DatabaseServerRepo_Expecter { return &DatabaseServerRepo_Expecter{mock: &_m.Mock} } -// CheckServer provides a mock function with given fields: server -func (_m *DatabaseServerRepo) CheckServer(server *biz.DatabaseServer) bool { - ret := _m.Called(server) +// CheckServer provides a mock function with given fields: ctx, server +func (_m *DatabaseServerRepo) CheckServer(ctx context.Context, server *biz.DatabaseServer) bool { + ret := _m.Called(ctx, server) if len(ret) == 0 { panic("no return value specified for CheckServer") } var r0 bool - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) bool); ok { - r0 = rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) bool); ok { + r0 = rf(ctx, server) } else { r0 = ret.Get(0).(bool) } @@ -48,14 +51,15 @@ type DatabaseServerRepo_CheckServer_Call struct { } // CheckServer is a helper method to define mock.On call +// - ctx context.Context // - server *biz.DatabaseServer -func (_e *DatabaseServerRepo_Expecter) CheckServer(server interface{}) *DatabaseServerRepo_CheckServer_Call { - return &DatabaseServerRepo_CheckServer_Call{Call: _e.mock.On("CheckServer", server)} +func (_e *DatabaseServerRepo_Expecter) CheckServer(ctx interface{}, server interface{}) *DatabaseServerRepo_CheckServer_Call { + return &DatabaseServerRepo_CheckServer_Call{Call: _e.mock.On("CheckServer", ctx, server)} } -func (_c *DatabaseServerRepo_CheckServer_Call) Run(run func(server *biz.DatabaseServer)) *DatabaseServerRepo_CheckServer_Call { +func (_c *DatabaseServerRepo_CheckServer_Call) Run(run func(ctx context.Context, server *biz.DatabaseServer)) *DatabaseServerRepo_CheckServer_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*biz.DatabaseServer)) + run(args[0].(context.Context), args[1].(*biz.DatabaseServer)) }) return _c } @@ -65,7 +69,7 @@ func (_c *DatabaseServerRepo_CheckServer_Call) Return(_a0 bool) *DatabaseServerR return _c } -func (_c *DatabaseServerRepo_CheckServer_Call) RunAndReturn(run func(*biz.DatabaseServer) bool) *DatabaseServerRepo_CheckServer_Call { +func (_c *DatabaseServerRepo_CheckServer_Call) RunAndReturn(run func(context.Context, *biz.DatabaseServer) bool) *DatabaseServerRepo_CheckServer_Call { _c.Call.Return(run) return _c } @@ -309,9 +313,9 @@ func (_c *DatabaseServerRepo_Delete_Call) RunAndReturn(run func(uint) error) *Da return _c } -// Get provides a mock function with given fields: id -func (_m *DatabaseServerRepo) Get(id uint) (*biz.DatabaseServer, error) { - ret := _m.Called(id) +// Get provides a mock function with given fields: ctx, id +func (_m *DatabaseServerRepo) Get(ctx context.Context, id uint) (*biz.DatabaseServer, error) { + ret := _m.Called(ctx, id) if len(ret) == 0 { panic("no return value specified for Get") @@ -319,19 +323,19 @@ func (_m *DatabaseServerRepo) Get(id uint) (*biz.DatabaseServer, error) { var r0 *biz.DatabaseServer var r1 error - if rf, ok := ret.Get(0).(func(uint) (*biz.DatabaseServer, error)); ok { - return rf(id) + if rf, ok := ret.Get(0).(func(context.Context, uint) (*biz.DatabaseServer, error)); ok { + return rf(ctx, id) } - if rf, ok := ret.Get(0).(func(uint) *biz.DatabaseServer); ok { - r0 = rf(id) + if rf, ok := ret.Get(0).(func(context.Context, uint) *biz.DatabaseServer); ok { + r0 = rf(ctx, id) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*biz.DatabaseServer) } } - if rf, ok := ret.Get(1).(func(uint) error); ok { - r1 = rf(id) + if rf, ok := ret.Get(1).(func(context.Context, uint) error); ok { + r1 = rf(ctx, id) } else { r1 = ret.Error(1) } @@ -345,14 +349,15 @@ type DatabaseServerRepo_Get_Call struct { } // Get is a helper method to define mock.On call +// - ctx context.Context // - id uint -func (_e *DatabaseServerRepo_Expecter) Get(id interface{}) *DatabaseServerRepo_Get_Call { - return &DatabaseServerRepo_Get_Call{Call: _e.mock.On("Get", id)} +func (_e *DatabaseServerRepo_Expecter) Get(ctx interface{}, id interface{}) *DatabaseServerRepo_Get_Call { + return &DatabaseServerRepo_Get_Call{Call: _e.mock.On("Get", ctx, id)} } -func (_c *DatabaseServerRepo_Get_Call) Run(run func(id uint)) *DatabaseServerRepo_Get_Call { +func (_c *DatabaseServerRepo_Get_Call) Run(run func(ctx context.Context, id uint)) *DatabaseServerRepo_Get_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(uint)) + run(args[0].(context.Context), args[1].(uint)) }) return _c } @@ -362,14 +367,14 @@ func (_c *DatabaseServerRepo_Get_Call) Return(_a0 *biz.DatabaseServer, _a1 error return _c } -func (_c *DatabaseServerRepo_Get_Call) RunAndReturn(run func(uint) (*biz.DatabaseServer, error)) *DatabaseServerRepo_Get_Call { +func (_c *DatabaseServerRepo_Get_Call) RunAndReturn(run func(context.Context, uint) (*biz.DatabaseServer, error)) *DatabaseServerRepo_Get_Call { _c.Call.Return(run) return _c } -// GetByName provides a mock function with given fields: name -func (_m *DatabaseServerRepo) GetByName(name string) (*biz.DatabaseServer, error) { - ret := _m.Called(name) +// GetByName provides a mock function with given fields: ctx, name +func (_m *DatabaseServerRepo) GetByName(ctx context.Context, name string) (*biz.DatabaseServer, error) { + ret := _m.Called(ctx, name) if len(ret) == 0 { panic("no return value specified for GetByName") @@ -377,19 +382,19 @@ func (_m *DatabaseServerRepo) GetByName(name string) (*biz.DatabaseServer, error var r0 *biz.DatabaseServer var r1 error - if rf, ok := ret.Get(0).(func(string) (*biz.DatabaseServer, error)); ok { - return rf(name) + if rf, ok := ret.Get(0).(func(context.Context, string) (*biz.DatabaseServer, error)); ok { + return rf(ctx, name) } - if rf, ok := ret.Get(0).(func(string) *biz.DatabaseServer); ok { - r0 = rf(name) + if rf, ok := ret.Get(0).(func(context.Context, string) *biz.DatabaseServer); ok { + r0 = rf(ctx, name) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*biz.DatabaseServer) } } - if rf, ok := ret.Get(1).(func(string) error); ok { - r1 = rf(name) + if rf, ok := ret.Get(1).(func(context.Context, string) error); ok { + r1 = rf(ctx, name) } else { r1 = ret.Error(1) } @@ -403,14 +408,15 @@ type DatabaseServerRepo_GetByName_Call struct { } // GetByName is a helper method to define mock.On call +// - ctx context.Context // - name string -func (_e *DatabaseServerRepo_Expecter) GetByName(name interface{}) *DatabaseServerRepo_GetByName_Call { - return &DatabaseServerRepo_GetByName_Call{Call: _e.mock.On("GetByName", name)} +func (_e *DatabaseServerRepo_Expecter) GetByName(ctx interface{}, name interface{}) *DatabaseServerRepo_GetByName_Call { + return &DatabaseServerRepo_GetByName_Call{Call: _e.mock.On("GetByName", ctx, name)} } -func (_c *DatabaseServerRepo_GetByName_Call) Run(run func(name string)) *DatabaseServerRepo_GetByName_Call { +func (_c *DatabaseServerRepo_GetByName_Call) Run(run func(ctx context.Context, name string)) *DatabaseServerRepo_GetByName_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(string)) + run(args[0].(context.Context), args[1].(string)) }) return _c } @@ -420,14 +426,14 @@ func (_c *DatabaseServerRepo_GetByName_Call) Return(_a0 *biz.DatabaseServer, _a1 return _c } -func (_c *DatabaseServerRepo_GetByName_Call) RunAndReturn(run func(string) (*biz.DatabaseServer, error)) *DatabaseServerRepo_GetByName_Call { +func (_c *DatabaseServerRepo_GetByName_Call) RunAndReturn(run func(context.Context, string) (*biz.DatabaseServer, error)) *DatabaseServerRepo_GetByName_Call { _c.Call.Return(run) return _c } -// List provides a mock function with given fields: page, limit, typ -func (_m *DatabaseServerRepo) List(page uint, limit uint, typ string) ([]*biz.DatabaseServer, int64, error) { - ret := _m.Called(page, limit, typ) +// List provides a mock function with given fields: ctx, page, limit, typ +func (_m *DatabaseServerRepo) List(ctx context.Context, page uint, limit uint, typ string) ([]*biz.DatabaseServer, int64, error) { + ret := _m.Called(ctx, page, limit, typ) if len(ret) == 0 { panic("no return value specified for List") @@ -436,25 +442,25 @@ func (_m *DatabaseServerRepo) List(page uint, limit uint, typ string) ([]*biz.Da var r0 []*biz.DatabaseServer var r1 int64 var r2 error - if rf, ok := ret.Get(0).(func(uint, uint, string) ([]*biz.DatabaseServer, int64, error)); ok { - return rf(page, limit, typ) + if rf, ok := ret.Get(0).(func(context.Context, uint, uint, string) ([]*biz.DatabaseServer, int64, error)); ok { + return rf(ctx, page, limit, typ) } - if rf, ok := ret.Get(0).(func(uint, uint, string) []*biz.DatabaseServer); ok { - r0 = rf(page, limit, typ) + if rf, ok := ret.Get(0).(func(context.Context, uint, uint, string) []*biz.DatabaseServer); ok { + r0 = rf(ctx, page, limit, typ) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*biz.DatabaseServer) } } - if rf, ok := ret.Get(1).(func(uint, uint, string) int64); ok { - r1 = rf(page, limit, typ) + if rf, ok := ret.Get(1).(func(context.Context, uint, uint, string) int64); ok { + r1 = rf(ctx, page, limit, typ) } else { r1 = ret.Get(1).(int64) } - if rf, ok := ret.Get(2).(func(uint, uint, string) error); ok { - r2 = rf(page, limit, typ) + if rf, ok := ret.Get(2).(func(context.Context, uint, uint, string) error); ok { + r2 = rf(ctx, page, limit, typ) } else { r2 = ret.Error(2) } @@ -468,16 +474,17 @@ type DatabaseServerRepo_List_Call struct { } // List is a helper method to define mock.On call +// - ctx context.Context // - page uint // - limit uint // - typ string -func (_e *DatabaseServerRepo_Expecter) List(page interface{}, limit interface{}, typ interface{}) *DatabaseServerRepo_List_Call { - return &DatabaseServerRepo_List_Call{Call: _e.mock.On("List", page, limit, typ)} +func (_e *DatabaseServerRepo_Expecter) List(ctx interface{}, page interface{}, limit interface{}, typ interface{}) *DatabaseServerRepo_List_Call { + return &DatabaseServerRepo_List_Call{Call: _e.mock.On("List", ctx, page, limit, typ)} } -func (_c *DatabaseServerRepo_List_Call) Run(run func(page uint, limit uint, typ string)) *DatabaseServerRepo_List_Call { +func (_c *DatabaseServerRepo_List_Call) Run(run func(ctx context.Context, page uint, limit uint, typ string)) *DatabaseServerRepo_List_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(uint), args[1].(uint), args[2].(string)) + run(args[0].(context.Context), args[1].(uint), args[2].(uint), args[3].(string)) }) return _c } @@ -487,7 +494,7 @@ func (_c *DatabaseServerRepo_List_Call) Return(_a0 []*biz.DatabaseServer, _a1 in return _c } -func (_c *DatabaseServerRepo_List_Call) RunAndReturn(run func(uint, uint, string) ([]*biz.DatabaseServer, int64, error)) *DatabaseServerRepo_List_Call { +func (_c *DatabaseServerRepo_List_Call) RunAndReturn(run func(context.Context, uint, uint, string) ([]*biz.DatabaseServer, int64, error)) *DatabaseServerRepo_List_Call { _c.Call.Return(run) return _c } @@ -550,9 +557,9 @@ func (_c *DatabaseServerRepo_ListUsers_Call) RunAndReturn(run func(uint) ([]*biz return _c } -// Operator provides a mock function with given fields: server -func (_m *DatabaseServerRepo) Operator(server *biz.DatabaseServer) (db.Operator, error) { - ret := _m.Called(server) +// Operator provides a mock function with given fields: ctx, server +func (_m *DatabaseServerRepo) Operator(ctx context.Context, server *biz.DatabaseServer) (db.Operator, error) { + ret := _m.Called(ctx, server) if len(ret) == 0 { panic("no return value specified for Operator") @@ -560,19 +567,19 @@ func (_m *DatabaseServerRepo) Operator(server *biz.DatabaseServer) (db.Operator, var r0 db.Operator var r1 error - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) (db.Operator, error)); ok { - return rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) (db.Operator, error)); ok { + return rf(ctx, server) } - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) db.Operator); ok { - r0 = rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) db.Operator); ok { + r0 = rf(ctx, server) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(db.Operator) } } - if rf, ok := ret.Get(1).(func(*biz.DatabaseServer) error); ok { - r1 = rf(server) + if rf, ok := ret.Get(1).(func(context.Context, *biz.DatabaseServer) error); ok { + r1 = rf(ctx, server) } else { r1 = ret.Error(1) } @@ -586,14 +593,15 @@ type DatabaseServerRepo_Operator_Call struct { } // Operator is a helper method to define mock.On call +// - ctx context.Context // - server *biz.DatabaseServer -func (_e *DatabaseServerRepo_Expecter) Operator(server interface{}) *DatabaseServerRepo_Operator_Call { - return &DatabaseServerRepo_Operator_Call{Call: _e.mock.On("Operator", server)} +func (_e *DatabaseServerRepo_Expecter) Operator(ctx interface{}, server interface{}) *DatabaseServerRepo_Operator_Call { + return &DatabaseServerRepo_Operator_Call{Call: _e.mock.On("Operator", ctx, server)} } -func (_c *DatabaseServerRepo_Operator_Call) Run(run func(server *biz.DatabaseServer)) *DatabaseServerRepo_Operator_Call { +func (_c *DatabaseServerRepo_Operator_Call) Run(run func(ctx context.Context, server *biz.DatabaseServer)) *DatabaseServerRepo_Operator_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*biz.DatabaseServer)) + run(args[0].(context.Context), args[1].(*biz.DatabaseServer)) }) return _c } @@ -603,7 +611,7 @@ func (_c *DatabaseServerRepo_Operator_Call) Return(_a0 db.Operator, _a1 error) * return _c } -func (_c *DatabaseServerRepo_Operator_Call) RunAndReturn(run func(*biz.DatabaseServer) (db.Operator, error)) *DatabaseServerRepo_Operator_Call { +func (_c *DatabaseServerRepo_Operator_Call) RunAndReturn(run func(context.Context, *biz.DatabaseServer) (db.Operator, error)) *DatabaseServerRepo_Operator_Call { _c.Call.Return(run) return _c } diff --git a/mocks/biz/DatabaseUserRepo.go b/mocks/biz/DatabaseUserRepo.go index a14f96ec..b6c2a6b1 100644 --- a/mocks/biz/DatabaseUserRepo.go +++ b/mocks/biz/DatabaseUserRepo.go @@ -3,7 +3,10 @@ package biz import ( + context "context" + biz "github.com/acepanel/panel/v3/internal/biz" + db "github.com/acepanel/panel/v3/pkg/db" mock "github.com/stretchr/testify/mock" @@ -172,9 +175,9 @@ func (_c *DatabaseUserRepo_DeleteByServerNames_Call) RunAndReturn(run func(uint, return _c } -// Get provides a mock function with given fields: id -func (_m *DatabaseUserRepo) Get(id uint) (*biz.DatabaseUser, error) { - ret := _m.Called(id) +// Get provides a mock function with given fields: ctx, id +func (_m *DatabaseUserRepo) Get(ctx context.Context, id uint) (*biz.DatabaseUser, error) { + ret := _m.Called(ctx, id) if len(ret) == 0 { panic("no return value specified for Get") @@ -182,19 +185,19 @@ func (_m *DatabaseUserRepo) Get(id uint) (*biz.DatabaseUser, error) { var r0 *biz.DatabaseUser var r1 error - if rf, ok := ret.Get(0).(func(uint) (*biz.DatabaseUser, error)); ok { - return rf(id) + if rf, ok := ret.Get(0).(func(context.Context, uint) (*biz.DatabaseUser, error)); ok { + return rf(ctx, id) } - if rf, ok := ret.Get(0).(func(uint) *biz.DatabaseUser); ok { - r0 = rf(id) + if rf, ok := ret.Get(0).(func(context.Context, uint) *biz.DatabaseUser); ok { + r0 = rf(ctx, id) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*biz.DatabaseUser) } } - if rf, ok := ret.Get(1).(func(uint) error); ok { - r1 = rf(id) + if rf, ok := ret.Get(1).(func(context.Context, uint) error); ok { + r1 = rf(ctx, id) } else { r1 = ret.Error(1) } @@ -208,14 +211,15 @@ type DatabaseUserRepo_Get_Call struct { } // Get is a helper method to define mock.On call +// - ctx context.Context // - id uint -func (_e *DatabaseUserRepo_Expecter) Get(id interface{}) *DatabaseUserRepo_Get_Call { - return &DatabaseUserRepo_Get_Call{Call: _e.mock.On("Get", id)} +func (_e *DatabaseUserRepo_Expecter) Get(ctx interface{}, id interface{}) *DatabaseUserRepo_Get_Call { + return &DatabaseUserRepo_Get_Call{Call: _e.mock.On("Get", ctx, id)} } -func (_c *DatabaseUserRepo_Get_Call) Run(run func(id uint)) *DatabaseUserRepo_Get_Call { +func (_c *DatabaseUserRepo_Get_Call) Run(run func(ctx context.Context, id uint)) *DatabaseUserRepo_Get_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(uint)) + run(args[0].(context.Context), args[1].(uint)) }) return _c } @@ -225,14 +229,14 @@ func (_c *DatabaseUserRepo_Get_Call) Return(_a0 *biz.DatabaseUser, _a1 error) *D return _c } -func (_c *DatabaseUserRepo_Get_Call) RunAndReturn(run func(uint) (*biz.DatabaseUser, error)) *DatabaseUserRepo_Get_Call { +func (_c *DatabaseUserRepo_Get_Call) RunAndReturn(run func(context.Context, uint) (*biz.DatabaseUser, error)) *DatabaseUserRepo_Get_Call { _c.Call.Return(run) return _c } -// List provides a mock function with given fields: page, limit, typ -func (_m *DatabaseUserRepo) List(page uint, limit uint, typ string) ([]*biz.DatabaseUser, int64, error) { - ret := _m.Called(page, limit, typ) +// List provides a mock function with given fields: ctx, page, limit, typ +func (_m *DatabaseUserRepo) List(ctx context.Context, page uint, limit uint, typ string) ([]*biz.DatabaseUser, int64, error) { + ret := _m.Called(ctx, page, limit, typ) if len(ret) == 0 { panic("no return value specified for List") @@ -241,25 +245,25 @@ func (_m *DatabaseUserRepo) List(page uint, limit uint, typ string) ([]*biz.Data var r0 []*biz.DatabaseUser var r1 int64 var r2 error - if rf, ok := ret.Get(0).(func(uint, uint, string) ([]*biz.DatabaseUser, int64, error)); ok { - return rf(page, limit, typ) + if rf, ok := ret.Get(0).(func(context.Context, uint, uint, string) ([]*biz.DatabaseUser, int64, error)); ok { + return rf(ctx, page, limit, typ) } - if rf, ok := ret.Get(0).(func(uint, uint, string) []*biz.DatabaseUser); ok { - r0 = rf(page, limit, typ) + if rf, ok := ret.Get(0).(func(context.Context, uint, uint, string) []*biz.DatabaseUser); ok { + r0 = rf(ctx, page, limit, typ) } else { if ret.Get(0) != nil { r0 = ret.Get(0).([]*biz.DatabaseUser) } } - if rf, ok := ret.Get(1).(func(uint, uint, string) int64); ok { - r1 = rf(page, limit, typ) + if rf, ok := ret.Get(1).(func(context.Context, uint, uint, string) int64); ok { + r1 = rf(ctx, page, limit, typ) } else { r1 = ret.Get(1).(int64) } - if rf, ok := ret.Get(2).(func(uint, uint, string) error); ok { - r2 = rf(page, limit, typ) + if rf, ok := ret.Get(2).(func(context.Context, uint, uint, string) error); ok { + r2 = rf(ctx, page, limit, typ) } else { r2 = ret.Error(2) } @@ -273,16 +277,17 @@ type DatabaseUserRepo_List_Call struct { } // List is a helper method to define mock.On call +// - ctx context.Context // - page uint // - limit uint // - typ string -func (_e *DatabaseUserRepo_Expecter) List(page interface{}, limit interface{}, typ interface{}) *DatabaseUserRepo_List_Call { - return &DatabaseUserRepo_List_Call{Call: _e.mock.On("List", page, limit, typ)} +func (_e *DatabaseUserRepo_Expecter) List(ctx interface{}, page interface{}, limit interface{}, typ interface{}) *DatabaseUserRepo_List_Call { + return &DatabaseUserRepo_List_Call{Call: _e.mock.On("List", ctx, page, limit, typ)} } -func (_c *DatabaseUserRepo_List_Call) Run(run func(page uint, limit uint, typ string)) *DatabaseUserRepo_List_Call { +func (_c *DatabaseUserRepo_List_Call) Run(run func(ctx context.Context, page uint, limit uint, typ string)) *DatabaseUserRepo_List_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(uint), args[1].(uint), args[2].(string)) + run(args[0].(context.Context), args[1].(uint), args[2].(uint), args[3].(string)) }) return _c } @@ -292,7 +297,7 @@ func (_c *DatabaseUserRepo_List_Call) Return(_a0 []*biz.DatabaseUser, _a1 int64, return _c } -func (_c *DatabaseUserRepo_List_Call) RunAndReturn(run func(uint, uint, string) ([]*biz.DatabaseUser, int64, error)) *DatabaseUserRepo_List_Call { +func (_c *DatabaseUserRepo_List_Call) RunAndReturn(run func(context.Context, uint, uint, string) ([]*biz.DatabaseUser, int64, error)) *DatabaseUserRepo_List_Call { _c.Call.Return(run) return _c } @@ -356,9 +361,9 @@ func (_c *DatabaseUserRepo_ListByNames_Call) RunAndReturn(run func(uint, []strin return _c } -// Operator provides a mock function with given fields: server -func (_m *DatabaseUserRepo) Operator(server *biz.DatabaseServer) (db.Operator, error) { - ret := _m.Called(server) +// Operator provides a mock function with given fields: ctx, server +func (_m *DatabaseUserRepo) Operator(ctx context.Context, server *biz.DatabaseServer) (db.Operator, error) { + ret := _m.Called(ctx, server) if len(ret) == 0 { panic("no return value specified for Operator") @@ -366,19 +371,19 @@ func (_m *DatabaseUserRepo) Operator(server *biz.DatabaseServer) (db.Operator, e var r0 db.Operator var r1 error - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) (db.Operator, error)); ok { - return rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) (db.Operator, error)); ok { + return rf(ctx, server) } - if rf, ok := ret.Get(0).(func(*biz.DatabaseServer) db.Operator); ok { - r0 = rf(server) + if rf, ok := ret.Get(0).(func(context.Context, *biz.DatabaseServer) db.Operator); ok { + r0 = rf(ctx, server) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(db.Operator) } } - if rf, ok := ret.Get(1).(func(*biz.DatabaseServer) error); ok { - r1 = rf(server) + if rf, ok := ret.Get(1).(func(context.Context, *biz.DatabaseServer) error); ok { + r1 = rf(ctx, server) } else { r1 = ret.Error(1) } @@ -392,14 +397,15 @@ type DatabaseUserRepo_Operator_Call struct { } // Operator is a helper method to define mock.On call +// - ctx context.Context // - server *biz.DatabaseServer -func (_e *DatabaseUserRepo_Expecter) Operator(server interface{}) *DatabaseUserRepo_Operator_Call { - return &DatabaseUserRepo_Operator_Call{Call: _e.mock.On("Operator", server)} +func (_e *DatabaseUserRepo_Expecter) Operator(ctx interface{}, server interface{}) *DatabaseUserRepo_Operator_Call { + return &DatabaseUserRepo_Operator_Call{Call: _e.mock.On("Operator", ctx, server)} } -func (_c *DatabaseUserRepo_Operator_Call) Run(run func(server *biz.DatabaseServer)) *DatabaseUserRepo_Operator_Call { +func (_c *DatabaseUserRepo_Operator_Call) Run(run func(ctx context.Context, server *biz.DatabaseServer)) *DatabaseUserRepo_Operator_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*biz.DatabaseServer)) + run(args[0].(context.Context), args[1].(*biz.DatabaseServer)) }) return _c } @@ -409,7 +415,7 @@ func (_c *DatabaseUserRepo_Operator_Call) Return(_a0 db.Operator, _a1 error) *Da return _c } -func (_c *DatabaseUserRepo_Operator_Call) RunAndReturn(run func(*biz.DatabaseServer) (db.Operator, error)) *DatabaseUserRepo_Operator_Call { +func (_c *DatabaseUserRepo_Operator_Call) RunAndReturn(run func(context.Context, *biz.DatabaseServer) (db.Operator, error)) *DatabaseUserRepo_Operator_Call { _c.Call.Return(run) return _c } @@ -460,17 +466,17 @@ func (_c *DatabaseUserRepo_Save_Call) RunAndReturn(run func(*biz.DatabaseUser) e return _c } -// UpdateRemark provides a mock function with given fields: req -func (_m *DatabaseUserRepo) UpdateRemark(req *request.DatabaseUserUpdateRemark) error { - ret := _m.Called(req) +// UpdateRemark provides a mock function with given fields: ctx, req +func (_m *DatabaseUserRepo) UpdateRemark(ctx context.Context, req *request.DatabaseUserUpdateRemark) error { + ret := _m.Called(ctx, req) if len(ret) == 0 { panic("no return value specified for UpdateRemark") } var r0 error - if rf, ok := ret.Get(0).(func(*request.DatabaseUserUpdateRemark) error); ok { - r0 = rf(req) + if rf, ok := ret.Get(0).(func(context.Context, *request.DatabaseUserUpdateRemark) error); ok { + r0 = rf(ctx, req) } else { r0 = ret.Error(0) } @@ -484,14 +490,15 @@ type DatabaseUserRepo_UpdateRemark_Call struct { } // UpdateRemark is a helper method to define mock.On call +// - ctx context.Context // - req *request.DatabaseUserUpdateRemark -func (_e *DatabaseUserRepo_Expecter) UpdateRemark(req interface{}) *DatabaseUserRepo_UpdateRemark_Call { - return &DatabaseUserRepo_UpdateRemark_Call{Call: _e.mock.On("UpdateRemark", req)} +func (_e *DatabaseUserRepo_Expecter) UpdateRemark(ctx interface{}, req interface{}) *DatabaseUserRepo_UpdateRemark_Call { + return &DatabaseUserRepo_UpdateRemark_Call{Call: _e.mock.On("UpdateRemark", ctx, req)} } -func (_c *DatabaseUserRepo_UpdateRemark_Call) Run(run func(req *request.DatabaseUserUpdateRemark)) *DatabaseUserRepo_UpdateRemark_Call { +func (_c *DatabaseUserRepo_UpdateRemark_Call) Run(run func(ctx context.Context, req *request.DatabaseUserUpdateRemark)) *DatabaseUserRepo_UpdateRemark_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].(*request.DatabaseUserUpdateRemark)) + run(args[0].(context.Context), args[1].(*request.DatabaseUserUpdateRemark)) }) return _c } @@ -501,7 +508,7 @@ func (_c *DatabaseUserRepo_UpdateRemark_Call) Return(_a0 error) *DatabaseUserRep return _c } -func (_c *DatabaseUserRepo_UpdateRemark_Call) RunAndReturn(run func(*request.DatabaseUserUpdateRemark) error) *DatabaseUserRepo_UpdateRemark_Call { +func (_c *DatabaseUserRepo_UpdateRemark_Call) RunAndReturn(run func(context.Context, *request.DatabaseUserUpdateRemark) error) *DatabaseUserRepo_UpdateRemark_Call { _c.Call.Return(run) return _c } diff --git a/mocks/biz/NotifyChannelRepo.go b/mocks/biz/NotifyChannelRepo.go new file mode 100644 index 00000000..755c9ad4 --- /dev/null +++ b/mocks/biz/NotifyChannelRepo.go @@ -0,0 +1,412 @@ +// Code generated by mockery. DO NOT EDIT. + +package biz + +import ( + biz "github.com/acepanel/panel/v3/internal/biz" + mock "github.com/stretchr/testify/mock" +) + +// NotifyChannelRepo is an autogenerated mock type for the NotifyChannelRepo type +type NotifyChannelRepo struct { + mock.Mock +} + +type NotifyChannelRepo_Expecter struct { + mock *mock.Mock +} + +func (_m *NotifyChannelRepo) EXPECT() *NotifyChannelRepo_Expecter { + return &NotifyChannelRepo_Expecter{mock: &_m.Mock} +} + +// All provides a mock function with no fields +func (_m *NotifyChannelRepo) All() ([]*biz.NotifyChannel, error) { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for All") + } + + var r0 []*biz.NotifyChannel + var r1 error + if rf, ok := ret.Get(0).(func() ([]*biz.NotifyChannel, error)); ok { + return rf() + } + if rf, ok := ret.Get(0).(func() []*biz.NotifyChannel); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.NotifyChannel) + } + } + + if rf, ok := ret.Get(1).(func() error); ok { + r1 = rf() + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// NotifyChannelRepo_All_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'All' +type NotifyChannelRepo_All_Call struct { + *mock.Call +} + +// All is a helper method to define mock.On call +func (_e *NotifyChannelRepo_Expecter) All() *NotifyChannelRepo_All_Call { + return &NotifyChannelRepo_All_Call{Call: _e.mock.On("All")} +} + +func (_c *NotifyChannelRepo_All_Call) Run(run func()) *NotifyChannelRepo_All_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *NotifyChannelRepo_All_Call) Return(_a0 []*biz.NotifyChannel, _a1 error) *NotifyChannelRepo_All_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *NotifyChannelRepo_All_Call) RunAndReturn(run func() ([]*biz.NotifyChannel, error)) *NotifyChannelRepo_All_Call { + _c.Call.Return(run) + return _c +} + +// Create provides a mock function with given fields: channel +func (_m *NotifyChannelRepo) Create(channel *biz.NotifyChannel) error { + ret := _m.Called(channel) + + if len(ret) == 0 { + panic("no return value specified for Create") + } + + var r0 error + if rf, ok := ret.Get(0).(func(*biz.NotifyChannel) error); ok { + r0 = rf(channel) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// NotifyChannelRepo_Create_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Create' +type NotifyChannelRepo_Create_Call struct { + *mock.Call +} + +// Create is a helper method to define mock.On call +// - channel *biz.NotifyChannel +func (_e *NotifyChannelRepo_Expecter) Create(channel interface{}) *NotifyChannelRepo_Create_Call { + return &NotifyChannelRepo_Create_Call{Call: _e.mock.On("Create", channel)} +} + +func (_c *NotifyChannelRepo_Create_Call) Run(run func(channel *biz.NotifyChannel)) *NotifyChannelRepo_Create_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(*biz.NotifyChannel)) + }) + return _c +} + +func (_c *NotifyChannelRepo_Create_Call) Return(_a0 error) *NotifyChannelRepo_Create_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *NotifyChannelRepo_Create_Call) RunAndReturn(run func(*biz.NotifyChannel) error) *NotifyChannelRepo_Create_Call { + _c.Call.Return(run) + return _c +} + +// Delete provides a mock function with given fields: id +func (_m *NotifyChannelRepo) Delete(id uint) error { + ret := _m.Called(id) + + if len(ret) == 0 { + panic("no return value specified for Delete") + } + + var r0 error + if rf, ok := ret.Get(0).(func(uint) error); ok { + r0 = rf(id) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// NotifyChannelRepo_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' +type NotifyChannelRepo_Delete_Call struct { + *mock.Call +} + +// Delete is a helper method to define mock.On call +// - id uint +func (_e *NotifyChannelRepo_Expecter) Delete(id interface{}) *NotifyChannelRepo_Delete_Call { + return &NotifyChannelRepo_Delete_Call{Call: _e.mock.On("Delete", id)} +} + +func (_c *NotifyChannelRepo_Delete_Call) Run(run func(id uint)) *NotifyChannelRepo_Delete_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint)) + }) + return _c +} + +func (_c *NotifyChannelRepo_Delete_Call) Return(_a0 error) *NotifyChannelRepo_Delete_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *NotifyChannelRepo_Delete_Call) RunAndReturn(run func(uint) error) *NotifyChannelRepo_Delete_Call { + _c.Call.Return(run) + return _c +} + +// Get provides a mock function with given fields: id +func (_m *NotifyChannelRepo) Get(id uint) (*biz.NotifyChannel, error) { + ret := _m.Called(id) + + if len(ret) == 0 { + panic("no return value specified for Get") + } + + var r0 *biz.NotifyChannel + var r1 error + if rf, ok := ret.Get(0).(func(uint) (*biz.NotifyChannel, error)); ok { + return rf(id) + } + if rf, ok := ret.Get(0).(func(uint) *biz.NotifyChannel); ok { + r0 = rf(id) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*biz.NotifyChannel) + } + } + + if rf, ok := ret.Get(1).(func(uint) error); ok { + r1 = rf(id) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// NotifyChannelRepo_Get_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Get' +type NotifyChannelRepo_Get_Call struct { + *mock.Call +} + +// Get is a helper method to define mock.On call +// - id uint +func (_e *NotifyChannelRepo_Expecter) Get(id interface{}) *NotifyChannelRepo_Get_Call { + return &NotifyChannelRepo_Get_Call{Call: _e.mock.On("Get", id)} +} + +func (_c *NotifyChannelRepo_Get_Call) Run(run func(id uint)) *NotifyChannelRepo_Get_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint)) + }) + return _c +} + +func (_c *NotifyChannelRepo_Get_Call) Return(_a0 *biz.NotifyChannel, _a1 error) *NotifyChannelRepo_Get_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *NotifyChannelRepo_Get_Call) RunAndReturn(run func(uint) (*biz.NotifyChannel, error)) *NotifyChannelRepo_Get_Call { + _c.Call.Return(run) + return _c +} + +// GetByIDs provides a mock function with given fields: ids +func (_m *NotifyChannelRepo) GetByIDs(ids []uint) ([]*biz.NotifyChannel, error) { + ret := _m.Called(ids) + + if len(ret) == 0 { + panic("no return value specified for GetByIDs") + } + + var r0 []*biz.NotifyChannel + var r1 error + if rf, ok := ret.Get(0).(func([]uint) ([]*biz.NotifyChannel, error)); ok { + return rf(ids) + } + if rf, ok := ret.Get(0).(func([]uint) []*biz.NotifyChannel); ok { + r0 = rf(ids) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.NotifyChannel) + } + } + + if rf, ok := ret.Get(1).(func([]uint) error); ok { + r1 = rf(ids) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + +// NotifyChannelRepo_GetByIDs_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'GetByIDs' +type NotifyChannelRepo_GetByIDs_Call struct { + *mock.Call +} + +// GetByIDs is a helper method to define mock.On call +// - ids []uint +func (_e *NotifyChannelRepo_Expecter) GetByIDs(ids interface{}) *NotifyChannelRepo_GetByIDs_Call { + return &NotifyChannelRepo_GetByIDs_Call{Call: _e.mock.On("GetByIDs", ids)} +} + +func (_c *NotifyChannelRepo_GetByIDs_Call) Run(run func(ids []uint)) *NotifyChannelRepo_GetByIDs_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].([]uint)) + }) + return _c +} + +func (_c *NotifyChannelRepo_GetByIDs_Call) Return(_a0 []*biz.NotifyChannel, _a1 error) *NotifyChannelRepo_GetByIDs_Call { + _c.Call.Return(_a0, _a1) + return _c +} + +func (_c *NotifyChannelRepo_GetByIDs_Call) RunAndReturn(run func([]uint) ([]*biz.NotifyChannel, error)) *NotifyChannelRepo_GetByIDs_Call { + _c.Call.Return(run) + return _c +} + +// List provides a mock function with given fields: page, limit +func (_m *NotifyChannelRepo) List(page uint, limit uint) ([]*biz.NotifyChannel, int64, error) { + ret := _m.Called(page, limit) + + if len(ret) == 0 { + panic("no return value specified for List") + } + + var r0 []*biz.NotifyChannel + var r1 int64 + var r2 error + if rf, ok := ret.Get(0).(func(uint, uint) ([]*biz.NotifyChannel, int64, error)); ok { + return rf(page, limit) + } + if rf, ok := ret.Get(0).(func(uint, uint) []*biz.NotifyChannel); ok { + r0 = rf(page, limit) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*biz.NotifyChannel) + } + } + + if rf, ok := ret.Get(1).(func(uint, uint) int64); ok { + r1 = rf(page, limit) + } else { + r1 = ret.Get(1).(int64) + } + + if rf, ok := ret.Get(2).(func(uint, uint) error); ok { + r2 = rf(page, limit) + } else { + r2 = ret.Error(2) + } + + return r0, r1, r2 +} + +// NotifyChannelRepo_List_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'List' +type NotifyChannelRepo_List_Call struct { + *mock.Call +} + +// List is a helper method to define mock.On call +// - page uint +// - limit uint +func (_e *NotifyChannelRepo_Expecter) List(page interface{}, limit interface{}) *NotifyChannelRepo_List_Call { + return &NotifyChannelRepo_List_Call{Call: _e.mock.On("List", page, limit)} +} + +func (_c *NotifyChannelRepo_List_Call) Run(run func(page uint, limit uint)) *NotifyChannelRepo_List_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(uint), args[1].(uint)) + }) + return _c +} + +func (_c *NotifyChannelRepo_List_Call) Return(_a0 []*biz.NotifyChannel, _a1 int64, _a2 error) *NotifyChannelRepo_List_Call { + _c.Call.Return(_a0, _a1, _a2) + return _c +} + +func (_c *NotifyChannelRepo_List_Call) RunAndReturn(run func(uint, uint) ([]*biz.NotifyChannel, int64, error)) *NotifyChannelRepo_List_Call { + _c.Call.Return(run) + return _c +} + +// Update provides a mock function with given fields: channel +func (_m *NotifyChannelRepo) Update(channel *biz.NotifyChannel) error { + ret := _m.Called(channel) + + if len(ret) == 0 { + panic("no return value specified for Update") + } + + var r0 error + if rf, ok := ret.Get(0).(func(*biz.NotifyChannel) error); ok { + r0 = rf(channel) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// NotifyChannelRepo_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' +type NotifyChannelRepo_Update_Call struct { + *mock.Call +} + +// Update is a helper method to define mock.On call +// - channel *biz.NotifyChannel +func (_e *NotifyChannelRepo_Expecter) Update(channel interface{}) *NotifyChannelRepo_Update_Call { + return &NotifyChannelRepo_Update_Call{Call: _e.mock.On("Update", channel)} +} + +func (_c *NotifyChannelRepo_Update_Call) Run(run func(channel *biz.NotifyChannel)) *NotifyChannelRepo_Update_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(*biz.NotifyChannel)) + }) + return _c +} + +func (_c *NotifyChannelRepo_Update_Call) Return(_a0 error) *NotifyChannelRepo_Update_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *NotifyChannelRepo_Update_Call) RunAndReturn(run func(*biz.NotifyChannel) error) *NotifyChannelRepo_Update_Call { + _c.Call.Return(run) + return _c +} + +// NewNotifyChannelRepo creates a new instance of NotifyChannelRepo. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. +// The first argument is typically a *testing.T value. +func NewNotifyChannelRepo(t interface { + mock.TestingT + Cleanup(func()) +}) *NotifyChannelRepo { + mock := &NotifyChannelRepo{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/pkg/db/clickhouse.go b/pkg/db/clickhouse.go index 27a2d3e7..cfad6423 100644 --- a/pkg/db/clickhouse.go +++ b/pkg/db/clickhouse.go @@ -1,6 +1,7 @@ package db import ( + "context" "database/sql" "fmt" "strings" @@ -17,7 +18,7 @@ type ClickHouse struct { } // NewClickHouse 创建 ClickHouse 连接(HTTP API) -func NewClickHouse(username, password, address string) (*ClickHouse, error) { +func NewClickHouse(ctx context.Context, username, password, address string) (*ClickHouse, error) { client := resty.New() client.SetBaseURL(fmt.Sprintf("http://%s", address)) client.SetTimeout(10 * 1000 * 1000 * 1000) // 10s @@ -30,7 +31,7 @@ func NewClickHouse(username, password, address string) (*ClickHouse, error) { } // 测试连接 - if err := ch.Ping(); err != nil { + if err := ch.ping(ctx); err != nil { _ = client.Close() return nil, fmt.Errorf("connect to clickhouse failed: %w", err) } @@ -43,8 +44,25 @@ func (r *ClickHouse) Close() { } func (r *ClickHouse) Ping() error { - _, err := r.exec("SELECT 1") - return err + return r.ping(context.Background()) +} + +// ping 带 context 的连通性检查,供构造时使用 +func (r *ClickHouse) ping(ctx context.Context) error { + resp, err := r.client.R(). + SetContext(ctx). + SetQueryParam("user", r.username). + SetQueryParam("password", r.password). + SetBody("SELECT 1"). + Post("/") + if err != nil { + return fmt.Errorf("clickhouse query failed: %w", err) + } + if resp.StatusCode() != 200 { + return fmt.Errorf("clickhouse query error: %s", strings.TrimSpace(resp.String())) + } + + return nil } func (r *ClickHouse) Query(query string, args ...any) (*sql.Rows, error) { diff --git a/pkg/db/elasticsearch.go b/pkg/db/elasticsearch.go index 2d71be15..4481e1ef 100644 --- a/pkg/db/elasticsearch.go +++ b/pkg/db/elasticsearch.go @@ -1,6 +1,7 @@ package db import ( + "context" "encoding/json" "fmt" "strings" @@ -30,7 +31,7 @@ type ESDocument struct { } // NewElasticsearch 创建 Elasticsearch 连接 -func NewElasticsearch(address, username, password string) (*Elasticsearch, error) { +func NewElasticsearch(ctx context.Context, address, username, password string) (*Elasticsearch, error) { client := resty.New() client.SetBaseURL(fmt.Sprintf("http://%s", address)) client.SetTimeout(10 * 1000 * 1000 * 1000) // 10s @@ -39,7 +40,7 @@ func NewElasticsearch(address, username, password string) (*Elasticsearch, error } es := &Elasticsearch{client: client} - if err := es.Ping(); err != nil { + if err := es.ping(ctx); err != nil { _ = client.Close() return nil, fmt.Errorf("connect to elasticsearch failed: %w", err) } @@ -52,13 +53,19 @@ func (r *Elasticsearch) Close() { } func (r *Elasticsearch) Ping() error { - resp, err := r.client.R().Get("/") + return r.ping(context.Background()) +} + +// ping 带 context 的连通性检查,供构造时使用 +func (r *Elasticsearch) ping(ctx context.Context) error { + resp, err := r.client.R().SetContext(ctx).Get("/") if err != nil { return err } if resp.StatusCode() != 200 { return fmt.Errorf("elasticsearch ping failed: %s", resp.String()) } + return nil } diff --git a/pkg/db/mongodb.go b/pkg/db/mongodb.go index c2363a07..4527aa47 100644 --- a/pkg/db/mongodb.go +++ b/pkg/db/mongodb.go @@ -1,6 +1,7 @@ package db import ( + "context" "encoding/json" "fmt" "strings" @@ -16,14 +17,14 @@ type MongoDB struct { } // NewMongoDB 创建 MongoDB 连接 -func NewMongoDB(username, password, address string) (*MongoDB, error) { +func NewMongoDB(ctx context.Context, username, password, address string) (*MongoDB, error) { m := &MongoDB{ username: username, password: password, address: address, } - if err := m.Ping(); err != nil { + if err := m.ping(ctx); err != nil { return nil, fmt.Errorf("connect to mongodb failed: %w", err) } @@ -33,7 +34,12 @@ func NewMongoDB(username, password, address string) (*MongoDB, error) { func (r *MongoDB) Close() {} func (r *MongoDB) Ping() error { - _, err := r.mongosh(`db.runCommand({ping:1})`) + return r.ping(context.Background()) +} + +// ping 带 context 的连通性检查,供构造时使用 +func (r *MongoDB) ping(ctx context.Context) error { + _, err := r.mongoshContext(ctx, `db.runCommand({ping:1})`) return err } @@ -149,14 +155,21 @@ func (r *MongoDB) Users() ([]MongoUser, error) { // mongosh 执行 mongosh 命令 func (r *MongoDB) mongosh(eval string) (string, error) { - cmd := fmt.Sprintf(`mongosh --quiet --eval "%s" mongodb://%s:%s@%s/admin 2>/dev/null`, + return r.mongoshContext(context.Background(), eval) +} + +// mongoshContext 执行 mongosh 命令,ctx 取消时终止进程 +func (r *MongoDB) mongoshContext(ctx context.Context, eval string) (string, error) { + // serverSelectionTimeoutMS 限制建连耗时,避免不可达地址长时间挂起 + cmd := fmt.Sprintf(`mongosh --quiet --eval "%s" "mongodb://%s:%s@%s/admin?serverSelectionTimeoutMS=10000" 2>/dev/null`, strings.ReplaceAll(eval, `"`, `\"`), r.username, r.password, r.address, ) - raw, err := shell.Execf(cmd) + raw, err := shell.ExecfWithContext(ctx, cmd) if err != nil { return "", fmt.Errorf("mongosh error: %w", err) } + return strings.TrimSpace(raw), nil } diff --git a/pkg/db/mysql.go b/pkg/db/mysql.go index 173757a2..ba090e07 100644 --- a/pkg/db/mysql.go +++ b/pkg/db/mysql.go @@ -1,6 +1,7 @@ package db import ( + "context" "database/sql" "fmt" "regexp" @@ -16,16 +17,18 @@ type MySQL struct { address string } -func NewMySQL(username, password, address string, typ ...string) (Operator, error) { - dsn := fmt.Sprintf("%s:%s@tcp(%s)/", username, password, address) +func NewMySQL(ctx context.Context, username, password, address string, typ ...string) (Operator, error) { + // 限制建连与读写超时,避免不可达地址阻塞调用方 + dsn := fmt.Sprintf("%s:%s@tcp(%s)/?timeout=5s&readTimeout=10s&writeTimeout=10s", username, password, address) if len(typ) > 0 && typ[0] == "unix" { - dsn = fmt.Sprintf("%s:%s@unix(%s)/", username, password, address) + dsn = fmt.Sprintf("%s:%s@unix(%s)/?timeout=5s&readTimeout=10s&writeTimeout=10s", username, password, address) } db, err := sql.Open("mysql", dsn) if err != nil { return nil, fmt.Errorf("init mysql connection failed: %w", err) } - if err = db.Ping(); err != nil { + if err = db.PingContext(ctx); err != nil { + _ = db.Close() return nil, fmt.Errorf("connect to mysql failed: %w", err) } return &MySQL{ diff --git a/pkg/db/postgres.go b/pkg/db/postgres.go index cd5edda4..867b1916 100644 --- a/pkg/db/postgres.go +++ b/pkg/db/postgres.go @@ -1,6 +1,7 @@ package db import ( + "context" "database/sql" "fmt" "slices" @@ -17,21 +18,23 @@ type Postgres struct { port uint } -func NewPostgres(username, password, address string, port uint) (Operator, error) { +func NewPostgres(ctx context.Context, username, password, address string, port uint) (Operator, error) { username = strings.ReplaceAll(username, `'`, `\'`) password = strings.ReplaceAll(password, `'`, `\'`) - dsn := fmt.Sprintf(`host=%s port=%d user='%s' password='%s' dbname=postgres sslmode=disable`, address, port, username, password) + // connect_timeout 限制建连耗时,避免不可达地址阻塞调用方 + dsn := fmt.Sprintf(`host=%s port=%d user='%s' password='%s' dbname=postgres sslmode=disable connect_timeout=5`, address, port, username, password) if password == "" { if username == "" { username = "postgres" } - dsn = fmt.Sprintf(`host=%s port=%d user='%s' dbname=postgres sslmode=disable`, address, port, username) + dsn = fmt.Sprintf(`host=%s port=%d user='%s' dbname=postgres sslmode=disable connect_timeout=5`, address, port, username) } db, err := sql.Open("postgres", dsn) if err != nil { return nil, fmt.Errorf("init postgres connection failed: %w", err) } - if err = db.Ping(); err != nil { + if err = db.PingContext(ctx); err != nil { + _ = db.Close() return nil, fmt.Errorf("connect to postgres failed: %w", err) } return &Postgres{ diff --git a/pkg/db/redis.go b/pkg/db/redis.go index e20f79b0..26c29114 100644 --- a/pkg/db/redis.go +++ b/pkg/db/redis.go @@ -1,6 +1,7 @@ package db import ( + "context" "encoding/json" "fmt" "time" @@ -26,8 +27,15 @@ type Redis struct { address string } -func NewRedis(username, password, address string) (*Redis, error) { - conn, err := redis.Dial("tcp", address, redis.DialUsername(username), redis.DialPassword(password)) +func NewRedis(ctx context.Context, username, password, address string) (*Redis, error) { + // 限制建连与读写超时,避免不可达地址阻塞调用方 + conn, err := redis.DialContext(ctx, "tcp", address, + redis.DialUsername(username), + redis.DialPassword(password), + redis.DialConnectTimeout(5*time.Second), + redis.DialReadTimeout(10*time.Second), + redis.DialWriteTimeout(10*time.Second), + ) if err != nil { return nil, err } diff --git a/pkg/notify/notify.go b/pkg/notify/notify.go new file mode 100644 index 00000000..08a54452 --- /dev/null +++ b/pkg/notify/notify.go @@ -0,0 +1,34 @@ +// Package notify 提供面板通知渠道的统一抽象 +package notify + +import ( + "context" + "encoding/json" + "fmt" +) + +// 通知渠道类型 +const ( + TypeSMTP = "smtp" +) + +// Message 通知消息 +type Message struct { + Subject string + Body string // HTML 正文 +} + +// Notifier 通知渠道 +type Notifier interface { + Send(ctx context.Context, msg *Message) error +} + +// New 按渠道类型构造通知器,config 为渠道的 JSON 配置 +func New(typ string, config json.RawMessage) (Notifier, error) { + switch typ { + case TypeSMTP: + return NewSMTP(config) + default: + return nil, fmt.Errorf("unsupported notify channel type: %s", typ) + } +} diff --git a/pkg/notify/smtp.go b/pkg/notify/smtp.go new file mode 100644 index 00000000..f747b9c9 --- /dev/null +++ b/pkg/notify/smtp.go @@ -0,0 +1,113 @@ +package notify + +import ( + "context" + "crypto/tls" + "encoding/json" + "errors" + "fmt" + "slices" + "time" + + "github.com/wneessen/go-mail" +) + +// SMTP 加密方式 +const ( + EncryptionNone = "none" + EncryptionSSL = "ssl" + EncryptionSTARTTLS = "starttls" +) + +// SMTPConfig SMTP 渠道配置 +type SMTPConfig struct { + Host string `json:"host"` + Port int `json:"port"` + Encryption string `json:"encryption"` // none / ssl / starttls + Username string `json:"username"` + Password string `json:"password"` + From string `json:"from"` + FromName string `json:"from_name"` + To []string `json:"to"` + SkipVerify bool `json:"skip_verify"` +} + +type smtpNotifier struct { + conf SMTPConfig +} + +// NewSMTP 构造 SMTP 通知器 +func NewSMTP(config json.RawMessage) (Notifier, error) { + var conf SMTPConfig + if err := json.Unmarshal(config, &conf); err != nil { + return nil, err + } + if conf.Host == "" || conf.Port <= 0 { + return nil, errors.New("smtp host and port are required") + } + if conf.From == "" { + conf.From = conf.Username + } + if conf.From == "" { + return nil, errors.New("smtp sender address is required") + } + if len(conf.To) == 0 { + return nil, errors.New("smtp recipients are required") + } + // 未知取值不能落到明文分支,否则拼写错误会让预期加密的连接静默降级 + if !slices.Contains([]string{EncryptionNone, EncryptionSSL, EncryptionSTARTTLS}, conf.Encryption) { + return nil, fmt.Errorf("unsupported smtp encryption: %s", conf.Encryption) + } + + return &smtpNotifier{conf: conf}, nil +} + +func (s *smtpNotifier) Send(ctx context.Context, msg *Message) error { + m := mail.NewMsg() + if err := m.FromFormat(s.conf.FromName, s.conf.From); err != nil { + return err + } + if err := m.To(s.conf.To...); err != nil { + return err + } + m.Subject(msg.Subject) + m.SetBodyString(mail.TypeTextHTML, msg.Body) + + options := []mail.Option{ + mail.WithPort(s.conf.Port), + mail.WithTimeout(30 * time.Second), + mail.WithTLSConfig(&tls.Config{ + ServerName: s.conf.Host, + InsecureSkipVerify: s.conf.SkipVerify, // nolint:gosec + MinVersion: tls.VersionTLS11, + }), + } + + switch s.conf.Encryption { + case EncryptionSSL: + options = append(options, mail.WithSSL()) + case EncryptionSTARTTLS: + options = append(options, mail.WithTLSPolicy(mail.TLSMandatory)) + default: + options = append(options, mail.WithTLSPolicy(mail.NoTLS)) + } + + // 无用户名视为匿名投递 + if s.conf.Username != "" { + options = append(options, + mail.WithSMTPAuth(mail.SMTPAuthAutoDiscover), + mail.WithUsername(s.conf.Username), + mail.WithPassword(s.conf.Password), + ) + } else { + options = append(options, mail.WithSMTPAuth(mail.SMTPAuthNoAuth)) + } + + client, err := mail.NewClient(s.conf.Host, options...) + if err != nil { + return err + } + defer func() { _ = client.Close() }() + + return client.DialAndSendWithContext(ctx, m) +} diff --git a/pkg/shell/exec.go b/pkg/shell/exec.go index d1938813..693e5811 100644 --- a/pkg/shell/exec.go +++ b/pkg/shell/exec.go @@ -85,6 +85,29 @@ func Execf(shell string, args ...any) (string, error) { return strings.TrimSpace(stdout.String()), nil } +// ExecfWithContext 安全执行 shell 命令,ctx 取消时终止进程 +func ExecfWithContext(ctx context.Context, shell string, args ...any) (string, error) { + if !preCheckArg(args) { + return "", errors.New("command contains illegal characters") + } + if len(args) > 0 { + shell = fmt.Sprintf(shell, args...) + } + + _ = os.Setenv("LC_ALL", "C") + cmd := exec.CommandContext(ctx, "bash", "-c", shell) + + var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + if err := cmd.Run(); err != nil { + return strings.TrimSpace(stdout.String()), fmt.Errorf("run %s failed, err: %w, stderr: %s", shell, err, strings.TrimSpace(stderr.String())) + } + + return strings.TrimSpace(stdout.String()), nil +} + // ExecfAsync 异步执行 shell 命令 func ExecfAsync(shell string, args ...any) error { if !preCheckArg(args) { diff --git a/pkg/sshlog/sshlog.go b/pkg/sshlog/sshlog.go new file mode 100644 index 00000000..ad7ce3fe --- /dev/null +++ b/pkg/sshlog/sshlog.go @@ -0,0 +1,114 @@ +// Package sshlog 解析 sshd 登录日志 +package sshlog + +import ( + "bufio" + "bytes" + "regexp" + "strings" + "time" + + "github.com/acepanel/panel/v3/pkg/types" +) + +// 登录记录状态 +const ( + StatusAccepted = "accepted" + StatusFailed = "failed" + StatusInvalidUser = "invalid_user" + StatusDisconnected = "disconnected" +) + +var ( + accepted = regexp.MustCompile(`Accepted\s+(\S+)\s+for\s+(\S+)\s+from\s+(\S+)\s+port\s+(\d+)`) + failed = regexp.MustCompile(`Failed\s+(\S+)\s+for\s+(?:invalid user\s+)?(\S+)\s+from\s+(\S+)\s+port\s+(\d+)`) + invalidUser = regexp.MustCompile(`Invalid user\s+(\S+)\s+from\s+(\S+)\s+port\s+(\d+)`) + disconnect = regexp.MustCompile(`Disconnected from\s+(?:authenticating\s+)?user\s+(\S+)\s+(\S+)\s+port\s+(\d+)`) +) + +// ParseMessage 从日志消息中提取 SSH 登录信息,无法识别时返回 nil +func ParseMessage(msg string) *types.SSHLoginLog { + if m := accepted.FindStringSubmatch(msg); m != nil { + return &types.SSHLoginLog{ + Method: m[1], + User: m[2], + IP: m[3], + Port: m[4], + Status: StatusAccepted, + } + } + if m := failed.FindStringSubmatch(msg); m != nil { + return &types.SSHLoginLog{ + Method: m[1], + User: m[2], + IP: m[3], + Port: m[4], + Status: StatusFailed, + } + } + if m := invalidUser.FindStringSubmatch(msg); m != nil { + return &types.SSHLoginLog{ + User: m[1], + IP: m[2], + Port: m[3], + Method: "-", + Status: StatusInvalidUser, + } + } + if m := disconnect.FindStringSubmatch(msg); m != nil { + return &types.SSHLoginLog{ + User: m[1], + IP: m[2], + Port: m[3], + Method: "-", + Status: StatusDisconnected, + } + } + + return nil +} + +// ParseChunk 从连续日志字节中解析 SSH 登录记录 +func ParseChunk(data []byte) []types.SSHLoginLog { + if len(data) == 0 { + return nil + } + + var logs []types.SSHLoginLog + scanner := bufio.NewScanner(bytes.NewReader(data)) + scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) + for scanner.Scan() { + line := scanner.Text() + if !strings.Contains(line, "sshd[") { + continue + } + record := ParseMessage(line) + if record == nil { + continue + } + record.Time = ParseTime(line) + logs = append(logs, *record) + } + + return logs +} + +// ParseTime 从 syslog 格式行中解析时间,失败返回 "-" +func ParseTime(line string) string { + // syslog 格式:Mon DD HH:MM:SS(前 15 个字符) + if len(line) < 15 { + return "-" + } + + ts := line[:15] + // 使用当前年份补全 + t, err := time.Parse("Jan 2 15:04:05", ts) + if err != nil { + t, err = time.Parse("Jan 2 15:04:05", ts) + if err != nil { + return "-" + } + } + + return t.AddDate(time.Now().Year(), 0, 0).Format(time.DateTime) +} diff --git a/web/src/api/panel/alert/index.ts b/web/src/api/panel/alert/index.ts new file mode 100644 index 00000000..3edd6181 --- /dev/null +++ b/web/src/api/panel/alert/index.ts @@ -0,0 +1,17 @@ +import { http } from '@/utils' + +export default { + // 告警规则列表 + rules: (page: number, limit: number): any => http.Get('/alert/rule', { params: { page, limit } }), + // 新增告警规则 + createRule: (data: any): any => http.Post('/alert/rule', data), + // 更新告警规则 + updateRule: (id: number, data: any): any => http.Put(`/alert/rule/${id}`, data), + // 删除告警规则 + deleteRule: (id: number): any => http.Delete(`/alert/rule/${id}`), + // 告警记录列表 + records: (page: number, limit: number): any => + http.Get('/alert/record', { params: { page, limit } }), + // 清空告警记录 + clearRecords: (): any => http.Post('/alert/record/clear'), +} diff --git a/web/src/api/panel/container/index.ts b/web/src/api/panel/container/index.ts index eb2586de..59928a4c 100644 --- a/web/src/api/panel/container/index.ts +++ b/web/src/api/panel/container/index.ts @@ -9,8 +9,7 @@ export default { // 添加容器 containerCreate: (config: any): any => http.Post('/container/container', config), // 更新容器(删除重建) - containerUpdate: (id: string, config: any): any => - http.Put(`/container/container/${id}`, config), + containerUpdate: (id: string, config: any): any => http.Put(`/container/container/${id}`, config), // 删除容器 containerRemove: (id: string): any => http.Delete(`/container/container/${id}`), // 启动容器 diff --git a/web/src/api/panel/monitor/index.ts b/web/src/api/panel/monitor/index.ts index 489b142a..53e5f7a8 100644 --- a/web/src/api/panel/monitor/index.ts +++ b/web/src/api/panel/monitor/index.ts @@ -4,8 +4,7 @@ export default { // 开关 setting: (): any => http.Get('/monitor/setting'), // 保存设置 - updateSetting: (enabled: boolean, days: number, interval: number): any => - http.Post('/monitor/setting', { enabled, days, interval }), + updateSetting: (data: any): any => http.Post('/monitor/setting', data), // 清空监控记录 clear: (): any => http.Post('/monitor/clear'), // 监控记录 diff --git a/web/src/api/panel/notify/index.ts b/web/src/api/panel/notify/index.ts new file mode 100644 index 00000000..62afaa53 --- /dev/null +++ b/web/src/api/panel/notify/index.ts @@ -0,0 +1,21 @@ +import { http } from '@/utils' + +export default { + // 通知渠道列表 + channels: (page: number, limit: number): any => + http.Get('/notify/channel', { params: { page, limit } }), + // 全部通知渠道 + allChannels: (): any => http.Get('/notify/channel/all'), + // 新增通知渠道 + createChannel: (data: any): any => http.Post('/notify/channel', data), + // 更新通知渠道 + updateChannel: (id: number, data: any): any => http.Put(`/notify/channel/${id}`, data), + // 删除通知渠道 + deleteChannel: (id: number): any => http.Delete(`/notify/channel/${id}`), + // 测试通知渠道 + testChannel: (id: number): any => http.Post(`/notify/channel/${id}/test`), + // 事件通知设置 + setting: (): any => http.Get('/notify/setting'), + // 保存事件通知设置 + updateSetting: (data: any): any => http.Post('/notify/setting', data), +} diff --git a/web/src/api/panel/tamper/index.ts b/web/src/api/panel/tamper/index.ts index 3b0e8160..2c35745d 100644 --- a/web/src/api/panel/tamper/index.ts +++ b/web/src/api/panel/tamper/index.ts @@ -23,8 +23,7 @@ export default { // 删除规则 deleteRule: (id: number): any => http.Delete(`/tamper/rule/${id}`), // 拦截日志 - logs: (page: number, limit: number): any => - http.Get('/tamper/log', { params: { page, limit } }), + logs: (page: number, limit: number): any => http.Get('/tamper/log', { params: { page, limit } }), // 清空日志 clearLogs: (): any => http.Delete('/tamper/log'), } diff --git a/web/src/components/system/HealthBanner.vue b/web/src/components/system/HealthBanner.vue index 65d8be4a..7d3070cd 100644 --- a/web/src/components/system/HealthBanner.vue +++ b/web/src/components/system/HealthBanner.vue @@ -1,4 +1,3 @@ -