diff --git a/pkg/common/adapters/database/bun.go b/pkg/common/adapters/database/bun.go index 043a81b..79eb32e 100644 --- a/pkg/common/adapters/database/bun.go +++ b/pkg/common/adapters/database/bun.go @@ -1507,6 +1507,13 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery { return b } +func (b *BunInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery { + if len(columns) > 0 { + b.query = b.query.ExcludeColumn(columns...) + } + return b +} + func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery { if len(columns) > 0 { b.query = b.query.Returning(strings.Join(columns, ", ")) @@ -1619,6 +1626,13 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer return b } +func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery { + if len(columns) > 0 { + b.query = b.query.ExcludeColumn(columns...) + } + return b +} + func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery { b.query = b.query.Where(query, args...) return b diff --git a/pkg/common/adapters/database/gorm.go b/pkg/common/adapters/database/gorm.go index 2fd2975..a288627 100644 --- a/pkg/common/adapters/database/gorm.go +++ b/pkg/common/adapters/database/gorm.go @@ -751,6 +751,13 @@ func (g *GormInsertQuery) OnConflict(action string) common.InsertQuery { return g } +func (g *GormInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery { + if len(columns) > 0 { + g.db = g.db.Omit(columns...) + } + return g +} + func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery { g.returningColumns = columns return g @@ -930,6 +937,13 @@ func (g *GormUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQue return g } +func (g *GormUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery { + if len(columns) > 0 { + g.db = g.db.Omit(columns...) + } + return g +} + func (g *GormUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery { g.db = g.db.Where(query, args...) return g diff --git a/pkg/common/adapters/database/pgsql.go b/pkg/common/adapters/database/pgsql.go index aea1d3d..f4addc0 100644 --- a/pkg/common/adapters/database/pgsql.go +++ b/pkg/common/adapters/database/pgsql.go @@ -691,6 +691,13 @@ func (p *PgSQLInsertQuery) OnConflict(action string) common.InsertQuery { return p } +func (p *PgSQLInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery { + for _, col := range columns { + delete(p.values, col) + } + return p +} + func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery { p.returning = columns return p @@ -850,6 +857,13 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu return p } +func (p *PgSQLUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery { + for _, col := range columns { + delete(p.sets, col) + } + return p +} + func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuery { pkName := "" if p.model != nil { diff --git a/pkg/common/interfaces.go b/pkg/common/interfaces.go index 9729134..a9c4516 100644 --- a/pkg/common/interfaces.go +++ b/pkg/common/interfaces.go @@ -81,6 +81,8 @@ type InsertQuery interface { Table(table string) InsertQuery Value(column string, value interface{}) InsertQuery OnConflict(action string) InsertQuery + // ExcludeColumn omits columns from a Model()-based INSERT (e.g. generated columns). + ExcludeColumn(columns ...string) InsertQuery Returning(columns ...string) InsertQuery // Execution @@ -94,6 +96,8 @@ type UpdateQuery interface { Table(table string) UpdateQuery Set(column string, value interface{}) UpdateQuery SetMap(values map[string]interface{}) UpdateQuery + // ExcludeColumn omits columns from a Model()-based UPDATE (e.g. generated columns). + ExcludeColumn(columns ...string) UpdateQuery Where(query string, args ...interface{}) UpdateQuery Returning(columns ...string) UpdateQuery diff --git a/pkg/common/recursive_crud.go b/pkg/common/recursive_crud.go index a6e4fcd..dcd0fde 100644 --- a/pkg/common/recursive_crud.go +++ b/pkg/common/recursive_crud.go @@ -116,7 +116,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD( case "insert", "create", "add": // Only perform insert if we have data to insert if hasData { - id, err := p.processInsert(ctx, regularData, tableName) + id, err := p.processInsert(ctx, regularData, model, tableName) if err != nil { logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err) return nil, fmt.Errorf("insert failed: %w", err) @@ -148,7 +148,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD( return result, nil } if hasData { - rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName]) + rows, err := p.processUpdate(ctx, regularData, model, tableName, data[pkName]) if err != nil { logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err) return nil, fmt.Errorf("update failed: %w", err) @@ -295,10 +295,12 @@ func (p *NestedCUDProcessor) injectForeignKeys(data map[string]interface{}, mode func (p *NestedCUDProcessor) processInsert( ctx context.Context, data map[string]interface{}, + model interface{}, tableName string, ) (interface{}, error) { logger.Debug("Inserting into %s with data: %+v", tableName, data) + reflection.RemoveNonWritableColumns(model, data) query := p.db.NewInsert().Table(tableName) for key, value := range data { @@ -335,6 +337,7 @@ func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string func (p *NestedCUDProcessor) processUpdate( ctx context.Context, data map[string]interface{}, + model interface{}, tableName string, id interface{}, ) (int64, error) { @@ -345,6 +348,7 @@ func (p *NestedCUDProcessor) processUpdate( logger.Debug("Updating %s with ID %v, data: %+v", tableName, id, data) + reflection.RemoveNonWritableColumns(model, data) query := p.db.NewUpdate().Table(tableName).SetMap(data).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id) result, err := query.Exec(ctx) diff --git a/pkg/common/recursive_crud_test.go b/pkg/common/recursive_crud_test.go index ce23b6c..3fb144f 100644 --- a/pkg/common/recursive_crud_test.go +++ b/pkg/common/recursive_crud_test.go @@ -99,6 +99,7 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery { return m } func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m } +func (m *mockInsertQuery) ExcludeColumn(columns ...string) InsertQuery { return m } func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m } func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) { m.db.insertCalls = append(m.db.insertCalls, m.values) @@ -131,6 +132,7 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery { return m } func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m } +func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m } func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m } func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) { // Record the update call diff --git a/pkg/mqttspec/handler.go b/pkg/mqttspec/handler.go index 98185cf..995b998 100644 --- a/pkg/mqttspec/handler.go +++ b/pkg/mqttspec/handler.go @@ -895,6 +895,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { // Insert record query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + if generated := reflection.NonWritableColumns(hookCtx.Model); len(generated) > 0 { + query = query.ExcludeColumn(generated...) + } if _, err := query.Exec(hookCtx.Context); err != nil { return nil, fmt.Errorf("failed to create record: %w", err) } @@ -924,6 +927,8 @@ func (h *Handler) update(hookCtx *HookContext) error { // the stored value unless disallowNulls is set, in which case null is skipped. values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls) + reflection.RemoveNonWritableColumns(hookCtx.Model, values) + if len(values) > 0 { query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) diff --git a/pkg/reflection/model_utils.go b/pkg/reflection/model_utils.go index 604028c..6f7bf9d 100644 --- a/pkg/reflection/model_utils.go +++ b/pkg/reflection/model_utils.go @@ -656,7 +656,7 @@ func isColumnWritableInType(typ reflect.Type, columnName string) (found bool, wr // Check bun tag for scanonly bunTag := field.Tag.Get("bun") if bunTag != "" { - if isBunFieldScanOnly(bunTag) { + if isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag) { return true, false } } @@ -689,6 +689,70 @@ func isBunFieldScanOnly(tag string) bool { return false } +// isBunFieldGenerated checks if a bun tag marks the column as database-generated +// (GENERATED ALWAYS AS ... STORED), which can be read but never written. +// Example: "email_normalized,generated" -> true +func isBunFieldGenerated(tag string) bool { + for _, part := range strings.Split(tag, ",") { + if strings.TrimSpace(part) == "generated" { + return true + } + } + return false +} + +// RemoveNonWritableColumns deletes from values every key that maps to a +// non-writable model column (bun scanonly/generated, gorm read-only). Used +// before writing a read-merged record back with UPDATE ... SET. +func RemoveNonWritableColumns(model any, values map[string]interface{}) { + for key := range values { + if !IsColumnWritable(model, key) { + delete(values, key) + } + } +} + +// NonWritableColumns returns the column names of the model that cannot be +// written (bun scanonly/generated, gorm read-only), including embedded structs. +func NonWritableColumns(model any) []string { + t := reflect.TypeOf(model) + for t != nil && (t.Kind() == reflect.Pointer || t.Kind() == reflect.Slice || t.Kind() == reflect.Array) { + t = t.Elem() + } + if t == nil || t.Kind() != reflect.Struct { + return nil + } + var cols []string + collectNonWritable(t, &cols) + return cols +} + +func collectNonWritable(typ reflect.Type, cols *[]string) { + for i := 0; i < typ.NumField(); i++ { + field := typ.Field(i) + if field.Anonymous { + ft := field.Type + if ft.Kind() == reflect.Pointer { + ft = ft.Elem() + } + if ft.Kind() == reflect.Struct { + collectNonWritable(ft, cols) + continue + } + } + bunTag, gormTag := field.Tag.Get("bun"), field.Tag.Get("gorm") + if bunTag == "-" || gormTag == "-" { + continue + } + if (bunTag != "" && (isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag))) || + (gormTag != "" && isGormFieldReadOnly(gormTag)) { + if name := getColumnNameFromField(field); name != "" { + *cols = append(*cols, name) + } + } + } +} + // isGormFieldReadOnly checks if a gorm tag indicates the field is read-only // Examples: // - "<-:false" -> true (no writes allowed) diff --git a/pkg/reflection/model_utils_test.go b/pkg/reflection/model_utils_test.go index 2ef7706..3b4b46b 100644 --- a/pkg/reflection/model_utils_test.go +++ b/pkg/reflection/model_utils_test.go @@ -1920,3 +1920,46 @@ func TestMapToStruct_Errors(t *testing.T) { }) } } + +func TestRemoveNonWritableColumns_Generated(t *testing.T) { + type m struct { + ID int `bun:"id,pk"` + Email string `bun:"email"` + Norm string `bun:"email_normalized,generated"` + Scan string `bun:"scan_col,scanonly"` + } + vals := map[string]interface{}{"id": 1, "email": "A", "email_normalized": "a", "scan_col": "x", "dynamic": 1} + RemoveNonWritableColumns(&m{}, vals) + if _, ok := vals["email_normalized"]; ok { + t.Error("generated column not removed") + } + if _, ok := vals["scan_col"]; ok { + t.Error("scanonly column not removed") + } + if len(vals) != 3 { + t.Errorf("unexpected keys: %v", vals) + } +} + +func TestNonWritableColumns(t *testing.T) { + type base struct { + Created string `bun:"created_at,scanonly"` + } + type m struct { + base + ID int `bun:"id,pk"` + Email string `bun:"email"` + Norm string `bun:"email_normalized,generated"` + Ro string `gorm:"column:ro;->"` + } + got := NonWritableColumns(&m{}) + want := map[string]bool{"created_at": true, "email_normalized": true, "ro": true} + if len(got) != len(want) { + t.Fatalf("got %v", got) + } + for _, c := range got { + if !want[c] { + t.Errorf("unexpected %s", c) + } + } +} diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index 9e8622b..c9e82ba 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -559,6 +559,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data if len(cols) == 0 { return invalidArg("no writable fields in data") } + reflection.RemoveNonWritableColumns(model, cols) q := tx.NewInsert().Table(tableName) for key, value := range cols { q = q.Value(key, value) @@ -726,6 +727,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, existingMap[key] = v } + reflection.RemoveNonWritableColumns(model, setCols) q := tx.NewUpdate().Table(tableName).SetMap(setCols). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) res, err := q.Exec(ctx) diff --git a/pkg/resolvemcp/writewhere.go b/pkg/resolvemcp/writewhere.go index 826bd3a..1cd07a9 100644 --- a/pkg/resolvemcp/writewhere.go +++ b/pkg/resolvemcp/writewhere.go @@ -194,6 +194,7 @@ func (h *Handler) executeWhere(ctx context.Context, req whereRequest) (_ *whereR cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", ")) var affected int64 if req.op == "update" { + reflection.RemoveNonWritableColumns(model, setCols) r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx) if err != nil { return fmt.Errorf("error updating records: %w", err) diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index 47b0a4f..c1b397c 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -824,6 +824,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat } responseData = v + reflection.RemoveNonWritableColumns(model, v) query := tx.NewInsert().Table(tableName) for key, value := range v { query = query.Value(key, common.ConvertSliceForBun(value)) @@ -971,6 +972,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat item = modifiedData } + reflection.RemoveNonWritableColumns(model, item) txQuery := tx.NewInsert().Table(tableName) for key, value := range item { txQuery = txQuery.Value(key, common.ConvertSliceForBun(value)) @@ -1127,6 +1129,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat itemMap = modifiedData } + reflection.RemoveNonWritableColumns(model, itemMap) txQuery := tx.NewInsert().Table(tableName) for key, value := range itemMap { txQuery = txQuery.Value(key, common.ConvertSliceForBun(value)) @@ -1322,6 +1325,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Overwrite with every key present in the request (including "" and null unless disallowed) common.MergeUpdateValues(existingMap, updates, h.disallowNulls) + reflection.RemoveNonWritableColumns(model, existingMap) // Build update query with merged data query := tx.NewUpdate().Table(tableName).SetMap(existingMap) @@ -1507,6 +1511,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Overwrite with every key present in the request (including "" and null unless disallowed) common.MergeUpdateValues(existingMap, item, h.disallowNulls) + reflection.RemoveNonWritableColumns(model, existingMap) txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) if _, err := txQuery.Exec(ctx); err != nil { @@ -1662,6 +1667,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url // Overwrite with every key present in the request (including "" and null unless disallowed) common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls) + reflection.RemoveNonWritableColumns(model, existingMap) txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) if _, err := txQuery.Exec(ctx); err != nil { diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index a67dbb8..41f94ee 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -1410,6 +1410,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat if provider, ok := modelValue.(common.TableNameProvider); !ok || provider.TableName() == "" { query = query.Table(tableName) } + if generated := reflection.NonWritableColumns(model); len(generated) > 0 { + query = query.ExcludeColumn(generated...) + } fields := reflection.GetSQLModelColumns(model) query = query.Returning(fields...) @@ -1657,6 +1660,9 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id // Create update query using Model() to preserve custom types and driver.Valuer interfaces query := tx.NewUpdate().Model(modelInstance) + if generated := reflection.NonWritableColumns(model); len(generated) > 0 { + query = query.ExcludeColumn(generated...) + } query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID) // Execute BeforeScan hooks - pass query chain so hooks can modify it diff --git a/pkg/websocketspec/handler.go b/pkg/websocketspec/handler.go index 352d1d8..10aa79c 100644 --- a/pkg/websocketspec/handler.go +++ b/pkg/websocketspec/handler.go @@ -758,6 +758,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { // Insert record query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) + if generated := reflection.NonWritableColumns(hookCtx.Model); len(generated) > 0 { + query = query.ExcludeColumn(generated...) + } if _, err := query.Exec(hookCtx.Context); err != nil { return nil, fmt.Errorf("failed to create record: %w", err) } @@ -786,6 +789,8 @@ func (h *Handler) update(hookCtx *HookContext) error { // the stored value unless disallowNulls is set, in which case null is skipped. values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls) + reflection.RemoveNonWritableColumns(hookCtx.Model, values) + if len(values) > 0 { query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) diff --git a/pkg/websocketspec/handler_test.go b/pkg/websocketspec/handler_test.go index 69a47b7..1b7b12b 100644 --- a/pkg/websocketspec/handler_test.go +++ b/pkg/websocketspec/handler_test.go @@ -226,6 +226,11 @@ func (m *MockInsertQuery) OnConflict(action string) common.InsertQuery { return args.Get(0).(common.InsertQuery) } +func (m *MockInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery { + args := m.Called(columns) + return args.Get(0).(common.InsertQuery) +} + func (m *MockInsertQuery) Returning(columns ...string) common.InsertQuery { args := m.Called(columns) return args.Get(0).(common.InsertQuery) @@ -254,6 +259,11 @@ func (m *MockUpdateQuery) Model(model interface{}) common.UpdateQuery { return args.Get(0).(common.UpdateQuery) } +func (m *MockUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery { + args := m.Called(columns) + return args.Get(0).(common.UpdateQuery) +} + func (m *MockUpdateQuery) Table(table string) common.UpdateQuery { args := m.Called(table) return args.Get(0).(common.UpdateQuery)