Compare commits

..
2 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
4 changed files with 288 additions and 1 deletions
+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)
}
}