diff --git a/pkg/common/adapters/database/bun.go b/pkg/common/adapters/database/bun.go index 79eb32e..ddfa6f0 100644 --- a/pkg/common/adapters/database/bun.go +++ b/pkg/common/adapters/database/bun.go @@ -10,6 +10,7 @@ import ( "time" "github.com/uptrace/bun" + "github.com/uptrace/bun/schema" "github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/dbtrace" @@ -1507,8 +1508,30 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery { return b } +// bunWritableExcludes drops columns bun already leaves out of INSERT/UPDATE +// (scanonly fields) or does not know, since bun's ExcludeColumn errors with +// "can't find column" for anything that is not in the table's writable fields. +func bunWritableExcludes(model bun.Model, columns []string) []string { + tm, ok := model.(interface{ Table() *schema.Table }) + if !ok || tm.Table() == nil { + return columns + } + table := tm.Table() + writable := make(map[string]struct{}, len(table.Fields)) + for _, f := range table.Fields { + writable[f.Name] = struct{}{} + } + out := make([]string, 0, len(columns)) + for _, c := range columns { + if _, ok := writable[c]; ok || c == "*" { + out = append(out, c) + } + } + return out +} + func (b *BunInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery { - if len(columns) > 0 { + if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 { b.query = b.query.ExcludeColumn(columns...) } return b @@ -1627,7 +1650,7 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer } func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery { - if len(columns) > 0 { + if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 { b.query = b.query.ExcludeColumn(columns...) } return b diff --git a/pkg/common/adapters/database/bun_exclude_test.go b/pkg/common/adapters/database/bun_exclude_test.go new file mode 100644 index 0000000..54c39ad --- /dev/null +++ b/pkg/common/adapters/database/bun_exclude_test.go @@ -0,0 +1,100 @@ +package database + +import ( + "database/sql" + "strings" + "testing" + + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/pgdialect" + + "github.com/bitechdev/ResolveSpec/pkg/reflection" +) + +// adhocBuffer mirrors the real-world DBAdhocBuffer: scanonly fields with both +// bun and gorm read-only tags. +type adhocBuffer struct { + CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"` + CQL2 string `json:"cql2,omitempty" gorm:"->" bun:",scanonly"` + RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"` + RecordError string `json:"_error,omitempty" gorm:"-" bun:",scanonly"` +} + +type excludeModel struct { + bun.BaseModel `bun:"table:public.crmnote,alias:crmnote"` + ID int `json:"id" bun:"id,pk"` + Note string `json:"note" bun:"note,type:citext,"` + Norm string `json:"norm" bun:"norm,generated"` + + adhocBuffer `json:",omitempty" bun:",scanonly"` +} + +func newExcludeDB() *bun.DB { + return bun.NewDB(&sql.DB{}, pgdialect.New()) +} + +// TestBunExcludeColumnWithNonWritableColumns feeds the reflection output +// straight into the adapter, as the handlers do, for insert and update. +func TestBunExcludeColumnWithNonWritableColumns(t *testing.T) { + db := newExcludeDB() + m := &excludeModel{} + cols := reflection.NonWritableColumns(m) + if len(cols) == 0 { + t.Fatal("expected non-writable columns") + } + + ins := &BunInsertQuery{query: db.NewInsert().Model(m)} + ins.ExcludeColumn(cols...) + insSQL, err := ins.query.AppendQuery(db.QueryGen(), nil) + if err != nil { + t.Fatalf("insert: %v", err) + } + + upd := &BunUpdateQuery{query: db.NewUpdate().Model(m).Where("id = 1")} + upd.ExcludeColumn(cols...) + updSQL, err := upd.query.AppendQuery(db.QueryGen(), nil) + if err != nil { + t.Fatalf("update: %v", err) + } + + for name, q := range map[string]string{"insert": string(insSQL), "update": string(updSQL)} { + for _, bad := range []string{"cql1", "cql2", "_rownumber", "_error", "norm"} { + if strings.Contains(q, `"`+bad+`"`) { + t.Errorf("%s writes non-writable column %s: %s", name, bad, q) + } + } + if !strings.Contains(q, `"note"`) { + t.Errorf("%s dropped writable column note: %s", name, q) + } + } +} + +func TestBunExcludeColumnIgnoresUnknownAndKeepsWritable(t *testing.T) { + db := newExcludeDB() + m := &excludeModel{} + + ins := &BunInsertQuery{query: db.NewInsert().Model(m)} + ins.ExcludeColumn("does_not_exist", "note") + q, err := ins.query.AppendQuery(db.QueryGen(), nil) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(q), `"note"`) { + t.Errorf("writable column note should have been excluded: %s", q) + } +} + +func TestBunExcludeColumnOnlyNonWritable(t *testing.T) { + db := newExcludeDB() + ins := &BunInsertQuery{query: db.NewInsert().Model(&excludeModel{})} + ins.ExcludeColumn("cql1") // everything filtered out: must not error or panic + if _, err := ins.query.AppendQuery(db.QueryGen(), nil); err != nil { + t.Fatal(err) + } +} + +func TestBunExcludeColumnWithoutModel(t *testing.T) { + db := newExcludeDB() + ins := &BunInsertQuery{query: db.NewInsert()} + ins.ExcludeColumn("cql1") // no model yet: must not panic +} diff --git a/pkg/reflection/model_utils_test.go b/pkg/reflection/model_utils_test.go index 3b4b46b..50eec4d 100644 --- a/pkg/reflection/model_utils_test.go +++ b/pkg/reflection/model_utils_test.go @@ -497,13 +497,13 @@ func TestIsColumnWritableWithEmbedded(t *testing.T) { // Test models with relations for GetSQLModelColumns type User struct { - ID int `bun:"id,pk" json:"id"` - Name string `bun:"name" json:"name"` - Email string `bun:"email" json:"email"` - ProfileData string `json:"profile_data"` // No bun/gorm tag - Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"` - Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"` - RowNumber int64 `bun:",scanonly" json:"_rownumber"` + ID int `bun:"id,pk" json:"id"` + Name string `bun:"name" json:"name"` + Email string `bun:"email" json:"email"` + ProfileData string `json:"profile_data"` // No bun/gorm tag + Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"` + Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"` + RowNumber int64 `bun:",scanonly" json:"_rownumber"` } type Post struct { @@ -528,8 +528,8 @@ type Tag struct { // Model with scan-only embedded struct type EntityWithScanOnlyEmbedded struct { - ID int `bun:"id,pk" json:"id"` - Name string `bun:"name" json:"name"` + ID int `bun:"id,pk" json:"id"` + Name string `bun:"name" json:"name"` AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only } @@ -1086,17 +1086,17 @@ func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) { // Models for relation testing type Author struct { - ID int `bun:"id,pk" json:"id"` - Name string `bun:"name" json:"name"` - Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"` + ID int `bun:"id,pk" json:"id"` + Name string `bun:"name" json:"name"` + Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"` } type Book struct { - ID int `bun:"id,pk" json:"id"` - Title string `bun:"title" json:"title"` - AuthorID int `bun:"author_id" json:"author_id"` - Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"` - Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"` + ID int `bun:"id,pk" json:"id"` + Title string `bun:"title" json:"title"` + AuthorID int `bun:"author_id" json:"author_id"` + Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"` + Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"` } type Publisher struct { @@ -1106,9 +1106,9 @@ type Publisher struct { } type Student struct { - ID int `gorm:"column:id;primaryKey" json:"id"` - Name string `gorm:"column:name" json:"name"` - Courses []Course `gorm:"many2many:student_courses" json:"courses"` + ID int `gorm:"column:id;primaryKey" json:"id"` + Name string `gorm:"column:name" json:"name"` + Courses []Course `gorm:"many2many:student_courses" json:"courses"` } type Course struct { @@ -1119,11 +1119,11 @@ type Course struct { // Recursive relation model type Category struct { - ID int `bun:"id,pk" json:"id"` - Name string `bun:"name" json:"name"` - ParentID *int `bun:"parent_id" json:"parent_id"` - Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"` - Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"` + ID int `bun:"id,pk" json:"id"` + Name string `bun:"name" json:"name"` + ParentID *int `bun:"parent_id" json:"parent_id"` + Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"` + Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"` } func TestGetRelationType(t *testing.T) { @@ -1299,7 +1299,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) { expected: nil, }, { - name: "model without primary key tags - fallback to ID field", + name: "model without primary key tags - fallback to ID field", model: struct { ID int Name string @@ -1307,7 +1307,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) { expected: 99, }, { - name: "model without ID field", + name: "model without ID field", model: struct { Name string }{Name: "Test"}, @@ -1508,10 +1508,10 @@ func TestGetSQLModelColumns_EdgeCases(t *testing.T) { // Test models with table:, rel:, join: tags for ExtractColumnFromBunTag type BunSpecialTagsModel struct { - Table string `bun:"table:users"` - Relation []Post `bun:"rel:has-many"` - Join string `bun:"join:id=user_id"` - NormalCol string `bun:"normal_col"` + Table string `bun:"table:users"` + Relation []Post `bun:"rel:has-many"` + Join string `bun:"join:id=user_id"` + NormalCol string `bun:"normal_col"` } func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) { @@ -1592,8 +1592,8 @@ func TestGetRelationType_GORMFallback(t *testing.T) { func TestGetRelationType_AdditionalCases(t *testing.T) { // Test model with GORM has-one (pointer without foreignKey or with references) type Address struct { - ID int `gorm:"column:id;primaryKey"` - UserID int `gorm:"column:user_id"` + ID int `gorm:"column:id;primaryKey"` + UserID int `gorm:"column:user_id"` } type UserWithAddress struct { @@ -1609,7 +1609,7 @@ func TestGetRelationType_AdditionalCases(t *testing.T) { type Employee struct { ID int - Company Company // Single struct (not pointer, not slice) - belongs-to + Company Company // Single struct (not pointer, not slice) - belongs-to Coworkers []Employee // Slice without bun/gorm tags - has-many } @@ -1963,3 +1963,31 @@ func TestNonWritableColumns(t *testing.T) { } } } + +func TestNonWritableColumns_EmbeddedScanOnlyBuffer(t *testing.T) { + type buffer struct { + CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"` + RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"` + } + type m struct { + ID int `json:"id" bun:"id,pk"` + Note string `json:"note" bun:"note,type:citext,"` + buffer `json:",omitempty" bun:",scanonly"` + } + got := NonWritableColumns(&m{}) + has := map[string]bool{} + for _, c := range got { + has[c] = true + } + if !has["cql1"] { + t.Errorf("cql1 should be non-writable, got %v", got) + } + if has["id"] || has["note"] { + t.Errorf("writable columns reported as non-writable: %v", got) + } + vals := map[string]interface{}{"id": 1, "note": "x", "cql1": "y"} + RemoveNonWritableColumns(&m{}, vals) + if _, ok := vals["cql1"]; ok || len(vals) != 2 { + t.Errorf("unexpected values: %v", vals) + } +}