Files
gochat/backend/internal/security/gorm_encryption.go
T
Rogeeandrogee f719529d66 fix(security): harden auth and secret handling (HH-444) (#101)
* 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>
2026-08-22 15:45:06 +08:00

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)
}