480 lines
14 KiB
Go
480 lines
14 KiB
Go
package cabana
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.golem15.com/golem15/summercms/modules/backpack"
|
|
"git.golem15.com/golem15/summercms/modules/bouncer"
|
|
"git.golem15.com/golem15/summercms/modules/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"} }
|
|
|
|
// FilterOptions serves the grouped scope's choices (D-27).
|
|
func (queryRow) FilterOptions(scope string) []pact.Option {
|
|
if scope != "filterByGroup" {
|
|
return nil
|
|
}
|
|
return []pact.Option{{Value: "4", Label: "demo::lang.group_four"}, {Value: "9", Label: "Group nine"}}
|
|
}
|
|
|
|
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{}
|
|
var _ pact.FilterOptions = queryRow{}
|
|
|
|
// TestListSearchNonTextColumns pins WR-08: a searchable numeric or date column
|
|
// is searched as text instead of failing with an SQL error (PostgreSQL has no
|
|
// lower(integer) or lower(timestamp)).
|
|
func TestListSearchNonTextColumns(t *testing.T) {
|
|
svc, db := newListService(t)
|
|
seedListRows(t, db)
|
|
const columns = `columns:
|
|
name:
|
|
label: Name
|
|
searchable: true
|
|
group_id:
|
|
label: Group
|
|
searchable: true
|
|
created_at:
|
|
label: Created
|
|
type: datetime
|
|
searchable: true
|
|
active:
|
|
label: Active
|
|
type: switch
|
|
searchable: true
|
|
`
|
|
schema, err := CompileList("acme.demo", queryController{}, filterFS(queryListConfig, columns, queryFilters))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
svc.reg.byID["acme.demo.widgets"].List = schema
|
|
|
|
byGroup := decodeList(t, getList(t, svc, "search=9"), http.StatusOK)
|
|
if ids := rowIDs(byGroup.Data); len(ids) != 1 || ids[0] != 3 {
|
|
t.Fatalf("integer column search = %v, want [3]", ids)
|
|
}
|
|
byDate := decodeList(t, getList(t, svc, "search=2024-01"), http.StatusOK)
|
|
if byDate.Meta.Total != 2 {
|
|
t.Fatalf("timestamp column search total = %d, want 2", byDate.Meta.Total)
|
|
}
|
|
byName := decodeList(t, getList(t, svc, "search=other"), http.StatusOK)
|
|
if ids := rowIDs(byName.Data); len(ids) != 1 || ids[0] != 4 {
|
|
t.Fatalf("text column search = %v, want [4]", ids)
|
|
}
|
|
}
|