From 7f451c35723657dfcf6886fc9f6da12bece0f366 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Thu, 24 Sep 2026 18:58:56 +0200 Subject: [PATCH] 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 --- cabana/query_test.go | 428 +++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 428 insertions(+) create mode 100644 cabana/query_test.go diff --git a/cabana/query_test.go b/cabana/query_test.go new file mode 100644 index 0000000..b5484e7 --- /dev/null +++ b/cabana/query_test.go @@ -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{}