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{}