mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-06 13:26:28 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cff3bde85 | ||
|
|
3e6224698c | ||
|
|
aec87a81e7 | ||
|
|
9235292586 | ||
|
|
23f10387c5 |
@@ -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,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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user