test(09-04): add failing tests for deterministic list queries
- Search, sort, filters, and adjacent pages must return the D-11 envelope - Empty and single results keep an array and the requested page size - Unknown identifiers and injected values fail closed before unsafe SQL
This commit is contained in:
428
cabana/query_test.go
Normal file
428
cabana/query_test.go
Normal file
@@ -0,0 +1,428 @@
|
||||
package cabana
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.golem15.com/golem15/summercms/backpack"
|
||||
"git.golem15.com/golem15/summercms/bouncer"
|
||||
"git.golem15.com/golem15/summercms/pact"
|
||||
tcpostgres "github.com/testcontainers/testcontainers-go/modules/postgres"
|
||||
pgdriver "gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const queryListConfig = `modelClass: Widget
|
||||
list: ~/plugins/acme/demo/models/widget/columns.yaml
|
||||
recordsPerPage: 20
|
||||
perPageOptions:
|
||||
- 2
|
||||
- 10
|
||||
- 20
|
||||
showSorting: true
|
||||
defaultSort:
|
||||
column: name
|
||||
direction: asc
|
||||
filter: config_filter.yaml
|
||||
`
|
||||
|
||||
const queryColumns = `columns:
|
||||
name:
|
||||
label: Name
|
||||
searchable: true
|
||||
sortable: true
|
||||
active:
|
||||
label: Active
|
||||
type: switch
|
||||
sortable: true
|
||||
artist:
|
||||
label: Artist
|
||||
relation: Artist
|
||||
select: name
|
||||
searchable: true
|
||||
`
|
||||
|
||||
const queryFilters = `scopes:
|
||||
active:
|
||||
label: Active
|
||||
type: switch
|
||||
column: active
|
||||
created:
|
||||
label: Created
|
||||
type: daterange
|
||||
column: created_at
|
||||
grouped:
|
||||
label: Group
|
||||
modelClass: Group
|
||||
nameFrom: name
|
||||
scope: filterByGroup
|
||||
`
|
||||
|
||||
func TestListQueryContract(t *testing.T) {
|
||||
svc, db := newListService(t)
|
||||
seedListRows(t, db)
|
||||
rec := getList(t, svc, "search=Other")
|
||||
body := decodeList(t, rec, http.StatusOK)
|
||||
if body.Meta.Total != 1 || len(body.Data) != 1 || rowID(body.Data[0]) != 4 {
|
||||
t.Fatalf("search = total %d ids %v", body.Meta.Total, rowIDs(body.Data))
|
||||
}
|
||||
sorted := decodeList(t, getList(t, svc, "sort=name&dir=desc&per_page=2&page=1"), http.StatusOK)
|
||||
if ids := rowIDs(sorted.Data); len(ids) != 2 || ids[0] != 1 || ids[1] != 2 {
|
||||
t.Fatalf("desc page = %v, want [1 2]", ids)
|
||||
}
|
||||
related := decodeList(t, getList(t, svc, "search=Beatles"), http.StatusOK)
|
||||
if related.Meta.Total != 1 || len(related.Data) != 1 || rowID(related.Data[0]) != 1 {
|
||||
t.Fatalf("relation search = total %d ids %v", related.Meta.Total, rowIDs(related.Data))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListQueryEmpty(t *testing.T) {
|
||||
svc, db := newListService(t)
|
||||
if err := db.Exec("DELETE FROM cabana_list_query_rows").Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rec := getList(t, svc, "per_page=10")
|
||||
if !json.Valid(rec.Body.Bytes()) || rec.Body.String() == "" {
|
||||
t.Fatalf("body %s", rec.Body.String())
|
||||
}
|
||||
body := decodeList(t, rec, http.StatusOK)
|
||||
if body.Data == nil || len(body.Data) != 0 || body.Meta.Page != 1 || body.Meta.PerPage != 10 || body.Meta.Total != 0 || body.Meta.LastPage != 1 {
|
||||
t.Fatalf("empty envelope = %+v data %#v", body.Meta, body.Data)
|
||||
}
|
||||
if strings.Contains(rec.Body.String(), `"data":null`) {
|
||||
t.Fatalf("empty data was null: %s", rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestListQuerySingle(t *testing.T) {
|
||||
svc, db := newListService(t)
|
||||
resetListRows(t, db)
|
||||
if err := db.Create(&queryRow{ID: 7, Name: "Only", Active: true, CreatedAt: time.Date(2024, 1, 2, 12, 0, 0, 0, time.UTC)}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
body := decodeList(t, getList(t, svc, "per_page=10"), http.StatusOK)
|
||||
if len(body.Data) != 1 || body.Meta.Total != 1 || body.Meta.LastPage != 1 || body.Meta.PerPage != 10 || rowID(body.Data[0]) != 7 {
|
||||
t.Fatalf("single = %+v ids %v", body.Meta, rowIDs(body.Data))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListQueryAdjacent(t *testing.T) {
|
||||
svc, db := newListService(t)
|
||||
seedListRows(t, db)
|
||||
page1 := decodeList(t, getList(t, svc, "sort=name&dir=asc&per_page=2&page=1"), http.StatusOK)
|
||||
page2 := decodeList(t, getList(t, svc, "sort=name&dir=asc&per_page=2&page=2"), http.StatusOK)
|
||||
if ids := rowIDs(page1.Data); len(ids) != 2 || ids[0] != 4 || ids[1] != 1 {
|
||||
t.Fatalf("page 1 = %v, want [4 1]", ids)
|
||||
}
|
||||
if ids := rowIDs(page2.Data); len(ids) != 2 || ids[0] != 2 || ids[1] != 3 {
|
||||
t.Fatalf("page 2 = %v, want [2 3]", ids)
|
||||
}
|
||||
seen := map[int]struct{}{}
|
||||
for _, id := range append(rowIDs(page1.Data), rowIDs(page2.Data)...) {
|
||||
if _, dup := seen[id]; dup {
|
||||
t.Fatalf("duplicate id %d across adjacent pages", id)
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
func TestListQueryFilters(t *testing.T) {
|
||||
svc, db := newListService(t)
|
||||
seedListRows(t, db)
|
||||
active := decodeList(t, getList(t, svc, "filter[active]=true&per_page=10"), http.StatusOK)
|
||||
if ids := rowIDs(active.Data); !sameIDs(ids, []int{1, 3, 4}) {
|
||||
t.Fatalf("switch ids = %v", ids)
|
||||
}
|
||||
dates := decodeList(t, getList(t, svc, "filter[created]=2024-01-01..2024-01-31&per_page=10"), http.StatusOK)
|
||||
if ids := rowIDs(dates.Data); !sameIDs(ids, []int{1, 2}) {
|
||||
t.Fatalf("date ids = %v", ids)
|
||||
}
|
||||
scoped := decodeList(t, getList(t, svc, "filter[grouped]=9&per_page=10"), http.StatusOK)
|
||||
if ids := rowIDs(scoped.Data); !sameIDs(ids, []int{3}) {
|
||||
t.Fatalf("scope ids = %v", ids)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListQueryRejectsInjection(t *testing.T) {
|
||||
svc, db := newListService(t)
|
||||
seedListRows(t, db)
|
||||
before := countListRows(t, db)
|
||||
cases := []struct {
|
||||
name string
|
||||
query string
|
||||
field string
|
||||
}{
|
||||
{name: "sort injection", query: "sort=name%3Bdrop", field: "sort"},
|
||||
{name: "sort case", query: "sort=Name", field: "sort"},
|
||||
{name: "direction case", query: "sort=name&dir=DESC", field: "dir"},
|
||||
{name: "unknown filter", query: "filter[missing]=1", field: "filter"},
|
||||
{name: "page zero", query: "page=0", field: "page"},
|
||||
{name: "page text", query: "page=foo", field: "page"},
|
||||
{name: "per page cap", query: "per_page=999999", field: "per_page"},
|
||||
{name: "bad switch", query: "filter[active]=maybe", field: "filter"},
|
||||
{name: "bad range", query: "filter[created]=2024-02-01..2024-01-01", field: "filter"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
rec := getList(t, svc, tc.query)
|
||||
if rec.Code != http.StatusUnprocessableEntity {
|
||||
t.Fatalf("status=%d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var body listEnvelope
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if body.Error.Code != "validation_failed" {
|
||||
t.Fatalf("code=%s body=%s", body.Error.Code, rec.Body.String())
|
||||
}
|
||||
if _, ok := body.Error.Details[tc.field]; !ok {
|
||||
t.Fatalf("details missing %s: %s", tc.field, rec.Body.String())
|
||||
}
|
||||
if strings.Contains(rec.Body.String(), "drop table") {
|
||||
t.Fatalf("error body echoed SQL: %s", rec.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
payload := decodeList(t, getList(t, svc, "search=%25%27+OR+%271%27%3D%271"), http.StatusOK)
|
||||
if payload.Meta.Total != 0 {
|
||||
t.Fatalf("injected search matched %d rows", payload.Meta.Total)
|
||||
}
|
||||
if countListRows(t, db) != before {
|
||||
t.Fatalf("row count changed from %d to %d", before, countListRows(t, db))
|
||||
}
|
||||
|
||||
t.Run("permission before query", func(t *testing.T) {
|
||||
locked := *svc
|
||||
locked.reg = &Registry{byID: map[string]*CompiledController{
|
||||
"acme.demo.widgets": {
|
||||
Controller: queryController{perms: []string{"acme.demo.access"}},
|
||||
List: svc.reg.byID["acme.demo.widgets"].List,
|
||||
},
|
||||
}}
|
||||
rec := httptest.NewRecorder()
|
||||
req := controllerRequest(&bouncer.Principal{ID: 4})
|
||||
req.URL.RawQuery = "sort=name%3Bdrop"
|
||||
locked.list(rec, req)
|
||||
if rec.Code != http.StatusForbidden {
|
||||
t.Fatalf("status=%d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
type queryArtist struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
}
|
||||
|
||||
func (queryArtist) TableName() string { return "cabana_list_query_artists" }
|
||||
|
||||
type queryRow struct {
|
||||
ID uint `gorm:"column:id;primaryKey"`
|
||||
Name string `gorm:"column:name"`
|
||||
Active bool `gorm:"column:active"`
|
||||
CreatedAt time.Time `gorm:"column:created_at"`
|
||||
GroupID uint `gorm:"column:group_id"`
|
||||
ArtistID *uint `gorm:"column:artist_id"`
|
||||
Artist queryArtist `gorm:"foreignKey:ArtistID"`
|
||||
}
|
||||
|
||||
func (queryRow) TableName() string { return "cabana_list_query_rows" }
|
||||
|
||||
func (queryRow) FilterScopes() []string { return []string{"filterByGroup"} }
|
||||
|
||||
func (queryRow) FilterScope(name string, db *gorm.DB, value any) *gorm.DB {
|
||||
if db == nil || name != "filterByGroup" {
|
||||
return db
|
||||
}
|
||||
return db.Where("group_id = ?", value)
|
||||
}
|
||||
|
||||
type queryController struct{ perms []string }
|
||||
|
||||
func (queryController) ID() string { return "acme.demo.widgets" }
|
||||
func (queryController) ModelName() string { return "Widget" }
|
||||
func (queryController) ConfigDir() string { return "controllers/widgets" }
|
||||
func (queryController) NewRecord() any { return &queryRow{} }
|
||||
func (c queryController) RequiredPermissions() []string {
|
||||
return c.perms
|
||||
}
|
||||
|
||||
var (
|
||||
listPGOnce sync.Once
|
||||
listPGDSN string
|
||||
listPGErr error
|
||||
)
|
||||
|
||||
func newListService(t *testing.T) (*service, *gorm.DB) {
|
||||
t.Helper()
|
||||
listPGOnce.Do(func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
defer cancel()
|
||||
ctr, err := tcpostgres.Run(ctx, "postgres:16-alpine",
|
||||
tcpostgres.WithDatabase("cabana"),
|
||||
tcpostgres.WithUsername("cabana"),
|
||||
tcpostgres.WithPassword("cabana"),
|
||||
tcpostgres.BasicWaitStrategies(),
|
||||
)
|
||||
if err != nil {
|
||||
listPGErr = err
|
||||
return
|
||||
}
|
||||
dsn, err := ctr.ConnectionString(ctx, "sslmode=disable")
|
||||
if err != nil {
|
||||
listPGErr = err
|
||||
return
|
||||
}
|
||||
listPGDSN = dsn
|
||||
})
|
||||
if listPGErr != nil {
|
||||
t.Fatalf("postgres: %v", listPGErr)
|
||||
}
|
||||
db, err := gorm.Open(pgdriver.Open(listPGDSN), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Migrator().DropTable(&queryRow{}, &queryArtist{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.AutoMigrate(&queryArtist{}, &queryRow{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
schema, err := CompileList("acme.demo", queryController{}, filterFS(queryListConfig, queryColumns, queryFilters))
|
||||
if err != nil {
|
||||
t.Fatalf("schema: %v", err)
|
||||
}
|
||||
app := backpack.New(nil)
|
||||
if err := app.Publish(db); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := &service{app: app, reg: &Registry{byID: map[string]*CompiledController{
|
||||
"acme.demo.widgets": {Controller: queryController{}, List: schema},
|
||||
}}}
|
||||
return svc, db
|
||||
}
|
||||
|
||||
func seedListRows(t *testing.T, db *gorm.DB) {
|
||||
t.Helper()
|
||||
resetListRows(t, db)
|
||||
if err := db.Create(&queryArtist{ID: 8, Name: "Beatles"}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
artistID := uint(8)
|
||||
rows := []queryRow{
|
||||
{ID: 1, Name: "Same", Active: true, CreatedAt: time.Date(2024, 1, 2, 12, 0, 0, 0, time.UTC), GroupID: 4, ArtistID: &artistID},
|
||||
{ID: 2, Name: "Same", Active: false, CreatedAt: time.Date(2024, 1, 31, 12, 0, 0, 0, time.UTC), GroupID: 4},
|
||||
{ID: 3, Name: "Same", Active: true, CreatedAt: time.Date(2024, 2, 1, 12, 0, 0, 0, time.UTC), GroupID: 9},
|
||||
{ID: 4, Name: "Other", Active: true, CreatedAt: time.Date(2023, 12, 31, 12, 0, 0, 0, time.UTC), GroupID: 4},
|
||||
}
|
||||
if err := db.Create(&rows).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func resetListRows(t *testing.T, db *gorm.DB) {
|
||||
t.Helper()
|
||||
if err := db.Exec("DELETE FROM cabana_list_query_rows").Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := db.Exec("DELETE FROM cabana_list_query_artists").Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func countListRows(t *testing.T, db *gorm.DB) int64 {
|
||||
t.Helper()
|
||||
var n int64
|
||||
if err := db.Model(&queryRow{}).Count(&n).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func getList(t *testing.T, svc *service, rawQuery string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := controllerRequest(&bouncer.Principal{ID: 1, IsSuperuser: true})
|
||||
req.URL.RawQuery = rawQuery
|
||||
rec := httptest.NewRecorder()
|
||||
svc.list(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
type listEnvelope struct {
|
||||
Data []map[string]any `json:"data"`
|
||||
Meta struct {
|
||||
Page int `json:"page"`
|
||||
PerPage int `json:"per_page"`
|
||||
Total int64 `json:"total"`
|
||||
LastPage int `json:"last_page"`
|
||||
} `json:"meta"`
|
||||
Error struct {
|
||||
Code string `json:"code"`
|
||||
Details map[string]any `json:"details"`
|
||||
} `json:"error"`
|
||||
}
|
||||
|
||||
func decodeList(t *testing.T, rec *httptest.ResponseRecorder, status int) listEnvelope {
|
||||
t.Helper()
|
||||
if rec.Code != status {
|
||||
t.Fatalf("status=%d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var body listEnvelope
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil {
|
||||
t.Fatalf("json: %v body=%s", err, rec.Body.String())
|
||||
}
|
||||
return body
|
||||
}
|
||||
|
||||
func rowID(row map[string]any) int {
|
||||
switch n := row["id"].(type) {
|
||||
case float64:
|
||||
return int(n)
|
||||
case int:
|
||||
return n
|
||||
case uint:
|
||||
return int(n)
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func rowIDs(rows []map[string]any) []int {
|
||||
out := make([]int, len(rows))
|
||||
for i, row := range rows {
|
||||
out[i] = rowID(row)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func sameIDs(got, want []int) bool {
|
||||
if len(got) != len(want) {
|
||||
return false
|
||||
}
|
||||
seen := map[int]int{}
|
||||
for _, id := range got {
|
||||
seen[id]++
|
||||
}
|
||||
for _, id := range want {
|
||||
seen[id]--
|
||||
if seen[id] < 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
for _, n := range seen {
|
||||
if n != 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
var _ pact.AdminRecordSource = queryController{}
|
||||
var _ pact.AdminPermissioned = queryController{}
|
||||
var _ pact.FilterScope = queryRow{}
|
||||
Reference in New Issue
Block a user