package security import ( "reflect" "gorm.io/gorm" "gorm.io/gorm/schema" ) // RegisterGORMEncryption encrypts string fields tagged secure:"..." on writes and // decrypts them after reads. Existing plaintext remains readable for online migration. func RegisterGORMEncryption(db *gorm.DB, encryptor *Encryptor) error { if encryptor == nil || !encryptor.IsEnabled() { return nil } before := func(tx *gorm.DB) { transformStatement(tx, encryptor, true) } after := func(tx *gorm.DB) { transformStatement(tx, encryptor, false) } if err := db.Callback().Create().Before("gorm:create").Register("gochat:encrypt", before); err != nil { return err } if err := db.Callback().Create().After("gorm:create").Register("gochat:decrypt", after); err != nil { return err } if err := db.Callback().Update().Before("gorm:update").Register("gochat:encrypt", before); err != nil { return err } if err := db.Callback().Update().After("gorm:update").Register("gochat:decrypt", after); err != nil { return err } return db.Callback().Query().After("gorm:after_query").Register("gochat:decrypt", after) } func transformStatement(tx *gorm.DB, encryptor *Encryptor, encrypt bool) { if tx.Statement == nil || tx.Statement.Schema == nil { return } if values, ok := tx.Statement.Dest.(map[string]interface{}); ok { transformMap(tx, encryptor, values, encrypt) return } transformValue(tx, encryptor, tx.Statement.ReflectValue, tx.Statement.Schema, encrypt) } func transformMap(tx *gorm.DB, encryptor *Encryptor, values map[string]interface{}, encrypt bool) { for key, value := range values { field := tx.Statement.Schema.LookUpField(key) if field == nil || field.StructField.Tag.Get("secure") == "" { continue } text, ok := value.(string) if !ok { continue } transformed, err := transformSecret(encryptor, text, encrypt) if err != nil { tx.AddError(err) return } values[key] = transformed } } func transformValue(tx *gorm.DB, encryptor *Encryptor, value reflect.Value, modelSchema *schema.Schema, encrypt bool) { for value.IsValid() && (value.Kind() == reflect.Pointer || value.Kind() == reflect.Interface) { if value.IsNil() { return } value = value.Elem() } if value.Kind() == reflect.Slice || value.Kind() == reflect.Array { for i := 0; i < value.Len(); i++ { transformValue(tx, encryptor, value.Index(i), modelSchema, encrypt) } return } if value.Kind() != reflect.Struct { return } for _, field := range modelSchema.Fields { if field.StructField.Tag.Get("secure") == "" { continue } current, zero := field.ValueOf(tx.Statement.Context, value) text, ok := current.(string) if !ok || zero || text == "" { continue } transformed, err := transformSecret(encryptor, text, encrypt) if err != nil { tx.AddError(err) return } if err := field.Set(tx.Statement.Context, value, transformed); err != nil { tx.AddError(err) return } } } func transformSecret(encryptor *Encryptor, value string, encrypt bool) (string, error) { if encrypt { if IsEncrypted(value) { return value, nil } return encryptor.Encrypt(value) } if !IsEncrypted(value) { return value, nil } return encryptor.Decrypt(value) }