Compare commits

..
3 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
7 changed files with 475 additions and 37 deletions
+25 -2
View File
@@ -10,6 +10,7 @@ import (
"time"
"github.com/uptrace/bun"
"github.com/uptrace/bun/schema"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
@@ -1507,8 +1508,30 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery {
return b
}
// bunWritableExcludes drops columns bun already leaves out of INSERT/UPDATE
// (scanonly fields) or does not know, since bun's ExcludeColumn errors with
// "can't find column" for anything that is not in the table's writable fields.
func bunWritableExcludes(model bun.Model, columns []string) []string {
tm, ok := model.(interface{ Table() *schema.Table })
if !ok || tm.Table() == nil {
return columns
}
table := tm.Table()
writable := make(map[string]struct{}, len(table.Fields))
for _, f := range table.Fields {
writable[f.Name] = struct{}{}
}
out := make([]string, 0, len(columns))
for _, c := range columns {
if _, ok := writable[c]; ok || c == "*" {
out = append(out, c)
}
}
return out
}
func (b *BunInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
if len(columns) > 0 {
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
b.query = b.query.ExcludeColumn(columns...)
}
return b
@@ -1627,7 +1650,7 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
}
func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
if len(columns) > 0 {
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
b.query = b.query.ExcludeColumn(columns...)
}
return b
@@ -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
}
+28
View File
@@ -1963,3 +1963,31 @@ func TestNonWritableColumns(t *testing.T) {
}
}
}
func TestNonWritableColumns_EmbeddedScanOnlyBuffer(t *testing.T) {
type buffer struct {
CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"`
RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"`
}
type m struct {
ID int `json:"id" bun:"id,pk"`
Note string `json:"note" bun:"note,type:citext,"`
buffer `json:",omitempty" bun:",scanonly"`
}
got := NonWritableColumns(&m{})
has := map[string]bool{}
for _, c := range got {
has[c] = true
}
if !has["cql1"] {
t.Errorf("cql1 should be non-writable, got %v", got)
}
if has["id"] || has["note"] {
t.Errorf("writable columns reported as non-writable: %v", got)
}
vals := map[string]interface{}{"id": 1, "note": "x", "cql1": "y"}
RemoveNonWritableColumns(&m{}, vals)
if _, ok := vals["cql1"]; ok || len(vals) != 2 {
t.Errorf("unexpected values: %v", vals)
}
}
+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)
}
}
+9 -1
View File
@@ -417,6 +417,14 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
if id == "" {
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
@@ -726,7 +734,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
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 {
return applyUserConds(q).WhereOr(sanitizedOr)
})
+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)
}
}