diff --git a/management/server/migration/migration.go b/management/server/migration/migration.go index ba31ee459..60c923352 100644 --- a/management/server/migration/migration.go +++ b/management/server/migration/migration.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "net" "strings" log "github.com/sirupsen/logrus" @@ -99,3 +100,96 @@ func MigrateFieldFromGobToJSON[T any, S any](db *gorm.DB, fieldName string) erro return nil } + +// MigrateNetIPFieldFromBlobToJSON migrates a Net IP column from Blob encoding to JSON encoding. +// T is the type of the model that contains the field to be migrated. +// S is the type of the field to be migrated. +func MigrateNetIPFieldFromBlobToJSON[T any](db *gorm.DB, fieldName string, indexName string) error { + oldColumnName := fieldName + newColumnName := fieldName + "_tmp" + + var model T + + if !db.Migrator().HasTable(&model) { + log.Printf("Table for %T does not exist, no migration needed", model) + return nil + } + + stmt := &gorm.Statement{DB: db} + err := stmt.Parse(&model) + if err != nil { + return fmt.Errorf("parse model: %w", err) + } + tableName := stmt.Schema.Table + + var item string + if err := db.Model(&model).Select(oldColumnName).First(&item).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + log.Printf("No records in table %s, no migration needed", tableName) + return nil + } + return fmt.Errorf("fetch first record: %w", err) + } + if len(item) < 0 { + return fmt.Errorf("no records fetched") + } + + var js json.RawMessage + var syntaxError *json.SyntaxError + err = json.Unmarshal([]byte(item), &js) + if err == nil || !errors.As(err, &syntaxError) { + log.Debugf("No migration needed for %s, %s", tableName, fieldName) + return nil + } + + if err := db.Transaction(func(tx *gorm.DB) error { + if err := tx.Exec(fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s TEXT", tableName, newColumnName)).Error; err != nil { + return fmt.Errorf("add column %s: %w", newColumnName, err) + } + + var rows []map[string]any + if err := tx.Table(tableName).Select("id", oldColumnName).Find(&rows).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + log.Printf("No records in table %s, no migration needed", tableName) + return nil + } + return fmt.Errorf("find rows: %w", err) + } + + for _, row := range rows { + blob, ok := row[oldColumnName].(string) + if !ok { + return fmt.Errorf("type assertion failed") + } + + jsonValue, err := json.Marshal(net.IP(blob)) + if err != nil { + return fmt.Errorf("re-encode to JSON: %w", err) + } + + if err := tx.Table(tableName).Where("id = ?", row["id"]).Update(newColumnName, jsonValue).Error; err != nil { + return fmt.Errorf("update row: %w", err) + } + } + + if indexName != "" { + if err := tx.Migrator().DropIndex(&model, indexName); err != nil { + return fmt.Errorf("drop index %s: %w", indexName, err) + } + } + + if err := tx.Exec(fmt.Sprintf("ALTER TABLE %s DROP COLUMN %s", tableName, oldColumnName)).Error; err != nil { + return fmt.Errorf("drop column %s: %w", oldColumnName, err) + } + if err := tx.Exec(fmt.Sprintf("ALTER TABLE %s RENAME COLUMN %s TO %s", tableName, newColumnName, oldColumnName)).Error; err != nil { + return fmt.Errorf("rename column %s to %s: %w", newColumnName, oldColumnName, err) + } + return nil + }); err != nil { + return err + } + + log.Printf("Migration of %s.%s from blob to json completed", tableName, fieldName) + + return nil +}