* fix(security): harden auth and credential handling (HH-444) * fix(security): address HH-444 review blockers * fix(security): close remaining HH-444 review blockers --------- Co-authored-by: Rogee <rogee@ipao.vip>
112 lines
3.1 KiB
Go
112 lines
3.1 KiB
Go
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)
|
|
}
|