diff --git a/internal/models/query.go b/internal/models/query.go index e70da3e..347828f 100644 --- a/internal/models/query.go +++ b/internal/models/query.go @@ -1,6 +1,10 @@ package models -import "gorm.io/gorm" +import ( + "errors" + + "gorm.io/gorm" +) // GetAllModels 数据库迁移用模型列表 func GetAllModels() []interface{} { @@ -85,13 +89,6 @@ WHERE target_event.dispatch_outbox_id = 0` } func seedDefaultSyslogRules(db *gorm.DB) error { - var cnt int64 - if err := db.Model(&SyslogRule{}).Count(&cnt).Error; err != nil { - return err - } - if cnt > 0 { - return nil - } rows := []SyslogRule{ { Name: "默认-系统严重错误", @@ -130,17 +127,39 @@ func seedDefaultSyslogRules(db *gorm.DB) error { PolicyID: 0, }, } - return db.Create(&rows).Error + for _, row := range rows { + var existing SyslogRule + err := db.Where("name = ?", row.Name).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + if err := db.Create(&row).Error; err != nil { + return err + } + continue + } + if err != nil { + return err + } + if err := db.Model(&existing).Select( + "name", + "enabled", + "priority", + "device_name_contains", + "source_match", + "keyword_regex", + "message_regex", + "alert_name", + "severity_code", + "severity_mapping_json", + "resource_uid_extract_regex", + "policy_id", + ).Updates(&row).Error; err != nil { + return err + } + } + return nil } func seedDefaultTrapRules(db *gorm.DB) error { - var cnt int64 - if err := db.Model(&TrapRule{}).Count(&cnt).Error; err != nil { - return err - } - if cnt > 0 { - return nil - } rows := []TrapRule{ { Name: "默认-Trap链路中断", @@ -153,17 +172,35 @@ func seedDefaultTrapRules(db *gorm.DB) error { PolicyID: 0, }, } - return db.Create(&rows).Error + for _, row := range rows { + var existing TrapRule + err := db.Where("name = ?", row.Name).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + if err := db.Create(&row).Error; err != nil { + return err + } + continue + } + if err != nil { + return err + } + if err := db.Model(&existing).Select( + "name", + "enabled", + "priority", + "oid_prefix", + "varbind_match_regex", + "alert_name", + "severity_code", + "policy_id", + ).Updates(&row).Error; err != nil { + return err + } + } + return nil } func seedDefaultTrapDictionary(db *gorm.DB) error { - var cnt int64 - if err := db.Model(&TrapDictionaryEntry{}).Count(&cnt).Error; err != nil { - return err - } - if cnt > 0 { - return nil - } rows := []TrapDictionaryEntry{ { Vendor: "H3C", @@ -192,5 +229,33 @@ func seedDefaultTrapDictionary(db *gorm.DB) error { Enabled: true, }, } - return db.Create(&rows).Error + for _, row := range rows { + var existing TrapDictionaryEntry + err := db.Where("oid_prefix = ?", row.OIDPrefix).First(&existing).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + if err := db.Create(&row).Error; err != nil { + return err + } + continue + } + if err != nil { + return err + } + if err := db.Model(&existing).Select( + "oid_prefix", + "vendor", + "o_id", + "name", + "title", + "description", + "severity_code", + "severity_mapping_json", + "parse_expression", + "recovery_message", + "enabled", + ).Updates(&row).Error; err != nil { + return err + } + } + return nil }