发布 0.1.0:通知渠道、告警规则、令牌版本与安全加固
This commit is contained in:
@@ -14,6 +14,7 @@ import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"oci-portal/internal/aiwire"
|
||||
"oci-portal/internal/model"
|
||||
@@ -506,13 +507,36 @@ func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, error) {
|
||||
// LogCall 落一条调用日志(仅元数据与用量,永不含请求 / 响应正文),返回落库 ID 供内容日志关联(失败为 0)。
|
||||
func (s *AiGatewayService) LogCall(entry model.AiCallLog) uint {
|
||||
entry.ErrMsg = truncateErr(entry.ErrMsg)
|
||||
if err := s.db.Create(&entry).Error; err != nil {
|
||||
stored := false
|
||||
err := s.db.Transaction(func(tx *gorm.DB) error {
|
||||
ok, err := lockAiLogParent(tx, &model.AiChannel{}, entry.ChannelID)
|
||||
if err != nil || !ok {
|
||||
return err
|
||||
}
|
||||
stored = true
|
||||
return tx.Create(&entry).Error
|
||||
})
|
||||
if err != nil {
|
||||
log.Printf("ai call log: %v", err)
|
||||
return 0
|
||||
}
|
||||
if !stored {
|
||||
return 0
|
||||
}
|
||||
return entry.ID
|
||||
}
|
||||
|
||||
func lockAiLogParent(tx *gorm.DB, value any, id uint) (bool, error) {
|
||||
if id == 0 {
|
||||
return true, nil
|
||||
}
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(value, id).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
return err == nil, err
|
||||
}
|
||||
|
||||
// CallLogs 分页查询调用日志。
|
||||
func (s *AiGatewayService) CallLogs(ctx context.Context, page, size int) ([]model.AiCallLog, int64, error) {
|
||||
if page < 1 {
|
||||
@@ -555,9 +579,19 @@ func (s *AiGatewayService) UpdateKeyContentLog(ctx context.Context, id uint, hou
|
||||
|
||||
// LogContent 写一条内容日志(调用方已确认密钥开启且未过期);正文截断至 64KB。
|
||||
func (s *AiGatewayService) LogContent(entry model.AiContentLog) {
|
||||
if entry.CallLogID == 0 {
|
||||
return
|
||||
}
|
||||
entry.RequestBody = truncateBody(entry.RequestBody)
|
||||
entry.ResponseBody = truncateBody(entry.ResponseBody)
|
||||
if err := s.db.Create(&entry).Error; err != nil {
|
||||
err := s.db.Transaction(func(tx *gorm.DB) error {
|
||||
ok, err := lockAiLogParent(tx, &model.AiCallLog{}, entry.CallLogID)
|
||||
if err != nil || !ok {
|
||||
return err
|
||||
}
|
||||
return tx.Create(&entry).Error
|
||||
})
|
||||
if err != nil {
|
||||
log.Printf("ai content log: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -488,6 +488,7 @@ func TestAiContentLogSwitch(t *testing.T) {
|
||||
t.Error("超过 7 天上限应被拒绝")
|
||||
}
|
||||
// 写入与截断(带调用日志关联)
|
||||
mustCreate(t, gw.db, &model.AiCallLog{ID: 42, KeyID: key.ID, ChannelID: 0})
|
||||
gw.LogContent(model.AiContentLog{CallLogID: 42, KeyID: key.ID, KeyName: "k1", Endpoint: "openai", Model: "m", RequestBody: strings.Repeat("x", 70*1024)})
|
||||
rows, total, err := gw.ContentLogs(ctx, key.ID, 0, 1, 20)
|
||||
if err != nil || total != 1 || len(rows[0].RequestBody) != 64*1024 {
|
||||
|
||||
@@ -0,0 +1,305 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// 告警规则约束与命中记录保留期(窗口计数之外多留几天便于排查)。
|
||||
const (
|
||||
alertMaxThreshold = 100
|
||||
alertMaxWindowMin = 1440
|
||||
alertHitRetention = 7 * 24 * time.Hour
|
||||
alertSourceIPIn = "in"
|
||||
alertSourceIPNotIn = "notin"
|
||||
)
|
||||
|
||||
// ErrInvalidAlertRule 标记规则字段非法,api 层映射 400。
|
||||
var ErrInvalidAlertRule = fmt.Errorf("告警规则字段非法")
|
||||
|
||||
// ListAlertRules 返回全部告警规则(创建顺序)。
|
||||
func (s *LogEventService) ListAlertRules(ctx context.Context) ([]model.AlertRule, error) {
|
||||
var rules []model.AlertRule
|
||||
if err := s.db.WithContext(ctx).Order("id").Find(&rules).Error; err != nil {
|
||||
return nil, fmt.Errorf("list alert rules: %w", err)
|
||||
}
|
||||
return rules, nil
|
||||
}
|
||||
|
||||
// CreateAlertRule 校验并创建规则。
|
||||
func (s *LogEventService) CreateAlertRule(ctx context.Context, rule model.AlertRule) (model.AlertRule, error) {
|
||||
if err := validateAlertRule(&rule); err != nil {
|
||||
return model.AlertRule{}, err
|
||||
}
|
||||
rule.ID = 0
|
||||
if err := s.db.WithContext(ctx).Create(&rule).Error; err != nil {
|
||||
return model.AlertRule{}, fmt.Errorf("create alert rule: %w", err)
|
||||
}
|
||||
return rule, nil
|
||||
}
|
||||
|
||||
// UpdateAlertRule 校验并整体覆盖规则(含启停)。
|
||||
func (s *LogEventService) UpdateAlertRule(ctx context.Context, id uint, rule model.AlertRule) (model.AlertRule, error) {
|
||||
if err := validateAlertRule(&rule); err != nil {
|
||||
return model.AlertRule{}, err
|
||||
}
|
||||
var cur model.AlertRule
|
||||
if err := s.db.WithContext(ctx).First(&cur, id).Error; err != nil {
|
||||
return model.AlertRule{}, fmt.Errorf("find alert rule %d: %w", id, err)
|
||||
}
|
||||
rule.ID, rule.CreatedAt = cur.ID, cur.CreatedAt
|
||||
if err := s.db.WithContext(ctx).Save(&rule).Error; err != nil {
|
||||
return model.AlertRule{}, fmt.Errorf("update alert rule: %w", err)
|
||||
}
|
||||
return rule, nil
|
||||
}
|
||||
|
||||
// DeleteAlertRule 删除规则及其命中记录。
|
||||
func (s *LogEventService) DeleteAlertRule(ctx context.Context, id uint) error {
|
||||
if err := s.db.WithContext(ctx).Delete(&model.AlertRule{}, id).Error; err != nil {
|
||||
return fmt.Errorf("delete alert rule: %w", err)
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Where("rule_id = ?", id).Delete(&model.AlertRuleHit{}).Error; err != nil {
|
||||
return fmt.Errorf("delete alert rule hits: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateAlertRule 校验字段并归一化;非法时返回含具体原因的 ErrInvalidAlertRule 包装。
|
||||
func validateAlertRule(rule *model.AlertRule) error {
|
||||
rule.Name = strings.TrimSpace(rule.Name)
|
||||
if rule.Name == "" {
|
||||
return fmt.Errorf("%w: 名称必填", ErrInvalidAlertRule)
|
||||
}
|
||||
if rule.SourceIPMode == "" {
|
||||
rule.SourceIPMode = alertSourceIPIn
|
||||
}
|
||||
if rule.SourceIPMode != alertSourceIPIn && rule.SourceIPMode != alertSourceIPNotIn {
|
||||
return fmt.Errorf("%w: 来源 IP 模式须为 in/notin", ErrInvalidAlertRule)
|
||||
}
|
||||
if rule.Threshold < 1 || rule.Threshold > alertMaxThreshold {
|
||||
return fmt.Errorf("%w: 阈值须在 1-%d 之间", ErrInvalidAlertRule, alertMaxThreshold)
|
||||
}
|
||||
if rule.Threshold > 1 && (rule.WindowMinutes < 1 || rule.WindowMinutes > alertMaxWindowMin) {
|
||||
return fmt.Errorf("%w: 阈值>1 时窗口须在 1-%d 分钟之间", ErrInvalidAlertRule, alertMaxWindowMin)
|
||||
}
|
||||
if rule.EventTypes != "" {
|
||||
rule.EventTypes = normalizeCSV(rule.EventTypes)
|
||||
}
|
||||
return validateAlertRuleIPs(rule)
|
||||
}
|
||||
|
||||
// validateAlertRuleIPs 归一化并校验来源 IP 列表(裸 IP 或 CIDR)。
|
||||
func validateAlertRuleIPs(rule *model.AlertRule) error {
|
||||
if rule.SourceIPs == "" {
|
||||
return nil
|
||||
}
|
||||
rule.SourceIPs = normalizeCSV(rule.SourceIPs)
|
||||
for _, item := range strings.Split(rule.SourceIPs, ",") {
|
||||
if _, err := parseIPMatcher(item); err != nil {
|
||||
return fmt.Errorf("%w: 来源 IP %q 不是合法的 IP 或 CIDR", ErrInvalidAlertRule, item)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// normalizeCSV 去除各项空白与空项后重组逗号分隔串。
|
||||
func normalizeCSV(s string) string {
|
||||
parts := strings.Split(s, ",")
|
||||
out := parts[:0]
|
||||
for _, p := range parts {
|
||||
if p = strings.TrimSpace(p); p != "" {
|
||||
out = append(out, p)
|
||||
}
|
||||
}
|
||||
return strings.Join(out, ",")
|
||||
}
|
||||
|
||||
// parseIPMatcher 把裸 IP 或 CIDR 解析为前缀(裸 IP 视为单地址前缀)。
|
||||
func parseIPMatcher(item string) (netip.Prefix, error) {
|
||||
if strings.Contains(item, "/") {
|
||||
return netip.ParsePrefix(item)
|
||||
}
|
||||
addr, err := netip.ParseAddr(item)
|
||||
if err != nil {
|
||||
return netip.Prefix{}, err
|
||||
}
|
||||
return netip.PrefixFrom(addr, addr.BitLen()), nil
|
||||
}
|
||||
|
||||
// ipListMatch 报告 ip 是否命中列表中的任一前缀;ip 解析失败视为未命中。
|
||||
func ipListMatch(list, ip string) bool {
|
||||
addr, err := netip.ParseAddr(ip)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
for _, item := range strings.Split(list, ",") {
|
||||
if p, err := parseIPMatcher(item); err == nil && p.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ruleHits 报告事件是否命中规则的全部条件(AND 语义,空条件视为任意)。
|
||||
func ruleHits(rule model.AlertRule, e *model.LogEvent, p parsedEvent) bool {
|
||||
if rule.OciConfigID != 0 && rule.OciConfigID != e.OciConfigID {
|
||||
return false
|
||||
}
|
||||
name := relayEventShortName(p.EventType)
|
||||
if rule.EventTypes != "" && !slices.Contains(strings.Split(rule.EventTypes, ","), name) {
|
||||
return false
|
||||
}
|
||||
if rule.ResourceMatch != "" && !strings.Contains(p.ResourceName, rule.ResourceMatch) {
|
||||
return false
|
||||
}
|
||||
return ruleIPHits(rule, p.SourceIP)
|
||||
}
|
||||
|
||||
// ruleIPHits 按模式判定来源 IP 条件:in 命中列表告警;notin 不在列表才告警,
|
||||
// 事件缺 IP 字段时 notin 不告警(避免解析缺字段导致白名单误报)。
|
||||
func ruleIPHits(rule model.AlertRule, ip string) bool {
|
||||
if rule.SourceIPs == "" {
|
||||
return true
|
||||
}
|
||||
if rule.SourceIPMode == alertSourceIPNotIn {
|
||||
return ip != "" && !ipListMatch(rule.SourceIPs, ip)
|
||||
}
|
||||
return ipListMatch(rule.SourceIPs, ip)
|
||||
}
|
||||
|
||||
// matchAlertRules 对一条已解析事件执行全部启用规则;任何内部错误只记日志,不影响解析主流程。
|
||||
func (s *LogEventService) matchAlertRules(ctx context.Context, rules []model.AlertRule, e *model.LogEvent, p parsedEvent) {
|
||||
if s.notifier == nil {
|
||||
return
|
||||
}
|
||||
for _, rule := range rules {
|
||||
if !rule.Enabled || !ruleHits(rule, e, p) {
|
||||
continue
|
||||
}
|
||||
count, ok := s.recordAlertHit(ctx, rule, e)
|
||||
if !ok || count < rule.Threshold || !s.alertCooldownPass(rule) {
|
||||
continue
|
||||
}
|
||||
s.notifier.SendTemplateAsync("audit_alert", map[string]string{
|
||||
"rule": rule.Name, "tenant": s.configAlias(ctx, e.OciConfigID),
|
||||
"event": relayEventShortName(p.EventType), "resource": p.ResourceName,
|
||||
"ip": p.SourceIP, "count": fmt.Sprint(count),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// recordAlertHit 落一条命中并返回窗口内累计次数;阈值 1 的规则免计数直接触发。
|
||||
func (s *LogEventService) recordAlertHit(ctx context.Context, rule model.AlertRule, e *model.LogEvent) (int, bool) {
|
||||
count, err := s.recordAlertHitTx(ctx, rule, e)
|
||||
if err != nil {
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
log.Printf("alert rule hit record: %v", err)
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
return count, true
|
||||
}
|
||||
|
||||
func (s *LogEventService) recordAlertHitTx(ctx context.Context, rule model.AlertRule, event *model.LogEvent) (int, error) {
|
||||
count := 0
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var err error
|
||||
count, err = insertAlertHit(tx, rule, event)
|
||||
return err
|
||||
})
|
||||
return count, err
|
||||
}
|
||||
|
||||
// insertAlertHit 按 rule→event 锁顺序确认引用存在后插入并统计窗口命中。
|
||||
func insertAlertHit(tx *gorm.DB, rule model.AlertRule, event *model.LogEvent) (int, error) {
|
||||
if err := lockAlertRefs(tx, rule.ID, event.ID); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if rule.Threshold <= 1 {
|
||||
return 1, nil
|
||||
}
|
||||
now := time.Now()
|
||||
hit := model.AlertRuleHit{RuleID: rule.ID, LogEventID: event.ID, HitAt: now}
|
||||
if err := tx.Create(&hit).Error; err != nil {
|
||||
return 0, fmt.Errorf("create alert rule hit: %w", err)
|
||||
}
|
||||
var count int64
|
||||
cutoff := now.Add(-time.Duration(rule.WindowMinutes) * time.Minute)
|
||||
err := tx.Model(&model.AlertRuleHit{}).
|
||||
Where("rule_id = ? AND hit_at >= ?", rule.ID, cutoff).Count(&count).Error
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("count alert rule hits: %w", err)
|
||||
}
|
||||
return int(count), nil
|
||||
}
|
||||
|
||||
func lockAlertRefs(tx *gorm.DB, ruleID, eventID uint) error {
|
||||
var rule model.AlertRule
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&rule, ruleID).Error; err != nil {
|
||||
return fmt.Errorf("lock alert rule %d: %w", ruleID, err)
|
||||
}
|
||||
var event model.LogEvent
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&event, eventID).Error; err != nil {
|
||||
return fmt.Errorf("lock log event %d: %w", eventID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// alertCooldownPass 报告规则是否已过冷却窗口;通过即记录本次发送时刻。
|
||||
// 阈值 1 的规则无冷却(每次命中即时告警,与既有云端事件通知一致)。
|
||||
func (s *LogEventService) alertCooldownPass(rule model.AlertRule) bool {
|
||||
if rule.Threshold <= 1 {
|
||||
return true
|
||||
}
|
||||
s.alertMu.Lock()
|
||||
defer s.alertMu.Unlock()
|
||||
window := time.Duration(rule.WindowMinutes) * time.Minute
|
||||
if last, ok := s.alertSentAt[rule.ID]; ok && time.Since(last) < window {
|
||||
return false
|
||||
}
|
||||
if s.alertSentAt == nil {
|
||||
s.alertSentAt = map[uint]time.Time{}
|
||||
}
|
||||
s.alertSentAt[rule.ID] = time.Now()
|
||||
return true
|
||||
}
|
||||
|
||||
// ClearAlertCooldown 清除已删除租户规则的进程内冷却状态。
|
||||
func (s *LogEventService) ClearAlertCooldown(ruleIDs []uint) {
|
||||
s.alertMu.Lock()
|
||||
defer s.alertMu.Unlock()
|
||||
for _, id := range ruleIDs {
|
||||
delete(s.alertSentAt, id)
|
||||
}
|
||||
}
|
||||
|
||||
// loadEnabledAlertRules 载入启用中的规则;失败时返回空集并记日志(解析主流程照常)。
|
||||
func (s *LogEventService) loadEnabledAlertRules(ctx context.Context) []model.AlertRule {
|
||||
var rules []model.AlertRule
|
||||
err := s.db.WithContext(ctx).Where("enabled = ?", true).Order("id").Find(&rules).Error
|
||||
if err != nil {
|
||||
log.Printf("load alert rules: %v", err)
|
||||
return nil
|
||||
}
|
||||
return rules
|
||||
}
|
||||
|
||||
// cleanupAlertHits 删除保留期外的命中记录(随 cleanupOnce 周期执行)。
|
||||
func (s *LogEventService) cleanupAlertHits(ctx context.Context) {
|
||||
cutoff := time.Now().Add(-alertHitRetention)
|
||||
if err := s.db.WithContext(ctx).Where("hit_at < ?", cutoff).Delete(&model.AlertRuleHit{}).Error; err != nil {
|
||||
log.Printf("cleanup alert hits: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"oci-portal/internal/crypto"
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
func TestValidateAlertRule(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
rule model.AlertRule
|
||||
wantErr string
|
||||
check func(t *testing.T, r model.AlertRule)
|
||||
}{
|
||||
{name: "名称必填", rule: model.AlertRule{Threshold: 1}, wantErr: "名称"},
|
||||
{name: "模式非法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPMode: "any"}, wantErr: "in/notin"},
|
||||
{name: "阈值越界", rule: model.AlertRule{Name: "r", Threshold: 101}, wantErr: "阈值"},
|
||||
{name: "阈值>1须带窗口", rule: model.AlertRule{Name: "r", Threshold: 3}, wantErr: "窗口"},
|
||||
{name: "IP 非法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPs: "300.1.1.1"}, wantErr: "IP"},
|
||||
{name: "CIDR 合法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPs: "10.0.0.0/8, 1.2.3.4"},
|
||||
check: func(t *testing.T, r model.AlertRule) {
|
||||
if r.SourceIPs != "10.0.0.0/8,1.2.3.4" {
|
||||
t.Errorf("SourceIPs = %q, 应去空白归一化", r.SourceIPs)
|
||||
}
|
||||
if r.SourceIPMode != alertSourceIPIn {
|
||||
t.Errorf("SourceIPMode = %q, 应默认 in", r.SourceIPMode)
|
||||
}
|
||||
}},
|
||||
{name: "事件清单归一化", rule: model.AlertRule{Name: "r", Threshold: 1, EventTypes: " TerminateInstance , CreateApiKey ,"},
|
||||
check: func(t *testing.T, r model.AlertRule) {
|
||||
if r.EventTypes != "TerminateInstance,CreateApiKey" {
|
||||
t.Errorf("EventTypes = %q", r.EventTypes)
|
||||
}
|
||||
}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
rule := tt.rule
|
||||
err := validateAlertRule(&rule)
|
||||
if tt.wantErr == "" {
|
||||
if err != nil {
|
||||
t.Fatalf("validateAlertRule: %v", err)
|
||||
}
|
||||
if tt.check != nil {
|
||||
tt.check(t, rule)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("err = %v, want contains %q", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHits(t *testing.T) {
|
||||
base := model.AlertRule{Name: "r", Threshold: 1, SourceIPMode: alertSourceIPIn}
|
||||
ev := &model.LogEvent{OciConfigID: 7}
|
||||
parsed := parsedEvent{
|
||||
EventType: "com.oraclecloud.ComputeApi.TerminateInstance",
|
||||
SourceIP: "203.0.113.8",
|
||||
ResourceName: "web-server-1",
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
mod func(r *model.AlertRule)
|
||||
p *parsedEvent
|
||||
want bool
|
||||
}{
|
||||
{name: "空条件全命中", mod: func(r *model.AlertRule) {}, want: true},
|
||||
{name: "租户匹配", mod: func(r *model.AlertRule) { r.OciConfigID = 7 }, want: true},
|
||||
{name: "租户不匹配", mod: func(r *model.AlertRule) { r.OciConfigID = 8 }, want: false},
|
||||
{name: "事件短名命中", mod: func(r *model.AlertRule) { r.EventTypes = "LaunchInstance,TerminateInstance" }, want: true},
|
||||
{name: "事件不在清单", mod: func(r *model.AlertRule) { r.EventTypes = "CreateUser" }, want: false},
|
||||
{name: "资源子串命中", mod: func(r *model.AlertRule) { r.ResourceMatch = "web-" }, want: true},
|
||||
{name: "资源不含", mod: func(r *model.AlertRule) { r.ResourceMatch = "db-" }, want: false},
|
||||
{name: "IP in 命中 CIDR", mod: func(r *model.AlertRule) { r.SourceIPs = "203.0.113.0/24" }, want: true},
|
||||
{name: "IP in 未命中", mod: func(r *model.AlertRule) { r.SourceIPs = "10.0.0.0/8" }, want: false},
|
||||
{name: "IP notin 白名单外告警", mod: func(r *model.AlertRule) {
|
||||
r.SourceIPs, r.SourceIPMode = "10.0.0.0/8", alertSourceIPNotIn
|
||||
}, want: true},
|
||||
{name: "IP notin 白名单内不告警", mod: func(r *model.AlertRule) {
|
||||
r.SourceIPs, r.SourceIPMode = "203.0.113.8", alertSourceIPNotIn
|
||||
}, want: false},
|
||||
{name: "notin 事件缺 IP 不告警", mod: func(r *model.AlertRule) {
|
||||
r.SourceIPs, r.SourceIPMode = "10.0.0.0/8", alertSourceIPNotIn
|
||||
}, p: &parsedEvent{EventType: parsed.EventType}, want: false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
rule := base
|
||||
tt.mod(&rule)
|
||||
p := parsed
|
||||
if tt.p != nil {
|
||||
p = *tt.p
|
||||
}
|
||||
if got := ruleHits(rule, ev, p); got != tt.want {
|
||||
t.Errorf("ruleHits = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAlertRuleCRUD(t *testing.T) {
|
||||
svc, _, _ := newLogEventEnv(t)
|
||||
ctx := context.Background()
|
||||
|
||||
created, err := svc.CreateAlertRule(ctx, model.AlertRule{Name: "非白名单终止", Enabled: true, Threshold: 1})
|
||||
if err != nil {
|
||||
t.Fatalf("create: %v", err)
|
||||
}
|
||||
if created.ID == 0 {
|
||||
t.Fatal("create 未回填 ID")
|
||||
}
|
||||
if _, err := svc.CreateAlertRule(ctx, model.AlertRule{Threshold: 1}); err == nil {
|
||||
t.Fatal("空名称应校验失败")
|
||||
}
|
||||
|
||||
created.Enabled = false
|
||||
created.EventTypes = "TerminateInstance"
|
||||
updated, err := svc.UpdateAlertRule(ctx, created.ID, created)
|
||||
if err != nil {
|
||||
t.Fatalf("update: %v", err)
|
||||
}
|
||||
if updated.Enabled || updated.EventTypes != "TerminateInstance" {
|
||||
t.Fatalf("update 未生效: %+v", updated)
|
||||
}
|
||||
|
||||
rules, err := svc.ListAlertRules(ctx)
|
||||
if err != nil || len(rules) != 1 {
|
||||
t.Fatalf("list = %v, %v", rules, err)
|
||||
}
|
||||
if err := svc.DeleteAlertRule(ctx, created.ID); err != nil {
|
||||
t.Fatalf("delete: %v", err)
|
||||
}
|
||||
if rules, _ := svc.ListAlertRules(ctx); len(rules) != 0 {
|
||||
t.Fatalf("delete 后仍有 %d 条", len(rules))
|
||||
}
|
||||
}
|
||||
|
||||
// auditEventPayload 构造一条含资源与来源 IP 的 CloudEvents 审计消息。
|
||||
func auditEventPayload(event, resource, ip string) string {
|
||||
return fmt.Sprintf(`{"eventType":"com.oraclecloud.ComputeApi.%s","source":"ComputeApi",`+
|
||||
`"eventTime":"2026-07-10T08:00:00Z","data":{"resourceName":%q,"identity":{"ipAddress":%q}}}`,
|
||||
event, resource, ip)
|
||||
}
|
||||
|
||||
// newAlertNotifyEnv 组装带假 Telegram 通道的告警测试环境。
|
||||
func newAlertNotifyEnv(t *testing.T) (*LogEventService, *telegramCapture, func()) {
|
||||
t.Helper()
|
||||
svc, db, _ := newLogEventEnv(t)
|
||||
srv, rec := newFakeTelegram(t, `{"ok":true}`)
|
||||
cipher, err := crypto.NewCipher("test-data-key")
|
||||
if err != nil {
|
||||
t.Fatalf("new cipher: %v", err)
|
||||
}
|
||||
settings := NewSettingService(db, cipher)
|
||||
token := "123456:AAfake"
|
||||
if err := settings.UpdateTelegram(context.Background(),
|
||||
UpdateTelegramInput{Enabled: true, BotToken: &token, ChatID: "42"}); err != nil {
|
||||
t.Fatalf("update telegram: %v", err)
|
||||
}
|
||||
n := NewNotifier(settings)
|
||||
n.base = srv.URL
|
||||
svc.SetNotifier(n, settings)
|
||||
return svc, rec, n.Wait
|
||||
}
|
||||
|
||||
func TestMatchAlertRulesNotify(t *testing.T) {
|
||||
svc, rec, wait := newAlertNotifyEnv(t)
|
||||
ctx := context.Background()
|
||||
_, err := svc.CreateAlertRule(ctx, model.AlertRule{
|
||||
Name: "白名单外终止", Enabled: true, Threshold: 1,
|
||||
EventTypes: "TerminateInstance", SourceIPs: "10.0.0.0/8", SourceIPMode: alertSourceIPNotIn,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create rule: %v", err)
|
||||
}
|
||||
// 命中:白名单外 IP;不命中:白名单内 IP
|
||||
mustIngest(t, svc, "m1", auditEventPayload("TerminateInstance", "web-1", "203.0.113.8"))
|
||||
mustIngest(t, svc, "m2", auditEventPayload("TerminateInstance", "web-2", "10.1.2.3"))
|
||||
svc.parseOnce(ctx)
|
||||
wait()
|
||||
|
||||
alerts := auditAlerts(rec.snapshot())
|
||||
joined := strings.Join(alerts, "\n---\n")
|
||||
if !strings.Contains(joined, "白名单外终止") || !strings.Contains(joined, "web-1") {
|
||||
t.Fatalf("应收到含规则名与资源的告警,got %q", joined)
|
||||
}
|
||||
if strings.Contains(joined, "web-2") {
|
||||
t.Fatalf("白名单内事件不应告警,got %q", joined)
|
||||
}
|
||||
}
|
||||
|
||||
// auditAlerts 过滤出审计告警推送(排除既有 notifyCritical 的云端事件通知)。
|
||||
func auditAlerts(texts []string) []string {
|
||||
var out []string
|
||||
for _, s := range texts {
|
||||
if strings.Contains(s, "审计告警") {
|
||||
out = append(out, s)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func TestAlertThresholdWindow(t *testing.T) {
|
||||
svc, rec, wait := newAlertNotifyEnv(t)
|
||||
ctx := context.Background()
|
||||
_, err := svc.CreateAlertRule(ctx, model.AlertRule{
|
||||
Name: "登录风暴", Enabled: true, Threshold: 3, WindowMinutes: 5, EventTypes: "InteractiveLogin",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create rule: %v", err)
|
||||
}
|
||||
for i := 1; i <= 4; i++ {
|
||||
mustIngest(t, svc, fmt.Sprint("login-", i),
|
||||
auditEventPayload("InteractiveLogin", "user@x.com", "203.0.113.8"))
|
||||
}
|
||||
svc.parseOnce(ctx)
|
||||
wait()
|
||||
|
||||
alerts := auditAlerts(rec.snapshot())
|
||||
if len(alerts) != 1 {
|
||||
t.Fatalf("窗口内 4 次命中应只告警 1 次(第 3 次触发后冷却),got %d 条: %v", len(alerts), alerts)
|
||||
}
|
||||
if !strings.Contains(alerts[0], "3 次") {
|
||||
t.Errorf("告警文案应含累计次数,got %q", alerts[0])
|
||||
}
|
||||
}
|
||||
|
||||
// TestAlertRuleBadDataDoesNotBlockParse 验证规则表异常不影响解析主流程。
|
||||
func TestAlertRuleBadDataDoesNotBlockParse(t *testing.T) {
|
||||
svc, db, _ := newLogEventEnv(t)
|
||||
ctx := context.Background()
|
||||
// 直插一条绕过校验的坏规则(IP 列表非法)
|
||||
bad := model.AlertRule{Name: "bad", Enabled: true, Threshold: 1, SourceIPs: "not-an-ip"}
|
||||
if err := db.Create(&bad).Error; err != nil {
|
||||
t.Fatalf("insert bad rule: %v", err)
|
||||
}
|
||||
mustIngest(t, svc, "m1", auditEventPayload("TerminateInstance", "web-1", "1.2.3.4"))
|
||||
svc.parseOnce(ctx)
|
||||
|
||||
var e model.LogEvent
|
||||
if err := db.First(&e, "message_id = ?", "m1").Error; err != nil {
|
||||
t.Fatalf("find event: %v", err)
|
||||
}
|
||||
if !e.Processed {
|
||||
t.Fatal("坏规则不应阻塞事件解析")
|
||||
}
|
||||
}
|
||||
|
||||
// mustIngest 落一条回传事件,失败即终止测试。
|
||||
func mustIngest(t *testing.T, svc *LogEventService, msgID, payload string) {
|
||||
t.Helper()
|
||||
if err := svc.Ingest(context.Background(), 1, msgID, []byte(payload), false); err != nil {
|
||||
t.Fatalf("ingest %s: %v", msgID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCleanupAlertHits 验证过期命中记录随清理删除。
|
||||
func TestCleanupAlertHits(t *testing.T) {
|
||||
svc, db, _ := newLogEventEnv(t)
|
||||
old := model.AlertRuleHit{RuleID: 1, HitAt: time.Now().Add(-8 * 24 * time.Hour)}
|
||||
fresh := model.AlertRuleHit{RuleID: 1, HitAt: time.Now()}
|
||||
if err := db.Create(&old).Error; err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
if err := db.Create(&fresh).Error; err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
svc.cleanupAlertHits(context.Background())
|
||||
var count int64
|
||||
db.Model(&model.AlertRuleHit{}).Count(&count)
|
||||
if count != 1 {
|
||||
t.Fatalf("清理后应剩 1 条,got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClearAlertCooldown(t *testing.T) {
|
||||
svc := NewLogEventService(nil)
|
||||
rule1 := model.AlertRule{ID: 1, Threshold: 2, WindowMinutes: 10}
|
||||
rule2 := model.AlertRule{ID: 2, Threshold: 2, WindowMinutes: 10}
|
||||
if !svc.alertCooldownPass(rule1) || !svc.alertCooldownPass(rule2) {
|
||||
t.Fatal("首次命中应通过冷却检查")
|
||||
}
|
||||
|
||||
svc.ClearAlertCooldown([]uint{rule1.ID})
|
||||
if !svc.alertCooldownPass(rule1) {
|
||||
t.Fatal("已清理规则应重新通过冷却检查")
|
||||
}
|
||||
if svc.alertCooldownPass(rule2) {
|
||||
t.Fatal("未清理规则不应通过冷却检查")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordAlertHitRejectsMissingRefs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
deleteRule bool
|
||||
}{
|
||||
{name: "规则已删除", deleteRule: true},
|
||||
{name: "事件已删除"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
testRecordAlertHitMissingRef(t, tt.deleteRule)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testRecordAlertHitMissingRef(t *testing.T, deleteRule bool) {
|
||||
t.Helper()
|
||||
svc, db, cfgID := newLogEventEnv(t)
|
||||
rule := model.AlertRule{Name: "r", OciConfigID: cfgID, Threshold: 2, WindowMinutes: 5}
|
||||
event := model.LogEvent{OciConfigID: cfgID, MessageID: "m"}
|
||||
if err := db.Create(&rule).Error; err != nil {
|
||||
t.Fatalf("create rule: %v", err)
|
||||
}
|
||||
if err := db.Create(&event).Error; err != nil {
|
||||
t.Fatalf("create event: %v", err)
|
||||
}
|
||||
var err error
|
||||
if deleteRule {
|
||||
err = db.Delete(&rule).Error
|
||||
} else {
|
||||
err = db.Delete(&event).Error
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("delete reference: %v", err)
|
||||
}
|
||||
if count, ok := svc.recordAlertHit(context.Background(), rule, &event); ok || count != 0 {
|
||||
t.Fatalf("record missing refs = (%d,%v), want (0,false)", count, ok)
|
||||
}
|
||||
var hits int64
|
||||
db.Model(&model.AlertRuleHit{}).Count(&hits)
|
||||
if hits != 0 {
|
||||
t.Fatalf("orphan alert hits = %d, want 0", hits)
|
||||
}
|
||||
}
|
||||
+65
-50
@@ -2,17 +2,19 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"oci-portal/internal/oci"
|
||||
)
|
||||
|
||||
// 审计查询时间窗(小时)上下限。
|
||||
// 批式懒加载的单批条数上下限与缺省值。
|
||||
const (
|
||||
minAuditHours = 1
|
||||
maxAuditHours = 72
|
||||
auditLimitDefault = 100
|
||||
auditLimitMax = 200
|
||||
)
|
||||
|
||||
// 审计原始事件缓存:约 3KB/条,上限 2000 条 ≈ 6MB;TTL 内详情秒开。
|
||||
@@ -21,39 +23,32 @@ const (
|
||||
auditRawTTL = 10 * time.Minute
|
||||
)
|
||||
|
||||
// ErrInvalidAuditHours 表示审计时间窗参数越界,handler 据此返回 400。
|
||||
var ErrInvalidAuditHours = errors.New("audit events: hours must be between 1 and 72")
|
||||
|
||||
// ErrInvalidAuditWindow 表示续查的绝对时间窗非法(格式/顺序/跨度),handler 返回 400。
|
||||
var ErrInvalidAuditWindow = errors.New("audit events: invalid start/end window")
|
||||
// ErrInvalidAuditCursor 表示续查游标不可解析(过期格式/被篡改),handler 返回 400。
|
||||
var ErrInvalidAuditCursor = errors.New("audit events: invalid cursor, refresh to restart")
|
||||
|
||||
// ErrAuditEventGone 表示原始事件缓存过期且小窗重查未找回;handler 映射 404。
|
||||
var ErrAuditEventGone = errors.New("原始事件已不可取回,请刷新列表后重试")
|
||||
|
||||
// AuditQuery 是审计查询参数:Hours 为首查(相对当前时刻);
|
||||
// Start/End(RFC3339)为续查绝对窗,配合 Page 从截断游标断点续翻。
|
||||
// AuditQuery 是批式懒加载查询参数:Cursor 为空表示自当前时刻首查,
|
||||
// 非空则从上次响应的游标位置继续向更早回溯;Limit 为单批目标条数。
|
||||
type AuditQuery struct {
|
||||
Region string
|
||||
Hours int
|
||||
Start string
|
||||
End string
|
||||
Page string
|
||||
Cursor string
|
||||
Limit int
|
||||
}
|
||||
|
||||
// AuditEventsView 是审计查询响应:列表不含 raw(详情接口取回);
|
||||
// Start/End 回传本次实际使用的绝对窗,截断续查必须原样带回(游标绑定查询参数)。
|
||||
// AuditEventsView 是批式查询响应:列表不含 raw(详情接口取回);
|
||||
// Cursor 供下一批续查原样带回,空且 Exhausted 表示已到 365 天保留期尽头。
|
||||
type AuditEventsView struct {
|
||||
Items []oci.AuditEvent `json:"items"`
|
||||
Truncated bool `json:"truncated"`
|
||||
NextPage string `json:"nextPage,omitempty"`
|
||||
Start time.Time `json:"start"`
|
||||
End time.Time `json:"end"`
|
||||
Cursor string `json:"cursor,omitempty"`
|
||||
Exhausted bool `json:"exhausted"`
|
||||
}
|
||||
|
||||
// AuditEvents 实时查询租户 OCI 审计事件,纯透传不入库;region 为空时用配置
|
||||
// 默认区域。原始事件剥离进进程内缓存,响应体瘦身约 10 倍(调研④方案 A1+B1)。
|
||||
func (s *OciConfigService) AuditEvents(ctx context.Context, id uint, q AuditQuery) (AuditEventsView, error) {
|
||||
start, end, err := auditWindow(q)
|
||||
cur, err := decodeAuditCursor(q.Cursor)
|
||||
if err != nil {
|
||||
return AuditEventsView{}, err
|
||||
}
|
||||
@@ -61,53 +56,73 @@ func (s *OciConfigService) AuditEvents(ctx context.Context, id uint, q AuditQuer
|
||||
if err != nil {
|
||||
return AuditEventsView{}, err
|
||||
}
|
||||
res, err := s.client.ListAuditEvents(ctx, cred, q.Region, start, end, q.Page)
|
||||
res, err := s.client.ListAuditEventsBatch(ctx, cred, q.Region, cur, auditLimit(q.Limit))
|
||||
if err != nil {
|
||||
return AuditEventsView{}, err
|
||||
}
|
||||
return AuditEventsView{
|
||||
Items: s.stripAuditRaw(res.Items),
|
||||
Truncated: res.Truncated,
|
||||
NextPage: res.NextPage,
|
||||
// 与实际请求同粒度(分钟),续查回传时窗口逐字节一致
|
||||
Start: start.UTC().Truncate(time.Minute),
|
||||
End: end.UTC().Truncate(time.Minute),
|
||||
}, nil
|
||||
view := AuditEventsView{Items: s.stripAuditRaw(id, res.Items), Exhausted: res.Exhausted}
|
||||
if res.Cursor != nil {
|
||||
view.Cursor = encodeAuditCursor(*res.Cursor)
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
// auditWindow 解析查询窗口:带 start/end/page 走绝对窗校验,否则按 hours 相对窗。
|
||||
func auditWindow(q AuditQuery) (time.Time, time.Time, error) {
|
||||
if q.Start != "" || q.End != "" || q.Page != "" {
|
||||
start, err1 := time.Parse(time.RFC3339, q.Start)
|
||||
end, err2 := time.Parse(time.RFC3339, q.End)
|
||||
if err1 != nil || err2 != nil || !start.Before(end) ||
|
||||
end.Sub(start) > maxAuditHours*time.Hour+time.Minute {
|
||||
return time.Time{}, time.Time{}, ErrInvalidAuditWindow
|
||||
}
|
||||
return start, end, nil
|
||||
// auditLimit 归一单批条数:缺省 100,上限 200(响应体量与页预算的平衡)。
|
||||
func auditLimit(limit int) int {
|
||||
if limit <= 0 {
|
||||
return auditLimitDefault
|
||||
}
|
||||
if q.Hours < minAuditHours || q.Hours > maxAuditHours {
|
||||
return time.Time{}, time.Time{}, ErrInvalidAuditHours
|
||||
}
|
||||
end := time.Now().UTC()
|
||||
return end.Add(-time.Duration(q.Hours) * time.Hour), end, nil
|
||||
return min(limit, auditLimitMax)
|
||||
}
|
||||
|
||||
// stripAuditRaw 把每条原始事件按 eventId 放进缓存并从列表剥离。
|
||||
func (s *OciConfigService) stripAuditRaw(items []oci.AuditEvent) []oci.AuditEvent {
|
||||
// encodeAuditCursor 把游标序列化为不透明字符串(base64url JSON)。
|
||||
// 内容仅时间窗与 OCI 翻页令牌,伪造只影响自己的查询范围,无需签名。
|
||||
func encodeAuditCursor(cur oci.AuditCursor) string {
|
||||
b, _ := json.Marshal(cur)
|
||||
return base64.RawURLEncoding.EncodeToString(b)
|
||||
}
|
||||
|
||||
// decodeAuditCursor 解析续查游标;空串返回自当前时刻起的首查游标。
|
||||
func decodeAuditCursor(s string) (oci.AuditCursor, error) {
|
||||
if s == "" {
|
||||
return oci.NewAuditCursor(time.Now()), nil
|
||||
}
|
||||
b, err := base64.RawURLEncoding.DecodeString(s)
|
||||
if err != nil {
|
||||
return oci.AuditCursor{}, ErrInvalidAuditCursor
|
||||
}
|
||||
var cur oci.AuditCursor
|
||||
if err := json.Unmarshal(b, &cur); err != nil || cur.Start.IsZero() || !cur.Start.Before(cur.End) {
|
||||
return oci.AuditCursor{}, ErrInvalidAuditCursor
|
||||
}
|
||||
return cur, nil
|
||||
}
|
||||
|
||||
// auditRawKey 组装租户隔离的原始事件缓存键。
|
||||
func auditRawKey(configID uint, eventID string) string {
|
||||
return strconv.FormatUint(uint64(configID), 10) + ":" + eventID
|
||||
}
|
||||
|
||||
// stripAuditRaw 把每条原始事件按租户与 eventId 放进缓存并从列表剥离。
|
||||
func (s *OciConfigService) stripAuditRaw(configID uint, items []oci.AuditEvent) []oci.AuditEvent {
|
||||
for i := range items {
|
||||
if items[i].EventId != "" && items[i].Raw != nil {
|
||||
s.auditRaw.Set(items[i].EventId, items[i].Raw, auditRawTTL)
|
||||
s.auditRaw.Set(auditRawKey(configID, items[i].EventId), items[i].Raw, auditRawTTL)
|
||||
}
|
||||
items[i].Raw = nil
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
// InvalidateAuditCache 删除指定租户的全部审计原始事件缓存。
|
||||
func (s *OciConfigService) InvalidateAuditCache(configID uint) {
|
||||
s.auditRaw.DeletePrefix(strconv.FormatUint(uint64(configID), 10) + ":")
|
||||
}
|
||||
|
||||
// AuditEventDetail 取回单条事件的原始 JSON:缓存命中即回;miss 按事件时间
|
||||
// 所在分钟起 2 分钟小窗重查兜底(调研④方案 A2),仍未命中报 ErrAuditEventGone。
|
||||
func (s *OciConfigService) AuditEventDetail(ctx context.Context, id uint, region, eventID string, eventTime time.Time) (json.RawMessage, error) {
|
||||
if raw, ok := s.auditRaw.Get(eventID); ok {
|
||||
if raw, ok := s.auditRaw.Get(auditRawKey(id, eventID)); ok {
|
||||
return raw.(json.RawMessage), nil
|
||||
}
|
||||
cred, err := s.credentialsByID(ctx, id)
|
||||
@@ -124,7 +139,7 @@ func (s *OciConfigService) AuditEventDetail(ctx context.Context, id uint, region
|
||||
if ev.EventId == "" || ev.Raw == nil {
|
||||
continue
|
||||
}
|
||||
s.auditRaw.Set(ev.EventId, ev.Raw, auditRawTTL)
|
||||
s.auditRaw.Set(auditRawKey(id, ev.EventId), ev.Raw, auditRawTTL)
|
||||
if ev.EventId == eventID {
|
||||
found = ev.Raw
|
||||
}
|
||||
|
||||
+111
-78
@@ -13,115 +13,124 @@ import (
|
||||
// auditStubClient 覆写审计查询并记录透传参数,其余行为沿用 fakeClient。
|
||||
type auditStubClient struct {
|
||||
*fakeClient
|
||||
batch oci.AuditBatchResult
|
||||
result oci.AuditEventsResult
|
||||
gotRegion string
|
||||
gotCursor oci.AuditCursor
|
||||
gotLimit int
|
||||
gotStart time.Time
|
||||
gotEnd time.Time
|
||||
gotPage string
|
||||
calls int
|
||||
}
|
||||
|
||||
func (f *auditStubClient) ListAuditEventsBatch(ctx context.Context, cred oci.Credentials, region string, cur oci.AuditCursor, limit int) (oci.AuditBatchResult, error) {
|
||||
f.calls++
|
||||
f.gotRegion, f.gotCursor, f.gotLimit = region, cur, limit
|
||||
return f.batch, nil
|
||||
}
|
||||
|
||||
func (f *auditStubClient) ListAuditEvents(ctx context.Context, cred oci.Credentials, region string, start, end time.Time, page string) (oci.AuditEventsResult, error) {
|
||||
f.calls++
|
||||
f.gotRegion, f.gotStart, f.gotEnd, f.gotPage = region, start, end, page
|
||||
f.gotRegion, f.gotStart, f.gotEnd = region, start, end
|
||||
return f.result, nil
|
||||
}
|
||||
|
||||
func TestAuditEventsHoursValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
hours int
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "下界 1 小时", hours: 1},
|
||||
{name: "默认 24 小时", hours: 24},
|
||||
{name: "上界 72 小时", hours: 72},
|
||||
{name: "0 越界", hours: 0, wantErr: true},
|
||||
{name: "负数越界", hours: -3, wantErr: true},
|
||||
{name: "100 越界", hours: 100, wantErr: true},
|
||||
func TestAuditCursorCodec(t *testing.T) {
|
||||
cur := oci.AuditCursor{
|
||||
Start: time.Date(2026, 7, 9, 10, 0, 0, 0, time.UTC),
|
||||
End: time.Date(2026, 7, 10, 10, 0, 0, 0, time.UTC),
|
||||
Page: "tok-1",
|
||||
WindowHours: 48,
|
||||
}
|
||||
client := &auditStubClient{
|
||||
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
|
||||
result: oci.AuditEventsResult{Items: []oci.AuditEvent{{EventName: "GetInstance"}}},
|
||||
}
|
||||
svc := newTestService(t, client)
|
||||
cfg := importAliveConfig(t, svc)
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := svc.AuditEvents(context.Background(), cfg.ID, AuditQuery{Hours: tt.hours})
|
||||
if tt.wantErr {
|
||||
if !errors.Is(err, ErrInvalidAuditHours) {
|
||||
t.Fatalf("AuditEvents(hours=%d) error = %v, want ErrInvalidAuditHours", tt.hours, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("AuditEvents(hours=%d): %v", tt.hours, err)
|
||||
}
|
||||
if len(got.Items) != 1 {
|
||||
t.Errorf("items = %d, want 1", len(got.Items))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditEventsWindowAndRegion(t *testing.T) {
|
||||
client := &auditStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}}
|
||||
svc := newTestService(t, client)
|
||||
cfg := importAliveConfig(t, svc)
|
||||
|
||||
got, err := svc.AuditEvents(context.Background(), cfg.ID, AuditQuery{Region: "ap-tokyo-1", Hours: 6})
|
||||
got, err := decodeAuditCursor(encodeAuditCursor(cur))
|
||||
if err != nil {
|
||||
t.Fatalf("AuditEvents: %v", err)
|
||||
t.Fatalf("roundtrip: %v", err)
|
||||
}
|
||||
if client.calls != 1 {
|
||||
t.Fatalf("calls = %d, want 1", client.calls)
|
||||
if !got.Start.Equal(cur.Start) || !got.End.Equal(cur.End) || got.Page != cur.Page || got.WindowHours != cur.WindowHours {
|
||||
t.Fatalf("roundtrip = %+v, want %+v", got, cur)
|
||||
}
|
||||
if client.gotRegion != "ap-tokyo-1" {
|
||||
t.Errorf("region = %q, want %q", client.gotRegion, "ap-tokyo-1")
|
||||
|
||||
// 空串 → 自当前时刻首查:24h 窗、分钟粒度、无窗内游标
|
||||
first, err := decodeAuditCursor("")
|
||||
if err != nil {
|
||||
t.Fatalf("first cursor: %v", err)
|
||||
}
|
||||
if window := client.gotEnd.Sub(client.gotStart); window != 6*time.Hour {
|
||||
t.Errorf("window = %v, want %v", window, 6*time.Hour)
|
||||
if first.End.Sub(first.Start) != 24*time.Hour || first.Page != "" || first.End.Second() != 0 {
|
||||
t.Fatalf("first cursor = %+v, 应为 24h 分钟粒度首窗", first)
|
||||
}
|
||||
// 响应回传分钟粒度的绝对窗,续查据此原样带回
|
||||
if got.Start.Second() != 0 || !got.End.After(got.Start) {
|
||||
t.Errorf("响应窗口 = [%v, %v), 应为分钟粒度且有序", got.Start, got.End)
|
||||
|
||||
bad := []string{"!!!", "bm90LWpzb24", encodeAuditCursor(oci.AuditCursor{})}
|
||||
for i, s := range bad {
|
||||
if _, err := decodeAuditCursor(s); !errors.Is(err, ErrInvalidAuditCursor) {
|
||||
t.Errorf("bad[%d] err = %v, want ErrInvalidAuditCursor", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditEventsResume(t *testing.T) {
|
||||
func TestAuditEventsBatchParams(t *testing.T) {
|
||||
next := oci.AuditCursor{
|
||||
Start: time.Date(2026, 7, 8, 10, 0, 0, 0, time.UTC),
|
||||
End: time.Date(2026, 7, 9, 10, 0, 0, 0, time.UTC),
|
||||
WindowHours: 24,
|
||||
}
|
||||
client := &auditStubClient{
|
||||
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
|
||||
result: oci.AuditEventsResult{Items: []oci.AuditEvent{}, Truncated: true, NextPage: "tok-2"},
|
||||
batch: oci.AuditBatchResult{
|
||||
Items: []oci.AuditEvent{{EventId: "e1", EventName: "GetInstance"}},
|
||||
Cursor: &next,
|
||||
},
|
||||
}
|
||||
svc := newTestService(t, client)
|
||||
cfg := importAliveConfig(t, svc)
|
||||
ctx := context.Background()
|
||||
|
||||
// 续查:绝对窗 + 游标透传;NextPage 原样回传
|
||||
got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{
|
||||
Start: "2026-07-08T10:00:00Z", End: "2026-07-09T10:00:00Z", Page: "tok-1",
|
||||
})
|
||||
// limit 归一:0 → 100;超上限截到 200;region 透传
|
||||
cases := []struct {
|
||||
name string
|
||||
limit int
|
||||
wantLimit int
|
||||
}{
|
||||
{"缺省 100", 0, 100},
|
||||
{"正常透传", 50, 50},
|
||||
{"超限截断", 999, 200},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Region: "ap-tokyo-1", Limit: tc.limit})
|
||||
if err != nil {
|
||||
t.Fatalf("AuditEvents: %v", err)
|
||||
}
|
||||
if client.gotLimit != tc.wantLimit || client.gotRegion != "ap-tokyo-1" {
|
||||
t.Fatalf("limit = %d(want %d), region = %q", client.gotLimit, tc.wantLimit, client.gotRegion)
|
||||
}
|
||||
if got.Cursor == "" || got.Exhausted {
|
||||
t.Fatalf("响应应携带续查游标: %+v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// 续查:响应游标原样带回可解析,并透传到 oci 层
|
||||
got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{})
|
||||
if err != nil {
|
||||
t.Fatalf("首查: %v", err)
|
||||
}
|
||||
if _, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Cursor: got.Cursor}); err != nil {
|
||||
t.Fatalf("续查: %v", err)
|
||||
}
|
||||
if client.gotPage != "tok-1" || !got.Truncated || got.NextPage != "tok-2" {
|
||||
t.Errorf("游标透传 page=%q next=%q truncated=%v", client.gotPage, got.NextPage, got.Truncated)
|
||||
if !client.gotCursor.Start.Equal(next.Start) || client.gotCursor.WindowHours != 24 {
|
||||
t.Fatalf("续查游标透传 = %+v, want %+v", client.gotCursor, next)
|
||||
}
|
||||
if !client.gotStart.Equal(time.Date(2026, 7, 8, 10, 0, 0, 0, time.UTC)) {
|
||||
t.Errorf("续查未用绝对窗: start=%v", client.gotStart)
|
||||
|
||||
// 尽头:Cursor 为 nil → 响应空游标 + exhausted
|
||||
client.batch = oci.AuditBatchResult{Items: []oci.AuditEvent{}, Exhausted: true}
|
||||
end, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{})
|
||||
if err != nil || end.Cursor != "" || !end.Exhausted {
|
||||
t.Fatalf("尽头响应 = %+v, %v", end, err)
|
||||
}
|
||||
// 非法窗口逐项拒绝
|
||||
bad := []AuditQuery{
|
||||
{Start: "not-a-time", End: "2026-07-09T10:00:00Z"},
|
||||
{Start: "2026-07-09T10:00:00Z", End: "2026-07-08T10:00:00Z"}, // 倒序
|
||||
{Start: "2026-07-01T00:00:00Z", End: "2026-07-09T10:00:00Z"}, // 超 72h
|
||||
{Page: "tok-only"}, // 带游标缺窗口
|
||||
}
|
||||
for i, q := range bad {
|
||||
if _, err := svc.AuditEvents(ctx, cfg.ID, q); !errors.Is(err, ErrInvalidAuditWindow) {
|
||||
t.Errorf("bad[%d] err = %v, want ErrInvalidAuditWindow", i, err)
|
||||
}
|
||||
|
||||
// 非法游标 → ErrInvalidAuditCursor
|
||||
if _, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Cursor: "!!!"}); !errors.Is(err, ErrInvalidAuditCursor) {
|
||||
t.Fatalf("非法游标 err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -130,7 +139,7 @@ func TestAuditRawStrippedAndDetail(t *testing.T) {
|
||||
eventTime := time.Date(2026, 7, 9, 10, 30, 40, 0, time.UTC)
|
||||
client := &auditStubClient{
|
||||
fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}},
|
||||
result: oci.AuditEventsResult{Items: []oci.AuditEvent{
|
||||
batch: oci.AuditBatchResult{Items: []oci.AuditEvent{
|
||||
{EventId: "evt-1", EventTime: &eventTime, EventName: "TerminateInstance", Raw: raw},
|
||||
}},
|
||||
}
|
||||
@@ -138,7 +147,7 @@ func TestAuditRawStrippedAndDetail(t *testing.T) {
|
||||
cfg := importAliveConfig(t, svc)
|
||||
ctx := context.Background()
|
||||
|
||||
got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Hours: 1})
|
||||
got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{})
|
||||
if err != nil {
|
||||
t.Fatalf("AuditEvents: %v", err)
|
||||
}
|
||||
@@ -184,3 +193,27 @@ func TestAuditDetailRequeryFallback(t *testing.T) {
|
||||
t.Errorf("未找回 err = %v, want ErrAuditEventGone", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditRawCacheTenantIsolationAndInvalidation(t *testing.T) {
|
||||
svc := newTestService(t, &fakeClient{})
|
||||
raw1 := json.RawMessage(`{"tenant":1}`)
|
||||
raw2 := json.RawMessage(`{"tenant":2}`)
|
||||
svc.stripAuditRaw(1, []oci.AuditEvent{{EventId: "same", Raw: raw1}})
|
||||
svc.stripAuditRaw(2, []oci.AuditEvent{{EventId: "same", Raw: raw2}})
|
||||
|
||||
assertAuditRaw(t, svc, 1, raw1)
|
||||
assertAuditRaw(t, svc, 2, raw2)
|
||||
svc.InvalidateAuditCache(1)
|
||||
if _, ok := svc.auditRaw.Get(auditRawKey(1, "same")); ok {
|
||||
t.Fatal("tenant 1 raw cache still exists after invalidation")
|
||||
}
|
||||
assertAuditRaw(t, svc, 2, raw2)
|
||||
}
|
||||
|
||||
func assertAuditRaw(t *testing.T, svc *OciConfigService, configID uint, want json.RawMessage) {
|
||||
t.Helper()
|
||||
got, ok := svc.auditRaw.Get(auditRawKey(configID, "same"))
|
||||
if !ok || string(got.(json.RawMessage)) != string(want) {
|
||||
t.Fatalf("tenant %d raw = %v, want %s", configID, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
+58
-12
@@ -28,6 +28,13 @@ var ErrLoginLocked = errors.New("too many failed attempts, try again later")
|
||||
// tokenTTL 是登录令牌有效期。
|
||||
const tokenTTL = 24 * time.Hour
|
||||
|
||||
// authClaims 在标准声明外携带令牌版本;版本落后于账号当前值即失效。
|
||||
// 存量令牌无 ver 字段解析为 0,与存量账号的零值版本兼容(升级不强制登出)。
|
||||
type authClaims struct {
|
||||
jwt.RegisteredClaims
|
||||
Ver uint `json:"ver"`
|
||||
}
|
||||
|
||||
// dummyBcryptHash 是恒定失败的占位哈希("dummy-password"),
|
||||
// 用户不存在分支比对它以对齐耗时,防用户枚举与时序侧信道。
|
||||
const dummyBcryptHash = "$2a$10$N9qo8uLOickgx2ZMRZoMye3xW1Wq8p1zEIfQpXCXbXyE3xY5C6P6W"
|
||||
@@ -106,7 +113,7 @@ func (s *AuthService) createUser(username, password string) error {
|
||||
// 锁定期内一律 ErrLoginLocked(正确密码同样拒绝);阈值与时长取安全设置。
|
||||
// 已启用两步验证时:密码通过但缺验证码返回 ErrTotpRequired(不计失败),验证码错误计入守卫。
|
||||
func (s *AuthService) Login(ctx context.Context, username, password, clientIP, totpCode string) (string, time.Time, error) {
|
||||
key := clientIP + "|" + username
|
||||
key := guardKey(clientIP, username)
|
||||
now := time.Now()
|
||||
sec := securityOf(s.settings)
|
||||
lockFor := time.Duration(sec.LoginLockMinutes) * time.Minute
|
||||
@@ -138,7 +145,7 @@ func (s *AuthService) Login(ctx context.Context, username, password, clientIP, t
|
||||
}
|
||||
}
|
||||
s.guard.success(key)
|
||||
return s.signToken(user.Username)
|
||||
return s.signToken(user.Username, user.TokenVersion)
|
||||
}
|
||||
|
||||
// failLogin 记失败;达到阈值转锁定并推送告警(开关 login_lock,缺省开)。
|
||||
@@ -169,15 +176,18 @@ func (s *AuthService) notifyLock(username, clientIP string, sec SecuritySettings
|
||||
})
|
||||
}
|
||||
|
||||
func (s *AuthService) signToken(username string) (string, time.Time, error) {
|
||||
func (s *AuthService) signToken(username string, ver uint) (string, time.Time, error) {
|
||||
now := time.Now()
|
||||
expires := now.Add(tokenTTL)
|
||||
claims := jwt.RegisteredClaims{
|
||||
Subject: username,
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
ExpiresAt: jwt.NewNumericDate(expires),
|
||||
// jti:同一秒签发的令牌若无唯一 ID 字节全同,登出一个会连坐全部
|
||||
ID: newTokenID(),
|
||||
claims := authClaims{
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: username,
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
ExpiresAt: jwt.NewNumericDate(expires),
|
||||
// jti:同一秒签发的令牌若无唯一 ID 字节全同,登出一个会连坐全部
|
||||
ID: newTokenID(),
|
||||
},
|
||||
Ver: ver,
|
||||
}
|
||||
token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(s.jwtSecret)
|
||||
if err != nil {
|
||||
@@ -186,6 +196,34 @@ func (s *AuthService) signToken(username string) (string, time.Time, error) {
|
||||
return token, expires, nil
|
||||
}
|
||||
|
||||
// IssueToken 按账号当前令牌版本签发新 JWT;敏感操作递增版本后用它为
|
||||
// 操作者重签,避免操作者自身会话中断。
|
||||
func (s *AuthService) IssueToken(ctx context.Context, username string) (string, time.Time, error) {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
return s.signToken(user.Username, user.TokenVersion)
|
||||
}
|
||||
|
||||
// bumpTokenVersion 原子递增账号令牌版本,使所有已签发令牌立即失效。
|
||||
func (s *AuthService) bumpTokenVersion(ctx context.Context, username string) error {
|
||||
err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", username).
|
||||
UpdateColumn("token_version", gorm.Expr("token_version + 1")).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("bump token version: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeSessions 撤销账号全部会话(版本递增),并为操作者重签新令牌。
|
||||
func (s *AuthService) RevokeSessions(ctx context.Context, username string) (string, time.Time, error) {
|
||||
if err := s.bumpTokenVersion(ctx, username); err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
return s.IssueToken(ctx, username)
|
||||
}
|
||||
|
||||
// newTokenID 生成 128 位随机令牌 ID(crypto/rand 自 Go 1.24 起不会失败)。
|
||||
func newTokenID() string {
|
||||
b := make([]byte, 16)
|
||||
@@ -193,9 +231,10 @@ func newTokenID() string {
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
||||
// ParseToken 验证 JWT 签名与有效期,返回其中的用户名。
|
||||
func (s *AuthService) ParseToken(tokenString string) (string, error) {
|
||||
claims := &jwt.RegisteredClaims{}
|
||||
// ParseToken 验证 JWT 签名、有效期与令牌版本,返回其中的用户名。
|
||||
// 版本落后于账号当前值(凭据等已变更)按无效处理,不区分具体原因。
|
||||
func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (string, error) {
|
||||
claims := &authClaims{}
|
||||
_, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("unexpected signing method %v", t.Header["alg"])
|
||||
@@ -208,6 +247,13 @@ func (s *AuthService) ParseToken(tokenString string) (string, error) {
|
||||
if _, hit := s.revoked.Get(tokenHash(tokenString)); hit {
|
||||
return "", errors.New("token revoked")
|
||||
}
|
||||
user, err := s.findUser(ctx, claims.Subject)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("token subject: %w", err)
|
||||
}
|
||||
if claims.Ver != user.TokenVersion {
|
||||
return "", errors.New("token version outdated")
|
||||
}
|
||||
return claims.Subject, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -109,7 +109,7 @@ func TestTokenRoundTrip(t *testing.T) {
|
||||
if expires.IsZero() {
|
||||
t.Error("expires is zero, want future time")
|
||||
}
|
||||
username, err := auth.ParseToken(token)
|
||||
username, err := auth.ParseToken(context.Background(), token)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseToken: %v", err)
|
||||
}
|
||||
@@ -129,35 +129,92 @@ func TestParseTokenRejectsForged(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(forged); err == nil {
|
||||
if _, err := auth.ParseToken(context.Background(), forged); err == nil {
|
||||
t.Error("ParseToken(forged): got nil error, want failure")
|
||||
}
|
||||
if _, err := auth.ParseToken("not.a.token"); err == nil {
|
||||
if _, err := auth.ParseToken(context.Background(), "not.a.token"); err == nil {
|
||||
t.Error("ParseToken(garbage): got nil error, want failure")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogoutRevokesToken(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
token, _, err := auth.signToken("admin")
|
||||
// ParseToken 现校验令牌版本,须存在对应账号
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
token, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(token); err != nil {
|
||||
if _, err := auth.ParseToken(context.Background(), token); err != nil {
|
||||
t.Fatalf("ParseToken before logout: %v", err)
|
||||
}
|
||||
auth.Logout(token)
|
||||
if _, err := auth.ParseToken(token); err == nil {
|
||||
if _, err := auth.ParseToken(context.Background(), token); err == nil {
|
||||
t.Error("ParseToken after logout: got nil error, want revoked")
|
||||
}
|
||||
// 幂等:重复登出与无效令牌登出都不应 panic,也不影响其他令牌
|
||||
auth.Logout(token)
|
||||
auth.Logout("not.a.token")
|
||||
fresh, _, err := auth.signToken("admin")
|
||||
fresh, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken fresh: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(fresh); err != nil {
|
||||
if _, err := auth.ParseToken(context.Background(), fresh); err != nil {
|
||||
t.Errorf("ParseToken(fresh) after revoking old: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenVersionInvalidatesOldToken 凭据变更递增令牌版本,旧 JWT 立即失效,
|
||||
// 重签的新令牌可用(审计 S-02 动态复现序列的反向断言)。
|
||||
func TestTokenVersionInvalidatesOldToken(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
old, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "")
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
finalName, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
NewPassword: "changed-456", CurrentPassword: "pass123",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateCredentials: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(ctx, old); err == nil {
|
||||
t.Error("旧 token 在凭据变更后仍有效, want 失效")
|
||||
}
|
||||
fresh, _, err := auth.IssueToken(ctx, finalName)
|
||||
if err != nil {
|
||||
t.Fatalf("IssueToken: %v", err)
|
||||
}
|
||||
if name, err := auth.ParseToken(ctx, fresh); err != nil || name != "admin" {
|
||||
t.Errorf("重签 token 应有效: name=%q err=%v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestRevokeSessions 撤销全部会话:旧 token 失效,返回的新 token 有效。
|
||||
func TestRevokeSessions(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
old, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "")
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
fresh, _, err := auth.RevokeSessions(ctx, "admin")
|
||||
if err != nil {
|
||||
t.Fatalf("RevokeSessions: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(ctx, old); err == nil {
|
||||
t.Error("撤销后旧 token 仍有效, want 失效")
|
||||
}
|
||||
if name, err := auth.ParseToken(ctx, fresh); err != nil || name != "admin" {
|
||||
t.Errorf("撤销后新 token 应有效: name=%q err=%v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
@@ -33,42 +34,55 @@ type UpdateCredentialsInput struct {
|
||||
CurrentPassword string `json:"currentPassword" binding:"required"`
|
||||
}
|
||||
|
||||
// UpdateCredentials 修改用户名 / 密码:当前密码必验;改名后旧 JWT 的
|
||||
// sub 不再命中账号,前端应强制重新登录。
|
||||
func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput) error {
|
||||
// UpdateCredentials 修改用户名 / 密码:当前密码必验;成功后令牌版本递增
|
||||
// (全部旧 JWT 立即失效),返回最终用户名供调用方为操作者重签新令牌。
|
||||
func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput) (string, error) {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
return "", err
|
||||
}
|
||||
if bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(in.CurrentPassword)) != nil {
|
||||
return ErrCredentialConfirm
|
||||
return "", ErrCredentialConfirm
|
||||
}
|
||||
newName := strings.TrimSpace(in.NewUsername)
|
||||
if err := validateCredentialChange(user, newName, in.NewPassword); err != nil {
|
||||
return err
|
||||
return "", err
|
||||
}
|
||||
updates, finalName, err := s.credentialUpdates(ctx, user, newName, in.NewPassword)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// 同一条 UPDATE 里递增令牌版本,与凭据变更保持原子
|
||||
updates["token_version"] = gorm.Expr("token_version + 1")
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", user.ID).Updates(updates).Error; err != nil {
|
||||
return "", fmt.Errorf("update credentials: %w", err)
|
||||
}
|
||||
return finalName, nil
|
||||
}
|
||||
|
||||
// credentialUpdates 组装凭据变更字段并返回最终用户名。
|
||||
func (s *AuthService) credentialUpdates(ctx context.Context, user *model.User, newName, newPassword string) (map[string]any, string, error) {
|
||||
updates := map[string]any{}
|
||||
finalName := user.Username
|
||||
if newName != "" && newName != user.Username {
|
||||
taken, err := s.usernameTaken(ctx, newName, user.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, "", err
|
||||
}
|
||||
if taken {
|
||||
return fmt.Errorf("用户名已被占用: %w", ErrCredentialInvalid)
|
||||
return nil, "", fmt.Errorf("用户名已被占用: %w", ErrCredentialInvalid)
|
||||
}
|
||||
updates["username"] = newName
|
||||
finalName = newName
|
||||
}
|
||||
if in.NewPassword != "" {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(in.NewPassword), bcrypt.DefaultCost)
|
||||
if newPassword != "" {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return fmt.Errorf("hash password: %w", err)
|
||||
return nil, "", fmt.Errorf("hash password: %w", err)
|
||||
}
|
||||
updates["password_hash"] = string(hash)
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", user.ID).Updates(updates).Error; err != nil {
|
||||
return fmt.Errorf("update credentials: %w", err)
|
||||
}
|
||||
return nil
|
||||
return updates, finalName, nil
|
||||
}
|
||||
|
||||
// validateCredentialChange 校验改名 / 改密输入;两者均无实际变更时报非法。
|
||||
@@ -122,7 +136,11 @@ func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username str
|
||||
return ErrNeedIdentity
|
||||
}
|
||||
}
|
||||
return s.settings.SetPasswordLoginDisabled(ctx, disabled)
|
||||
if err := s.settings.SetPasswordLoginDisabled(ctx, disabled); err != nil {
|
||||
return err
|
||||
}
|
||||
// 登录策略属敏感变更:递增令牌版本,已签发会话全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
}
|
||||
|
||||
func (s *AuthService) identityCount(ctx context.Context, userID uint) (int64, error) {
|
||||
|
||||
@@ -26,22 +26,22 @@ func TestUpdateCredentials(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// 当前密码错误
|
||||
err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "newpass-123", CurrentPassword: "wrong"})
|
||||
_, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "newpass-123", CurrentPassword: "wrong"})
|
||||
if !errors.Is(err, ErrCredentialConfirm) {
|
||||
t.Fatalf("wrong current password: err = %v, want ErrCredentialConfirm", err)
|
||||
}
|
||||
// 短密码 / 无变更均拒绝
|
||||
err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "short", CurrentPassword: "pass123"})
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "short", CurrentPassword: "pass123"})
|
||||
if !errors.Is(err, ErrCredentialInvalid) {
|
||||
t.Fatalf("short password: err = %v, want ErrCredentialInvalid", err)
|
||||
}
|
||||
err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewUsername: "admin", CurrentPassword: "pass123"})
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewUsername: "admin", CurrentPassword: "pass123"})
|
||||
if !errors.Is(err, ErrCredentialInvalid) {
|
||||
t.Fatalf("no-op change: err = %v, want ErrCredentialInvalid", err)
|
||||
}
|
||||
|
||||
// 同时改名改密
|
||||
err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
NewUsername: "root", NewPassword: "newpass-123", CurrentPassword: "pass123",
|
||||
})
|
||||
if err != nil {
|
||||
|
||||
@@ -9,17 +9,17 @@ import (
|
||||
)
|
||||
|
||||
// IdentityProviders 列出域内 SAML 身份提供者。
|
||||
func (s *OciConfigService) IdentityProviders(ctx context.Context, id uint) ([]oci.IdentityProviderInfo, error) {
|
||||
func (s *OciConfigService) IdentityProviders(ctx context.Context, id uint, domainID string) ([]oci.IdentityProviderInfo, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.ListIdentityProviders(ctx, cred, homeRegion)
|
||||
return s.client.ListIdentityProviders(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// CreateIdentityProvider 创建禁用态 SAML IdP,激活另走 activate。
|
||||
// 映射与 JIT 字段缺省时填控制台默认值(normalizeIdpInput)。
|
||||
func (s *OciConfigService) CreateIdentityProvider(ctx context.Context, id uint, in oci.CreateIdpInput) (oci.IdentityProviderInfo, error) {
|
||||
func (s *OciConfigService) CreateIdentityProvider(ctx context.Context, id uint, domainID string, in oci.CreateIdpInput) (oci.IdentityProviderInfo, error) {
|
||||
if strings.TrimSpace(in.Name) == "" {
|
||||
return oci.IdentityProviderInfo{}, fmt.Errorf("create identity provider: name is required")
|
||||
}
|
||||
@@ -34,7 +34,7 @@ func (s *OciConfigService) CreateIdentityProvider(ctx context.Context, id uint,
|
||||
if err != nil {
|
||||
return oci.IdentityProviderInfo{}, err
|
||||
}
|
||||
return s.client.CreateSamlIdentityProvider(ctx, cred, homeRegion, normalizeIdpInput(in))
|
||||
return s.client.CreateSamlIdentityProvider(ctx, cred, homeRegion, domainID, normalizeIdpInput(in))
|
||||
}
|
||||
|
||||
// normalizeIdpInput 填充映射字段的控制台默认值:名称 ID 格式「无」、
|
||||
@@ -50,30 +50,30 @@ func normalizeIdpInput(in oci.CreateIdpInput) oci.CreateIdpInput {
|
||||
}
|
||||
|
||||
// SetIdentityProviderEnabled 激活/停用 IdP 并同步登录页显示。
|
||||
func (s *OciConfigService) SetIdentityProviderEnabled(ctx context.Context, id uint, idpID string, enabled bool) (oci.IdentityProviderInfo, error) {
|
||||
func (s *OciConfigService) SetIdentityProviderEnabled(ctx context.Context, id uint, domainID, idpID string, enabled bool) (oci.IdentityProviderInfo, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return oci.IdentityProviderInfo{}, err
|
||||
}
|
||||
return s.client.SetIdentityProviderEnabled(ctx, cred, homeRegion, idpID, enabled)
|
||||
return s.client.SetIdentityProviderEnabled(ctx, cred, homeRegion, domainID, idpID, enabled)
|
||||
}
|
||||
|
||||
// DeleteIdentityProvider 级联删除 IdP:先删关联的免 MFA 规则,再从登录页
|
||||
// 移除、停用、删除 IdP 本体。
|
||||
func (s *OciConfigService) DeleteIdentityProvider(ctx context.Context, id uint, idpID string) error {
|
||||
func (s *OciConfigService) DeleteIdentityProvider(ctx context.Context, id uint, domainID, idpID string) error {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.deleteIdpExemptions(ctx, cred, homeRegion, idpID); err != nil {
|
||||
if err := s.deleteIdpExemptions(ctx, cred, homeRegion, domainID, idpID); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.client.DeleteIdentityProvider(ctx, cred, homeRegion, idpID)
|
||||
return s.client.DeleteIdentityProvider(ctx, cred, homeRegion, domainID, idpID)
|
||||
}
|
||||
|
||||
// deleteIdpExemptions 删除 sign-on 策略中引用目标 IdP 的免 MFA 规则。
|
||||
func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Credentials, region, idpID string) error {
|
||||
rules, err := s.client.ListConsoleSignOnRules(ctx, cred, region)
|
||||
func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Credentials, region, domainID, idpID string) error {
|
||||
rules, err := s.client.ListConsoleSignOnRules(ctx, cred, region, domainID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("list sign-on rules before delete idp: %w", err)
|
||||
}
|
||||
@@ -81,7 +81,7 @@ func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Cre
|
||||
if r.BuiltIn || !strings.Contains(r.ConditionValue, idpID) {
|
||||
continue
|
||||
}
|
||||
if err := s.client.DeleteMfaExemptionRule(ctx, cred, region, r.ID); err != nil {
|
||||
if err := s.client.DeleteMfaExemptionRule(ctx, cred, region, domainID, r.ID); err != nil {
|
||||
return fmt.Errorf("delete exemption rule %s: %w", r.Name, err)
|
||||
}
|
||||
}
|
||||
@@ -89,25 +89,25 @@ func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Cre
|
||||
}
|
||||
|
||||
// DomainSamlMetadata 下载域的 SP SAML 元数据 XML(提供给 IdP 侧配置)。
|
||||
func (s *OciConfigService) DomainSamlMetadata(ctx context.Context, id uint) ([]byte, error) {
|
||||
func (s *OciConfigService) DomainSamlMetadata(ctx context.Context, id uint, domainID string) ([]byte, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.DownloadDomainSamlMetadata(ctx, cred, homeRegion)
|
||||
return s.client.DownloadDomainSamlMetadata(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// ConsoleSignOnRules 按优先级列出 OCI Console sign-on 策略的规则。
|
||||
func (s *OciConfigService) ConsoleSignOnRules(ctx context.Context, id uint) ([]oci.SignOnRuleInfo, error) {
|
||||
func (s *OciConfigService) ConsoleSignOnRules(ctx context.Context, id uint, domainID string) ([]oci.SignOnRuleInfo, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.ListConsoleSignOnRules(ctx, cred, homeRegion)
|
||||
return s.client.ListConsoleSignOnRules(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// CreateMfaExemption 为指定 IdP 创建免 MFA 规则并置为最高优先级。
|
||||
func (s *OciConfigService) CreateMfaExemption(ctx context.Context, id uint, idpID string) (oci.SignOnRuleInfo, error) {
|
||||
func (s *OciConfigService) CreateMfaExemption(ctx context.Context, id uint, domainID, idpID string) (oci.SignOnRuleInfo, error) {
|
||||
if strings.TrimSpace(idpID) == "" {
|
||||
return oci.SignOnRuleInfo{}, fmt.Errorf("create mfa exemption: identityProviderId is required")
|
||||
}
|
||||
@@ -115,16 +115,16 @@ func (s *OciConfigService) CreateMfaExemption(ctx context.Context, id uint, idpI
|
||||
if err != nil {
|
||||
return oci.SignOnRuleInfo{}, err
|
||||
}
|
||||
name, err := s.exemptionRuleName(ctx, cred, homeRegion, idpID)
|
||||
name, err := s.exemptionRuleName(ctx, cred, homeRegion, domainID, idpID)
|
||||
if err != nil {
|
||||
return oci.SignOnRuleInfo{}, err
|
||||
}
|
||||
return s.client.CreateMfaExemptionRule(ctx, cred, homeRegion, idpID, name)
|
||||
return s.client.CreateMfaExemptionRule(ctx, cred, homeRegion, domainID, idpID, name)
|
||||
}
|
||||
|
||||
// exemptionRuleName 校验 IdP 存在并用其名称生成规则名。
|
||||
func (s *OciConfigService) exemptionRuleName(ctx context.Context, cred oci.Credentials, region, idpID string) (string, error) {
|
||||
idps, err := s.client.ListIdentityProviders(ctx, cred, region)
|
||||
func (s *OciConfigService) exemptionRuleName(ctx context.Context, cred oci.Credentials, region, domainID, idpID string) (string, error) {
|
||||
idps, err := s.client.ListIdentityProviders(ctx, cred, region, domainID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -137,10 +137,10 @@ func (s *OciConfigService) exemptionRuleName(ctx context.Context, cred oci.Crede
|
||||
}
|
||||
|
||||
// DeleteMfaExemption 删除免 MFA 规则并恢复其余规则优先级。
|
||||
func (s *OciConfigService) DeleteMfaExemption(ctx context.Context, id uint, ruleID string) error {
|
||||
func (s *OciConfigService) DeleteMfaExemption(ctx context.Context, id uint, domainID, ruleID string) error {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.client.DeleteMfaExemptionRule(ctx, cred, homeRegion, ruleID)
|
||||
return s.client.DeleteMfaExemptionRule(ctx, cred, homeRegion, domainID, ruleID)
|
||||
}
|
||||
|
||||
+133
-40
@@ -57,6 +57,9 @@ type LogEventService struct {
|
||||
|
||||
relayPollTick time.Duration // 订阅确认轮询间隔,零值用默认
|
||||
relayPollTimeout time.Duration // 订阅确认轮询上限,零值用默认
|
||||
|
||||
alertMu sync.Mutex // 保护告警规则冷却表
|
||||
alertSentAt map[uint]time.Time // 规则 ID → 上次告警时刻(阈值型规则冷却)
|
||||
}
|
||||
|
||||
// NewLogEventService 组装依赖;调用 StartParser / StartCleanup 后台协程后生效。
|
||||
@@ -80,37 +83,62 @@ type LogWebhookInfo struct {
|
||||
|
||||
// EnsureSecret 为配置生成(或幂等返回)回传 secret。
|
||||
func (s *LogEventService) EnsureSecret(ctx context.Context, cfgID uint) (LogWebhookInfo, error) {
|
||||
if err := s.requireConfig(ctx, cfgID); err != nil {
|
||||
secret, err := generateWebhookSecret()
|
||||
if err != nil {
|
||||
return LogWebhookInfo{}, err
|
||||
}
|
||||
if info, ok, err := s.SecretInfo(ctx, cfgID); err != nil || ok {
|
||||
return info, err
|
||||
var info LogWebhookInfo
|
||||
err = s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var err error
|
||||
info, err = ensureSecretTx(tx, cfgID, secret)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return LogWebhookInfo{}, fmt.Errorf("ensure webhook secret: %w", err)
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func generateWebhookSecret() (string, error) {
|
||||
buf := make([]byte, 32)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return LogWebhookInfo{}, fmt.Errorf("generate webhook secret: %w", err)
|
||||
return "", fmt.Errorf("generate webhook secret: %w", err)
|
||||
}
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
func ensureSecretTx(tx *gorm.DB, cfgID uint, secret string) (LogWebhookInfo, error) {
|
||||
if err := lockOciConfig(tx, cfgID); err != nil {
|
||||
return LogWebhookInfo{}, err
|
||||
}
|
||||
if info, ok, err := secretInfoTx(tx, cfgID); err != nil || ok {
|
||||
return info, err
|
||||
}
|
||||
secret := hex.EncodeToString(buf)
|
||||
st := model.Setting{Key: secretKey(cfgID), Value: secret, UpdatedAt: time.Now()}
|
||||
if err := s.db.WithContext(ctx).Save(&st).Error; err != nil {
|
||||
if err := tx.Create(&st).Error; err != nil {
|
||||
return LogWebhookInfo{}, fmt.Errorf("save webhook secret: %w", err)
|
||||
}
|
||||
return webhookInfo(secret, st.UpdatedAt), nil
|
||||
}
|
||||
|
||||
// requireConfig 校验配置存在,不存在透传 gorm.ErrRecordNotFound(api 层映射 404)。
|
||||
func (s *LogEventService) requireConfig(ctx context.Context, cfgID uint) error {
|
||||
var cfg model.OciConfig
|
||||
if err := s.db.WithContext(ctx).Select("id").First(&cfg, cfgID).Error; err != nil {
|
||||
return fmt.Errorf("find oci config %d: %w", cfgID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SecretInfo 查询配置是否已生成 secret;未生成时 ok 为 false。
|
||||
func (s *LogEventService) SecretInfo(ctx context.Context, cfgID uint) (LogWebhookInfo, bool, error) {
|
||||
var info LogWebhookInfo
|
||||
var ok bool
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := lockOciConfig(tx, cfgID); err != nil {
|
||||
return err
|
||||
}
|
||||
var err error
|
||||
info, ok, err = secretInfoTx(tx, cfgID)
|
||||
return err
|
||||
})
|
||||
return info, ok, err
|
||||
}
|
||||
|
||||
func secretInfoTx(tx *gorm.DB, cfgID uint) (LogWebhookInfo, bool, error) {
|
||||
var st model.Setting
|
||||
err := s.db.WithContext(ctx).First(&st, "key = ?", secretKey(cfgID)).Error
|
||||
err := tx.First(&st, "key = ?", secretKey(cfgID)).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return LogWebhookInfo{}, false, nil
|
||||
@@ -154,19 +182,61 @@ func (s *LogEventService) ResolveSecret(ctx context.Context, secret string) (uin
|
||||
return 0, false
|
||||
}
|
||||
for _, row := range rows {
|
||||
if subtle.ConstantTimeCompare([]byte(row.Value), []byte(secret)) == 1 {
|
||||
id, err := strconv.ParseUint(strings.TrimPrefix(row.Key, logWebhookSecretPrefix), 10, 64)
|
||||
if err != nil {
|
||||
return 0, false
|
||||
}
|
||||
return uint(id), true
|
||||
if subtle.ConstantTimeCompare([]byte(row.Value), []byte(secret)) != 1 {
|
||||
continue
|
||||
}
|
||||
id, ok := webhookConfigID(row.Key)
|
||||
if !ok || !s.configExists(ctx, id) {
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
// webhookConfigID 从 secret Setting 键解析租户配置 ID。
|
||||
func webhookConfigID(key string) (uint, bool) {
|
||||
id, err := strconv.ParseUint(strings.TrimPrefix(key, logWebhookSecretPrefix), 10, 64)
|
||||
return uint(id), err == nil && id > 0
|
||||
}
|
||||
|
||||
// configExists 确认 secret 对应租户仍存在;数据库异常按认证失败处理。
|
||||
func (s *LogEventService) configExists(ctx context.Context, id uint) bool {
|
||||
var cfg model.OciConfig
|
||||
err := s.db.WithContext(ctx).Select("id").First(&cfg, id).Error
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
log.Printf("resolve webhook config: %v", err)
|
||||
}
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// Ingest 落一条回传事件;MessageID 唯一索引冲突即静默忽略(at-least-once 幂等)。
|
||||
func (s *LogEventService) Ingest(ctx context.Context, cfgID uint, messageID string, payload []byte, truncated bool) error {
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := lockOciConfig(tx, cfgID); err != nil {
|
||||
return err
|
||||
}
|
||||
return createLogEvent(tx, cfgID, messageID, payload, truncated)
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("ingest log event: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// lockOciConfig 与租户删除共用行锁,避免并发清理后写入孤儿数据。
|
||||
func lockOciConfig(tx *gorm.DB, cfgID uint) error {
|
||||
var cfg model.OciConfig
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Select("id").First(&cfg, cfgID).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("lock oci config %d: %w", cfgID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// createLogEvent 在已锁定租户的事务内幂等写入事件。
|
||||
func createLogEvent(tx *gorm.DB, cfgID uint, messageID string, payload []byte, truncated bool) error {
|
||||
event := model.LogEvent{
|
||||
OciConfigID: cfgID,
|
||||
MessageID: messageID,
|
||||
@@ -174,13 +244,9 @@ func (s *LogEventService) Ingest(ctx context.Context, cfgID uint, messageID stri
|
||||
Truncated: truncated,
|
||||
ReceivedAt: time.Now(),
|
||||
}
|
||||
err := s.db.WithContext(ctx).
|
||||
Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "message_id"}}, DoNothing: true}).
|
||||
Create(&event).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("ingest log event: %w", err)
|
||||
}
|
||||
return nil
|
||||
return tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "message_id"}}, DoNothing: true,
|
||||
}).Create(&event).Error
|
||||
}
|
||||
|
||||
// LogEventQuery 是回传事件查询参数;CfgID 为 0 表示全部租户。
|
||||
@@ -304,18 +370,44 @@ func (s *LogEventService) parseOnce(ctx context.Context) {
|
||||
log.Printf("log event parse load: %v", err)
|
||||
return
|
||||
}
|
||||
for i := range events {
|
||||
e := &events[i]
|
||||
parsed := parseLogEvent([]byte(e.Payload))
|
||||
e.EventType, e.Source, e.SourceIP, e.EventTime =
|
||||
parsed.EventType, parsed.Source, parsed.SourceIP, parsed.EventTime
|
||||
e.Processed = true
|
||||
if err := s.db.WithContext(ctx).Save(e).Error; err != nil {
|
||||
log.Printf("log event parse save %d: %v", e.ID, err)
|
||||
return
|
||||
}
|
||||
s.notifyCritical(ctx, e, parsed)
|
||||
if len(events) == 0 {
|
||||
return
|
||||
}
|
||||
rules := s.loadEnabledAlertRules(ctx)
|
||||
for i := range events {
|
||||
s.processLogEvent(ctx, &events[i], rules)
|
||||
}
|
||||
}
|
||||
|
||||
// processLogEvent 仅在条件更新命中原行后触发通知,删除并发胜出时静默跳过。
|
||||
func (s *LogEventService) processLogEvent(ctx context.Context, event *model.LogEvent, rules []model.AlertRule) {
|
||||
parsed := parseLogEvent([]byte(event.Payload))
|
||||
if !s.updateParsedEvent(ctx, event, parsed) {
|
||||
return
|
||||
}
|
||||
s.notifyCritical(ctx, event, parsed)
|
||||
s.matchAlertRules(ctx, rules, event, parsed)
|
||||
}
|
||||
|
||||
// updateParsedEvent 用条件 UPDATE 禁止 Save 在删除后隐式重建事件。
|
||||
func (s *LogEventService) updateParsedEvent(ctx context.Context, event *model.LogEvent, parsed parsedEvent) bool {
|
||||
updates := map[string]any{
|
||||
"event_type": parsed.EventType, "source": parsed.Source,
|
||||
"source_ip": parsed.SourceIP, "event_time": parsed.EventTime, "processed": true,
|
||||
}
|
||||
res := s.db.WithContext(ctx).Model(&model.LogEvent{}).
|
||||
Where("id = ? AND processed = ?", event.ID, false).Updates(updates)
|
||||
if res.Error != nil {
|
||||
log.Printf("log event parse update %d: %v", event.ID, res.Error)
|
||||
return false
|
||||
}
|
||||
if res.RowsAffected != 1 {
|
||||
return false
|
||||
}
|
||||
event.EventType, event.Source, event.SourceIP, event.EventTime =
|
||||
parsed.EventType, parsed.Source, parsed.SourceIP, parsed.EventTime
|
||||
event.Processed = true
|
||||
return true
|
||||
}
|
||||
|
||||
// onsEnvelope 覆盖 ONS 消息与 CloudEvents 审计事件的常见字段;
|
||||
@@ -443,6 +535,7 @@ func (s *LogEventService) cleanupOnce(ctx context.Context) {
|
||||
if err := s.cleanup(ctx, logEventRetention, logEventMaxRows); err != nil {
|
||||
log.Printf("log event cleanup: %v", err)
|
||||
}
|
||||
s.cleanupAlertHits(ctx)
|
||||
}
|
||||
|
||||
// cleanup 先删过期记录,再对超量部分删最旧;阈值参数化便于测试。
|
||||
|
||||
@@ -2,7 +2,9 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -28,7 +30,8 @@ func newLogEventEnv(t *testing.T) (*LogEventService, *gorm.DB, uint) {
|
||||
t.Fatalf("db handle: %v", err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}, &model.LogEvent{}); err != nil {
|
||||
if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}, &model.LogEvent{},
|
||||
&model.AlertRule{}, &model.AlertRuleHit{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
cfg := model.OciConfig{Alias: "测试租户"}
|
||||
@@ -98,9 +101,103 @@ func TestEnsureAndResolveSecret(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestEnsureSecretRejectsUnknownConfig(t *testing.T) {
|
||||
svc, _, _ := newLogEventEnv(t)
|
||||
if _, err := svc.EnsureSecret(context.Background(), 9999); err == nil {
|
||||
t.Fatal("ensure secret for unknown config succeeded, want error")
|
||||
svc, db, _ := newLogEventEnv(t)
|
||||
if _, err := svc.EnsureSecret(context.Background(), 9999); !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("ensure unknown config error = %v, want record not found", err)
|
||||
}
|
||||
var count int64
|
||||
db.Model(&model.Setting{}).Where("key = ?", secretKey(9999)).Count(&count)
|
||||
if count != 0 {
|
||||
t.Fatalf("unknown config secret rows = %d, want 0", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureSecretConcurrentDeleteNoOrphan(t *testing.T) {
|
||||
svc, db := newConcurrentSecretEnv(t)
|
||||
for i := 0; i < 8; i++ {
|
||||
cfg := model.OciConfig{Alias: fmt.Sprintf("race-%d", i)}
|
||||
if err := db.Create(&cfg).Error; err != nil {
|
||||
t.Fatalf("create config: %v", err)
|
||||
}
|
||||
ensureErr, deleteErr := runEnsureDeleteRace(svc, db, cfg.ID)
|
||||
if deleteErr != nil {
|
||||
t.Fatalf("delete config %d: %v", cfg.ID, deleteErr)
|
||||
}
|
||||
if ensureErr != nil && !errors.Is(ensureErr, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("ensure config %d: %v", cfg.ID, ensureErr)
|
||||
}
|
||||
assertNoSecretSetting(t, db, cfg.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func newConcurrentSecretEnv(t *testing.T) (*LogEventService, *gorm.DB) {
|
||||
t.Helper()
|
||||
dsn := filepath.Join(t.TempDir(), "secret.db") + "?_txlock=immediate&_pragma=busy_timeout%3d5000"
|
||||
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
t.Fatalf("db handle: %v", err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(4)
|
||||
t.Cleanup(func() { _ = sqlDB.Close() })
|
||||
if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
return NewLogEventService(db), db
|
||||
}
|
||||
|
||||
func runEnsureDeleteRace(svc *LogEventService, db *gorm.DB, cfgID uint) (error, error) {
|
||||
start := make(chan struct{})
|
||||
ensureDone, deleteDone := make(chan error, 1), make(chan error, 1)
|
||||
go func() {
|
||||
<-start
|
||||
_, err := svc.EnsureSecret(context.Background(), cfgID)
|
||||
ensureDone <- err
|
||||
}()
|
||||
go func() {
|
||||
<-start
|
||||
deleteDone <- deleteConfigAndSecret(db, cfgID)
|
||||
}()
|
||||
close(start)
|
||||
return <-ensureDone, <-deleteDone
|
||||
}
|
||||
|
||||
func deleteConfigAndSecret(db *gorm.DB, cfgID uint) error {
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := lockOciConfig(tx, cfgID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Delete(&model.Setting{}, "key = ?", secretKey(cfgID)).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.OciConfig{}, cfgID).Error
|
||||
})
|
||||
}
|
||||
|
||||
func assertNoSecretSetting(t *testing.T, db *gorm.DB, cfgID uint) {
|
||||
t.Helper()
|
||||
var count int64
|
||||
db.Model(&model.Setting{}).Where("key = ?", secretKey(cfgID)).Count(&count)
|
||||
if count != 0 {
|
||||
t.Fatalf("orphan secret rows for config %d = %d", cfgID, count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveSecretRejectsOrphanSetting(t *testing.T) {
|
||||
svc, db, _ := newLogEventEnv(t)
|
||||
orphan := model.Setting{Key: secretKey(9999), Value: "orphan-secret"}
|
||||
if err := db.Create(&orphan).Error; err != nil {
|
||||
t.Fatalf("create orphan secret: %v", err)
|
||||
}
|
||||
if id, ok := svc.ResolveSecret(context.Background(), orphan.Value); ok || id != 0 {
|
||||
t.Fatalf("resolve orphan secret = (%d,%v), want (0,false)", id, ok)
|
||||
}
|
||||
info, ok, err := svc.SecretInfo(context.Background(), 9999)
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) || ok || info.Secret != "" {
|
||||
t.Fatalf("orphan secret info = (%q,%v,%v), want empty,false,record not found", info.Secret, ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -127,6 +224,19 @@ func TestIngestIdempotent(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestIngestRejectsUnknownConfig(t *testing.T) {
|
||||
svc, db, _ := newLogEventEnv(t)
|
||||
err := svc.Ingest(context.Background(), 9999, "orphan", []byte(`{}`), false)
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("ingest unknown config error = %v, want record not found", err)
|
||||
}
|
||||
var count int64
|
||||
db.Model(&model.LogEvent{}).Where("message_id = ?", "orphan").Count(&count)
|
||||
if count != 0 {
|
||||
t.Fatalf("orphan event rows = %d, want 0", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseLogEvent(t *testing.T) {
|
||||
ts := "2026-07-07T08:00:00Z"
|
||||
tests := []struct {
|
||||
@@ -209,6 +319,40 @@ func TestParseOnceMarksProcessed(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateParsedEventDoesNotRecreateDeleted(t *testing.T) {
|
||||
svc, db, cfgID := newLogEventEnv(t)
|
||||
event := model.LogEvent{OciConfigID: cfgID, MessageID: "deleted", Payload: `{}`}
|
||||
if err := db.Create(&event).Error; err != nil {
|
||||
t.Fatalf("create event: %v", err)
|
||||
}
|
||||
if err := db.Delete(&model.LogEvent{}, event.ID).Error; err != nil {
|
||||
t.Fatalf("delete event: %v", err)
|
||||
}
|
||||
if svc.updateParsedEvent(context.Background(), &event, parsedEvent{EventType: "x"}) {
|
||||
t.Fatal("deleted event update reported success")
|
||||
}
|
||||
var count int64
|
||||
db.Model(&model.LogEvent{}).Where("id = ?", event.ID).Count(&count)
|
||||
if count != 0 {
|
||||
t.Fatalf("deleted event was recreated, rows = %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateParsedEventOnlyOnce(t *testing.T) {
|
||||
svc, db, cfgID := newLogEventEnv(t)
|
||||
event := model.LogEvent{OciConfigID: cfgID, MessageID: "once", Payload: `{}`}
|
||||
if err := db.Create(&event).Error; err != nil {
|
||||
t.Fatalf("create event: %v", err)
|
||||
}
|
||||
parsed := parsedEvent{EventType: "x"}
|
||||
if !svc.updateParsedEvent(context.Background(), &event, parsed) {
|
||||
t.Fatal("first conditional update should succeed")
|
||||
}
|
||||
if svc.updateParsedEvent(context.Background(), &event, parsed) {
|
||||
t.Fatal("processed event should not update twice")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogEventCleanup(t *testing.T) {
|
||||
svc, db, cfgID := newLogEventEnv(t)
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -1,10 +1,18 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// maxGuardEntries 是 failures / lockedAt 各自的条目上限(S-03):
|
||||
// 攻击者持续提交唯一用户名时先清过期、仍满驱逐最旧,内存有界。
|
||||
const maxGuardEntries = 4096
|
||||
|
||||
// guardUserMax 是计入守卫 key 的用户名字节上限,超长截断防高基数撑爆。
|
||||
const guardUserMax = 32
|
||||
|
||||
// loginGuard 按「IP|用户名」双维度做滑动窗口失败计数与锁定;
|
||||
// 窗口与锁定时长同值、阈值由调用方按安全设置传入(探索文档主题三);
|
||||
// 内存实现(单体面板),重启清零可接受。
|
||||
@@ -21,6 +29,16 @@ func newLoginGuard() *loginGuard {
|
||||
}
|
||||
}
|
||||
|
||||
// guardKey 规范化守卫键:用户名去首尾空白并按字节截断,
|
||||
// 高基数/超长输入不会产生无界的独立条目。
|
||||
func guardKey(clientIP, username string) string {
|
||||
name := strings.TrimSpace(username)
|
||||
if len(name) > guardUserMax {
|
||||
name = name[:guardUserMax]
|
||||
}
|
||||
return clientIP + "|" + name
|
||||
}
|
||||
|
||||
// locked 判定 key 是否处于锁定期;过期锁惰性清除。
|
||||
func (g *loginGuard) locked(key string, now time.Time, lockFor time.Duration) bool {
|
||||
g.mu.Lock()
|
||||
@@ -37,9 +55,13 @@ func (g *loginGuard) locked(key string, now time.Time, lockFor time.Duration) bo
|
||||
}
|
||||
|
||||
// fail 记一次失败并裁剪窗口;达到阈值转入锁定并返回 true(仅触发那一次)。
|
||||
// 新建条目前先保证容量(过期清理 + 最旧驱逐),防高基数 key 单向增长。
|
||||
func (g *loginGuard) fail(key string, now time.Time, limit int, window time.Duration) bool {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
if _, exists := g.failures[key]; !exists {
|
||||
g.ensureCapacity(now, window)
|
||||
}
|
||||
kept := g.failures[key][:0]
|
||||
for _, t := range g.failures[key] {
|
||||
if now.Sub(t) < window {
|
||||
@@ -49,6 +71,9 @@ func (g *loginGuard) fail(key string, now time.Time, limit int, window time.Dura
|
||||
kept = append(kept, now)
|
||||
if len(kept) >= limit {
|
||||
delete(g.failures, key)
|
||||
// 转入锁定同样要保证容量:该写入发生在已有 failures 条目的
|
||||
// 第 N 次失败,不经过新建条目路径的 ensureCapacity
|
||||
g.ensureLockCapacity(now, window)
|
||||
g.lockedAt[key] = now
|
||||
return true
|
||||
}
|
||||
@@ -56,6 +81,77 @@ func (g *loginGuard) fail(key string, now time.Time, limit int, window time.Dura
|
||||
return false
|
||||
}
|
||||
|
||||
// ensureLockCapacity 在写入 lockedAt 前保证容量:先清已过锁定期的锁,
|
||||
// 仍满驱逐锁定最早的条目(其剩余锁定时间最短,提前解锁的影响最小)。
|
||||
func (g *loginGuard) ensureLockCapacity(now time.Time, lockFor time.Duration) {
|
||||
if len(g.lockedAt) < maxGuardEntries {
|
||||
return
|
||||
}
|
||||
for k, at := range g.lockedAt {
|
||||
if now.Sub(at) >= lockFor {
|
||||
delete(g.lockedAt, k)
|
||||
}
|
||||
}
|
||||
if len(g.lockedAt) >= maxGuardEntries {
|
||||
g.evictOldestLock()
|
||||
}
|
||||
}
|
||||
|
||||
// evictOldestLock 驱逐锁定时间最早的 lockedAt 条目(调用方须持锁)。
|
||||
func (g *loginGuard) evictOldestLock() {
|
||||
var oldestKey string
|
||||
var oldestAt time.Time
|
||||
for k, at := range g.lockedAt {
|
||||
if oldestKey == "" || at.Before(oldestAt) {
|
||||
oldestKey, oldestAt = k, at
|
||||
}
|
||||
}
|
||||
if oldestKey != "" {
|
||||
delete(g.lockedAt, oldestKey)
|
||||
}
|
||||
}
|
||||
|
||||
// ensureCapacity 在新建 failures 条目前保证容量:先清出窗口外的失败记录
|
||||
// 与已过锁定期的锁,仍满则驱逐最后失败时间最早的条目(调用方须持锁)。
|
||||
func (g *loginGuard) ensureCapacity(now time.Time, window time.Duration) {
|
||||
if len(g.failures) < maxGuardEntries && len(g.lockedAt) < maxGuardEntries {
|
||||
return
|
||||
}
|
||||
for k, times := range g.failures {
|
||||
if len(times) == 0 || now.Sub(times[len(times)-1]) >= window {
|
||||
delete(g.failures, k)
|
||||
}
|
||||
}
|
||||
for k, at := range g.lockedAt {
|
||||
if now.Sub(at) >= window {
|
||||
delete(g.lockedAt, k)
|
||||
}
|
||||
}
|
||||
if len(g.failures) >= maxGuardEntries {
|
||||
g.evictOldest()
|
||||
}
|
||||
}
|
||||
|
||||
// evictOldest 驱逐最后失败时间最早的 failures 条目(调用方须持锁);
|
||||
// 仅在容量兜底时触发,O(n) 扫描可接受。
|
||||
func (g *loginGuard) evictOldest() {
|
||||
var oldestKey string
|
||||
var oldestAt time.Time
|
||||
for k, times := range g.failures {
|
||||
if len(times) == 0 {
|
||||
delete(g.failures, k)
|
||||
continue
|
||||
}
|
||||
last := times[len(times)-1]
|
||||
if oldestKey == "" || last.Before(oldestAt) {
|
||||
oldestKey, oldestAt = k, last
|
||||
}
|
||||
}
|
||||
if oldestKey != "" {
|
||||
delete(g.failures, oldestKey)
|
||||
}
|
||||
}
|
||||
|
||||
// success 清空 key 的失败计数(成功登录重置窗口)。
|
||||
func (g *loginGuard) success(key string) {
|
||||
g.mu.Lock()
|
||||
|
||||
@@ -0,0 +1,105 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TestGuardKeyNormalization 用户名去空白并按字节截断,高基数超长输入收敛。
|
||||
func TestGuardKeyNormalization(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ip string
|
||||
user string
|
||||
want string
|
||||
}{
|
||||
{name: "常规", ip: "1.2.3.4", user: "admin", want: "1.2.3.4|admin"},
|
||||
{name: "去空白", ip: "1.2.3.4", user: " admin ", want: "1.2.3.4|admin"},
|
||||
{name: "超长截断", ip: "1.2.3.4", user: strings.Repeat("a", 100), want: "1.2.3.4|" + strings.Repeat("a", guardUserMax)},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := guardKey(tt.ip, tt.user); got != tt.want {
|
||||
t.Errorf("guardKey = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardCapacityBounded 高基数唯一 key 持续失败,条目数不超过上限。
|
||||
func TestGuardCapacityBounded(t *testing.T) {
|
||||
g := newLoginGuard()
|
||||
now := time.Now()
|
||||
for i := 0; i < maxGuardEntries*2; i++ {
|
||||
g.fail(fmt.Sprintf("1.2.3.4|user-%d", i), now, 100, time.Hour)
|
||||
}
|
||||
if n := len(g.failures); n > maxGuardEntries {
|
||||
t.Errorf("failures 条目 = %d, want ≤ %d", n, maxGuardEntries)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardCapacitySweepExpired 容量满时优先清理窗口外条目,而非驱逐活跃条目。
|
||||
func TestGuardCapacitySweepExpired(t *testing.T) {
|
||||
g := newLoginGuard()
|
||||
base := time.Now()
|
||||
window := time.Hour
|
||||
// 填满:全部为「早已过窗」的旧失败
|
||||
for i := 0; i < maxGuardEntries; i++ {
|
||||
g.fail(fmt.Sprintf("old|%d", i), base.Add(-2*window), 100, window)
|
||||
}
|
||||
// 新失败触发清理:过期条目被清,新条目正常记录
|
||||
g.fail("fresh|user", base, 100, window)
|
||||
if _, ok := g.failures["fresh|user"]; !ok {
|
||||
t.Error("新条目未记录")
|
||||
}
|
||||
if n := len(g.failures); n > maxGuardEntries {
|
||||
t.Errorf("清理后条目 = %d, want ≤ %d", n, maxGuardEntries)
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardEvictOldestKeepsNewest 无过期可清时驱逐最旧,最新条目保留。
|
||||
func TestGuardEvictOldestKeepsNewest(t *testing.T) {
|
||||
g := newLoginGuard()
|
||||
base := time.Now()
|
||||
window := 24 * time.Hour
|
||||
for i := 0; i < maxGuardEntries; i++ {
|
||||
// 均在窗口内,时间递增;user-0 最旧
|
||||
g.fail(fmt.Sprintf("ip|user-%d", i), base.Add(time.Duration(i)*time.Second), 100, window)
|
||||
}
|
||||
g.fail("ip|newest", base.Add(time.Duration(maxGuardEntries+1)*time.Second), 100, window)
|
||||
if _, ok := g.failures["ip|newest"]; !ok {
|
||||
t.Error("最新条目未记录")
|
||||
}
|
||||
if _, ok := g.failures["ip|user-0"]; ok {
|
||||
t.Error("最旧条目未被驱逐")
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardLockCapacityBounded 高基数唯一 key 全部打满阈值转入锁定,
|
||||
// 锁定期内(无可清理的过期锁)条目数仍不超过上限,最新锁保留、最早锁被驱逐。
|
||||
func TestGuardLockCapacityBounded(t *testing.T) {
|
||||
g := newLoginGuard()
|
||||
base := time.Now()
|
||||
window := 15 * time.Minute
|
||||
total := maxGuardEntries + 100
|
||||
for i := 0; i < total; i++ {
|
||||
key := fmt.Sprintf("1.2.3.4|user-%d", i)
|
||||
// 每个 key 打满 2 次(阈值)即转入锁定;时间递增,i 越小锁定越早
|
||||
at := base.Add(time.Duration(i) * time.Millisecond)
|
||||
g.fail(key, at, 2, window)
|
||||
if !g.fail(key, at, 2, window) {
|
||||
t.Fatalf("key %s 未按阈值转入锁定", key)
|
||||
}
|
||||
}
|
||||
if n := len(g.lockedAt); n > maxGuardEntries {
|
||||
t.Errorf("lockedAt 条目 = %d, want ≤ %d", n, maxGuardEntries)
|
||||
}
|
||||
if _, ok := g.lockedAt[fmt.Sprintf("1.2.3.4|user-%d", total-1)]; !ok {
|
||||
t.Error("最新锁定条目未保留")
|
||||
}
|
||||
if _, ok := g.lockedAt["1.2.3.4|user-0"]; ok {
|
||||
t.Error("最早锁定条目未被驱逐")
|
||||
}
|
||||
}
|
||||
+134
-18
@@ -24,12 +24,12 @@ const notifyTextLimit = 3800
|
||||
// telegramAPIBase 是 Telegram Bot API 官方根地址。
|
||||
const telegramAPIBase = "https://api.telegram.org"
|
||||
|
||||
// Notifier 通过 Telegram Bot API 推送通知;未启用或未配置时静默跳过。
|
||||
// 云外 HTTP 调用,不属于 internal/oci 的职责范围。
|
||||
// Notifier 向 Telegram 与全部启用的通知渠道(webhook/ntfy/bark/smtp)推送;
|
||||
// 未启用或未配置的渠道静默跳过。云外 HTTP 调用,不属于 internal/oci 的职责范围。
|
||||
type Notifier struct {
|
||||
settings *SettingService
|
||||
client *http.Client
|
||||
base string // API 根地址,测试时注入 httptest 服务
|
||||
base string // Telegram API 根地址,测试时注入 httptest 服务
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
@@ -42,8 +42,28 @@ func NewNotifier(settings *SettingService) *Notifier {
|
||||
}
|
||||
}
|
||||
|
||||
// Send 读取配置并发送一条消息;未启用或未配置时直接返回 nil。
|
||||
// Send 向全部启用渠道发送一条纯文本消息;无渠道可用时返回 nil。
|
||||
func (n *Notifier) Send(ctx context.Context, text string) error {
|
||||
return n.dispatch(ctx, notifyMsg{Title: "oci-portal 通知", Text: text})
|
||||
}
|
||||
|
||||
// dispatch 把消息投递到 Telegram 与全部启用渠道;渠道间互不影响,错误合并返回供日志。
|
||||
func (n *Notifier) dispatch(ctx context.Context, msg notifyMsg) error {
|
||||
msg.Text = truncateText(msg.Text, notifyTextLimit)
|
||||
var errs []error
|
||||
if err := n.sendTelegram(ctx, msg); err != nil {
|
||||
errs = append(errs, fmt.Errorf("telegram: %w", err))
|
||||
}
|
||||
for _, chType := range NotifyChannelTypes {
|
||||
if err := n.sendChannel(ctx, chType, msg); err != nil {
|
||||
errs = append(errs, fmt.Errorf("%s: %w", chType, err))
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// sendTelegram 走既有 Telegram 路径;未启用或未配置时静默跳过。
|
||||
func (n *Notifier) sendTelegram(ctx context.Context, msg notifyMsg) error {
|
||||
cfg, err := n.settings.TelegramConfig(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -51,7 +71,38 @@ func (n *Notifier) Send(ctx context.Context, text string) error {
|
||||
if !cfg.Enabled || cfg.BotToken == "" || cfg.ChatID == "" {
|
||||
return nil
|
||||
}
|
||||
return n.post(ctx, cfg, text)
|
||||
return n.postHTML(ctx, cfg, mdToTelegramHTML(msg.Text))
|
||||
}
|
||||
|
||||
// sendChannel 向单个新增渠道投递;未启用或关键字段缺失时静默跳过。
|
||||
func (n *Notifier) sendChannel(ctx context.Context, chType string, msg notifyMsg) error {
|
||||
switch chType {
|
||||
case ChannelWebhook:
|
||||
cfg, err := n.settings.webhookChannel(ctx)
|
||||
if err != nil || !cfg.Enabled || cfg.URL == "" {
|
||||
return err
|
||||
}
|
||||
return sendWebhook(ctx, n.client, cfg, msg)
|
||||
case ChannelNtfy:
|
||||
cfg, err := n.settings.ntfyChannel(ctx)
|
||||
if err != nil || !cfg.Enabled || cfg.Topic == "" {
|
||||
return err
|
||||
}
|
||||
return sendNtfy(ctx, n.client, cfg, msg)
|
||||
case ChannelBark:
|
||||
cfg, err := n.settings.barkChannel(ctx)
|
||||
if err != nil || !cfg.Enabled || cfg.DeviceKey == "" {
|
||||
return err
|
||||
}
|
||||
return sendBark(ctx, n.client, cfg, msg)
|
||||
case ChannelSMTP:
|
||||
cfg, err := n.settings.smtpChannel(ctx)
|
||||
if err != nil || !cfg.Enabled || cfg.Host == "" {
|
||||
return err
|
||||
}
|
||||
return sendSMTP(ctx, cfg, msg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SendAsync 异步发送并自带超时,失败只记日志,绝不影响调用链路。
|
||||
@@ -62,13 +113,13 @@ func (n *Notifier) SendAsync(text string) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), notifySendTimeout)
|
||||
defer cancel()
|
||||
if err := n.Send(ctx, text); err != nil {
|
||||
log.Printf("telegram notify: %v", err)
|
||||
log.Printf("notify: %v", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// SendTemplateAsync 按 kind 的生效模板(自定义或默认)渲染变量后异步发送;
|
||||
// 模板支持轻量 Markdown,按 Telegram HTML 语义投递。
|
||||
// 模板支持轻量 Markdown,Telegram 按 HTML 语义投递,其余渠道按纯文本投递。
|
||||
func (n *Notifier) SendTemplateAsync(kind string, vars map[string]string) {
|
||||
n.wg.Add(1)
|
||||
go func() {
|
||||
@@ -76,7 +127,7 @@ func (n *Notifier) SendTemplateAsync(kind string, vars map[string]string) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), notifySendTimeout)
|
||||
defer cancel()
|
||||
if err := n.sendTemplate(ctx, kind, "", vars); err != nil {
|
||||
log.Printf("telegram notify %s: %v", kind, err)
|
||||
log.Printf("notify %s: %v", kind, err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -88,12 +139,12 @@ func (n *Notifier) SendTemplateTest(ctx context.Context, kind, tplOverride strin
|
||||
if !ok {
|
||||
return fmt.Errorf("未知通知类型 %q", kind)
|
||||
}
|
||||
cfg, err := n.settings.TelegramConfig(ctx)
|
||||
enabled, err := n.anyChannelEnabled(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !cfg.Enabled || cfg.BotToken == "" || cfg.ChatID == "" {
|
||||
return fmt.Errorf("Telegram 通知未启用或未配置,请先在「通知方式」保存配置")
|
||||
if !enabled {
|
||||
return fmt.Errorf("没有已启用的通知渠道,请先在「通知方式」配置并启用")
|
||||
}
|
||||
vars := map[string]string{}
|
||||
for k, v := range def.Sample {
|
||||
@@ -102,15 +153,29 @@ func (n *Notifier) SendTemplateTest(ctx context.Context, kind, tplOverride strin
|
||||
return n.sendTemplate(ctx, kind, tplOverride, vars)
|
||||
}
|
||||
|
||||
// sendTemplate 渲染并发送:取生效模板 → 变量替换 → 截断 → Markdown 转 Telegram HTML。
|
||||
func (n *Notifier) sendTemplate(ctx context.Context, kind, tplOverride string, vars map[string]string) error {
|
||||
cfg, err := n.settings.TelegramConfig(ctx)
|
||||
// anyChannelEnabled 报告是否存在至少一个已启用且配置齐全的渠道(含 Telegram)。
|
||||
func (n *Notifier) anyChannelEnabled(ctx context.Context) (bool, error) {
|
||||
tg, err := n.settings.TelegramConfig(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
return false, err
|
||||
}
|
||||
if !cfg.Enabled || cfg.BotToken == "" || cfg.ChatID == "" {
|
||||
return nil
|
||||
if tg.Enabled && tg.BotToken != "" && tg.ChatID != "" {
|
||||
return true, nil
|
||||
}
|
||||
views, err := n.settings.NotifyChannels(ctx)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
for _, v := range views {
|
||||
if v.Enabled {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// sendTemplate 渲染并分发:取生效模板 → 变量替换 → 投递到全部启用渠道。
|
||||
func (n *Notifier) sendTemplate(ctx context.Context, kind, tplOverride string, vars map[string]string) error {
|
||||
tpl := tplOverride
|
||||
if strings.TrimSpace(tpl) == "" {
|
||||
custom, err := n.settings.NotifyTemplate(ctx, kind)
|
||||
@@ -120,7 +185,58 @@ func (n *Notifier) sendTemplate(ctx context.Context, kind, tplOverride string, v
|
||||
tpl = notifyTplText(custom, kind)
|
||||
}
|
||||
text := renderNotifyTemplate(tpl, vars)
|
||||
return n.postHTML(ctx, cfg, mdToTelegramHTML(truncateText(text, notifyTextLimit)))
|
||||
return n.dispatch(ctx, notifyMsg{Title: notifyTplLabel(kind), Text: text})
|
||||
}
|
||||
|
||||
// SendChannelTest 向指定渠道同步发送测试消息;不看启用开关,便于保存后即时验证。
|
||||
func (n *Notifier) SendChannelTest(ctx context.Context, chType string) error {
|
||||
msg := notifyMsg{Title: "oci-portal 测试消息", Text: "✅ 通知渠道配置成功:" + chType}
|
||||
switch chType {
|
||||
case ChannelWebhook:
|
||||
cfg, err := n.settings.webhookChannel(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.URL == "" {
|
||||
return fmt.Errorf("请先保存 Webhook URL")
|
||||
}
|
||||
return sendWebhook(ctx, n.client, cfg, msg)
|
||||
case ChannelNtfy:
|
||||
cfg, err := n.settings.ntfyChannel(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.Topic == "" {
|
||||
return fmt.Errorf("请先保存 ntfy topic")
|
||||
}
|
||||
return sendNtfy(ctx, n.client, cfg, msg)
|
||||
}
|
||||
return n.sendChannelTestRest(ctx, chType, msg)
|
||||
}
|
||||
|
||||
// sendChannelTestRest 承接 SendChannelTest 的 bark/smtp 分支(拆分控制函数长度)。
|
||||
func (n *Notifier) sendChannelTestRest(ctx context.Context, chType string, msg notifyMsg) error {
|
||||
switch chType {
|
||||
case ChannelBark:
|
||||
cfg, err := n.settings.barkChannel(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.DeviceKey == "" {
|
||||
return fmt.Errorf("请先保存 Bark device key")
|
||||
}
|
||||
return sendBark(ctx, n.client, cfg, msg)
|
||||
case ChannelSMTP:
|
||||
cfg, err := n.settings.smtpChannel(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.Host == "" || cfg.From == "" || cfg.To == "" {
|
||||
return fmt.Errorf("请先保存 SMTP 主机、发件人与收件人")
|
||||
}
|
||||
return sendSMTP(ctx, cfg, msg)
|
||||
}
|
||||
return ErrUnknownNotifyChannel
|
||||
}
|
||||
|
||||
// Wait 等待全部在途异步通知发完,供退出前收尾与测试同步。
|
||||
|
||||
@@ -0,0 +1,218 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"mime"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/smtp"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// 通知渠道类型;telegram 走既有独立路径,不纳入此枚举。
|
||||
const (
|
||||
ChannelWebhook = "webhook"
|
||||
ChannelNtfy = "ntfy"
|
||||
ChannelBark = "bark"
|
||||
ChannelSMTP = "smtp"
|
||||
)
|
||||
|
||||
// NotifyChannelTypes 是新增渠道的稳定顺序(视图输出与遍历发送共用)。
|
||||
var NotifyChannelTypes = []string{ChannelWebhook, ChannelNtfy, ChannelBark, ChannelSMTP}
|
||||
|
||||
// notifyMsg 是渲染完成的渠道无关载荷:Title 取通知类型显示名,Text 为模板渲染文本。
|
||||
type notifyMsg struct {
|
||||
Title string
|
||||
Text string
|
||||
}
|
||||
|
||||
// defaultWebhookBody 是通用 Webhook 的缺省 body 模板(兼容 Slack incoming webhook)。
|
||||
const defaultWebhookBody = `{"text": "{{title}}\n{{text}}"}`
|
||||
|
||||
// 渠道服务端缺省地址。
|
||||
const (
|
||||
defaultNtfyServer = "https://ntfy.sh"
|
||||
defaultBarkServer = "https://api.day.app"
|
||||
)
|
||||
|
||||
// WebhookChannel 是通用 Webhook 渠道配置;BodyTemplate 空串按缺省模板发送。
|
||||
type WebhookChannel struct {
|
||||
Enabled bool
|
||||
URL string
|
||||
BodyTemplate string
|
||||
}
|
||||
|
||||
// NtfyChannel 是 ntfy 渠道配置;Server 空串用官方 ntfy.sh,Token 可选。
|
||||
type NtfyChannel struct {
|
||||
Enabled bool
|
||||
Server string
|
||||
Topic string
|
||||
Token string
|
||||
}
|
||||
|
||||
// BarkChannel 是 Bark(iOS)渠道配置;Server 空串用官方 api.day.app。
|
||||
type BarkChannel struct {
|
||||
Enabled bool
|
||||
Server string
|
||||
DeviceKey string
|
||||
}
|
||||
|
||||
// SMTPChannel 是邮件渠道配置;465 端口走隐式 TLS,其余端口 STARTTLS,Username 空则免认证。
|
||||
type SMTPChannel struct {
|
||||
Enabled bool
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
Password string
|
||||
From string
|
||||
To string
|
||||
}
|
||||
|
||||
// validNotifyURL 校验渠道出站地址:仅接受 http/https 且 host 非空。
|
||||
func validNotifyURL(raw string) error {
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
|
||||
return fmt.Errorf("地址须为 http(s):// 开头的合法 URL")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// renderWebhookBody 渲染 body 模板:占位值先做 JSON 字符串转义(去掉包裹引号)再替换,
|
||||
// 保证换行/引号嵌入 JSON 字符串后仍合法。
|
||||
func renderWebhookBody(tpl string, msg notifyMsg) string {
|
||||
if strings.TrimSpace(tpl) == "" {
|
||||
tpl = defaultWebhookBody
|
||||
}
|
||||
r := strings.NewReplacer("{{title}}", jsonEscape(msg.Title), "{{text}}", jsonEscape(msg.Text))
|
||||
return r.Replace(tpl)
|
||||
}
|
||||
|
||||
// jsonEscape 返回 s 作为 JSON 字符串字面量的内容(不含首尾引号)。
|
||||
func jsonEscape(s string) string {
|
||||
b, _ := json.Marshal(s)
|
||||
return string(b[1 : len(b)-1])
|
||||
}
|
||||
|
||||
// sendWebhook 向自定义 endpoint POST 渲染后的 JSON body。
|
||||
func sendWebhook(ctx context.Context, client *http.Client, cfg WebhookChannel, msg notifyMsg) error {
|
||||
body := renderWebhookBody(cfg.BodyTemplate, msg)
|
||||
return postJSON(ctx, client, cfg.URL, "", []byte(body))
|
||||
}
|
||||
|
||||
// sendNtfy 以 JSON publish 模式向 ntfy 服务端根路径发布(UTF-8 标题无兼容问题)。
|
||||
func sendNtfy(ctx context.Context, client *http.Client, cfg NtfyChannel, msg notifyMsg) error {
|
||||
server := cfg.Server
|
||||
if server == "" {
|
||||
server = defaultNtfyServer
|
||||
}
|
||||
body, err := json.Marshal(map[string]string{"topic": cfg.Topic, "title": msg.Title, "message": msg.Text})
|
||||
if err != nil {
|
||||
return fmt.Errorf("ntfy send: %w", err)
|
||||
}
|
||||
return postJSON(ctx, client, strings.TrimRight(server, "/"), cfg.Token, body)
|
||||
}
|
||||
|
||||
// sendBark 调用 Bark v2 push 接口。
|
||||
func sendBark(ctx context.Context, client *http.Client, cfg BarkChannel, msg notifyMsg) error {
|
||||
server := cfg.Server
|
||||
if server == "" {
|
||||
server = defaultBarkServer
|
||||
}
|
||||
body, err := json.Marshal(map[string]string{"title": msg.Title, "body": msg.Text, "device_key": cfg.DeviceKey})
|
||||
if err != nil {
|
||||
return fmt.Errorf("bark send: %w", err)
|
||||
}
|
||||
return postJSON(ctx, client, strings.TrimRight(server, "/")+"/push", "", body)
|
||||
}
|
||||
|
||||
// postJSON 发送 JSON POST;token 非空时附 Bearer 认证,非 2xx 一律视为失败。
|
||||
func postJSON(ctx context.Context, client *http.Client, u, token string, body []byte) error {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return sanitizeURLError(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return sanitizeURLError(err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("http %d", resp.StatusCode)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// sendSMTP 组装邮件报文并投递;主题按 RFC2047 编码,正文 UTF-8 纯文本。
|
||||
func sendSMTP(ctx context.Context, cfg SMTPChannel, msg notifyMsg) error {
|
||||
var b strings.Builder
|
||||
fmt.Fprintf(&b, "From: %s\r\n", cfg.From)
|
||||
fmt.Fprintf(&b, "To: %s\r\n", cfg.To)
|
||||
fmt.Fprintf(&b, "Subject: %s\r\n", mime.QEncoding.Encode("utf-8", msg.Title))
|
||||
b.WriteString("MIME-Version: 1.0\r\nContent-Type: text/plain; charset=utf-8\r\n\r\n")
|
||||
b.WriteString(strings.ReplaceAll(msg.Text, "\n", "\r\n"))
|
||||
return smtpDeliver(ctx, cfg, []byte(b.String()))
|
||||
}
|
||||
|
||||
// smtpDeliver 建立 SMTP 会话:465 端口隐式 TLS,其余端口连通后尝试 STARTTLS。
|
||||
func smtpDeliver(ctx context.Context, cfg SMTPChannel, mail []byte) error {
|
||||
addr := net.JoinHostPort(cfg.Host, fmt.Sprint(cfg.Port))
|
||||
d := &net.Dialer{}
|
||||
conn, err := d.DialContext(ctx, "tcp", addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("smtp dial: %w", err)
|
||||
}
|
||||
if cfg.Port == 465 {
|
||||
conn = tls.Client(conn, &tls.Config{ServerName: cfg.Host})
|
||||
}
|
||||
c, err := smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return fmt.Errorf("smtp handshake: %w", err)
|
||||
}
|
||||
defer c.Close()
|
||||
return smtpWrite(c, cfg, mail)
|
||||
}
|
||||
|
||||
// smtpWrite 在既有会话上完成 STARTTLS/认证/收发件与报文写入。
|
||||
func smtpWrite(c *smtp.Client, cfg SMTPChannel, mail []byte) error {
|
||||
if cfg.Port != 465 {
|
||||
if ok, _ := c.Extension("STARTTLS"); ok {
|
||||
if err := c.StartTLS(&tls.Config{ServerName: cfg.Host}); err != nil {
|
||||
return fmt.Errorf("smtp starttls: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
if cfg.Username != "" {
|
||||
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
|
||||
if err := c.Auth(auth); err != nil {
|
||||
return fmt.Errorf("smtp auth: %w", err)
|
||||
}
|
||||
}
|
||||
if err := c.Mail(cfg.From); err != nil {
|
||||
return fmt.Errorf("smtp from: %w", err)
|
||||
}
|
||||
if err := c.Rcpt(cfg.To); err != nil {
|
||||
return fmt.Errorf("smtp rcpt: %w", err)
|
||||
}
|
||||
w, err := c.Data()
|
||||
if err != nil {
|
||||
return fmt.Errorf("smtp data: %w", err)
|
||||
}
|
||||
if _, err := w.Write(mail); err != nil {
|
||||
w.Close()
|
||||
return fmt.Errorf("smtp write: %w", err)
|
||||
}
|
||||
if err := w.Close(); err != nil {
|
||||
return fmt.Errorf("smtp close: %w", err)
|
||||
}
|
||||
return c.Quit()
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRenderWebhookBody(t *testing.T) {
|
||||
msg := notifyMsg{Title: "任务失败", Text: "line1\nline\"2\""}
|
||||
tests := []struct {
|
||||
name string
|
||||
tpl string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "空模板用缺省 Slack 形状",
|
||||
tpl: "",
|
||||
want: `{"text": "任务失败\nline1\nline\"2\""}`,
|
||||
},
|
||||
{
|
||||
name: "自定义模板占位替换并转义",
|
||||
tpl: `{"msgtype":"text","text":{"content":"{{text}}"}}`,
|
||||
want: `{"msgtype":"text","text":{"content":"line1\nline\"2\""}}`,
|
||||
},
|
||||
{
|
||||
name: "无占位原样输出",
|
||||
tpl: `{"fixed":1}`,
|
||||
want: `{"fixed":1}`,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := renderWebhookBody(tt.tpl, msg)
|
||||
if got != tt.want {
|
||||
t.Errorf("renderWebhookBody = %s, want %s", got, tt.want)
|
||||
}
|
||||
if !json.Valid([]byte(got)) {
|
||||
t.Errorf("渲染结果不是合法 JSON: %s", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// jsonCapture 收集一次 HTTP 请求的路径、认证头与 JSON body。
|
||||
type jsonCapture struct {
|
||||
mu sync.Mutex
|
||||
path string
|
||||
auth string
|
||||
fields map[string]string
|
||||
}
|
||||
|
||||
// newCaptureServer 起假渠道服务端,固定返回 status。
|
||||
func newCaptureServer(t *testing.T, status int) (*httptest.Server, *jsonCapture) {
|
||||
t.Helper()
|
||||
rec := &jsonCapture{}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
rec.mu.Lock()
|
||||
defer rec.mu.Unlock()
|
||||
rec.path = r.URL.Path
|
||||
rec.auth = r.Header.Get("Authorization")
|
||||
rec.fields = map[string]string{}
|
||||
_ = json.NewDecoder(r.Body).Decode(&rec.fields)
|
||||
w.WriteHeader(status)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv, rec
|
||||
}
|
||||
|
||||
func TestSendNtfy(t *testing.T) {
|
||||
srv, rec := newCaptureServer(t, http.StatusOK)
|
||||
cfg := NtfyChannel{Server: srv.URL, Topic: "oci", Token: "tk-1"}
|
||||
msg := notifyMsg{Title: "标题", Text: "正文"}
|
||||
if err := sendNtfy(context.Background(), srv.Client(), cfg, msg); err != nil {
|
||||
t.Fatalf("sendNtfy: %v", err)
|
||||
}
|
||||
if rec.auth != "Bearer tk-1" {
|
||||
t.Errorf("Authorization = %q, want Bearer tk-1", rec.auth)
|
||||
}
|
||||
want := map[string]string{"topic": "oci", "title": "标题", "message": "正文"}
|
||||
for k, v := range want {
|
||||
if rec.fields[k] != v {
|
||||
t.Errorf("body[%s] = %q, want %q", k, rec.fields[k], v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendBark(t *testing.T) {
|
||||
srv, rec := newCaptureServer(t, http.StatusOK)
|
||||
cfg := BarkChannel{Server: srv.URL, DeviceKey: "dk-1"}
|
||||
if err := sendBark(context.Background(), srv.Client(), cfg, notifyMsg{Title: "标题", Text: "正文"}); err != nil {
|
||||
t.Fatalf("sendBark: %v", err)
|
||||
}
|
||||
if rec.path != "/push" {
|
||||
t.Errorf("path = %q, want /push", rec.path)
|
||||
}
|
||||
want := map[string]string{"title": "标题", "body": "正文", "device_key": "dk-1"}
|
||||
for k, v := range want {
|
||||
if rec.fields[k] != v {
|
||||
t.Errorf("body[%s] = %q, want %q", k, rec.fields[k], v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostJSONNon2xx(t *testing.T) {
|
||||
srv, _ := newCaptureServer(t, http.StatusBadGateway)
|
||||
err := postJSON(context.Background(), srv.Client(), srv.URL, "", []byte(`{}`))
|
||||
if err == nil || !strings.Contains(err.Error(), "502") {
|
||||
t.Fatalf("postJSON err = %v, want contains 502", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidNotifyURL(t *testing.T) {
|
||||
tests := []struct {
|
||||
raw string
|
||||
wantErr bool
|
||||
}{
|
||||
{"https://example.com/hook", false},
|
||||
{"http://10.0.0.1:8080/x", false},
|
||||
{"ftp://example.com", true},
|
||||
{"not-a-url", true},
|
||||
{"", true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if err := validNotifyURL(tt.raw); (err != nil) != tt.wantErr {
|
||||
t.Errorf("validNotifyURL(%q) err = %v, wantErr %v", tt.raw, err, tt.wantErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// fakeSMTP 起一个最小 SMTP 服务端,按标准会话应答并收集 DATA 报文。
|
||||
func fakeSMTP(t *testing.T) (addr string, got *strings.Builder) {
|
||||
t.Helper()
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { ln.Close() })
|
||||
got = &strings.Builder{}
|
||||
go func() {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
fmt.Fprint(conn, "220 fake ESMTP\r\n")
|
||||
r := bufio.NewReader(conn)
|
||||
for {
|
||||
line, err := r.ReadString('\n')
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
cmd := strings.ToUpper(strings.TrimSpace(line))
|
||||
switch {
|
||||
case strings.HasPrefix(cmd, "EHLO"), strings.HasPrefix(cmd, "HELO"):
|
||||
fmt.Fprint(conn, "250-fake\r\n250 AUTH PLAIN\r\n")
|
||||
case strings.HasPrefix(cmd, "AUTH"):
|
||||
fmt.Fprint(conn, "235 ok\r\n")
|
||||
case strings.HasPrefix(cmd, "MAIL"), strings.HasPrefix(cmd, "RCPT"):
|
||||
fmt.Fprint(conn, "250 ok\r\n")
|
||||
case cmd == "DATA":
|
||||
fmt.Fprint(conn, "354 go\r\n")
|
||||
readSMTPData(r, got)
|
||||
fmt.Fprint(conn, "250 ok\r\n")
|
||||
case cmd == "QUIT":
|
||||
fmt.Fprint(conn, "221 bye\r\n")
|
||||
return
|
||||
default:
|
||||
fmt.Fprint(conn, "250 ok\r\n")
|
||||
}
|
||||
}
|
||||
}()
|
||||
return ln.Addr().String(), got
|
||||
}
|
||||
|
||||
// readSMTPData 读取 DATA 段直到单独一行 "."。
|
||||
func readSMTPData(r *bufio.Reader, got *strings.Builder) {
|
||||
for {
|
||||
line, err := r.ReadString('\n')
|
||||
if err != nil || strings.TrimRight(line, "\r\n") == "." {
|
||||
return
|
||||
}
|
||||
got.WriteString(line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendSMTP(t *testing.T) {
|
||||
addr, got := fakeSMTP(t)
|
||||
host, portStr, _ := net.SplitHostPort(addr)
|
||||
var port int
|
||||
fmt.Sscanf(portStr, "%d", &port)
|
||||
cfg := SMTPChannel{
|
||||
Host: host, Port: port,
|
||||
Username: "u", Password: "p",
|
||||
From: "a@x.com", To: "b@y.com",
|
||||
}
|
||||
err := sendSMTP(context.Background(), cfg, notifyMsg{Title: "任务失败", Text: "第一行\n第二行"})
|
||||
if err != nil {
|
||||
t.Fatalf("sendSMTP: %v", err)
|
||||
}
|
||||
mail := got.String()
|
||||
for _, want := range []string{"From: a@x.com", "To: b@y.com", "Subject: =?utf-8?q?", "第一行\r\n第二行"} {
|
||||
if !strings.Contains(mail, want) {
|
||||
t.Errorf("mail 缺少 %q:\n%s", want, mail)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestDispatchFanout 验证多渠道并发投递:一个渠道失败不影响其他渠道送达。
|
||||
func TestDispatchFanout(t *testing.T) {
|
||||
okSrv, okRec := newCaptureServer(t, http.StatusOK)
|
||||
badSrv, badRec := newCaptureServer(t, http.StatusInternalServerError)
|
||||
settings, _ := newSettingEnv(t)
|
||||
ctx := context.Background()
|
||||
mustUpdateChannel(t, settings, ChannelWebhook, UpdateNotifyChannelInput{Enabled: true, URL: badSrv.URL})
|
||||
mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Enabled: true, Server: okSrv.URL, Topic: "oci"})
|
||||
n := NewNotifier(settings)
|
||||
|
||||
err := n.dispatch(ctx, notifyMsg{Title: "T", Text: "body"})
|
||||
if err == nil || !strings.Contains(err.Error(), "webhook") {
|
||||
t.Fatalf("dispatch err = %v, want webhook 失败", err)
|
||||
}
|
||||
if badRec.fields == nil {
|
||||
t.Error("webhook 渠道未收到请求")
|
||||
}
|
||||
if okRec.fields["message"] != "body" {
|
||||
t.Errorf("ntfy 渠道未正常送达: %v", okRec.fields)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDispatchSkipsDisabled 验证未启用/未配置的渠道静默跳过且不报错。
|
||||
func TestDispatchSkipsDisabled(t *testing.T) {
|
||||
srv, rec := newCaptureServer(t, http.StatusOK)
|
||||
settings, _ := newSettingEnv(t)
|
||||
mustUpdateChannel(t, settings, ChannelWebhook, UpdateNotifyChannelInput{Enabled: false, URL: srv.URL})
|
||||
n := NewNotifier(settings)
|
||||
if err := n.dispatch(context.Background(), notifyMsg{Title: "T", Text: "x"}); err != nil {
|
||||
t.Fatalf("dispatch: %v", err)
|
||||
}
|
||||
if rec.fields != nil {
|
||||
t.Error("未启用渠道不应收到请求")
|
||||
}
|
||||
}
|
||||
|
||||
// mustUpdateChannel 保存渠道配置,失败即终止测试。
|
||||
func mustUpdateChannel(t *testing.T, s *SettingService, chType string, in UpdateNotifyChannelInput) {
|
||||
t.Helper()
|
||||
if err := s.UpdateNotifyChannel(context.Background(), chType, in); err != nil {
|
||||
t.Fatalf("update channel %s: %v", chType, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateNotifyChannelValidation(t *testing.T) {
|
||||
secret := "sec-123456"
|
||||
tests := []struct {
|
||||
name string
|
||||
chType string
|
||||
in UpdateNotifyChannelInput
|
||||
wantErr string
|
||||
}{
|
||||
{name: "未知类型", chType: "pigeon", wantErr: "未知通知渠道"},
|
||||
{name: "启用 webhook 缺 URL", chType: ChannelWebhook, in: UpdateNotifyChannelInput{Enabled: true}, wantErr: "URL"},
|
||||
{name: "webhook 非法协议", chType: ChannelWebhook, in: UpdateNotifyChannelInput{Enabled: true, URL: "ftp://x"}, wantErr: "http(s)"},
|
||||
{name: "启用 ntfy 缺 topic", chType: ChannelNtfy, in: UpdateNotifyChannelInput{Enabled: true}, wantErr: "topic"},
|
||||
{name: "启用 smtp 缺主机", chType: ChannelSMTP, in: UpdateNotifyChannelInput{Enabled: true, Port: 587}, wantErr: "SMTP"},
|
||||
{name: "smtp 端口越界", chType: ChannelSMTP, in: UpdateNotifyChannelInput{Enabled: true, Host: "h", Port: 70000, From: "a@x", To: "b@y"}, wantErr: "端口"},
|
||||
{name: "未启用可存草稿", chType: ChannelNtfy, in: UpdateNotifyChannelInput{Enabled: false}},
|
||||
{name: "配置齐全可启用", chType: ChannelBark, in: UpdateNotifyChannelInput{Enabled: true, Secret: &secret}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
settings, _ := newSettingEnv(t)
|
||||
err := settings.UpdateNotifyChannel(context.Background(), tt.chType, tt.in)
|
||||
if tt.wantErr == "" && err != nil {
|
||||
t.Fatalf("UpdateNotifyChannel: %v", err)
|
||||
}
|
||||
if tt.wantErr != "" && (err == nil || !strings.Contains(err.Error(), tt.wantErr)) {
|
||||
t.Fatalf("err = %v, want contains %q", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyChannelSecretLifecycle 验证密文字段的存取语义:
|
||||
// 保存后视图只回 set/tail,nil 沿用旧值,空串清除。
|
||||
func TestNotifyChannelSecretLifecycle(t *testing.T) {
|
||||
settings, _ := newSettingEnv(t)
|
||||
ctx := context.Background()
|
||||
secret := "ntfy-token-9876"
|
||||
mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Topic: "t", Secret: &secret})
|
||||
|
||||
view := mustChannelView(t, settings, ChannelNtfy)
|
||||
if !view.SecretSet || view.SecretTail != "9876" {
|
||||
t.Fatalf("view = %+v, want SecretSet 且尾号 9876", view)
|
||||
}
|
||||
cfg, err := settings.ntfyChannel(ctx)
|
||||
if err != nil || cfg.Token != secret {
|
||||
t.Fatalf("ntfyChannel token = %q err = %v, want 原文", cfg.Token, err)
|
||||
}
|
||||
|
||||
// nil 沿用
|
||||
mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Topic: "t2"})
|
||||
if cfg, _ := settings.ntfyChannel(ctx); cfg.Token != secret || cfg.Topic != "t2" {
|
||||
t.Fatalf("nil Secret 应沿用旧 token,got %+v", cfg)
|
||||
}
|
||||
|
||||
// 空串清除
|
||||
empty := ""
|
||||
mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Topic: "t2", Secret: &empty})
|
||||
if cfg, _ := settings.ntfyChannel(ctx); cfg.Token != "" {
|
||||
t.Fatalf("空串 Secret 应清除,got %q", cfg.Token)
|
||||
}
|
||||
if view := mustChannelView(t, settings, ChannelNtfy); view.SecretSet {
|
||||
t.Fatal("清除后 SecretSet 应为 false")
|
||||
}
|
||||
}
|
||||
|
||||
// mustChannelView 取单渠道脱敏视图。
|
||||
func mustChannelView(t *testing.T, s *SettingService, chType string) NotifyChannelView {
|
||||
t.Helper()
|
||||
views, err := s.NotifyChannels(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("NotifyChannels: %v", err)
|
||||
}
|
||||
for _, v := range views {
|
||||
if v.Type == chType {
|
||||
return v
|
||||
}
|
||||
}
|
||||
t.Fatalf("view %s not found", chType)
|
||||
return NotifyChannelView{}
|
||||
}
|
||||
@@ -21,6 +21,7 @@ var notifyTplOrder = []string{
|
||||
"task_fail", "task_recover", "snatch_success", "tenant_dead", "task_stop",
|
||||
"login_lock", "model_deprecated",
|
||||
"log_event_instance", "log_event_identity", "log_event_policy", "log_event_region", "log_event_login",
|
||||
"audit_alert",
|
||||
}
|
||||
|
||||
var notifyTplDefs = map[string]notifyTplDef{
|
||||
@@ -38,9 +39,13 @@ var notifyTplDefs = map[string]notifyTplDef{
|
||||
},
|
||||
"snatch_success": {
|
||||
Label: "抢机成功",
|
||||
Default: "🎉 抢机成功:{{task_name}}\n{{message}}",
|
||||
Vars: []string{"task_name", "message"},
|
||||
Sample: map[string]string{"task_name": "示例抢机任务", "message": "已创建 1/1 台实例"},
|
||||
Default: "🎉 抢机成功:{{task_name}}\n租户:{{tenant}}\n区域:{{region}}\n类型:{{shape}}\n配置:{{spec}}\n镜像:{{image}}\nIP:{{ip}}\n密码:{{root_password}}",
|
||||
Vars: []string{"task_name", "tenant", "region", "shape", "spec", "image", "ip", "root_password", "message"},
|
||||
Sample: map[string]string{
|
||||
"task_name": "示例抢机任务", "tenant": "免费01", "region": "Japan Central (Osaka)",
|
||||
"shape": "VM.Standard.A1.Flex", "spec": "4C/24G", "image": "Canonical-Ubuntu-24.04-aarch64-2025.05.20-0",
|
||||
"ip": "129.150.32.10", "root_password": "Xy3#kP9m-2Qw@7Zn", "message": "created 1: ocid1.instance…",
|
||||
},
|
||||
},
|
||||
"tenant_dead": {
|
||||
Label: "租户失联",
|
||||
@@ -96,6 +101,13 @@ var notifyTplDefs = map[string]notifyTplDef{
|
||||
Vars: []string{"tenant", "event", "resource"},
|
||||
Sample: map[string]string{"tenant": "免费01", "event": "InteractiveLogin", "resource": "demo@example.com"},
|
||||
},
|
||||
"audit_alert": {
|
||||
Label: "审计告警",
|
||||
Default: "🔔 审计告警:{{rule}}\n租户 {{tenant}} · {{event}} {{resource}}\n来源 {{ip}} · 窗口内 {{count}} 次",
|
||||
Vars: []string{"rule", "tenant", "event", "resource", "ip", "count"},
|
||||
Sample: map[string]string{"rule": "非白名单终止实例", "tenant": "免费01",
|
||||
"event": "TerminateInstance", "resource": "web-server-1", "ip": "203.0.113.8", "count": "1"},
|
||||
},
|
||||
}
|
||||
|
||||
var notifyVarPattern = regexp.MustCompile(`\{\{(\w+)\}\}`)
|
||||
@@ -161,3 +173,11 @@ func notifyTplText(custom, kind string) string {
|
||||
}
|
||||
return notifyTplDefs[kind].Default
|
||||
}
|
||||
|
||||
// notifyTplLabel 返回 kind 的显示名,供渠道消息标题(ntfy/bark/邮件主题)使用。
|
||||
func notifyTplLabel(kind string) string {
|
||||
if def, ok := notifyTplDefs[kind]; ok {
|
||||
return def.Label
|
||||
}
|
||||
return "oci-portal 通知"
|
||||
}
|
||||
|
||||
@@ -197,7 +197,12 @@ func (o *OAuthService) HandleCallback(ctx context.Context, provider, state, code
|
||||
return "", "", p.mode, err
|
||||
}
|
||||
if p.mode == "bind" {
|
||||
return "", ident.Display, p.mode, o.bind(ctx, p.username, provider, ident)
|
||||
if err := o.bind(ctx, p.username, provider, ident); err != nil {
|
||||
return "", ident.Display, p.mode, err
|
||||
}
|
||||
// 绑定属敏感变更:版本递增使旧令牌失效,同时为操作者签新令牌随回跳带回
|
||||
token, _, err := o.auth.RevokeSessions(ctx, p.username)
|
||||
return token, ident.Display, p.mode, err
|
||||
}
|
||||
token, display, err = o.loginByIdentity(ctx, provider, ident)
|
||||
return token, display, p.mode, err
|
||||
@@ -305,7 +310,7 @@ func (o *OAuthService) loginByIdentity(ctx context.Context, provider string, ide
|
||||
if err := o.db.WithContext(ctx).First(&user, row.UserID).Error; err != nil {
|
||||
return "", "", fmt.Errorf("find bound user: %w", err)
|
||||
}
|
||||
token, _, err := o.auth.signToken(user.Username)
|
||||
token, _, err := o.auth.signToken(user.Username, user.TokenVersion)
|
||||
return token, ident.Display, err
|
||||
}
|
||||
|
||||
@@ -345,5 +350,6 @@ func (o *OAuthService) Unbind(ctx context.Context, username string, id uint) err
|
||||
if res.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
// 解绑属敏感变更:递增令牌版本,已签发会话全部失效
|
||||
return o.auth.bumpTokenVersion(ctx, username)
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -15,10 +16,12 @@ import (
|
||||
|
||||
// OciConfigService 管理 API Key 配置的导入、测活与快照同步。
|
||||
type OciConfigService struct {
|
||||
db *gorm.DB
|
||||
cipher *crypto.Cipher
|
||||
client oci.Client
|
||||
// auditRaw 暂存审计原始事件(eventId → raw,TTL 10 分钟),
|
||||
db *gorm.DB
|
||||
cipher *crypto.Cipher
|
||||
client oci.Client
|
||||
cleanupTasks *TaskService
|
||||
cleanupEvents *LogEventService
|
||||
// auditRaw 按 configId:eventId 暂存审计原始事件(TTL 10 分钟),
|
||||
// 列表响应剥离 raw 后详情接口据此秒开;miss 走小窗重查兜底。
|
||||
auditRaw *cache.Cache
|
||||
}
|
||||
@@ -28,6 +31,12 @@ func NewOciConfigService(db *gorm.DB, cipher *crypto.Cipher, client oci.Client)
|
||||
return &OciConfigService{db: db, cipher: cipher, client: client, auditRaw: cache.New(auditRawMax)}
|
||||
}
|
||||
|
||||
// SetTenantCleanupDeps 注入租户删除提交后的任务与告警内存状态同步依赖。
|
||||
func (s *OciConfigService) SetTenantCleanupDeps(tasks *TaskService, events *LogEventService) {
|
||||
s.cleanupTasks = tasks
|
||||
s.cleanupEvents = events
|
||||
}
|
||||
|
||||
// ImportInput 是导入一份 API Key 的输入:
|
||||
// ConfigINI 与显式字段二选一(ConfigINI 优先),私钥内容必填。
|
||||
type ImportInput struct {
|
||||
@@ -263,17 +272,18 @@ func (s *OciConfigService) Get(ctx context.Context, id uint) (*model.OciConfig,
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
// Delete 删除配置及其区域 / 区间缓存。
|
||||
// Delete 在单一事务中删除配置及全部租户级本地关联数据。
|
||||
func (s *OciConfigService) Delete(ctx context.Context, id uint) error {
|
||||
res := s.db.WithContext(ctx).Delete(&model.OciConfig{}, id)
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("delete oci config %d: %w", id, res.Error)
|
||||
unlock := func() {}
|
||||
if s.cleanupTasks != nil {
|
||||
unlock = s.cleanupTasks.lockTenantCleanup()
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return fmt.Errorf("delete oci config %d: %w", id, gorm.ErrRecordNotFound)
|
||||
defer unlock()
|
||||
result, err := s.deleteTenant(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.db.WithContext(ctx).Where("oci_config_id = ?", id).Delete(&model.RegionCache{})
|
||||
s.db.WithContext(ctx).Where("oci_config_id = ?", id).Delete(&model.CompartmentCache{})
|
||||
s.afterTenantDelete(ctx, result)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -294,12 +304,43 @@ func (s *OciConfigService) refresh(ctx context.Context, cfg *model.OciConfig, cr
|
||||
s.applyProfile(ctx, cfg, cred)
|
||||
s.syncScopeCaches(ctx, cfg, cred)
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Save(cfg).Error; err != nil {
|
||||
return fmt.Errorf("save oci config snapshot: %w", err)
|
||||
s.applySuspension(ctx, cfg, cred)
|
||||
return s.saveSnapshot(ctx, cfg)
|
||||
}
|
||||
|
||||
// saveSnapshot 只更新当前版本,避免在途测活把已删租户重新插入。
|
||||
func (s *OciConfigService) saveSnapshot(ctx context.Context, cfg *model.OciConfig) error {
|
||||
res := s.db.WithContext(ctx).Model(&model.OciConfig{}).
|
||||
Where("id = ? AND updated_at = ?", cfg.ID, cfg.UpdatedAt).
|
||||
Select("*").Omit("id", "created_at").Updates(cfg)
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("save oci config snapshot: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected != 1 {
|
||||
return fmt.Errorf("save oci config snapshot: %w", gorm.ErrRecordNotFound)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// applySuspension 查询账户能力接口:云端标记暂停的租户覆盖测活结论为 suspended。
|
||||
// 该接口对暂停/已终止租户仍可访问(常规 API 此时多被拒,会误判失联);查询失败不改变结论。
|
||||
func (s *OciConfigService) applySuspension(ctx context.Context, cfg *model.OciConfig, cred oci.Credentials) {
|
||||
caps, err := s.client.GetAccountCapabilities(ctx, cred, homeRegionOf(cfg))
|
||||
if err != nil || !caps.Suspended {
|
||||
return
|
||||
}
|
||||
cfg.AliveStatus = model.AliveStatusSuspended
|
||||
cfg.LastError = strings.TrimSpace("账户已被云端暂停 " + capsStatusNote(caps))
|
||||
}
|
||||
|
||||
// capsStatusNote 生成暂停详情备注(账户状态,如 terminated)。
|
||||
func capsStatusNote(caps oci.AccountCapabilities) string {
|
||||
if caps.AccountStatus == "" {
|
||||
return ""
|
||||
}
|
||||
return "(account status: " + caps.AccountStatus + ")"
|
||||
}
|
||||
|
||||
// applyProfile 拉取订阅信息;失败只把类别降级为 unknown,不影响测活结论。
|
||||
func (s *OciConfigService) applyProfile(ctx context.Context, cfg *model.OciConfig, cred oci.Credentials) {
|
||||
profile, err := s.client.FetchAccountProfile(ctx, cred)
|
||||
|
||||
@@ -25,6 +25,8 @@ type fakeClient struct {
|
||||
tenancyErr error
|
||||
profile oci.AccountProfile
|
||||
profileErr error
|
||||
caps oci.AccountCapabilities
|
||||
capsErr error
|
||||
|
||||
regionSubs []oci.RegionSubscription
|
||||
regionSubsErr error
|
||||
@@ -43,6 +45,10 @@ type fakeClient struct {
|
||||
instancesErr error
|
||||
costItems []oci.CostItem
|
||||
costErr error
|
||||
instance oci.Instance
|
||||
instanceErr error
|
||||
image oci.Image
|
||||
imageErr error
|
||||
|
||||
subscribedHomeRegion string
|
||||
subscribedKey string
|
||||
@@ -55,6 +61,10 @@ func (f *fakeClient) ValidateKey(ctx context.Context, cred oci.Credentials) (oci
|
||||
return f.tenancy, f.tenancyErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) GetAccountCapabilities(ctx context.Context, cred oci.Credentials, region string) (oci.AccountCapabilities, error) {
|
||||
return f.caps, f.capsErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) FetchAccountProfile(ctx context.Context, cred oci.Credentials) (oci.AccountProfile, error) {
|
||||
return f.profile, f.profileErr
|
||||
}
|
||||
@@ -107,6 +117,14 @@ func (f *fakeClient) ListInstances(ctx context.Context, cred oci.Credentials, re
|
||||
return f.instances, f.instancesErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) GetInstance(ctx context.Context, cred oci.Credentials, region, instanceID string) (oci.Instance, error) {
|
||||
return f.instance, f.instanceErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) GetImage(ctx context.Context, cred oci.Credentials, region, imageID string) (oci.Image, error) {
|
||||
return f.image, f.imageErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) SummarizeCosts(ctx context.Context, cred oci.Credentials, q oci.CostQuery) ([]oci.CostItem, error) {
|
||||
return f.costItems, f.costErr
|
||||
}
|
||||
@@ -312,6 +330,7 @@ func TestVerifyMissingConfig(t *testing.T) {
|
||||
|
||||
func TestDelete(t *testing.T) {
|
||||
svc := newTestService(t, &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}})
|
||||
migrateTenantDeleteModels(t, svc.db)
|
||||
cfg, err := svc.Import(context.Background(), trialImportInput())
|
||||
if err != nil {
|
||||
t.Fatalf("Import: %v", err)
|
||||
@@ -458,3 +477,39 @@ func TestListReturnsSummaries(t *testing.T) {
|
||||
t.Errorf("items[1] alias/group = %q/%q, want 生产号/empty", items[1].Alias, items[1].Group)
|
||||
}
|
||||
}
|
||||
|
||||
// TestVerifySuspendedOverridesStatus 云端暂停标记应覆盖测活结论为 suspended;
|
||||
// 能力接口失败或未标记时不改变原结论。
|
||||
func TestVerifySuspendedOverridesStatus(t *testing.T) {
|
||||
client := &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}
|
||||
svc := newTestService(t, client)
|
||||
cfg := importAliveConfig(t, svc)
|
||||
ctx := context.Background()
|
||||
|
||||
// 标记暂停:即使 ValidateKey 成功也应判 suspended
|
||||
client.caps = oci.AccountCapabilities{Suspended: true, AccountStatus: "terminated"}
|
||||
got, _, err := svc.Verify(ctx, cfg.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("Verify: %v", err)
|
||||
}
|
||||
if got.AliveStatus != model.AliveStatusSuspended {
|
||||
t.Fatalf("AliveStatus = %q, want suspended", got.AliveStatus)
|
||||
}
|
||||
if got.LastError == "" {
|
||||
t.Error("暂停时 LastError 应携带账户状态备注")
|
||||
}
|
||||
|
||||
// 解除暂停:恢复 alive
|
||||
client.caps = oci.AccountCapabilities{}
|
||||
got, _, err = svc.Verify(ctx, cfg.ID)
|
||||
if err != nil || got.AliveStatus != model.AliveStatusAlive {
|
||||
t.Fatalf("解除暂停后 = %q, %v; want alive", got.AliveStatus, err)
|
||||
}
|
||||
|
||||
// 能力接口失败:不影响测活结论
|
||||
client.capsErr = errors.New("boom")
|
||||
got, _, err = svc.Verify(ctx, cfg.ID)
|
||||
if err != nil || got.AliveStatus != model.AliveStatusAlive {
|
||||
t.Fatalf("能力接口失败后 = %q, %v; want alive", got.AliveStatus, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -89,7 +89,7 @@ func (s *OciConfigService) overviewTenants(ctx context.Context, out *Overview) e
|
||||
switch cfg.AliveStatus {
|
||||
case model.AliveStatusAlive:
|
||||
out.Tenants.Alive++
|
||||
case model.AliveStatusDead:
|
||||
case model.AliveStatusDead, model.AliveStatusSuspended:
|
||||
out.Tenants.Dead++
|
||||
}
|
||||
}
|
||||
|
||||
@@ -39,6 +39,9 @@ func (s *OciConfigService) saveRegionCache(ctx context.Context, cfgID uint, subs
|
||||
})
|
||||
}
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := lockOciConfig(tx, cfgID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("oci_config_id = ?", cfgID).Delete(&model.RegionCache{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -58,6 +61,9 @@ func (s *OciConfigService) saveCompartmentCache(ctx context.Context, cfgID uint,
|
||||
})
|
||||
}
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := lockOciConfig(tx, cfgID); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("oci_config_id = ?", cfgID).Delete(&model.CompartmentCache{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,309 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// 通知渠道配置键,前缀 notify_ch_<type>_;标注「密文」的键 AES-GCM 加密落库。
|
||||
const (
|
||||
settingChWebhookEnabled = "notify_ch_webhook_enabled" // "1"/"0"
|
||||
settingChWebhookURL = "notify_ch_webhook_url"
|
||||
settingChWebhookBodyTpl = "notify_ch_webhook_body_tpl"
|
||||
settingChNtfyEnabled = "notify_ch_ntfy_enabled"
|
||||
settingChNtfyServer = "notify_ch_ntfy_server"
|
||||
settingChNtfyTopic = "notify_ch_ntfy_topic"
|
||||
settingChNtfyToken = "notify_ch_ntfy_token" // 密文
|
||||
settingChBarkEnabled = "notify_ch_bark_enabled"
|
||||
settingChBarkServer = "notify_ch_bark_server"
|
||||
settingChBarkDeviceKey = "notify_ch_bark_device_key" // 密文
|
||||
settingChSMTPEnabled = "notify_ch_smtp_enabled"
|
||||
settingChSMTPHost = "notify_ch_smtp_host"
|
||||
settingChSMTPPort = "notify_ch_smtp_port"
|
||||
settingChSMTPUsername = "notify_ch_smtp_username"
|
||||
settingChSMTPPassword = "notify_ch_smtp_password" // 密文
|
||||
settingChSMTPFrom = "notify_ch_smtp_from"
|
||||
settingChSMTPTo = "notify_ch_smtp_to"
|
||||
)
|
||||
|
||||
// ErrUnknownNotifyChannel 标记未知渠道类型,api 层映射 400。
|
||||
var ErrUnknownNotifyChannel = fmt.Errorf("未知通知渠道类型")
|
||||
|
||||
// NotifyChannelView 是单渠道的脱敏视图:密文字段只回是否已配置与尾 4 位。
|
||||
type NotifyChannelView struct {
|
||||
Type string `json:"type"`
|
||||
Enabled bool `json:"enabled"`
|
||||
// webhook
|
||||
URL string `json:"url,omitempty"`
|
||||
BodyTemplate string `json:"bodyTemplate,omitempty"`
|
||||
// ntfy / bark 共用 Server
|
||||
Server string `json:"server,omitempty"`
|
||||
Topic string `json:"topic,omitempty"`
|
||||
// smtp
|
||||
Host string `json:"host,omitempty"`
|
||||
Port int `json:"port,omitempty"`
|
||||
Username string `json:"username,omitempty"`
|
||||
From string `json:"from,omitempty"`
|
||||
To string `json:"to,omitempty"`
|
||||
// 密文字段状态(ntfy token / bark device key / smtp password)
|
||||
SecretSet bool `json:"secretSet"`
|
||||
SecretTail string `json:"secretTail,omitempty"`
|
||||
}
|
||||
|
||||
// UpdateNotifyChannelInput 是单渠道更新入参;Secret 为 nil 沿用已存密文,空串清除。
|
||||
type UpdateNotifyChannelInput struct {
|
||||
Enabled bool
|
||||
URL string
|
||||
BodyTemplate string
|
||||
Server string
|
||||
Topic string
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
From string
|
||||
To string
|
||||
Secret *string
|
||||
}
|
||||
|
||||
// notifyChannelKeys 返回渠道的(启用键, 明文键列表, 密文键);未知类型返回启用键空串。
|
||||
func notifyChannelKeys(chType string) (enabled string, plain []string, secret string) {
|
||||
switch chType {
|
||||
case ChannelWebhook:
|
||||
return settingChWebhookEnabled, []string{settingChWebhookURL, settingChWebhookBodyTpl}, ""
|
||||
case ChannelNtfy:
|
||||
return settingChNtfyEnabled, []string{settingChNtfyServer, settingChNtfyTopic}, settingChNtfyToken
|
||||
case ChannelBark:
|
||||
return settingChBarkEnabled, []string{settingChBarkServer}, settingChBarkDeviceKey
|
||||
case ChannelSMTP:
|
||||
return settingChSMTPEnabled, []string{settingChSMTPHost, settingChSMTPPort,
|
||||
settingChSMTPUsername, settingChSMTPFrom, settingChSMTPTo}, settingChSMTPPassword
|
||||
}
|
||||
return "", nil, ""
|
||||
}
|
||||
|
||||
// NotifyChannels 返回全部新增渠道的脱敏视图(不含 telegram,其沿用独立接口)。
|
||||
func (s *SettingService) NotifyChannels(ctx context.Context) ([]NotifyChannelView, error) {
|
||||
out := make([]NotifyChannelView, 0, len(NotifyChannelTypes))
|
||||
for _, t := range NotifyChannelTypes {
|
||||
v, err := s.notifyChannelView(ctx, t)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, v)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// notifyChannelView 组装单渠道脱敏视图。
|
||||
func (s *SettingService) notifyChannelView(ctx context.Context, chType string) (NotifyChannelView, error) {
|
||||
enabledKey, plain, secretKey := notifyChannelKeys(chType)
|
||||
vals, err := s.getMany(ctx, append(plain, enabledKey)...)
|
||||
if err != nil {
|
||||
return NotifyChannelView{}, err
|
||||
}
|
||||
view := NotifyChannelView{Type: chType, Enabled: vals[enabledKey] == "1"}
|
||||
fillChannelView(&view, chType, vals)
|
||||
if secretKey == "" {
|
||||
return view, nil
|
||||
}
|
||||
secret, err := s.decryptSetting(ctx, secretKey)
|
||||
if err != nil {
|
||||
return NotifyChannelView{}, err
|
||||
}
|
||||
view.SecretSet = secret != ""
|
||||
view.SecretTail = tokenTail(secret)
|
||||
return view, nil
|
||||
}
|
||||
|
||||
// fillChannelView 把明文键值填入视图对应字段。
|
||||
func fillChannelView(v *NotifyChannelView, chType string, vals map[string]string) {
|
||||
switch chType {
|
||||
case ChannelWebhook:
|
||||
v.URL, v.BodyTemplate = vals[settingChWebhookURL], vals[settingChWebhookBodyTpl]
|
||||
case ChannelNtfy:
|
||||
v.Server, v.Topic = vals[settingChNtfyServer], vals[settingChNtfyTopic]
|
||||
case ChannelBark:
|
||||
v.Server = vals[settingChBarkServer]
|
||||
case ChannelSMTP:
|
||||
v.Host, v.Username = vals[settingChSMTPHost], vals[settingChSMTPUsername]
|
||||
v.From, v.To = vals[settingChSMTPFrom], vals[settingChSMTPTo]
|
||||
v.Port, _ = strconv.Atoi(vals[settingChSMTPPort])
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateNotifyChannel 保存单渠道配置;字段先经 validateChannelInput 校验。
|
||||
func (s *SettingService) UpdateNotifyChannel(ctx context.Context, chType string, in UpdateNotifyChannelInput) error {
|
||||
enabledKey, _, secretKey := notifyChannelKeys(chType)
|
||||
if enabledKey == "" {
|
||||
return ErrUnknownNotifyChannel
|
||||
}
|
||||
if err := validateChannelInput(chType, in); err != nil {
|
||||
return err
|
||||
}
|
||||
pairs := channelPlainPairs(chType, in)
|
||||
pairs[enabledKey] = boolSetting(in.Enabled)
|
||||
for k, v := range pairs {
|
||||
if err := s.set(ctx, k, v); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if secretKey == "" || in.Secret == nil {
|
||||
return nil
|
||||
}
|
||||
return s.encryptSetting(ctx, secretKey, *in.Secret)
|
||||
}
|
||||
|
||||
// channelPlainPairs 返回渠道明文键与入参值的对应关系。
|
||||
func channelPlainPairs(chType string, in UpdateNotifyChannelInput) map[string]string {
|
||||
switch chType {
|
||||
case ChannelWebhook:
|
||||
return map[string]string{settingChWebhookURL: in.URL, settingChWebhookBodyTpl: in.BodyTemplate}
|
||||
case ChannelNtfy:
|
||||
return map[string]string{settingChNtfyServer: in.Server, settingChNtfyTopic: in.Topic}
|
||||
case ChannelBark:
|
||||
return map[string]string{settingChBarkServer: in.Server}
|
||||
case ChannelSMTP:
|
||||
return map[string]string{settingChSMTPHost: in.Host, settingChSMTPPort: strconv.Itoa(in.Port),
|
||||
settingChSMTPUsername: in.Username, settingChSMTPFrom: in.From, settingChSMTPTo: in.To}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateChannelInput 按渠道校验入参;仅在启用时强校验必填,便于「先存草稿后启用」。
|
||||
func validateChannelInput(chType string, in UpdateNotifyChannelInput) error {
|
||||
if !in.Enabled {
|
||||
return validateChannelURLs(chType, in)
|
||||
}
|
||||
switch chType {
|
||||
case ChannelWebhook:
|
||||
if strings.TrimSpace(in.URL) == "" {
|
||||
return fmt.Errorf("启用 Webhook 须填写 URL")
|
||||
}
|
||||
case ChannelNtfy:
|
||||
if strings.TrimSpace(in.Topic) == "" {
|
||||
return fmt.Errorf("启用 ntfy 须填写 topic")
|
||||
}
|
||||
case ChannelSMTP:
|
||||
if in.Host == "" || in.Port <= 0 || in.Port > 65535 || in.From == "" || in.To == "" {
|
||||
return fmt.Errorf("启用 SMTP 须填写主机、端口(1-65535)、发件人与收件人")
|
||||
}
|
||||
}
|
||||
return validateChannelURLs(chType, in)
|
||||
}
|
||||
|
||||
// validateChannelURLs 校验渠道内的 URL 形字段(非空才校验)。
|
||||
func validateChannelURLs(chType string, in UpdateNotifyChannelInput) error {
|
||||
if chType == ChannelWebhook && in.URL != "" {
|
||||
return validNotifyURL(in.URL)
|
||||
}
|
||||
if (chType == ChannelNtfy || chType == ChannelBark) && in.Server != "" {
|
||||
return validNotifyURL(in.Server)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// boolSetting 把布尔序列化为设置值 "1"/"0"。
|
||||
func boolSetting(b bool) string {
|
||||
if b {
|
||||
return "1"
|
||||
}
|
||||
return "0"
|
||||
}
|
||||
|
||||
// encryptSetting 加密保存设置值;空串直接落空值表示清除。
|
||||
func (s *SettingService) encryptSetting(ctx context.Context, key, value string) error {
|
||||
if value == "" {
|
||||
return s.set(ctx, key, "")
|
||||
}
|
||||
enc, err := s.cipher.EncryptString(value)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt setting %s: %w", key, err)
|
||||
}
|
||||
return s.set(ctx, key, enc)
|
||||
}
|
||||
|
||||
// decryptSetting 读取并解密设置值;未配置返回空串。
|
||||
func (s *SettingService) decryptSetting(ctx context.Context, key string) (string, error) {
|
||||
enc, err := s.get(ctx, key)
|
||||
if err != nil || enc == "" {
|
||||
return "", err
|
||||
}
|
||||
val, err := s.cipher.DecryptString(enc)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decrypt setting %s: %w", key, err)
|
||||
}
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// webhookChannel 读取解密后的 Webhook 配置,供发送端消费。
|
||||
func (s *SettingService) webhookChannel(ctx context.Context) (WebhookChannel, error) {
|
||||
vals, err := s.getMany(ctx, settingChWebhookEnabled, settingChWebhookURL, settingChWebhookBodyTpl)
|
||||
if err != nil {
|
||||
return WebhookChannel{}, err
|
||||
}
|
||||
return WebhookChannel{
|
||||
Enabled: vals[settingChWebhookEnabled] == "1",
|
||||
URL: vals[settingChWebhookURL],
|
||||
BodyTemplate: vals[settingChWebhookBodyTpl],
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ntfyChannel 读取解密后的 ntfy 配置。
|
||||
func (s *SettingService) ntfyChannel(ctx context.Context) (NtfyChannel, error) {
|
||||
vals, err := s.getMany(ctx, settingChNtfyEnabled, settingChNtfyServer, settingChNtfyTopic)
|
||||
if err != nil {
|
||||
return NtfyChannel{}, err
|
||||
}
|
||||
token, err := s.decryptSetting(ctx, settingChNtfyToken)
|
||||
if err != nil {
|
||||
return NtfyChannel{}, err
|
||||
}
|
||||
return NtfyChannel{
|
||||
Enabled: vals[settingChNtfyEnabled] == "1",
|
||||
Server: vals[settingChNtfyServer],
|
||||
Topic: vals[settingChNtfyTopic],
|
||||
Token: token,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// barkChannel 读取解密后的 Bark 配置。
|
||||
func (s *SettingService) barkChannel(ctx context.Context) (BarkChannel, error) {
|
||||
vals, err := s.getMany(ctx, settingChBarkEnabled, settingChBarkServer)
|
||||
if err != nil {
|
||||
return BarkChannel{}, err
|
||||
}
|
||||
key, err := s.decryptSetting(ctx, settingChBarkDeviceKey)
|
||||
if err != nil {
|
||||
return BarkChannel{}, err
|
||||
}
|
||||
return BarkChannel{
|
||||
Enabled: vals[settingChBarkEnabled] == "1",
|
||||
Server: vals[settingChBarkServer],
|
||||
DeviceKey: key,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// smtpChannel 读取解密后的 SMTP 配置。
|
||||
func (s *SettingService) smtpChannel(ctx context.Context) (SMTPChannel, error) {
|
||||
vals, err := s.getMany(ctx, settingChSMTPEnabled, settingChSMTPHost, settingChSMTPPort,
|
||||
settingChSMTPUsername, settingChSMTPFrom, settingChSMTPTo)
|
||||
if err != nil {
|
||||
return SMTPChannel{}, err
|
||||
}
|
||||
password, err := s.decryptSetting(ctx, settingChSMTPPassword)
|
||||
if err != nil {
|
||||
return SMTPChannel{}, err
|
||||
}
|
||||
port, _ := strconv.Atoi(vals[settingChSMTPPort])
|
||||
return SMTPChannel{
|
||||
Enabled: vals[settingChSMTPEnabled] == "1",
|
||||
Host: vals[settingChSMTPHost],
|
||||
Port: port,
|
||||
Username: vals[settingChSMTPUsername],
|
||||
Password: password,
|
||||
From: vals[settingChSMTPFrom],
|
||||
To: vals[settingChSMTPTo],
|
||||
}, nil
|
||||
}
|
||||
+240
-12
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -22,6 +23,13 @@ const taskRunTimeout = 10 * time.Minute
|
||||
// taskLogKeep 是每个任务保留的执行日志条数。
|
||||
const taskLogKeep = 100
|
||||
|
||||
// defaultSnatchIPWait 是抢机成功后等待公网 IPv4 分配的总预算(全部实例共享,
|
||||
// 实例从创建到 VNIC 就绪通常 1~2 分钟);taskRunTimeout 余量内,超时降级不阻塞通知。
|
||||
const defaultSnatchIPWait = 3 * time.Minute
|
||||
|
||||
// snatchIPPollInterval 是等待公网 IP 的轮询间隔。
|
||||
const snatchIPPollInterval = 10 * time.Second
|
||||
|
||||
// TaskService 管理后台任务的存储、cron 调度与执行。
|
||||
type TaskService struct {
|
||||
db *gorm.DB
|
||||
@@ -31,21 +39,32 @@ type TaskService struct {
|
||||
cron *cron.Cron
|
||||
// aiGateway 供 AI 探测任务执行渠道探测,由 main 装配(可为 nil)
|
||||
aiGateway *AiGatewayService
|
||||
// snatchIPWait 是抢机成功后等待公网 IP 的预算,测试置 0 跳过等待
|
||||
snatchIPWait time.Duration
|
||||
|
||||
mu sync.Mutex
|
||||
entries map[uint]cron.EntryID
|
||||
// runMu 让租户删除等待在途执行结束,并阻止新执行读取删除前的任务快照。
|
||||
runMu sync.RWMutex
|
||||
|
||||
// snatchVarsMu 保护 snatchVars:runSnatch 抢满时写入成功通知变量,
|
||||
// execute 组装执行后快照时取走(一次性)
|
||||
snatchVarsMu sync.Mutex
|
||||
snatchVars map[uint]map[string]string
|
||||
}
|
||||
|
||||
// NewTaskService 组装依赖;notifier 传 nil 表示整体关闭通知,
|
||||
// settings 供发送前按事件类型过滤(nil 视为全开)。调用 Start 后开始调度。
|
||||
func NewTaskService(db *gorm.DB, configs *OciConfigService, notifier *Notifier, settings *SettingService) *TaskService {
|
||||
return &TaskService{
|
||||
db: db,
|
||||
configs: configs,
|
||||
notifier: notifier,
|
||||
settings: settings,
|
||||
cron: cron.New(),
|
||||
entries: map[uint]cron.EntryID{},
|
||||
db: db,
|
||||
configs: configs,
|
||||
notifier: notifier,
|
||||
settings: settings,
|
||||
cron: cron.New(),
|
||||
entries: map[uint]cron.EntryID{},
|
||||
snatchIPWait: defaultSnatchIPWait,
|
||||
snatchVars: map[uint]map[string]string{},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -371,9 +390,37 @@ func (s *TaskService) unschedule(taskID uint) {
|
||||
}
|
||||
}
|
||||
|
||||
// ApplyTenantCleanup 在租户删除事务提交后注销已删除任务并对齐 AI 探测任务。
|
||||
func (s *TaskService) ApplyTenantCleanup(ctx context.Context, taskIDs []uint, channelsGone bool) {
|
||||
for _, id := range taskIDs {
|
||||
s.unschedule(id)
|
||||
s.forgetSnatchVars(id)
|
||||
}
|
||||
if channelsGone {
|
||||
s.SyncAiProbeTask(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *TaskService) forgetSnatchVars(taskID uint) {
|
||||
s.snatchVarsMu.Lock()
|
||||
delete(s.snatchVars, taskID)
|
||||
s.snatchVarsMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *TaskService) lockTenantCleanup() func() {
|
||||
s.runMu.Lock()
|
||||
return s.runMu.Unlock
|
||||
}
|
||||
|
||||
// execute 执行一次任务:加载 → 分派 → 更新任务状态并写日志,
|
||||
// 前后状态交给通知判定,只在状态变化时推送。
|
||||
func (s *TaskService) execute(taskID uint) *model.TaskLog {
|
||||
s.runMu.RLock()
|
||||
defer s.runMu.RUnlock()
|
||||
return s.executeLocked(taskID)
|
||||
}
|
||||
|
||||
func (s *TaskService) executeLocked(taskID uint) *model.TaskLog {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), taskRunTimeout)
|
||||
defer cancel()
|
||||
task, err := s.GetTask(ctx, taskID)
|
||||
@@ -391,12 +438,40 @@ func (s *TaskService) execute(taskID uint) *model.TaskLog {
|
||||
task.LastError = oci.CompactError(runErr)
|
||||
message = task.LastError
|
||||
}
|
||||
s.db.Save(task)
|
||||
cur := taskSnapshot{Name: task.Name, Status: task.Status, LastError: task.LastError, Message: message}
|
||||
if !s.storeTaskRun(ctx, task) {
|
||||
return nil
|
||||
}
|
||||
cur := taskSnapshot{Name: task.Name, Status: task.Status, LastError: task.LastError,
|
||||
Message: message, SnatchVars: s.takeSnatchVars(task.ID)}
|
||||
s.notify(notifyEvents(prev, cur))
|
||||
return s.appendLog(task.ID, runErr == nil, message, time.Since(start))
|
||||
}
|
||||
|
||||
func (s *TaskService) storeTaskRun(ctx context.Context, task *model.Task) bool {
|
||||
stored, err := s.persistTaskRun(ctx, task)
|
||||
if err != nil {
|
||||
log.Printf("persist task run: %v", err)
|
||||
return false
|
||||
}
|
||||
return stored
|
||||
}
|
||||
|
||||
func (s *TaskService) persistTaskRun(ctx context.Context, task *model.Task) (bool, error) {
|
||||
updates := map[string]any{
|
||||
"status": task.Status, "last_run_at": task.LastRunAt,
|
||||
"last_error": task.LastError, "run_count": task.RunCount,
|
||||
}
|
||||
if task.Type == model.TaskTypeSnatch {
|
||||
updates["payload"] = task.Payload
|
||||
}
|
||||
res := s.db.WithContext(ctx).Model(&model.Task{}).
|
||||
Where("id = ? AND updated_at = ?", task.ID, task.UpdatedAt).Updates(updates)
|
||||
if res.Error != nil {
|
||||
return false, fmt.Errorf("persist task %d run: %w", task.ID, res.Error)
|
||||
}
|
||||
return res.RowsAffected == 1, nil
|
||||
}
|
||||
|
||||
// run 按类型分派任务执行。
|
||||
func (s *TaskService) run(ctx context.Context, task *model.Task) (string, error) {
|
||||
switch task.Type {
|
||||
@@ -650,12 +725,158 @@ func (s *TaskService) runSnatch(ctx context.Context, task *model.Task) (string,
|
||||
writeSnatchPayload(task, &p)
|
||||
return fmt.Sprintf("created %d (%s)%s, %d remaining", len(instances), strings.Join(ids, ","), adNote, remaining), nil
|
||||
}
|
||||
task.Status = model.TaskStatusSucceeded
|
||||
writeSnatchPayload(task, &p)
|
||||
s.unschedule(task.ID)
|
||||
s.snatchSucceed(ctx, task, &p, in, instances)
|
||||
return fmt.Sprintf("created %d: %s%s", len(instances), strings.Join(ids, ","), adNote), nil
|
||||
}
|
||||
|
||||
// snatchSucceed 处理抢满收尾:标记成功、停止调度,并组装成功通知变量暂存,
|
||||
// 供 execute 发通知时合并进 snatch_success 事件。
|
||||
func (s *TaskService) snatchSucceed(ctx context.Context, task *model.Task, p *snatchPayload, in oci.CreateInstanceInput, instances []oci.Instance) {
|
||||
task.Status = model.TaskStatusSucceeded
|
||||
writeSnatchPayload(task, p)
|
||||
s.unschedule(task.ID)
|
||||
s.storeSnatchVars(task.ID, s.snatchSuccessVars(ctx, p, in, instances))
|
||||
}
|
||||
|
||||
// snatchSuccessVars 组装抢机成功通知的模板变量;单项查询失败以占位值回退,
|
||||
// 绝不阻塞通知发送。密码只进通知变量,不落任务日志。
|
||||
func (s *TaskService) snatchSuccessVars(ctx context.Context, p *snatchPayload, in oci.CreateInstanceInput, instances []oci.Instance) map[string]string {
|
||||
tenant, region := "-", in.Region
|
||||
if cfg, err := s.configs.Get(ctx, p.OciConfigID); err == nil {
|
||||
tenant = cfg.Alias
|
||||
if region == "" {
|
||||
region = cfg.Region
|
||||
}
|
||||
}
|
||||
ips, pwds := s.snatchInstanceNet(ctx, p.OciConfigID, region, in, instances)
|
||||
return map[string]string{
|
||||
"tenant": tenant,
|
||||
"region": regionLabel(region),
|
||||
"shape": in.Shape,
|
||||
"spec": specLabel(in),
|
||||
"image": s.snatchImageLabel(ctx, p.OciConfigID, region, in),
|
||||
"ip": ips,
|
||||
"root_password": pwds,
|
||||
}
|
||||
}
|
||||
|
||||
// snatchInstanceNet 逐台等待公网 IPv4 并取回 root 密码;全部实例共享
|
||||
// snatchIPWait 总预算,多台结果按创建顺序以「、」连接。
|
||||
func (s *TaskService) snatchInstanceNet(ctx context.Context, cfgID uint, region string, in oci.CreateInstanceInput, instances []oci.Instance) (string, string) {
|
||||
deadline := time.Now().Add(s.snatchIPWait)
|
||||
ips := make([]string, 0, len(instances))
|
||||
pwds := make([]string, 0, len(instances))
|
||||
for _, inst := range instances {
|
||||
ips = append(ips, s.waitPublicIPv4(ctx, cfgID, region, inst.ID, deadline))
|
||||
pwds = append(pwds, rootPasswordOf(in, inst))
|
||||
}
|
||||
return strings.Join(ips, "、"), strings.Join(pwds, "、")
|
||||
}
|
||||
|
||||
// waitPublicIPv4 轮询实例公网 IPv4 直至分配或超过 deadline;
|
||||
// 超时回退内网 IPv4(标注内网),再无则「待分配」。
|
||||
func (s *TaskService) waitPublicIPv4(ctx context.Context, cfgID uint, region, instanceID string, deadline time.Time) string {
|
||||
private := ""
|
||||
for {
|
||||
inst, err := s.configs.Instance(ctx, cfgID, region, instanceID)
|
||||
if err == nil {
|
||||
if inst.PublicIP != "" {
|
||||
return inst.PublicIP
|
||||
}
|
||||
if inst.PrivateIP != "" {
|
||||
private = inst.PrivateIP
|
||||
}
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
break
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return waitIPFallback(private)
|
||||
case <-time.After(snatchIPPollInterval):
|
||||
}
|
||||
}
|
||||
return waitIPFallback(private)
|
||||
}
|
||||
|
||||
// waitIPFallback 是公网 IP 等待超时后的展示回退。
|
||||
func waitIPFallback(private string) string {
|
||||
if private != "" {
|
||||
return private + "(内网)"
|
||||
}
|
||||
return "待分配"
|
||||
}
|
||||
|
||||
// rootPasswordOf 取单台实例的 root 密码:随机生成模式从创建时写入的
|
||||
// FreeformTags 回读(每台独立),显式密码模式各台相同;SSH 密钥模式无密码。
|
||||
func rootPasswordOf(in oci.CreateInstanceInput, inst oci.Instance) string {
|
||||
if pwd := inst.FreeformTags["RootPassword"]; pwd != "" {
|
||||
return pwd
|
||||
}
|
||||
if in.RootPassword != "" {
|
||||
return in.RootPassword
|
||||
}
|
||||
return "—(SSH 密钥登录)"
|
||||
}
|
||||
|
||||
// specLabel 拼 Flex 规格,如 4C/24G;固定规格 shape(无 ocpus 参数)为 -。
|
||||
func specLabel(in oci.CreateInstanceInput) string {
|
||||
if in.Ocpus <= 0 {
|
||||
return "-"
|
||||
}
|
||||
return fmt.Sprintf("%gC/%gG", in.Ocpus, in.MemoryInGBs)
|
||||
}
|
||||
|
||||
// regionLabel 把区域标识换成控制台风格别名展示,如「Japan Central (Osaka)」;
|
||||
// 区域名与三字码均可匹配,表中未收录时回退原值。
|
||||
func regionLabel(region string) string {
|
||||
if region == "" {
|
||||
return "-"
|
||||
}
|
||||
regions, err := oci.AllRegions()
|
||||
if err != nil {
|
||||
return region
|
||||
}
|
||||
for _, r := range regions {
|
||||
if strings.EqualFold(r.Name, region) || strings.EqualFold(r.Key, region) {
|
||||
return r.Alias
|
||||
}
|
||||
}
|
||||
return region
|
||||
}
|
||||
|
||||
// snatchImageLabel 返回启动源展示名:镜像查显示名(失败回退 OCID),
|
||||
// 引导卷启动源固定文案。
|
||||
func (s *TaskService) snatchImageLabel(ctx context.Context, cfgID uint, region string, in oci.CreateInstanceInput) string {
|
||||
if in.BootVolumeID != "" {
|
||||
return "引导卷"
|
||||
}
|
||||
if in.ImageID == "" {
|
||||
return "-"
|
||||
}
|
||||
img, err := s.configs.Image(ctx, cfgID, region, in.ImageID)
|
||||
if err != nil || img.DisplayName == "" {
|
||||
return in.ImageID
|
||||
}
|
||||
return img.DisplayName
|
||||
}
|
||||
|
||||
// storeSnatchVars 暂存抢机成功通知变量,由 takeSnatchVars 取走。
|
||||
func (s *TaskService) storeSnatchVars(taskID uint, vars map[string]string) {
|
||||
s.snatchVarsMu.Lock()
|
||||
defer s.snatchVarsMu.Unlock()
|
||||
s.snatchVars[taskID] = vars
|
||||
}
|
||||
|
||||
// takeSnatchVars 取走并清除暂存的成功通知变量(无则 nil)。
|
||||
func (s *TaskService) takeSnatchVars(taskID uint) map[string]string {
|
||||
s.snatchVarsMu.Lock()
|
||||
defer s.snatchVarsMu.Unlock()
|
||||
vars := s.snatchVars[taskID]
|
||||
delete(s.snatchVars, taskID)
|
||||
return vars
|
||||
}
|
||||
|
||||
// snatchInstanceInput 组装本次创建参数:可用域显式指定时原样使用;
|
||||
// 留空(自动)时按执行序号轮询区域全部可用域——ad-1、ad-2、ad-3 依次循环,
|
||||
// 分摊单可用域容量不足。附加说明串供执行日志展示本次所用可用域。
|
||||
@@ -735,6 +956,9 @@ type taskSnapshot struct {
|
||||
Status string
|
||||
LastError string
|
||||
Message string // 本次执行结果摘要,仅执行后快照填写
|
||||
// SnatchVars 是抢机成功通知的补充变量(租户/区域/实例明细),
|
||||
// 仅抢满那次执行后快照携带
|
||||
SnatchVars map[string]string
|
||||
}
|
||||
|
||||
// notifyKind 是通知事件类型,与设置页「通知管理」开关一一对应。
|
||||
@@ -760,7 +984,11 @@ type notifyEvent struct {
|
||||
func notifyEvents(prev, cur taskSnapshot) []notifyEvent {
|
||||
var events []notifyEvent
|
||||
if prev.Status != model.TaskStatusSucceeded && cur.Status == model.TaskStatusSucceeded {
|
||||
events = append(events, notifyEvent{notifySnatchSuccess, map[string]string{"task_name": cur.Name, "message": cur.Message}})
|
||||
vars := map[string]string{"task_name": cur.Name, "message": cur.Message}
|
||||
for k, v := range cur.SnatchVars {
|
||||
vars[k] = v
|
||||
}
|
||||
events = append(events, notifyEvent{notifySnatchSuccess, vars})
|
||||
}
|
||||
switch {
|
||||
case prev.Status != model.TaskStatusFailed && cur.Status == model.TaskStatusFailed:
|
||||
|
||||
@@ -80,7 +80,10 @@ func newTaskEnv(t *testing.T, client oci.Client) (*TaskService, *OciConfigServic
|
||||
t.Fatalf("new cipher: %v", err)
|
||||
}
|
||||
configs := NewOciConfigService(db, cipher, client)
|
||||
return NewTaskService(db, configs, nil, NewSettingService(db, cipher)), configs, db
|
||||
tasks := NewTaskService(db, configs, nil, NewSettingService(db, cipher))
|
||||
// 测试不等待公网 IP 分配:预算置 0,首查后立即回退
|
||||
tasks.snatchIPWait = 0
|
||||
return tasks, configs, db
|
||||
}
|
||||
|
||||
func TestCreateTaskValidation(t *testing.T) {
|
||||
@@ -699,3 +702,102 @@ func TestUpdateTaskFailedToActiveResetsAuthFail(t *testing.T) {
|
||||
t.Error("重新启用后任务未回到调度")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSnatchSuccessVars 锁定抢机成功通知变量的组装:租户/区域/类型/镜像,
|
||||
// 以及 root 密码三种凭据模式与 IP 逐级回退;单项查询失败不阻塞其余字段。
|
||||
func TestSnatchSuccessVars(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
client *fakeClient
|
||||
in oci.CreateInstanceInput
|
||||
insts []oci.Instance
|
||||
want map[string]string
|
||||
}{
|
||||
{
|
||||
name: "生成密码模式全字段",
|
||||
client: &fakeClient{
|
||||
tenancy: oci.TenancyInfo{Name: "mytenancy", HomeRegionKey: "FRA"},
|
||||
instance: oci.Instance{PublicIP: "203.0.113.7"},
|
||||
image: oci.Image{DisplayName: "Ubuntu-24.04"},
|
||||
},
|
||||
in: oci.CreateInstanceInput{Region: "ap-osaka-1", Shape: "VM.Standard.A1.Flex",
|
||||
Ocpus: 4, MemoryInGBs: 24, ImageID: "ocid1.image.oc1..img"},
|
||||
insts: []oci.Instance{{ID: "i1", FreeformTags: map[string]string{"RootPassword": "pw-from-tag"}}},
|
||||
want: map[string]string{
|
||||
"tenant": "试用期",
|
||||
"region": "Japan Central (Osaka)",
|
||||
"shape": "VM.Standard.A1.Flex",
|
||||
"spec": "4C/24G",
|
||||
"image": "Ubuntu-24.04",
|
||||
"ip": "203.0.113.7",
|
||||
"root_password": "pw-from-tag",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "显式密码_未知区域_IP回退内网_镜像失败回退OCID",
|
||||
client: &fakeClient{
|
||||
instance: oci.Instance{PrivateIP: "10.0.0.5"},
|
||||
imageErr: fmt.Errorf("boom"),
|
||||
},
|
||||
in: oci.CreateInstanceInput{Region: "xx-nowhere-1", Shape: "VM.Standard.E2.1.Micro",
|
||||
ImageID: "ocid1.image.oc1..img", RootPassword: "explicit-pw"},
|
||||
insts: []oci.Instance{{ID: "i1"}},
|
||||
want: map[string]string{
|
||||
"tenant": "试用期",
|
||||
"region": "xx-nowhere-1",
|
||||
"shape": "VM.Standard.E2.1.Micro",
|
||||
"spec": "-",
|
||||
"image": "ocid1.image.oc1..img",
|
||||
"ip": "10.0.0.5(内网)",
|
||||
"root_password": "explicit-pw",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "SSH模式_引导卷_区域取配置默认_多台待分配",
|
||||
client: &fakeClient{instanceErr: fmt.Errorf("boom")},
|
||||
in: oci.CreateInstanceInput{Shape: "s",
|
||||
BootVolumeID: "ocid1.bootvolume.oc1..bv", SSHPublicKey: "ssh-ed25519 AAAA"},
|
||||
insts: []oci.Instance{{ID: "i1"}, {ID: "i2"}},
|
||||
want: map[string]string{
|
||||
"region": "Germany Central (Frankfurt)",
|
||||
"image": "引导卷",
|
||||
"ip": "待分配、待分配",
|
||||
"root_password": "—(SSH 密钥登录)、—(SSH 密钥登录)",
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
tasks, configs, _ := newTaskEnv(t, tt.client)
|
||||
cfg, err := configs.Import(context.Background(), trialImportInput())
|
||||
if err != nil {
|
||||
t.Fatalf("import: %v", err)
|
||||
}
|
||||
got := tasks.snatchSuccessVars(context.Background(),
|
||||
&snatchPayload{OciConfigID: cfg.ID}, tt.in, tt.insts)
|
||||
for k, want := range tt.want {
|
||||
if got[k] != want {
|
||||
t.Errorf("%s = %q, want %q", k, got[k], want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyEventsSnatchVarsMerged 抢满成功事件应合并暂存的补充变量。
|
||||
func TestNotifyEventsSnatchVarsMerged(t *testing.T) {
|
||||
prev := taskSnapshot{Name: "抢机", Status: model.TaskStatusActive}
|
||||
cur := taskSnapshot{Name: "抢机", Status: model.TaskStatusSucceeded, Message: "created 1",
|
||||
SnatchVars: map[string]string{"tenant": "免费01·t", "ip": "203.0.113.7", "root_password": "pw"}}
|
||||
events := notifyEvents(prev, cur)
|
||||
if len(events) != 1 || events[0].Kind != notifySnatchSuccess {
|
||||
t.Fatalf("events = %+v, want 1 条 snatch_success", events)
|
||||
}
|
||||
want := map[string]string{"task_name": "抢机", "message": "created 1",
|
||||
"tenant": "免费01·t", "ip": "203.0.113.7", "root_password": "pw"}
|
||||
for k, v := range want {
|
||||
if events[0].Vars[k] != v {
|
||||
t.Errorf("vars[%s] = %q, want %q", k, events[0].Vars[k], v)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -31,27 +31,36 @@ func (s *OciConfigService) credentialsAndHomeRegion(ctx context.Context, id uint
|
||||
return cred, homeRegionOf(cfg), nil
|
||||
}
|
||||
|
||||
// TenantUsers 列出租户 IAM 用户。
|
||||
func (s *OciConfigService) TenantUsers(ctx context.Context, id uint) ([]oci.TenantUser, error) {
|
||||
cred, err := s.credentialsByID(ctx, id)
|
||||
// IdentityDomains 列出租户全部 ACTIVE 身份域(域选择器数据源)。
|
||||
func (s *OciConfigService) IdentityDomains(ctx context.Context, id uint) ([]oci.IdentityDomain, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.ListTenantUsers(ctx, cred)
|
||||
return s.client.ListIdentityDomains(ctx, cred, homeRegion)
|
||||
}
|
||||
|
||||
// TenantUsers 列出租户 IAM 用户;domainID 非空只列该身份域。
|
||||
func (s *OciConfigService) TenantUsers(ctx context.Context, id uint, domainID string) ([]oci.TenantUser, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.ListTenantUsers(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// TenantUserDetail 查用户的域档案与管理员状态(编辑表单预填充用)。
|
||||
func (s *OciConfigService) TenantUserDetail(ctx context.Context, id uint, userID string) (oci.TenantUserDetail, error) {
|
||||
func (s *OciConfigService) TenantUserDetail(ctx context.Context, id uint, domainID, userID string) (oci.TenantUserDetail, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return oci.TenantUserDetail{}, err
|
||||
}
|
||||
return s.client.GetTenantUserDetail(ctx, cred, homeRegion, userID)
|
||||
return s.client.GetTenantUserDetail(ctx, cred, homeRegion, domainID, userID)
|
||||
}
|
||||
|
||||
// CreateTenantUser 新增用户;名字与姓氏必填(域用户档案要求),
|
||||
// 描述缺省用「名 姓」。
|
||||
func (s *OciConfigService) CreateTenantUser(ctx context.Context, id uint, in oci.CreateTenantUserInput) (oci.TenantUser, error) {
|
||||
func (s *OciConfigService) CreateTenantUser(ctx context.Context, id uint, domainID string, in oci.CreateTenantUserInput) (oci.TenantUser, error) {
|
||||
if strings.TrimSpace(in.Name) == "" {
|
||||
return oci.TenantUser{}, fmt.Errorf("create tenant user: name is required")
|
||||
}
|
||||
@@ -65,11 +74,11 @@ func (s *OciConfigService) CreateTenantUser(ctx context.Context, id uint, in oci
|
||||
if err != nil {
|
||||
return oci.TenantUser{}, err
|
||||
}
|
||||
return s.client.CreateTenantUser(ctx, cred, homeRegion, in)
|
||||
return s.client.CreateTenantUser(ctx, cred, homeRegion, domainID, in)
|
||||
}
|
||||
|
||||
// UpdateTenantUser 编辑用户资料;nil 字段不修改,全 nil 报错。
|
||||
func (s *OciConfigService) UpdateTenantUser(ctx context.Context, id uint, userID string, in oci.UpdateTenantUserInput) (oci.TenantUser, error) {
|
||||
func (s *OciConfigService) UpdateTenantUser(ctx context.Context, id uint, domainID, userID string, in oci.UpdateTenantUserInput) (oci.TenantUser, error) {
|
||||
hasField := in.Description != nil || in.Email != nil || in.GivenName != nil || in.FamilyName != nil
|
||||
hasAdmin := in.GrantDomainAdmin != nil || in.AddToAdminGroup != nil
|
||||
if !hasField && !hasAdmin {
|
||||
@@ -79,29 +88,29 @@ func (s *OciConfigService) UpdateTenantUser(ctx context.Context, id uint, userID
|
||||
if err != nil {
|
||||
return oci.TenantUser{}, err
|
||||
}
|
||||
return s.client.UpdateTenantUser(ctx, cred, homeRegion, userID, in)
|
||||
return s.client.UpdateTenantUser(ctx, cred, homeRegion, domainID, userID, in)
|
||||
}
|
||||
|
||||
// IdentitySetting 读取域身份设置(当前只透出主邮箱必填开关)。
|
||||
func (s *OciConfigService) IdentitySetting(ctx context.Context, id uint) (oci.IdentitySettingInfo, error) {
|
||||
func (s *OciConfigService) IdentitySetting(ctx context.Context, id uint, domainID string) (oci.IdentitySettingInfo, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return oci.IdentitySettingInfo{}, err
|
||||
}
|
||||
return s.client.GetIdentitySetting(ctx, cred, homeRegion)
|
||||
return s.client.GetIdentitySetting(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// UpdateIdentitySetting 修改「用户需要提供主电子邮件地址」开关。
|
||||
func (s *OciConfigService) UpdateIdentitySetting(ctx context.Context, id uint, primaryEmailRequired bool) (oci.IdentitySettingInfo, error) {
|
||||
func (s *OciConfigService) UpdateIdentitySetting(ctx context.Context, id uint, domainID string, primaryEmailRequired bool) (oci.IdentitySettingInfo, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return oci.IdentitySettingInfo{}, err
|
||||
}
|
||||
return s.client.UpdateIdentitySetting(ctx, cred, homeRegion, primaryEmailRequired)
|
||||
return s.client.UpdateIdentitySetting(ctx, cred, homeRegion, domainID, primaryEmailRequired)
|
||||
}
|
||||
|
||||
// DeleteTenantUser 删除 IAM 用户;拒绝删除当前配置正在使用的用户。
|
||||
func (s *OciConfigService) DeleteTenantUser(ctx context.Context, id uint, userID string) error {
|
||||
func (s *OciConfigService) DeleteTenantUser(ctx context.Context, id uint, domainID, userID string) error {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -109,25 +118,25 @@ func (s *OciConfigService) DeleteTenantUser(ctx context.Context, id uint, userID
|
||||
if userID == cred.UserOCID {
|
||||
return fmt.Errorf("delete tenant user: refusing to delete the user this config signs requests with")
|
||||
}
|
||||
return s.client.DeleteTenantUser(ctx, cred, homeRegion, userID)
|
||||
return s.client.DeleteTenantUser(ctx, cred, homeRegion, domainID, userID)
|
||||
}
|
||||
|
||||
// ResetTenantUserPassword 重置用户控制台密码,返回一次性新密码。
|
||||
func (s *OciConfigService) ResetTenantUserPassword(ctx context.Context, id uint, userID string) (string, error) {
|
||||
func (s *OciConfigService) ResetTenantUserPassword(ctx context.Context, id uint, domainID, userID string) (string, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return s.client.ResetTenantUserPassword(ctx, cred, homeRegion, userID)
|
||||
return s.client.ResetTenantUserPassword(ctx, cred, homeRegion, domainID, userID)
|
||||
}
|
||||
|
||||
// DeleteTenantUserMfa 清除用户全部 MFA,返回删除的经典 TOTP 设备数。
|
||||
func (s *OciConfigService) DeleteTenantUserMfa(ctx context.Context, id uint, userID string) (int, error) {
|
||||
func (s *OciConfigService) DeleteTenantUserMfa(ctx context.Context, id uint, domainID, userID string) (int, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return s.client.DeleteTenantUserMfaDevices(ctx, cred, homeRegion, userID)
|
||||
return s.client.DeleteTenantUserMfaDevices(ctx, cred, homeRegion, domainID, userID)
|
||||
}
|
||||
|
||||
// DeleteTenantUserApiKeys 删除用户的 API Key;默认保留当前配置使用的指纹。
|
||||
@@ -140,17 +149,17 @@ func (s *OciConfigService) DeleteTenantUserApiKeys(ctx context.Context, id uint,
|
||||
}
|
||||
|
||||
// NotificationRecipients 查询域通知收件人设置。
|
||||
func (s *OciConfigService) NotificationRecipients(ctx context.Context, id uint) (oci.NotificationRecipients, error) {
|
||||
func (s *OciConfigService) NotificationRecipients(ctx context.Context, id uint, domainID string) (oci.NotificationRecipients, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return oci.NotificationRecipients{}, err
|
||||
}
|
||||
return s.client.GetNotificationRecipients(ctx, cred, homeRegion)
|
||||
return s.client.GetNotificationRecipients(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// UpdateNotificationRecipients 把域通知改为只发给指定收件人;
|
||||
// 收件人为空表示关闭 test mode 恢复默认发送。
|
||||
func (s *OciConfigService) UpdateNotificationRecipients(ctx context.Context, id uint, emails []string) (oci.NotificationRecipients, error) {
|
||||
func (s *OciConfigService) UpdateNotificationRecipients(ctx context.Context, id uint, domainID string, emails []string) (oci.NotificationRecipients, error) {
|
||||
if err := validateEmails(emails); err != nil {
|
||||
return oci.NotificationRecipients{}, err
|
||||
}
|
||||
@@ -158,7 +167,7 @@ func (s *OciConfigService) UpdateNotificationRecipients(ctx context.Context, id
|
||||
if err != nil {
|
||||
return oci.NotificationRecipients{}, err
|
||||
}
|
||||
return s.client.UpdateNotificationRecipients(ctx, cred, homeRegion, emails)
|
||||
return s.client.UpdateNotificationRecipients(ctx, cred, homeRegion, domainID, emails)
|
||||
}
|
||||
|
||||
func validateEmails(emails []string) error {
|
||||
@@ -171,16 +180,16 @@ func validateEmails(emails []string) error {
|
||||
}
|
||||
|
||||
// PasswordPolicies 列出域密码策略。
|
||||
func (s *OciConfigService) PasswordPolicies(ctx context.Context, id uint) ([]oci.PasswordPolicyInfo, error) {
|
||||
func (s *OciConfigService) PasswordPolicies(ctx context.Context, id uint, domainID string) ([]oci.PasswordPolicyInfo, error) {
|
||||
cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.client.ListPasswordPolicies(ctx, cred, homeRegion)
|
||||
return s.client.ListPasswordPolicies(ctx, cred, homeRegion, domainID)
|
||||
}
|
||||
|
||||
// UpdatePasswordPolicy 修改域密码策略的过期天数或最小长度。
|
||||
func (s *OciConfigService) UpdatePasswordPolicy(ctx context.Context, id uint, policyID string, in oci.UpdatePasswordPolicyInput) (oci.PasswordPolicyInfo, error) {
|
||||
func (s *OciConfigService) UpdatePasswordPolicy(ctx context.Context, id uint, domainID, policyID string, in oci.UpdatePasswordPolicyInput) (oci.PasswordPolicyInfo, error) {
|
||||
if in.PasswordExpiresAfter == nil && in.MinLength == nil {
|
||||
return oci.PasswordPolicyInfo{}, fmt.Errorf("update password policy: nothing to update")
|
||||
}
|
||||
@@ -188,5 +197,5 @@ func (s *OciConfigService) UpdatePasswordPolicy(ctx context.Context, id uint, po
|
||||
if err != nil {
|
||||
return oci.PasswordPolicyInfo{}, err
|
||||
}
|
||||
return s.client.UpdatePasswordPolicy(ctx, cred, homeRegion, policyID, in)
|
||||
return s.client.UpdatePasswordPolicy(ctx, cred, homeRegion, domainID, policyID, in)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,457 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
type tenantDeleteResult struct {
|
||||
config model.OciConfig
|
||||
deletedTaskIDs []uint
|
||||
alertRuleIDs []uint
|
||||
channelsGone bool
|
||||
}
|
||||
|
||||
type tenantTaskAction struct {
|
||||
task model.Task
|
||||
deleteTask bool
|
||||
payload string
|
||||
}
|
||||
|
||||
type tenancyCacheInvalidator interface {
|
||||
InvalidateTenancy(tenancyOCID string)
|
||||
}
|
||||
|
||||
func (s *OciConfigService) deleteTenant(ctx context.Context, id uint) (*tenantDeleteResult, error) {
|
||||
result := &tenantDeleteResult{}
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return s.deleteTenantInTx(tx, id, result)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("delete oci config %d: %w", id, err)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *OciConfigService) deleteTenantInTx(tx *gorm.DB, id uint, result *tenantDeleteResult) error {
|
||||
if err := lockTenant(tx, id, &result.config); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.deleteTenantTasks(tx, id, result); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteTenantEvents(tx, id, result); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteTenantAI(tx, id, result); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteTenantSnapshots(tx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
return deleteTenantConfig(tx, id)
|
||||
}
|
||||
|
||||
func lockTenant(tx *gorm.DB, id uint, cfg *model.OciConfig) error {
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(cfg, id).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("load tenant: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *OciConfigService) deleteTenantTasks(tx *gorm.DB, id uint, result *tenantDeleteResult) error {
|
||||
var tasks []model.Task
|
||||
err := tx.Where("type IN ?", []string{
|
||||
model.TaskTypeSnatch, model.TaskTypeHealthCheck, model.TaskTypeCost,
|
||||
}).Find(&tasks).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("load tenant tasks: %w", err)
|
||||
}
|
||||
for i := range tasks {
|
||||
action, ok, err := planTenantTask(tasks[i], id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if err := applyTenantTask(tx, action, result); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func planTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
|
||||
switch task.Type {
|
||||
case model.TaskTypeSnatch:
|
||||
return planSnatchTask(task, id)
|
||||
case model.TaskTypeHealthCheck, model.TaskTypeCost:
|
||||
return planMultiTenantTask(task, id)
|
||||
default:
|
||||
return tenantTaskAction{}, false, nil
|
||||
}
|
||||
}
|
||||
|
||||
func planSnatchTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
|
||||
var payload snatchPayload
|
||||
if err := json.Unmarshal([]byte(task.Payload), &payload); err != nil {
|
||||
return tenantTaskAction{}, false, fmt.Errorf("parse task %d payload: %w", task.ID, err)
|
||||
}
|
||||
if payload.OciConfigID != id {
|
||||
return tenantTaskAction{}, false, nil
|
||||
}
|
||||
return tenantTaskAction{task: task, deleteTask: true}, true, nil
|
||||
}
|
||||
|
||||
func planMultiTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) {
|
||||
payload, err := decodeConfigIDs(task)
|
||||
if err != nil {
|
||||
return tenantTaskAction{}, false, err
|
||||
}
|
||||
if len(payload.OciConfigIDs) == 0 {
|
||||
return tenantTaskAction{task: task, payload: task.Payload}, true, nil
|
||||
}
|
||||
remaining, found := removeConfigID(payload.OciConfigIDs, id)
|
||||
if !found {
|
||||
return tenantTaskAction{}, false, nil
|
||||
}
|
||||
if len(remaining) == 0 {
|
||||
return tenantTaskAction{task: task, deleteTask: true}, true, nil
|
||||
}
|
||||
payload.OciConfigIDs = remaining
|
||||
raw, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return tenantTaskAction{}, false, fmt.Errorf("encode task %d payload: %w", task.ID, err)
|
||||
}
|
||||
return tenantTaskAction{task: task, payload: string(raw)}, true, nil
|
||||
}
|
||||
|
||||
func decodeConfigIDs(task model.Task) (healthCheckPayload, error) {
|
||||
var payload healthCheckPayload
|
||||
if strings.TrimSpace(task.Payload) == "" {
|
||||
return payload, nil
|
||||
}
|
||||
if err := json.Unmarshal([]byte(task.Payload), &payload); err != nil {
|
||||
return payload, fmt.Errorf("parse task %d payload: %w", task.ID, err)
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func removeConfigID(ids []uint, id uint) ([]uint, bool) {
|
||||
remaining := make([]uint, 0, len(ids))
|
||||
found := false
|
||||
for _, candidate := range ids {
|
||||
if candidate == id {
|
||||
found = true
|
||||
continue
|
||||
}
|
||||
remaining = append(remaining, candidate)
|
||||
}
|
||||
return remaining, found
|
||||
}
|
||||
|
||||
func applyTenantTask(tx *gorm.DB, action tenantTaskAction, result *tenantDeleteResult) error {
|
||||
if err := deleteTaskHistory(tx, action.task.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
if action.deleteTask {
|
||||
if err := tx.Delete(&model.Task{}, action.task.ID).Error; err != nil {
|
||||
return fmt.Errorf("delete task %d: %w", action.task.ID, err)
|
||||
}
|
||||
result.deletedTaskIDs = append(result.deletedTaskIDs, action.task.ID)
|
||||
return nil
|
||||
}
|
||||
return resetTenantTask(tx, action)
|
||||
}
|
||||
|
||||
func deleteTaskHistory(tx *gorm.DB, taskID uint) error {
|
||||
if err := tx.Where("task_id = ?", taskID).Delete(&model.TaskLog{}).Error; err != nil {
|
||||
return fmt.Errorf("delete task %d logs: %w", taskID, err)
|
||||
}
|
||||
key := deadAliasKey(taskID)
|
||||
if err := tx.Where("key = ?", key).Delete(&model.Setting{}).Error; err != nil {
|
||||
return fmt.Errorf("delete task %d state: %w", taskID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func resetTenantTask(tx *gorm.DB, action tenantTaskAction) error {
|
||||
updates := map[string]any{
|
||||
"payload": action.payload, "last_run_at": nil,
|
||||
"last_error": "", "run_count": 0,
|
||||
}
|
||||
err := tx.Model(&model.Task{}).Where("id = ?", action.task.ID).Updates(updates).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("reset task %d: %w", action.task.ID, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteTenantEvents(tx *gorm.DB, id uint, result *tenantDeleteResult) error {
|
||||
ruleIDs, eventIDs, affectedRules, err := loadTenantEventRefs(tx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := deleteAlertHits(tx, ruleIDs, eventIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
result.alertRuleIDs = mergeIDs(ruleIDs, affectedRules)
|
||||
if err := deleteWhere(tx, &model.AlertRule{}, "oci_config_id = ?", id); err != nil {
|
||||
return fmt.Errorf("delete tenant alert rules: %w", err)
|
||||
}
|
||||
if err := deleteWhere(tx, &model.LogEvent{}, "oci_config_id = ?", id); err != nil {
|
||||
return fmt.Errorf("delete tenant log events: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadTenantEventRefs(tx *gorm.DB, id uint) ([]uint, []uint, []uint, error) {
|
||||
ruleIDs, err := lockedTenantRuleIDs(tx, id)
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("load tenant alert rules: %w", err)
|
||||
}
|
||||
eventIDs, err := lockedTenantEventIDs(tx, id)
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("load tenant log events: %w", err)
|
||||
}
|
||||
affected, err := alertHitRuleIDs(tx, eventIDs)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
return ruleIDs, eventIDs, affected, nil
|
||||
}
|
||||
|
||||
func lockedTenantRuleIDs(tx *gorm.DB, id uint) ([]uint, error) {
|
||||
var rows []model.AlertRule
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").
|
||||
Where("oci_config_id = ?", id).Order("id").Find(&rows).Error
|
||||
ids := make([]uint, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
ids = append(ids, row.ID)
|
||||
}
|
||||
return ids, err
|
||||
}
|
||||
|
||||
func lockedTenantEventIDs(tx *gorm.DB, id uint) ([]uint, error) {
|
||||
var rows []model.LogEvent
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").
|
||||
Where("oci_config_id = ?", id).Order("id").Find(&rows).Error
|
||||
ids := make([]uint, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
ids = append(ids, row.ID)
|
||||
}
|
||||
return ids, err
|
||||
}
|
||||
|
||||
func alertHitRuleIDs(tx *gorm.DB, eventIDs []uint) ([]uint, error) {
|
||||
if len(eventIDs) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
ids := make([]uint, 0)
|
||||
err := tx.Model(&model.AlertRuleHit{}).Where("log_event_id IN ?", eventIDs).
|
||||
Distinct().Pluck("rule_id", &ids).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load affected alert rules: %w", err)
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func mergeIDs(groups ...[]uint) []uint {
|
||||
seen := make(map[uint]struct{})
|
||||
out := make([]uint, 0)
|
||||
for _, ids := range groups {
|
||||
for _, id := range ids {
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
out = append(out, id)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func deleteAlertHits(tx *gorm.DB, ruleIDs, eventIDs []uint) error {
|
||||
query := tx.Model(&model.AlertRuleHit{})
|
||||
switch {
|
||||
case len(ruleIDs) > 0 && len(eventIDs) > 0:
|
||||
query = query.Where("rule_id IN ? OR log_event_id IN ?", ruleIDs, eventIDs)
|
||||
case len(ruleIDs) > 0:
|
||||
query = query.Where("rule_id IN ?", ruleIDs)
|
||||
case len(eventIDs) > 0:
|
||||
query = query.Where("log_event_id IN ?", eventIDs)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
if err := query.Delete(&model.AlertRuleHit{}).Error; err != nil {
|
||||
return fmt.Errorf("delete tenant alert hits: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteTenantAI(tx *gorm.DB, id uint, result *tenantDeleteResult) error {
|
||||
channelIDs, err := lockedAiChannelIDs(tx, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("load tenant AI channels: %w", err)
|
||||
}
|
||||
if len(channelIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
callIDs, err := lockedAiCallIDs(tx, channelIDs)
|
||||
if err != nil {
|
||||
return fmt.Errorf("load tenant AI calls: %w", err)
|
||||
}
|
||||
if err := deleteTenantAIRows(tx, channelIDs, callIDs); err != nil {
|
||||
return err
|
||||
}
|
||||
result.channelsGone = true
|
||||
return reconcileAiProbeRows(tx, result)
|
||||
}
|
||||
|
||||
func lockedAiChannelIDs(tx *gorm.DB, id uint) ([]uint, error) {
|
||||
var rows []model.AiChannel
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").
|
||||
Where("oci_config_id = ?", id).Order("id").Find(&rows).Error
|
||||
ids := make([]uint, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
ids = append(ids, row.ID)
|
||||
}
|
||||
return ids, err
|
||||
}
|
||||
|
||||
func lockedAiCallIDs(tx *gorm.DB, channelIDs []uint) ([]uint, error) {
|
||||
var rows []model.AiCallLog
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").
|
||||
Where("channel_id IN ?", channelIDs).Order("id").Find(&rows).Error
|
||||
ids := make([]uint, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
ids = append(ids, row.ID)
|
||||
}
|
||||
return ids, err
|
||||
}
|
||||
|
||||
func deleteTenantAIRows(tx *gorm.DB, channelIDs, callIDs []uint) error {
|
||||
if len(callIDs) > 0 {
|
||||
if err := deleteWhere(tx, &model.AiContentLog{}, "call_log_id IN ?", callIDs); err != nil {
|
||||
return fmt.Errorf("delete tenant AI content logs: %w", err)
|
||||
}
|
||||
}
|
||||
steps := []struct {
|
||||
name string
|
||||
value any
|
||||
}{
|
||||
{"AI call logs", &model.AiCallLog{}},
|
||||
{"AI model cache", &model.AiModelCache{}},
|
||||
{"AI channels", &model.AiChannel{}},
|
||||
}
|
||||
for _, step := range steps {
|
||||
if err := deleteAIChannelRows(tx, step.value, channelIDs); err != nil {
|
||||
return fmt.Errorf("delete tenant %s: %w", step.name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteAIChannelRows(tx *gorm.DB, value any, channelIDs []uint) error {
|
||||
column := "channel_id"
|
||||
if _, ok := value.(*model.AiChannel); ok {
|
||||
column = "id"
|
||||
}
|
||||
return deleteWhere(tx, value, column+" IN ?", channelIDs)
|
||||
}
|
||||
|
||||
func reconcileAiProbeRows(tx *gorm.DB, result *tenantDeleteResult) error {
|
||||
var channelCount int64
|
||||
if err := tx.Model(&model.AiChannel{}).Count(&channelCount).Error; err != nil {
|
||||
return fmt.Errorf("count remaining AI channels: %w", err)
|
||||
}
|
||||
var ids []uint
|
||||
err := tx.Model(&model.Task{}).Where("type = ?", model.TaskTypeAiProbe).Pluck("id", &ids).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("load AI probe task: %w", err)
|
||||
}
|
||||
for _, id := range ids {
|
||||
if err := deleteTaskHistory(tx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
if channelCount == 0 {
|
||||
if err := tx.Delete(&model.Task{}, id).Error; err != nil {
|
||||
return fmt.Errorf("delete AI probe task %d: %w", id, err)
|
||||
}
|
||||
result.deletedTaskIDs = append(result.deletedTaskIDs, id)
|
||||
continue
|
||||
}
|
||||
if err := resetTaskHistoryFields(tx, id); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func resetTaskHistoryFields(tx *gorm.DB, id uint) error {
|
||||
updates := map[string]any{"last_run_at": nil, "last_error": "", "run_count": 0}
|
||||
if err := tx.Model(&model.Task{}).Where("id = ?", id).Updates(updates).Error; err != nil {
|
||||
return fmt.Errorf("reset task %d history: %w", id, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteTenantSnapshots(tx *gorm.DB, id uint) error {
|
||||
steps := []struct {
|
||||
name string
|
||||
value any
|
||||
}{
|
||||
{"check snapshots", &model.CheckSnapshot{}},
|
||||
{"cost snapshots", &model.CostSnapshot{}},
|
||||
{"region cache", &model.RegionCache{}},
|
||||
{"compartment cache", &model.CompartmentCache{}},
|
||||
}
|
||||
for _, step := range steps {
|
||||
if err := deleteWhere(tx, step.value, "oci_config_id = ?", id); err != nil {
|
||||
return fmt.Errorf("delete tenant %s: %w", step.name, err)
|
||||
}
|
||||
}
|
||||
if err := tx.Where("key = ?", secretKey(id)).Delete(&model.Setting{}).Error; err != nil {
|
||||
return fmt.Errorf("delete tenant webhook secret: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteTenantConfig(tx *gorm.DB, id uint) error {
|
||||
res := tx.Delete(&model.OciConfig{}, id)
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("delete tenant row: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteWhere(tx *gorm.DB, value any, query string, args ...any) error {
|
||||
return tx.Where(query, args...).Delete(value).Error
|
||||
}
|
||||
|
||||
func (s *OciConfigService) afterTenantDelete(ctx context.Context, result *tenantDeleteResult) {
|
||||
// 数据已提交,客户端断开不应中断必需的运行时对齐。
|
||||
ctx = context.WithoutCancel(ctx)
|
||||
s.InvalidateAuditCache(result.config.ID)
|
||||
if client, ok := s.client.(tenancyCacheInvalidator); ok {
|
||||
client.InvalidateTenancy(result.config.TenancyOCID)
|
||||
}
|
||||
if s.cleanupEvents != nil {
|
||||
s.cleanupEvents.ClearAlertCooldown(result.alertRuleIDs)
|
||||
}
|
||||
if s.cleanupTasks != nil {
|
||||
s.cleanupTasks.ApplyTenantCleanup(ctx, result.deletedTaskIDs, result.channelsGone)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,609 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"oci-portal/internal/crypto"
|
||||
"oci-portal/internal/model"
|
||||
"oci-portal/internal/oci"
|
||||
)
|
||||
|
||||
type invalidatingClient struct {
|
||||
*fakeClient
|
||||
invalidated []string
|
||||
}
|
||||
|
||||
type blockingClient struct {
|
||||
*fakeClient
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (c *blockingClient) ValidateKey(context.Context, oci.Credentials) (oci.TenancyInfo, error) {
|
||||
c.started <- struct{}{}
|
||||
<-c.release
|
||||
return oci.TenancyInfo{Name: "target"}, nil
|
||||
}
|
||||
|
||||
func (c *invalidatingClient) InvalidateTenancy(id string) {
|
||||
c.invalidated = append(c.invalidated, id)
|
||||
}
|
||||
|
||||
func newTenantDeleteEnv(t *testing.T, client oci.Client) (*OciConfigService, *TaskService, *gorm.DB) {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
t.Fatalf("database handle: %v", err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
migrateTenantDeleteModels(t, db)
|
||||
cipher, err := crypto.NewCipher("test-data-key")
|
||||
if err != nil {
|
||||
t.Fatalf("new cipher: %v", err)
|
||||
}
|
||||
configs := NewOciConfigService(db, cipher, client)
|
||||
tasks := NewTaskService(db, configs, nil, nil)
|
||||
configs.SetTenantCleanupDeps(tasks, nil)
|
||||
return configs, tasks, db
|
||||
}
|
||||
|
||||
func migrateTenantDeleteModels(t *testing.T, db *gorm.DB) {
|
||||
t.Helper()
|
||||
err := db.AutoMigrate(
|
||||
&model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.Setting{},
|
||||
&model.CheckSnapshot{}, &model.CostSnapshot{}, &model.RegionCache{}, &model.CompartmentCache{},
|
||||
&model.LogEvent{}, &model.AlertRule{}, &model.AlertRuleHit{},
|
||||
&model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}, &model.AiContentLog{},
|
||||
&model.Proxy{}, &model.AiKey{}, &model.SystemLog{},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlanTenantTask(t *testing.T) {
|
||||
for _, tt := range tenantTaskCases() {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
action, ok, err := planTenantTask(tt.task, 1)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if ok != tt.wantOK || action.deleteTask != tt.wantDelete {
|
||||
t.Errorf("result = (ok=%v, delete=%v), want (%v, %v)", ok, action.deleteTask, tt.wantOK, tt.wantDelete)
|
||||
}
|
||||
if action.payload != tt.wantPayload {
|
||||
t.Errorf("payload = %q, want %q", action.payload, tt.wantPayload)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type tenantTaskCase struct {
|
||||
name string
|
||||
task model.Task
|
||||
wantOK bool
|
||||
wantDelete bool
|
||||
wantPayload string
|
||||
wantErr bool
|
||||
}
|
||||
|
||||
func tenantTaskCases() []tenantTaskCase {
|
||||
return []tenantTaskCase{
|
||||
{name: "抢机命中", task: taskOf(model.TaskTypeSnatch, `{"ociConfigId":1}`), wantOK: true, wantDelete: true},
|
||||
{name: "抢机未命中", task: taskOf(model.TaskTypeSnatch, `{"ociConfigId":2}`)},
|
||||
{name: "测活全局", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[]}`), wantOK: true, wantPayload: `{"ociConfigIds":[]}`},
|
||||
{name: "测活单租户", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[1]}`), wantOK: true, wantDelete: true},
|
||||
{name: "测活多租户", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`},
|
||||
{name: "成本去重命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[1,1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`},
|
||||
{name: "成本未命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[2]}`)},
|
||||
{name: "非法 JSON", task: taskOf(model.TaskTypeCost, `{`), wantErr: true},
|
||||
}
|
||||
}
|
||||
|
||||
func taskOf(taskType, payload string) model.Task {
|
||||
return model.Task{ID: 10, Type: taskType, Payload: payload}
|
||||
}
|
||||
|
||||
func TestDeleteTenantCleansRelatedRows(t *testing.T) {
|
||||
client := &invalidatingClient{fakeClient: &fakeClient{}}
|
||||
configs, _, db := newTenantDeleteEnv(t, client)
|
||||
target, other := seedDeleteTenants(t, db)
|
||||
seedTenantSnapshots(t, db, target.ID, other.ID)
|
||||
seedTenantEvents(t, db, target.ID, other.ID)
|
||||
seedTenantAI(t, db, target.ID, other.ID)
|
||||
seedRetainedGlobals(t, db)
|
||||
if err := configs.Delete(context.Background(), target.ID); err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
assertTenantRowsGone(t, db, target.ID)
|
||||
assertOtherTenantRowsRemain(t, db, other.ID)
|
||||
assertRetainedGlobals(t, db)
|
||||
if len(client.invalidated) != 1 || client.invalidated[0] != target.TenancyOCID {
|
||||
t.Errorf("invalidated = %v, want [%s]", client.invalidated, target.TenancyOCID)
|
||||
}
|
||||
}
|
||||
|
||||
func seedDeleteTenants(t *testing.T, db *gorm.DB) (model.OciConfig, model.OciConfig) {
|
||||
t.Helper()
|
||||
target := model.OciConfig{Alias: "target", TenancyOCID: "ocid1.tenancy.target"}
|
||||
other := model.OciConfig{Alias: "other", TenancyOCID: "ocid1.tenancy.other"}
|
||||
mustCreate(t, db, &target)
|
||||
mustCreate(t, db, &other)
|
||||
return target, other
|
||||
}
|
||||
|
||||
func seedTenantSnapshots(t *testing.T, db *gorm.DB, target, other uint) {
|
||||
t.Helper()
|
||||
for _, id := range []uint{target, other} {
|
||||
mustCreate(t, db, &model.CheckSnapshot{OciConfigID: id})
|
||||
mustCreate(t, db, &model.CostSnapshot{OciConfigID: id, Day: "2026-07-10"})
|
||||
mustCreate(t, db, &model.RegionCache{OciConfigID: id, Key: "PHX"})
|
||||
mustCreate(t, db, &model.CompartmentCache{OciConfigID: id, OCID: fmt.Sprintf("comp-%d", id)})
|
||||
mustCreate(t, db, &model.Setting{Key: secretKey(id), Value: fmt.Sprintf("secret-%d", id)})
|
||||
}
|
||||
}
|
||||
|
||||
func seedTenantEvents(t *testing.T, db *gorm.DB, target, other uint) {
|
||||
t.Helper()
|
||||
targetEvent := model.LogEvent{OciConfigID: target, MessageID: "target-event"}
|
||||
otherEvent := model.LogEvent{OciConfigID: other, MessageID: "other-event"}
|
||||
targetRule := model.AlertRule{Name: "target-rule", OciConfigID: target}
|
||||
globalRule := model.AlertRule{Name: "global-rule", OciConfigID: 0}
|
||||
for _, value := range []any{&targetEvent, &otherEvent, &targetRule, &globalRule} {
|
||||
mustCreate(t, db, value)
|
||||
}
|
||||
hits := []model.AlertRuleHit{
|
||||
{RuleID: targetRule.ID, LogEventID: otherEvent.ID},
|
||||
{RuleID: globalRule.ID, LogEventID: targetEvent.ID},
|
||||
{RuleID: globalRule.ID, LogEventID: otherEvent.ID},
|
||||
}
|
||||
for i := range hits {
|
||||
mustCreate(t, db, &hits[i])
|
||||
}
|
||||
}
|
||||
|
||||
func seedTenantAI(t *testing.T, db *gorm.DB, target, other uint) {
|
||||
t.Helper()
|
||||
for _, id := range []uint{target, other} {
|
||||
channel := model.AiChannel{Name: fmt.Sprintf("channel-%d", id), OciConfigID: id, Region: "us-phoenix-1"}
|
||||
mustCreate(t, db, &channel)
|
||||
mustCreate(t, db, &model.AiModelCache{ChannelID: channel.ID, ModelOcid: fmt.Sprintf("model-%d", id)})
|
||||
call := model.AiCallLog{ChannelID: channel.ID, ChannelName: channel.Name}
|
||||
mustCreate(t, db, &call)
|
||||
mustCreate(t, db, &model.AiContentLog{CallLogID: call.ID, RequestBody: "sensitive"})
|
||||
}
|
||||
}
|
||||
|
||||
func seedRetainedGlobals(t *testing.T, db *gorm.DB) {
|
||||
t.Helper()
|
||||
mustCreate(t, db, &model.Proxy{Name: "shared", Type: "http"})
|
||||
mustCreate(t, db, &model.AiKey{Name: "global-key", KeyHash: "hash", Tail: "hash"})
|
||||
mustCreate(t, db, &model.SystemLog{Method: "DELETE", Path: "/api/v1/oci-configs/1"})
|
||||
mustCreate(t, db, &model.Setting{Key: "notify_channels", Value: "[]"})
|
||||
}
|
||||
|
||||
func assertTenantRowsGone(t *testing.T, db *gorm.DB, id uint) {
|
||||
t.Helper()
|
||||
rows := []any{
|
||||
&model.OciConfig{}, &model.CheckSnapshot{}, &model.CostSnapshot{},
|
||||
&model.RegionCache{}, &model.CompartmentCache{}, &model.LogEvent{},
|
||||
&model.AlertRule{}, &model.AiChannel{},
|
||||
}
|
||||
for _, value := range rows {
|
||||
column := "oci_config_id"
|
||||
if _, ok := value.(*model.OciConfig); ok {
|
||||
column = "id"
|
||||
}
|
||||
assertCount(t, db, value, column+" = ?", []any{id}, 0)
|
||||
}
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{secretKey(id)}, 0)
|
||||
assertCount(t, db, &model.AlertRuleHit{}, "", nil, 1)
|
||||
assertCount(t, db, &model.AiModelCache{}, "", nil, 1)
|
||||
assertCount(t, db, &model.AiCallLog{}, "", nil, 1)
|
||||
assertCount(t, db, &model.AiContentLog{}, "", nil, 1)
|
||||
}
|
||||
|
||||
func assertOtherTenantRowsRemain(t *testing.T, db *gorm.DB, id uint) {
|
||||
t.Helper()
|
||||
for _, value := range []any{
|
||||
&model.OciConfig{}, &model.CheckSnapshot{}, &model.CostSnapshot{},
|
||||
&model.RegionCache{}, &model.CompartmentCache{}, &model.LogEvent{},
|
||||
&model.AiChannel{},
|
||||
} {
|
||||
column := "oci_config_id"
|
||||
if _, ok := value.(*model.OciConfig); ok {
|
||||
column = "id"
|
||||
}
|
||||
assertCount(t, db, value, column+" = ?", []any{id}, 1)
|
||||
}
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{secretKey(id)}, 1)
|
||||
assertRemainingIndirectRows(t, db, id)
|
||||
}
|
||||
|
||||
func assertRemainingIndirectRows(t *testing.T, db *gorm.DB, otherID uint) {
|
||||
t.Helper()
|
||||
var channel model.AiChannel
|
||||
if err := db.Where("oci_config_id = ?", otherID).First(&channel).Error; err != nil {
|
||||
t.Fatalf("load other AI channel: %v", err)
|
||||
}
|
||||
assertCount(t, db, &model.AiModelCache{}, "channel_id = ?", []any{channel.ID}, 1)
|
||||
assertCount(t, db, &model.AiCallLog{}, "channel_id = ?", []any{channel.ID}, 1)
|
||||
var call model.AiCallLog
|
||||
if err := db.Where("channel_id = ?", channel.ID).First(&call).Error; err != nil {
|
||||
t.Fatalf("load other AI call: %v", err)
|
||||
}
|
||||
assertCount(t, db, &model.AiContentLog{}, "call_log_id = ?", []any{call.ID}, 1)
|
||||
assertRemainingAlertHit(t, db, otherID)
|
||||
}
|
||||
|
||||
func assertRemainingAlertHit(t *testing.T, db *gorm.DB, otherID uint) {
|
||||
t.Helper()
|
||||
var hit model.AlertRuleHit
|
||||
if err := db.First(&hit).Error; err != nil {
|
||||
t.Fatalf("load remaining alert hit: %v", err)
|
||||
}
|
||||
var rule model.AlertRule
|
||||
var event model.LogEvent
|
||||
if err := db.First(&rule, hit.RuleID).Error; err != nil {
|
||||
t.Fatalf("load remaining rule: %v", err)
|
||||
}
|
||||
if err := db.First(&event, hit.LogEventID).Error; err != nil {
|
||||
t.Fatalf("load remaining event: %v", err)
|
||||
}
|
||||
if rule.OciConfigID != 0 || event.OciConfigID != otherID {
|
||||
t.Errorf("remaining hit = rule cfg %d/event cfg %d, want global/other", rule.OciConfigID, event.OciConfigID)
|
||||
}
|
||||
}
|
||||
|
||||
func assertRetainedGlobals(t *testing.T, db *gorm.DB) {
|
||||
t.Helper()
|
||||
assertCount(t, db, &model.AlertRule{}, "oci_config_id = 0", nil, 1)
|
||||
assertCount(t, db, &model.Proxy{}, "", nil, 1)
|
||||
assertCount(t, db, &model.AiKey{}, "", nil, 1)
|
||||
assertCount(t, db, &model.SystemLog{}, "", nil, 1)
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{"notify_channels"}, 1)
|
||||
}
|
||||
|
||||
func TestDeleteTenantRewritesTasksAndCron(t *testing.T) {
|
||||
configs, tasks, db := newTenantDeleteEnv(t, &fakeClient{})
|
||||
target, other := seedDeleteTenants(t, db)
|
||||
created := seedTenantTasks(t, tasks, target.ID, other.ID)
|
||||
if err := configs.Delete(context.Background(), target.ID); err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
assertDeletedTasks(t, db, tasks, created[:2])
|
||||
assertRewrittenTasks(t, db, created[2:])
|
||||
}
|
||||
|
||||
func TestDeleteTenantReconcilesAiProbe(t *testing.T) {
|
||||
for _, keepOther := range []bool{false, true} {
|
||||
name := fmt.Sprintf("keepOther=%v", keepOther)
|
||||
t.Run(name, func(t *testing.T) {
|
||||
configs, tasks, db := newTenantDeleteEnv(t, &fakeClient{})
|
||||
target, other := seedDeleteTenants(t, db)
|
||||
seedProbeChannels(t, db, target.ID, other.ID, keepOther)
|
||||
tasks.SyncAiProbeTask(context.Background())
|
||||
probe := loadAiProbe(t, db)
|
||||
seedTaskHistory(t, db, probe.ID)
|
||||
if err := configs.Delete(context.Background(), target.ID); err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
assertAiProbeResult(t, db, tasks, probe.ID, keepOther)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func seedProbeChannels(t *testing.T, db *gorm.DB, target, other uint, keepOther bool) {
|
||||
t.Helper()
|
||||
mustCreate(t, db, &model.AiChannel{Name: "target", OciConfigID: target, Region: "r1"})
|
||||
if keepOther {
|
||||
mustCreate(t, db, &model.AiChannel{Name: "other", OciConfigID: other, Region: "r1"})
|
||||
}
|
||||
}
|
||||
|
||||
func loadAiProbe(t *testing.T, db *gorm.DB) model.Task {
|
||||
t.Helper()
|
||||
var task model.Task
|
||||
if err := db.Where("type = ?", model.TaskTypeAiProbe).First(&task).Error; err != nil {
|
||||
t.Fatalf("load AI probe: %v", err)
|
||||
}
|
||||
return task
|
||||
}
|
||||
|
||||
func seedTaskHistory(t *testing.T, db *gorm.DB, taskID uint) {
|
||||
t.Helper()
|
||||
updates := map[string]any{"last_error": "old", "run_count": 2, "last_run_at": time.Now()}
|
||||
if err := db.Model(&model.Task{}).Where("id = ?", taskID).Updates(updates).Error; err != nil {
|
||||
t.Fatalf("seed task history: %v", err)
|
||||
}
|
||||
mustCreate(t, db, &model.TaskLog{TaskID: taskID, Message: "old"})
|
||||
mustCreate(t, db, &model.Setting{Key: deadAliasKey(taskID), Value: "old"})
|
||||
}
|
||||
|
||||
func assertAiProbeResult(t *testing.T, db *gorm.DB, tasks *TaskService, id uint, keep bool) {
|
||||
t.Helper()
|
||||
want := int64(0)
|
||||
if keep {
|
||||
want = 1
|
||||
}
|
||||
assertCount(t, db, &model.Task{}, "id = ?", []any{id}, want)
|
||||
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{id}, 0)
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(id)}, 0)
|
||||
if !keep {
|
||||
if _, ok := tasks.entries[id]; ok {
|
||||
t.Errorf("AI probe %d still scheduled", id)
|
||||
}
|
||||
return
|
||||
}
|
||||
probe := loadAiProbe(t, db)
|
||||
if probe.RunCount != 0 || probe.LastRunAt != nil || probe.LastError != "" {
|
||||
t.Errorf("AI probe history not reset: %+v", probe)
|
||||
}
|
||||
}
|
||||
|
||||
func seedTenantTasks(t *testing.T, tasks *TaskService, target, other uint) []model.Task {
|
||||
t.Helper()
|
||||
inputs := []CreateTaskInput{
|
||||
{Name: "snatch", Type: model.TaskTypeSnatch, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigId":%d,"instance":{"displayName":"vm","region":"r","availabilityDomain":"a","subnetId":"s","shape":"x","imageId":"i"}}`, target))},
|
||||
{Name: "single", Type: model.TaskTypeCost, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigIds":[%d]}`, target))},
|
||||
{Name: "mixed", Type: model.TaskTypeHealthCheck, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigIds":[%d,%d]}`, target, other))},
|
||||
{Name: "global", Type: model.TaskTypeCost, CronExpr: "0 0 * * *", Payload: []byte(`{"ociConfigIds":[]}`)},
|
||||
{Name: "other", Type: model.TaskTypeHealthCheck, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigIds":[%d]}`, other))},
|
||||
}
|
||||
return createTasksWithHistory(t, tasks, inputs)
|
||||
}
|
||||
|
||||
func createTasksWithHistory(t *testing.T, tasks *TaskService, inputs []CreateTaskInput) []model.Task {
|
||||
t.Helper()
|
||||
out := make([]model.Task, 0, len(inputs))
|
||||
for _, input := range inputs {
|
||||
task, err := tasks.CreateTask(context.Background(), input)
|
||||
if err != nil {
|
||||
t.Fatalf("create task %s: %v", input.Name, err)
|
||||
}
|
||||
tasks.db.Model(task).Updates(map[string]any{"last_error": "old", "run_count": 3, "last_run_at": time.Now()})
|
||||
mustCreate(t, tasks.db, &model.TaskLog{TaskID: task.ID, Message: "old"})
|
||||
mustCreate(t, tasks.db, &model.Setting{Key: deadAliasKey(task.ID), Value: `["target"]`})
|
||||
out = append(out, *task)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func assertDeletedTasks(t *testing.T, db *gorm.DB, tasks *TaskService, deleted []model.Task) {
|
||||
t.Helper()
|
||||
for _, task := range deleted {
|
||||
assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 0)
|
||||
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 0)
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(task.ID)}, 0)
|
||||
if _, ok := tasks.entries[task.ID]; ok {
|
||||
t.Errorf("task %d still scheduled", task.ID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func assertRewrittenTasks(t *testing.T, db *gorm.DB, tasks []model.Task) {
|
||||
t.Helper()
|
||||
wantPayload := []string{`{"ociConfigIds":[2]}`, `{"ociConfigIds":[]}`, `{"ociConfigIds":[2]}`}
|
||||
for i, original := range tasks {
|
||||
var got model.Task
|
||||
if err := db.First(&got, original.ID).Error; err != nil {
|
||||
t.Fatalf("load task %d: %v", original.ID, err)
|
||||
}
|
||||
if got.Payload != wantPayload[i] {
|
||||
t.Errorf("task %d payload = %s, want %s", got.ID, got.Payload, wantPayload[i])
|
||||
}
|
||||
wantHistory := int64(0)
|
||||
if original.Name == "other" {
|
||||
wantHistory = 1
|
||||
} else if got.RunCount != 0 || got.LastRunAt != nil || got.LastError != "" {
|
||||
t.Errorf("task %d history fields not reset: %+v", got.ID, got)
|
||||
}
|
||||
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{got.ID}, wantHistory)
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(got.ID)}, wantHistory)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteTenantRollback(t *testing.T) {
|
||||
configs, tasks, db := newTenantDeleteEnv(t, &fakeClient{})
|
||||
target, _ := seedDeleteTenants(t, db)
|
||||
mustCreate(t, db, &model.CheckSnapshot{OciConfigID: target.ID})
|
||||
task := createHealthTask(t, tasks, target.ID)
|
||||
seedTaskHistory(t, db, task.ID)
|
||||
registerDeleteFailure(t, db, "cost_snapshots")
|
||||
err := configs.Delete(context.Background(), target.ID)
|
||||
if err == nil || !errors.Is(err, errInjectedTenantDelete) {
|
||||
t.Fatalf("delete error = %v, want injected failure", err)
|
||||
}
|
||||
assertCount(t, db, &model.OciConfig{}, "id = ?", []any{target.ID}, 1)
|
||||
assertCount(t, db, &model.CheckSnapshot{}, "oci_config_id = ?", []any{target.ID}, 1)
|
||||
assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 1)
|
||||
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 1)
|
||||
assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(task.ID)}, 1)
|
||||
}
|
||||
|
||||
func TestDeleteTenantWaitsForRunningTask(t *testing.T) {
|
||||
client := &blockingClient{fakeClient: &fakeClient{}, started: make(chan struct{}, 1), release: make(chan struct{})}
|
||||
configs, tasks, db := newTenantDeleteEnv(t, client)
|
||||
target, _ := seedDeleteTenants(t, db)
|
||||
setTenantPrivateKey(t, configs, target.ID)
|
||||
task := createHealthTask(t, tasks, target.ID)
|
||||
executed := make(chan struct{})
|
||||
go func() {
|
||||
tasks.execute(task.ID)
|
||||
close(executed)
|
||||
}()
|
||||
<-client.started
|
||||
deleted := make(chan error, 1)
|
||||
go func() { deleted <- configs.Delete(context.Background(), target.ID) }()
|
||||
assertDeleteBlocked(t, deleted)
|
||||
close(client.release)
|
||||
<-executed
|
||||
if err := <-deleted; err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 0)
|
||||
assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 0)
|
||||
}
|
||||
|
||||
func TestVerifyDoesNotResurrectDeletedTenant(t *testing.T) {
|
||||
client := &blockingClient{fakeClient: &fakeClient{}, started: make(chan struct{}, 1), release: make(chan struct{})}
|
||||
configs, _, db := newTenantDeleteEnv(t, client)
|
||||
target, _ := seedDeleteTenants(t, db)
|
||||
setTenantPrivateKey(t, configs, target.ID)
|
||||
verified := make(chan error, 1)
|
||||
go func() {
|
||||
_, _, err := configs.Verify(context.Background(), target.ID)
|
||||
verified <- err
|
||||
}()
|
||||
<-client.started
|
||||
if err := configs.Delete(context.Background(), target.ID); err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
close(client.release)
|
||||
if err := <-verified; !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("verify error = %v, want record not found", err)
|
||||
}
|
||||
assertCount(t, db, &model.OciConfig{}, "id = ?", []any{target.ID}, 0)
|
||||
}
|
||||
|
||||
func TestScopeCacheRejectsDeletedTenant(t *testing.T) {
|
||||
configs, _, db := newTenantDeleteEnv(t, &fakeClient{})
|
||||
target, _ := seedDeleteTenants(t, db)
|
||||
if err := configs.Delete(context.Background(), target.ID); err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
err := configs.saveRegionCache(context.Background(), target.ID,
|
||||
[]oci.RegionSubscription{{Key: "PHX", Name: "us-phoenix-1"}})
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("save region cache error = %v, want record not found", err)
|
||||
}
|
||||
err = configs.saveCompartmentCache(context.Background(), target.ID,
|
||||
[]oci.Compartment{{ID: "compartment", Name: "deleted"}})
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("save compartment cache error = %v, want record not found", err)
|
||||
}
|
||||
assertCount(t, db, &model.RegionCache{}, "oci_config_id = ?", []any{target.ID}, 0)
|
||||
assertCount(t, db, &model.CompartmentCache{}, "oci_config_id = ?", []any{target.ID}, 0)
|
||||
}
|
||||
|
||||
func setTenantPrivateKey(t *testing.T, configs *OciConfigService, id uint) {
|
||||
t.Helper()
|
||||
encrypted, err := configs.cipher.EncryptString("private-key")
|
||||
if err != nil {
|
||||
t.Fatalf("encrypt private key: %v", err)
|
||||
}
|
||||
if err := configs.db.Model(&model.OciConfig{}).Where("id = ?", id).
|
||||
Update("private_key_enc", encrypted).Error; err != nil {
|
||||
t.Fatalf("set private key: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func createHealthTask(t *testing.T, tasks *TaskService, cfgID uint) *model.Task {
|
||||
t.Helper()
|
||||
payload := []byte(fmt.Sprintf(`{"ociConfigIds":[%d]}`, cfgID))
|
||||
task, err := tasks.CreateTask(context.Background(), CreateTaskInput{
|
||||
Name: "running", Type: model.TaskTypeHealthCheck,
|
||||
CronExpr: "0 0 * * *", Payload: payload,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create task: %v", err)
|
||||
}
|
||||
return task
|
||||
}
|
||||
|
||||
func assertDeleteBlocked(t *testing.T, deleted <-chan error) {
|
||||
t.Helper()
|
||||
select {
|
||||
case err := <-deleted:
|
||||
t.Fatalf("delete returned before running task finished: %v", err)
|
||||
case <-time.After(50 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersistTaskRunDoesNotResurrectDeletedTask(t *testing.T) {
|
||||
_, tasks, db := newTenantDeleteEnv(t, &fakeClient{})
|
||||
task := &model.Task{Name: "stale", Type: model.TaskTypeCost, Status: model.TaskStatusActive}
|
||||
mustCreate(t, db, task)
|
||||
stale := *task
|
||||
if err := db.Delete(task).Error; err != nil {
|
||||
t.Fatalf("delete task: %v", err)
|
||||
}
|
||||
stored, err := tasks.persistTaskRun(context.Background(), &stale)
|
||||
if err != nil || stored {
|
||||
t.Fatalf("persist stale task = (%v, %v), want (false, nil)", stored, err)
|
||||
}
|
||||
assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 0)
|
||||
}
|
||||
|
||||
func TestAiLogsRejectDeletedTenantParents(t *testing.T) {
|
||||
configs, _, db := newTenantDeleteEnv(t, &fakeClient{})
|
||||
target, _ := seedDeleteTenants(t, db)
|
||||
channel := model.AiChannel{Name: "target", OciConfigID: target.ID, Region: "r1"}
|
||||
mustCreate(t, db, &channel)
|
||||
gw := NewAiGatewayService(db, configs, &fakeClient{})
|
||||
callID := gw.LogCall(model.AiCallLog{ChannelID: channel.ID, ChannelName: channel.Name})
|
||||
if callID == 0 {
|
||||
t.Fatal("initial call log was not created")
|
||||
}
|
||||
if err := configs.Delete(context.Background(), target.ID); err != nil {
|
||||
t.Fatalf("delete tenant: %v", err)
|
||||
}
|
||||
lateID := gw.LogCall(model.AiCallLog{ChannelID: channel.ID, ChannelName: channel.Name})
|
||||
if lateID != 0 {
|
||||
t.Errorf("late call ID = %d, want 0", lateID)
|
||||
}
|
||||
gw.LogContent(model.AiContentLog{CallLogID: callID, RequestBody: "late"})
|
||||
gw.LogContent(model.AiContentLog{CallLogID: 0, RequestBody: "orphan"})
|
||||
assertCount(t, db, &model.AiCallLog{}, "channel_id = ?", []any{channel.ID}, 0)
|
||||
assertCount(t, db, &model.AiContentLog{}, "", nil, 0)
|
||||
}
|
||||
|
||||
var errInjectedTenantDelete = errors.New("injected tenant delete failure")
|
||||
|
||||
func registerDeleteFailure(t *testing.T, db *gorm.DB, table string) {
|
||||
t.Helper()
|
||||
name := "test:tenant_delete_failure"
|
||||
err := db.Callback().Delete().Before("gorm:delete").Register(name, func(tx *gorm.DB) {
|
||||
if tx.Statement.Table == table {
|
||||
tx.AddError(errInjectedTenantDelete)
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("register callback: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Callback().Delete().Remove(name) })
|
||||
}
|
||||
|
||||
func mustCreate(t *testing.T, db *gorm.DB, value any) {
|
||||
t.Helper()
|
||||
if err := db.Create(value).Error; err != nil {
|
||||
t.Fatalf("create %T: %v", value, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCount(t *testing.T, db *gorm.DB, value any, query string, args []any, want int64) {
|
||||
t.Helper()
|
||||
q := db.Model(value)
|
||||
if query != "" {
|
||||
q = q.Where(query, args...)
|
||||
}
|
||||
var got int64
|
||||
if err := q.Count(&got).Error; err != nil {
|
||||
t.Fatalf("count %T: %v", value, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Errorf("count %T = %d, want %d", value, got, want)
|
||||
}
|
||||
}
|
||||
@@ -102,7 +102,8 @@ func (s *AuthService) ActivateTotp(ctx context.Context, username, code string) e
|
||||
s.totpMu.Lock()
|
||||
delete(s.totpPending, username)
|
||||
s.totpMu.Unlock()
|
||||
return nil
|
||||
// 两步验证形态变更:旧令牌全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
}
|
||||
|
||||
// DisableTotp 停用两步验证;需当前验证码或登录密码任一确认。
|
||||
@@ -122,7 +123,8 @@ func (s *AuthService) DisableTotp(ctx context.Context, username, password, code
|
||||
if err != nil {
|
||||
return fmt.Errorf("clear totp secret: %w", err)
|
||||
}
|
||||
return nil
|
||||
// 两步验证形态变更:旧令牌全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
}
|
||||
|
||||
// confirmDisable 校验停用凭证:验证码或密码任一通过即可。
|
||||
|
||||
@@ -176,7 +176,7 @@ func TestOAuthBindLoginUnbind(t *testing.T) {
|
||||
if display != "octocat" {
|
||||
t.Errorf("display = %q", display)
|
||||
}
|
||||
if username, err := auth.ParseToken(token); err != nil || username != "admin" {
|
||||
if username, err := auth.ParseToken(context.Background(), token); err != nil || username != "admin" {
|
||||
t.Errorf("token 应属 admin, got %q (%v)", username, err)
|
||||
}
|
||||
// 未绑定身份拒绝登录
|
||||
|
||||
Reference in New Issue
Block a user