Compare commits

...
5 Commits
Author SHA1 Message Date
Hein 8cff3bde85 test(wrap_bunrouter): add tests for route param preservation
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m28s
Tests / Unit Tests (push) Successful in 1m30s
Build , Vet Test, and Lint / Lint Code (push) Successful in 1m53s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m9s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m10s
Tests / Race Detector (push) Successful in 3m39s
2026-10-05 16:44:44 +02:00
Hein 3e6224698c fix(handler): enforce single record return for ID queries 2026-10-05 16:04:57 +02:00
Hein aec87a81e7 fix(bun): ignore scanonly columns in ExcludeColumn
Tests / Integration Tests (push) Skipped
Tests / Unit Tests (push) Successful in 1m37s
Tests / Race Detector (push) Successful in 3m52s
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m16s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m19s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m19s
Bun's ExcludeColumn errors with "can't find column" for scanonly fields
because they are not in the table's writable fields. Filter the exclude
list to writable bun fields so models with scanonly buffers can insert
and update again. Add tests for the adapter and reflection.
2026-10-05 14:10:52 +02:00
warkanum 9235292586 fix(crud): skip generated and read-only columns on insert and update
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m34s
Tests / Unit Tests (push) Successful in 1m42s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m11s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m25s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m28s
Tests / Race Detector (push) Successful in 4m6s
Read-merge-update wrote every model column back, so GENERATED ALWAYS
columns failed with SQLSTATE 428C9. Add a bun 'generated' tag option,
reflection.NonWritableColumns/RemoveNonWritableColumns, and apply them
in resolvespec, restheadspec, websocketspec, mqttspec, resolvemcp and
the nested CUD processor. Add ExcludeColumn to InsertQuery/UpdateQuery
for model-based writes.
2026-10-02 22:43:45 +02:00
warkanum 23f10387c5 ci(release): fix rust toolchain setup and dart publish validation warnings
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 3m15s
Tests / Unit Tests (push) Successful in 3m26s
Build , Vet Test, and Lint / Lint Code (push) Successful in 4m45s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 4m55s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 4m55s
Tests / Race Detector (push) Successful in 7m6s
2026-10-01 21:20:04 +02:00
21 changed files with 678 additions and 41 deletions
+7 -3
View File
@@ -129,9 +129,7 @@ jobs:
- uses: actions/checkout@v4 - uses: actions/checkout@v4
- name: Set up Rust - name: Set up Rust
run: | uses: dtolnay/rust-toolchain@stable
rustup toolchain install stable --profile minimal
rustup default stable
- name: Test - name: Test
run: cargo test run: cargo test
@@ -252,6 +250,12 @@ jobs:
run: | run: |
sed -i -E "s/^version: .*/version: ${VERSION}/" pubspec.yaml sed -i -E "s/^version: .*/version: ${VERSION}/" pubspec.yaml
sed -i -E "s#^publish_to: .*#publish_to: ${SERVER_URL}/api/packages/${OWNER}/pub#" pubspec.yaml sed -i -E "s#^publish_to: .*#publish_to: ${SERVER_URL}/api/packages/${OWNER}/pub#" pubspec.yaml
if ! grep -q "^## ${VERSION}\$" CHANGELOG.md; then
{ head -n 1 CHANGELOG.md; printf '\n## %s\n\n- Release %s.\n' "$VERSION" "$VERSION"; tail -n +2 CHANGELOG.md; } > CHANGELOG.tmp
mv CHANGELOG.tmp CHANGELOG.md
fi
# pub warns about a dirty git tree; commit the stamped files locally (never pushed)
git -c user.name=ci -c user.email=ci@localhost commit -q -am "ci: stamp dart version ${VERSION}"
- name: Dry run - name: Dry run
if: ${{ env.PUBLISH != 'true' }} if: ${{ env.PUBLISH != 'true' }}
+1
View File
@@ -1,6 +1,7 @@
name: resolvespec name: resolvespec
description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints. description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints.
version: 0.1.0 version: 0.1.0
repository: https://git.warky.dev/wdevs/ResolveSpec
publish_to: none publish_to: none
environment: environment:
+37
View File
@@ -10,6 +10,7 @@ import (
"time" "time"
"github.com/uptrace/bun" "github.com/uptrace/bun"
"github.com/uptrace/bun/schema"
"github.com/bitechdev/ResolveSpec/pkg/common" "github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace" "github.com/bitechdev/ResolveSpec/pkg/dbtrace"
@@ -1507,6 +1508,35 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery {
return b 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 columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
b.query = b.query.ExcludeColumn(columns...)
}
return b
}
func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery { func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery {
if len(columns) > 0 { if len(columns) > 0 {
b.query = b.query.Returning(strings.Join(columns, ", ")) b.query = b.query.Returning(strings.Join(columns, ", "))
@@ -1619,6 +1649,13 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
return b return b
} }
func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
b.query = b.query.ExcludeColumn(columns...)
}
return b
}
func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery { func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
b.query = b.query.Where(query, args...) b.query = b.query.Where(query, args...)
return b return b
@@ -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
}
+14
View File
@@ -751,6 +751,13 @@ func (g *GormInsertQuery) OnConflict(action string) common.InsertQuery {
return g 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 { func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery {
g.returningColumns = columns g.returningColumns = columns
return g return g
@@ -930,6 +937,13 @@ func (g *GormUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQue
return g 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 { func (g *GormUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
g.db = g.db.Where(query, args...) g.db = g.db.Where(query, args...)
return g return g
+14
View File
@@ -691,6 +691,13 @@ func (p *PgSQLInsertQuery) OnConflict(action string) common.InsertQuery {
return p 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 { func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery {
p.returning = columns p.returning = columns
return p return p
@@ -850,6 +857,13 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu
return p 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 { func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuery {
pkName := "" pkName := ""
if p.model != nil { if p.model != nil {
+4
View File
@@ -81,6 +81,8 @@ type InsertQuery interface {
Table(table string) InsertQuery Table(table string) InsertQuery
Value(column string, value interface{}) InsertQuery Value(column string, value interface{}) InsertQuery
OnConflict(action string) 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 Returning(columns ...string) InsertQuery
// Execution // Execution
@@ -94,6 +96,8 @@ type UpdateQuery interface {
Table(table string) UpdateQuery Table(table string) UpdateQuery
Set(column string, value interface{}) UpdateQuery Set(column string, value interface{}) UpdateQuery
SetMap(values map[string]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 Where(query string, args ...interface{}) UpdateQuery
Returning(columns ...string) UpdateQuery Returning(columns ...string) UpdateQuery
+6 -2
View File
@@ -116,7 +116,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
case "insert", "create", "add": case "insert", "create", "add":
// Only perform insert if we have data to insert // Only perform insert if we have data to insert
if hasData { if hasData {
id, err := p.processInsert(ctx, regularData, tableName) id, err := p.processInsert(ctx, regularData, model, tableName)
if err != nil { if err != nil {
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err) logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
return nil, fmt.Errorf("insert failed: %w", err) return nil, fmt.Errorf("insert failed: %w", err)
@@ -148,7 +148,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
return result, nil return result, nil
} }
if hasData { if hasData {
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName]) rows, err := p.processUpdate(ctx, regularData, model, tableName, data[pkName])
if err != nil { if err != nil {
logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err) 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) 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( func (p *NestedCUDProcessor) processInsert(
ctx context.Context, ctx context.Context,
data map[string]interface{}, data map[string]interface{},
model interface{},
tableName string, tableName string,
) (interface{}, error) { ) (interface{}, error) {
logger.Debug("Inserting into %s with data: %+v", tableName, data) logger.Debug("Inserting into %s with data: %+v", tableName, data)
reflection.RemoveNonWritableColumns(model, data)
query := p.db.NewInsert().Table(tableName) query := p.db.NewInsert().Table(tableName)
for key, value := range data { for key, value := range data {
@@ -335,6 +337,7 @@ func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string
func (p *NestedCUDProcessor) processUpdate( func (p *NestedCUDProcessor) processUpdate(
ctx context.Context, ctx context.Context,
data map[string]interface{}, data map[string]interface{},
model interface{},
tableName string, tableName string,
id interface{}, id interface{},
) (int64, error) { ) (int64, error) {
@@ -345,6 +348,7 @@ func (p *NestedCUDProcessor) processUpdate(
logger.Debug("Updating %s with ID %v, data: %+v", tableName, id, data) 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) query := p.db.NewUpdate().Table(tableName).SetMap(data).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id)
result, err := query.Exec(ctx) result, err := query.Exec(ctx)
+2
View File
@@ -99,6 +99,7 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
return m return m
} }
func (m *mockInsertQuery) OnConflict(action string) 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) Returning(columns ...string) InsertQuery { return m }
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) { func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
m.db.insertCalls = append(m.db.insertCalls, m.values) m.db.insertCalls = append(m.db.insertCalls, m.values)
@@ -131,6 +132,7 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
return m return m
} }
func (m *mockUpdateQuery) Where(condition string, args ...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) Returning(columns ...string) UpdateQuery { return m }
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) { func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
// Record the update call // Record the update call
+5
View File
@@ -895,6 +895,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// Insert record // Insert record
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) 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 { if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to create record: %w", err) 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. // 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) values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
if len(values) > 0 { if len(values) > 0 {
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
+65 -1
View File
@@ -656,7 +656,7 @@ func isColumnWritableInType(typ reflect.Type, columnName string) (found bool, wr
// Check bun tag for scanonly // Check bun tag for scanonly
bunTag := field.Tag.Get("bun") bunTag := field.Tag.Get("bun")
if bunTag != "" { if bunTag != "" {
if isBunFieldScanOnly(bunTag) { if isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag) {
return true, false return true, false
} }
} }
@@ -689,6 +689,70 @@ func isBunFieldScanOnly(tag string) bool {
return false 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 // isGormFieldReadOnly checks if a gorm tag indicates the field is read-only
// Examples: // Examples:
// - "<-:false" -> true (no writes allowed) // - "<-:false" -> true (no writes allowed)
+105 -34
View File
@@ -497,13 +497,13 @@ func TestIsColumnWritableWithEmbedded(t *testing.T) {
// Test models with relations for GetSQLModelColumns // Test models with relations for GetSQLModelColumns
type User struct { type User struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
Email string `bun:"email" json:"email"` Email string `bun:"email" json:"email"`
ProfileData string `json:"profile_data"` // No bun/gorm tag ProfileData string `json:"profile_data"` // No bun/gorm tag
Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"` Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"`
Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"` Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"`
RowNumber int64 `bun:",scanonly" json:"_rownumber"` RowNumber int64 `bun:",scanonly" json:"_rownumber"`
} }
type Post struct { type Post struct {
@@ -528,8 +528,8 @@ type Tag struct {
// Model with scan-only embedded struct // Model with scan-only embedded struct
type EntityWithScanOnlyEmbedded struct { type EntityWithScanOnlyEmbedded struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only
} }
@@ -1086,17 +1086,17 @@ func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) {
// Models for relation testing // Models for relation testing
type Author struct { type Author struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"` Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"`
} }
type Book struct { type Book struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Title string `bun:"title" json:"title"` Title string `bun:"title" json:"title"`
AuthorID int `bun:"author_id" json:"author_id"` AuthorID int `bun:"author_id" json:"author_id"`
Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"` Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"`
Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"` Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"`
} }
type Publisher struct { type Publisher struct {
@@ -1106,9 +1106,9 @@ type Publisher struct {
} }
type Student struct { type Student struct {
ID int `gorm:"column:id;primaryKey" json:"id"` ID int `gorm:"column:id;primaryKey" json:"id"`
Name string `gorm:"column:name" json:"name"` Name string `gorm:"column:name" json:"name"`
Courses []Course `gorm:"many2many:student_courses" json:"courses"` Courses []Course `gorm:"many2many:student_courses" json:"courses"`
} }
type Course struct { type Course struct {
@@ -1119,11 +1119,11 @@ type Course struct {
// Recursive relation model // Recursive relation model
type Category struct { type Category struct {
ID int `bun:"id,pk" json:"id"` ID int `bun:"id,pk" json:"id"`
Name string `bun:"name" json:"name"` Name string `bun:"name" json:"name"`
ParentID *int `bun:"parent_id" json:"parent_id"` ParentID *int `bun:"parent_id" json:"parent_id"`
Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"` Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"`
Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"` Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"`
} }
func TestGetRelationType(t *testing.T) { func TestGetRelationType(t *testing.T) {
@@ -1299,7 +1299,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
expected: nil, expected: nil,
}, },
{ {
name: "model without primary key tags - fallback to ID field", name: "model without primary key tags - fallback to ID field",
model: struct { model: struct {
ID int ID int
Name string Name string
@@ -1307,7 +1307,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
expected: 99, expected: 99,
}, },
{ {
name: "model without ID field", name: "model without ID field",
model: struct { model: struct {
Name string Name string
}{Name: "Test"}, }{Name: "Test"},
@@ -1508,10 +1508,10 @@ func TestGetSQLModelColumns_EdgeCases(t *testing.T) {
// Test models with table:, rel:, join: tags for ExtractColumnFromBunTag // Test models with table:, rel:, join: tags for ExtractColumnFromBunTag
type BunSpecialTagsModel struct { type BunSpecialTagsModel struct {
Table string `bun:"table:users"` Table string `bun:"table:users"`
Relation []Post `bun:"rel:has-many"` Relation []Post `bun:"rel:has-many"`
Join string `bun:"join:id=user_id"` Join string `bun:"join:id=user_id"`
NormalCol string `bun:"normal_col"` NormalCol string `bun:"normal_col"`
} }
func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) { func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) {
@@ -1592,8 +1592,8 @@ func TestGetRelationType_GORMFallback(t *testing.T) {
func TestGetRelationType_AdditionalCases(t *testing.T) { func TestGetRelationType_AdditionalCases(t *testing.T) {
// Test model with GORM has-one (pointer without foreignKey or with references) // Test model with GORM has-one (pointer without foreignKey or with references)
type Address struct { type Address struct {
ID int `gorm:"column:id;primaryKey"` ID int `gorm:"column:id;primaryKey"`
UserID int `gorm:"column:user_id"` UserID int `gorm:"column:user_id"`
} }
type UserWithAddress struct { type UserWithAddress struct {
@@ -1609,7 +1609,7 @@ func TestGetRelationType_AdditionalCases(t *testing.T) {
type Employee struct { type Employee struct {
ID int 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 Coworkers []Employee // Slice without bun/gorm tags - has-many
} }
@@ -1920,3 +1920,74 @@ 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)
}
}
}
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)
}
}
+2
View File
@@ -559,6 +559,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
if len(cols) == 0 { if len(cols) == 0 {
return invalidArg("no writable fields in data") return invalidArg("no writable fields in data")
} }
reflection.RemoveNonWritableColumns(model, cols)
q := tx.NewInsert().Table(tableName) q := tx.NewInsert().Table(tableName)
for key, value := range cols { for key, value := range cols {
q = q.Value(key, value) q = q.Value(key, value)
@@ -726,6 +727,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
existingMap[key] = v existingMap[key] = v
} }
reflection.RemoveNonWritableColumns(model, setCols)
q := tx.NewUpdate().Table(tableName).SetMap(setCols). q := tx.NewUpdate().Table(tableName).SetMap(setCols).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id) Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
res, err := q.Exec(ctx) res, err := q.Exec(ctx)
+1
View File
@@ -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, ", ")) cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", "))
var affected int64 var affected int64
if req.op == "update" { if req.op == "update" {
reflection.RemoveNonWritableColumns(model, setCols)
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx) r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
if err != nil { if err != nil {
return fmt.Errorf("error updating records: %w", err) return fmt.Errorf("error updating records: %w", err)
+6
View File
@@ -824,6 +824,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
} }
responseData = v responseData = v
reflection.RemoveNonWritableColumns(model, v)
query := tx.NewInsert().Table(tableName) query := tx.NewInsert().Table(tableName)
for key, value := range v { for key, value := range v {
query = query.Value(key, common.ConvertSliceForBun(value)) query = query.Value(key, common.ConvertSliceForBun(value))
@@ -971,6 +972,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
item = modifiedData item = modifiedData
} }
reflection.RemoveNonWritableColumns(model, item)
txQuery := tx.NewInsert().Table(tableName) txQuery := tx.NewInsert().Table(tableName)
for key, value := range item { for key, value := range item {
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value)) txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
@@ -1127,6 +1129,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
itemMap = modifiedData itemMap = modifiedData
} }
reflection.RemoveNonWritableColumns(model, itemMap)
txQuery := tx.NewInsert().Table(tableName) txQuery := tx.NewInsert().Table(tableName)
for key, value := range itemMap { for key, value := range itemMap {
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value)) 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) // Overwrite with every key present in the request (including "" and null unless disallowed)
common.MergeUpdateValues(existingMap, updates, h.disallowNulls) common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
reflection.RemoveNonWritableColumns(model, existingMap)
// Build update query with merged data // Build update query with merged data
query := tx.NewUpdate().Table(tableName).SetMap(existingMap) 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) // Overwrite with every key present in the request (including "" and null unless disallowed)
common.MergeUpdateValues(existingMap, item, h.disallowNulls) 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) txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if _, err := txQuery.Exec(ctx); err != nil { 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) // Overwrite with every key present in the request (including "" and null unless disallowed)
common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls) 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) txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if _, err := txQuery.Exec(ctx); err != nil { if _, err := txQuery.Exec(ctx); err != nil {
+61
View File
@@ -0,0 +1,61 @@
package resolvespec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/uptrace/bunrouter"
)
type wrapCtxKey struct{}
// The auth wrapper must hand the handler the middleware-enriched request
// without dropping the bunrouter route params.
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
var gotSchema, gotEntity, gotID string
var gotCtxVal any
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotSchema = req.Param("schema")
gotEntity = req.Param("entity")
gotID = req.Param("id")
gotCtxVal = req.Context().Value(wrapCtxKey{})
return nil
}
auth := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
})
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
}
if gotCtxVal != "enriched" {
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
}
}
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
var gotID string
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotID = req.Param("id")
return nil
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
if gotID != "7" {
t.Errorf("id = %q, want 7", gotID)
}
}
+15 -1
View File
@@ -417,6 +417,14 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
if id == "" { if id == "" {
options.SingleRecordAsObject = false options.SingleRecordAsObject = false
} else {
// The primary key is already filtered, so never return more than one
// record regardless of limit/offset/cursor headers or joins.
one := 1
options.Limit = &one
options.Offset = nil
options.CursorForward = ""
options.CursorBackward = ""
} }
// Validate and unwrap model type to get base struct // Validate and unwrap model type to get base struct
@@ -726,7 +734,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
sanitizedOr = common.EnsureOuterParentheses(sanitizedOr) sanitizedOr = common.EnsureOuterParentheses(sanitizedOr)
} }
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" && common.Hardening().SQLStrict { if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" {
query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery { query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery {
return applyUserConds(q).WhereOr(sanitizedOr) return applyUserConds(q).WhereOr(sanitizedOr)
}) })
@@ -1410,6 +1418,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
if provider, ok := modelValue.(common.TableNameProvider); !ok || provider.TableName() == "" { if provider, ok := modelValue.(common.TableNameProvider); !ok || provider.TableName() == "" {
query = query.Table(tableName) query = query.Table(tableName)
} }
if generated := reflection.NonWritableColumns(model); len(generated) > 0 {
query = query.ExcludeColumn(generated...)
}
fields := reflection.GetSQLModelColumns(model) fields := reflection.GetSQLModelColumns(model)
query = query.Returning(fields...) query = query.Returning(fields...)
@@ -1657,6 +1668,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 // Create update query using Model() to preserve custom types and driver.Valuer interfaces
query := tx.NewUpdate().Model(modelInstance) 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) query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
// Execute BeforeScan hooks - pass query chain so hooks can modify it // Execute BeforeScan hooks - pass query chain so hooks can modify it
+157
View File
@@ -0,0 +1,157 @@
package restheadspec
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"testing"
"github.com/DATA-DOG/go-sqlmock"
"github.com/uptrace/bun"
"github.com/uptrace/bun/dialect/pgdialect"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
// readCapturingSQL runs handleRead and returns every SELECT it issued.
func readCapturingSQL(t *testing.T, id string, options ExtendedRequestOptions) []string {
queries, _ := readCapturingSQLAndBody(t, id, options)
return queries
}
// readCapturingSQLAndBody is readCapturingSQL that also returns the response body.
// The mocked row carries the requested id so the body can be checked against it.
func readCapturingSQLAndBody(t *testing.T, id string, options ExtendedRequestOptions) ([]string, string) {
t.Helper()
resetTotalCache(t)
var queries []string
matcher := sqlmock.QueryMatcherFunc(func(_, actual string) error {
queries = append(queries, actual)
return nil
})
sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(matcher))
if err != nil {
t.Fatal(err)
}
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { _ = sqlDB.Close() })
h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry())
rowID, err := strconv.Atoi(id)
if err != nil {
rowID = 7
}
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(rowID, "a"))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectCommit()
rec := httptest.NewRecorder()
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil))
h.handleRead(itemCtx(t), w, id, options)
if rec.Code != http.StatusOK {
t.Fatalf("status %d body %s", rec.Code, rec.Body)
}
return queries, rec.Body.String()
}
func TestReadByIDIgnoresLimitOffsetAndCursor(t *testing.T) {
limit, offset := 50, 10
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{
Limit: &limit,
Offset: &offset,
},
})
last := queries[len(queries)-1]
if !strings.Contains(last, "LIMIT 1") || strings.Contains(last, "OFFSET") {
t.Fatalf("read by id must be LIMIT 1 with no OFFSET: %s", last)
}
}
func TestReadWithoutIDKeepsRequestedLimit(t *testing.T) {
limit := 50
queries := readCapturingSQL(t, "", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{Limit: &limit},
})
if last := queries[len(queries)-1]; !strings.Contains(last, "LIMIT 50") {
t.Fatalf("list read must keep its limit: %s", last)
}
}
// topLevelOr reports whether the WHERE clause has an OR outside any parentheses,
// i.e. one that would let rows bypass the AND-ed primary key condition.
func topLevelOr(sql string) bool {
where := sql[strings.Index(sql, "WHERE")+len("WHERE"):]
depth, inStr := 0, false
for i := 0; i < len(where); i++ {
switch c := where[i]; {
case c == '\'':
inStr = !inStr
case inStr:
case c == '(':
depth++
case c == ')':
depth--
case depth == 0 && strings.HasPrefix(where[i:], " OR "):
return true
}
}
return false
}
func TestReadByIDCustomSQLOrCannotEscapePrimaryKey(t *testing.T) {
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
RequestOptions: common.RequestOptions{
Filters: []common.FilterOption{{Column: "name", Operator: "eq", Value: "a"}},
},
CustomSQLOr: "name = 'x'",
})
last := queries[len(queries)-1]
if !strings.Contains(last, `"id" = '7'`) && !strings.Contains(last, `"id" = 7`) {
t.Fatalf("primary key filter missing: %s", last)
}
if topLevelOr(last) {
t.Fatalf("OR escapes the primary key filter: %s", last)
}
}
func TestReadByIDFiltersAndReturnsRequestedRecord(t *testing.T) {
queries, body := readCapturingSQLAndBody(t, "42", ExtendedRequestOptions{})
last := queries[len(queries)-1]
if !strings.Contains(last, `"items"."id" = '42'`) && !strings.Contains(last, `"items"."id" = 42`) {
t.Fatalf("query must filter the primary key to 42: %s", last)
}
if strings.Contains(last, "= 7") || strings.Contains(last, "= '7'") {
t.Fatalf("query filters a different id: %s", last)
}
// every query that touches rows (count and select) must carry the id filter
for _, q := range queries {
if strings.Contains(q, "FROM") && !strings.Contains(q, "42") {
t.Fatalf("query without the id filter: %s", q)
}
}
var rows []struct {
ID int `json:"id"`
}
data := body
if i := strings.Index(body, `"data"`); i >= 0 {
data = body[i+len(`"data"`):]
}
if i := strings.Index(data, "["); i >= 0 {
data = data[i:]
}
dec := json.NewDecoder(strings.NewReader(data))
if err := dec.Decode(&rows); err != nil {
t.Fatalf("decode %q: %v", body, err)
}
if len(rows) != 1 || rows[0].ID != 42 {
t.Fatalf("response must contain exactly the record with id 42: %s", body)
}
}
+61
View File
@@ -0,0 +1,61 @@
package restheadspec
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/uptrace/bunrouter"
)
type wrapCtxKey struct{}
// The auth wrapper must hand the handler the middleware-enriched request
// without dropping the bunrouter route params.
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
var gotSchema, gotEntity, gotID string
var gotCtxVal any
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotSchema = req.Param("schema")
gotEntity = req.Param("entity")
gotID = req.Param("id")
gotCtxVal = req.Context().Value(wrapCtxKey{})
return nil
}
auth := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
})
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
}
if gotCtxVal != "enriched" {
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
}
}
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
var gotID string
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
gotID = req.Param("id")
return nil
}
router := bunrouter.New()
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
if gotID != "7" {
t.Errorf("id = %q, want 7", gotID)
}
}
+5
View File
@@ -758,6 +758,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// Insert record // Insert record
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName) 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 { if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to create record: %w", err) 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. // 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) values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
if len(values) > 0 { if len(values) > 0 {
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values). query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID) Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
+10
View File
@@ -226,6 +226,11 @@ func (m *MockInsertQuery) OnConflict(action string) common.InsertQuery {
return args.Get(0).(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 { func (m *MockInsertQuery) Returning(columns ...string) common.InsertQuery {
args := m.Called(columns) args := m.Called(columns)
return args.Get(0).(common.InsertQuery) return args.Get(0).(common.InsertQuery)
@@ -254,6 +259,11 @@ func (m *MockUpdateQuery) Model(model interface{}) common.UpdateQuery {
return args.Get(0).(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 { func (m *MockUpdateQuery) Table(table string) common.UpdateQuery {
args := m.Called(table) args := m.Called(table)
return args.Get(0).(common.UpdateQuery) return args.Get(0).(common.UpdateQuery)