mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-06 05:16:27 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cff3bde85 | ||
|
|
3e6224698c | ||
|
|
aec87a81e7 |
@@ -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,8 +1508,30 @@ 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 {
|
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...)
|
b.query = b.query.ExcludeColumn(columns...)
|
||||||
}
|
}
|
||||||
return b
|
return b
|
||||||
@@ -1627,7 +1650,7 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
|
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...)
|
b.query = b.query.ExcludeColumn(columns...)
|
||||||
}
|
}
|
||||||
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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user