mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-06 13:26:28 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cff3bde85 | ||
|
|
3e6224698c |
@@ -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