- Search, sort, filters, and pagination use compiled selectors and bound values - Equal sort keys break ties on the primary key so adjacent pages do not overlap - Unknown identifiers return validation_failed before SQL
512 lines
13 KiB
Go
512 lines
13 KiB
Go
package cabana
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"reflect"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"git.golem15.com/golem15/summercms/lagoon"
|
|
"git.golem15.com/golem15/summercms/pact"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/clause"
|
|
)
|
|
|
|
// ListInput is the raw admin list query. Identifiers are matched later
|
|
// against the compiled schema; nothing here is SQL.
|
|
type ListInput struct {
|
|
Search string
|
|
Sort string
|
|
Dir string
|
|
Page string
|
|
PerPage string
|
|
Filters map[string]string
|
|
}
|
|
|
|
// ListMeta is the D-11 pagination block.
|
|
type ListMeta struct {
|
|
Page int `json:"page"`
|
|
PerPage int `json:"per_page"`
|
|
Total int64 `json:"total"`
|
|
LastPage int `json:"last_page"`
|
|
}
|
|
|
|
// ListResult is one page of projected rows.
|
|
type ListResult struct {
|
|
Data []map[string]any
|
|
Meta ListMeta
|
|
}
|
|
|
|
// ListValidationError is a D-10 validation_failed failure.
|
|
type ListValidationError struct {
|
|
Details map[string]any
|
|
}
|
|
|
|
func (e *ListValidationError) Error() string { return "validation_failed" }
|
|
|
|
func listInvalid(field, message string) *ListValidationError {
|
|
return &ListValidationError{Details: map[string]any{field: []string{message}}}
|
|
}
|
|
|
|
func listQueryFromRequest(r *http.Request) ListInput {
|
|
q := r.URL.Query()
|
|
filters := map[string]string{}
|
|
for key, vals := range q {
|
|
if !strings.HasPrefix(key, "filter[") || !strings.HasSuffix(key, "]") || len(key) <= len("filter[]") {
|
|
continue
|
|
}
|
|
name := key[len("filter[") : len(key)-1]
|
|
if len(vals) == 0 {
|
|
filters[name] = ""
|
|
continue
|
|
}
|
|
filters[name] = vals[len(vals)-1]
|
|
}
|
|
return ListInput{
|
|
Search: q.Get(listSearchTerm),
|
|
Sort: q.Get("sort"),
|
|
Dir: q.Get("dir"),
|
|
Page: q.Get("page"),
|
|
PerPage: q.Get("per_page"),
|
|
Filters: filters,
|
|
}
|
|
}
|
|
|
|
// ExecuteList runs an allowlisted search, sort, filter, and page against db.
|
|
func ExecuteList(ctx context.Context, db *gorm.DB, cc *CompiledController, in ListInput) (*ListResult, error) {
|
|
if cc == nil || cc.List == nil {
|
|
return nil, errors.New("cabana: list schema is missing")
|
|
}
|
|
if db == nil {
|
|
return nil, errors.New("cabana: database is not configured")
|
|
}
|
|
src, ok := cc.Controller.(pact.AdminRecordSource)
|
|
if !ok || src == nil || src.NewRecord() == nil {
|
|
return nil, errors.New("cabana: admin controller has no record source")
|
|
}
|
|
model := src.NewRecord()
|
|
mt := reflect.TypeOf(model)
|
|
if mt == nil || mt.Kind() != reflect.Pointer || mt.Elem().Kind() != reflect.Struct {
|
|
return nil, errors.New("cabana: admin model must be a pointer")
|
|
}
|
|
if ctx == nil {
|
|
ctx = context.Background()
|
|
}
|
|
page, per, err := normalizePage(cc.List, in)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sortCol, sortDesc, err := normalizeSort(cc.List, in)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
q := db.WithContext(ctx).Model(model)
|
|
if ext, ok := cc.Controller.(pact.ListExtendQuery); ok && ext != nil {
|
|
if next := ext.ListExtendQuery(ctx, q); next != nil {
|
|
q = next
|
|
}
|
|
}
|
|
q, err = applyListSearch(q, cc.List, model, in.Search)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
q, err = applyListFilters(q, cc.List, model, in.Filters)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var total int64
|
|
if err := q.Session(&gorm.Session{}).Count(&total).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
q = applyListOrder(q, model, sortCol, sortDesc)
|
|
for _, col := range cc.List.Columns {
|
|
if col.Relation != "" {
|
|
q = q.Preload(col.Relation)
|
|
}
|
|
}
|
|
offset := (page - 1) * per
|
|
slice := reflect.New(reflect.SliceOf(mt.Elem()))
|
|
if err := q.Offset(offset).Limit(per).Find(slice.Interface()).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
values := slice.Elem()
|
|
data := make([]map[string]any, 0, values.Len())
|
|
for i := 0; i < values.Len(); i++ {
|
|
data = append(data, projectRow(values.Index(i).Addr().Interface(), cc.List.Columns))
|
|
}
|
|
paged := lagoon.Paginate(data, page, per, total)
|
|
return &ListResult{
|
|
Data: paged.Data,
|
|
Meta: ListMeta{
|
|
Page: paged.Meta.CurrentPage,
|
|
PerPage: paged.Meta.PerPage,
|
|
Total: paged.Meta.Total,
|
|
LastPage: paged.Meta.LastPage,
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
func normalizePage(schema *ListSchema, in ListInput) (int, int, error) {
|
|
page := 1
|
|
if in.Page != "" {
|
|
n, err := strconv.Atoi(in.Page)
|
|
if err != nil || n < 1 {
|
|
return 0, 0, listInvalid("page", "must be a positive integer")
|
|
}
|
|
page = n
|
|
}
|
|
per := schema.RecordsPerPage
|
|
if per < 1 {
|
|
per = 20
|
|
}
|
|
if in.PerPage != "" {
|
|
n, err := strconv.Atoi(in.PerPage)
|
|
if err != nil || !pageSizeAllowed(schema, n) {
|
|
return 0, 0, listInvalid("per_page", "is not an allowed page size")
|
|
}
|
|
per = n
|
|
}
|
|
return page, per, nil
|
|
}
|
|
|
|
func pageSizeAllowed(schema *ListSchema, n int) bool {
|
|
if n < 1 {
|
|
return false
|
|
}
|
|
options := schema.PerPageOptions
|
|
if len(options) == 0 {
|
|
return n == schema.RecordsPerPage || (schema.RecordsPerPage < 1 && n == 20)
|
|
}
|
|
for _, option := range options {
|
|
if option == n {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func normalizeSort(schema *ListSchema, in ListInput) (string, bool, error) {
|
|
if in.Sort == "" {
|
|
if schema.DefaultSort == nil {
|
|
return "", false, nil
|
|
}
|
|
return schema.DefaultSort.Column, schema.DefaultSort.Direction == "desc", nil
|
|
}
|
|
if !schema.ShowSorting {
|
|
return "", false, listInvalid("sort", "is not sortable")
|
|
}
|
|
var match *ListColumn
|
|
for i := range schema.Columns {
|
|
if schema.Columns[i].Key == in.Sort {
|
|
match = &schema.Columns[i]
|
|
break
|
|
}
|
|
}
|
|
if match == nil || !match.Sortable || match.Relation != "" {
|
|
return "", false, listInvalid("sort", "is not a sortable column")
|
|
}
|
|
dir := in.Dir
|
|
if dir == "" {
|
|
dir = "asc"
|
|
}
|
|
if dir != "asc" && dir != "desc" {
|
|
return "", false, listInvalid("dir", "must be asc or desc")
|
|
}
|
|
return match.Key, dir == "desc", nil
|
|
}
|
|
|
|
func applyListOrder(db *gorm.DB, model any, column string, desc bool) *gorm.DB {
|
|
table := tableName(model)
|
|
pk := primaryColumn(model)
|
|
if column != "" {
|
|
db = db.Order(clause.OrderByColumn{
|
|
Column: clause.Column{Table: table, Name: column},
|
|
Desc: desc,
|
|
})
|
|
}
|
|
if column != pk {
|
|
db = db.Order(clause.OrderByColumn{Column: clause.Column{Table: table, Name: pk}})
|
|
}
|
|
return db
|
|
}
|
|
|
|
func applyListSearch(db *gorm.DB, schema *ListSchema, model any, term string) (*gorm.DB, error) {
|
|
term = strings.TrimSpace(term)
|
|
if term == "" {
|
|
return db, nil
|
|
}
|
|
var cols []ListColumn
|
|
for _, col := range schema.Columns {
|
|
if col.Searchable {
|
|
cols = append(cols, col)
|
|
}
|
|
}
|
|
if len(cols) == 0 {
|
|
return nil, listInvalid("search", "is not searchable")
|
|
}
|
|
pattern := "%" + escapeLike(strings.ToLower(term)) + "%"
|
|
joined := map[string]struct{}{}
|
|
parts := make([]string, 0, len(cols))
|
|
args := make([]any, 0, len(cols))
|
|
main := tableName(model)
|
|
for _, col := range cols {
|
|
table := main
|
|
column := col.Key
|
|
if col.Relation != "" {
|
|
if _, ok := relatedTable(model, col.Relation); !ok || !identifier(col.Select) {
|
|
return nil, listInvalid("search", "is not searchable")
|
|
}
|
|
// GORM aliases the joined association with the relation field name.
|
|
table = col.Relation
|
|
column = col.Select
|
|
if _, done := joined[col.Relation]; !done {
|
|
db = db.Joins(col.Relation)
|
|
joined[col.Relation] = struct{}{}
|
|
}
|
|
} else if !identifier(column) {
|
|
return nil, listInvalid("search", "is not searchable")
|
|
}
|
|
parts = append(parts, fmt.Sprintf("LOWER(%s) LIKE ? ESCAPE '\\'", qualifiedColumn(db, table, column)))
|
|
args = append(args, pattern)
|
|
}
|
|
return db.Where(strings.Join(parts, " OR "), args...), nil
|
|
}
|
|
|
|
func applyListFilters(db *gorm.DB, schema *ListSchema, model any, filters map[string]string) (*gorm.DB, error) {
|
|
if len(filters) == 0 {
|
|
return db, nil
|
|
}
|
|
byName := map[string]ListFilter{}
|
|
for _, filter := range schema.Filters {
|
|
byName[filter.Name] = filter
|
|
}
|
|
table := tableName(model)
|
|
var provider pact.FilterScope
|
|
if scope, ok := model.(pact.FilterScope); ok {
|
|
provider = scope
|
|
}
|
|
for name, raw := range filters {
|
|
filter, ok := byName[name]
|
|
if !ok || !identifier(name) {
|
|
return nil, listInvalid("filter", "is not a declared filter")
|
|
}
|
|
switch filter.Type {
|
|
case "switch":
|
|
value, ok := switchArgument(filter, raw)
|
|
if !ok {
|
|
return nil, listInvalid("filter", "is not a declared value")
|
|
}
|
|
db = db.Where(clause.Eq{Column: clause.Column{Table: table, Name: filter.Column}, Value: value})
|
|
case "daterange":
|
|
start, end, err := parseDateRange(raw)
|
|
if err != nil {
|
|
return nil, listInvalid("filter", "is not a valid date range")
|
|
}
|
|
db = db.Where(clause.Gte{Column: clause.Column{Table: table, Name: filter.Column}, Value: start})
|
|
db = db.Where(clause.Lt{Column: clause.Column{Table: table, Name: filter.Column}, Value: end})
|
|
case "scope":
|
|
arg, err := scopeArgument(raw)
|
|
if err != nil {
|
|
return nil, listInvalid("filter", "is not a declared value")
|
|
}
|
|
if provider == nil {
|
|
return nil, errors.New("cabana: filter scope provider is missing")
|
|
}
|
|
next := provider.FilterScope(filter.Scope, db, arg)
|
|
if next == nil {
|
|
return nil, errors.New("cabana: filter scope returned nil")
|
|
}
|
|
db = next
|
|
default:
|
|
return nil, listInvalid("filter", "is not a declared filter")
|
|
}
|
|
}
|
|
return db, nil
|
|
}
|
|
|
|
func switchArgument(filter ListFilter, raw string) (any, bool) {
|
|
if len(filter.Options) > 0 {
|
|
for _, opt := range filter.Options {
|
|
if string(opt.Value.raw) == raw {
|
|
return scalarValue(opt.Value), true
|
|
}
|
|
}
|
|
return nil, false
|
|
}
|
|
if filter.TrueValue != nil && string(filter.TrueValue.raw) == raw {
|
|
return scalarValue(*filter.TrueValue), true
|
|
}
|
|
if filter.FalseValue != nil && string(filter.FalseValue.raw) == raw {
|
|
return scalarValue(*filter.FalseValue), true
|
|
}
|
|
return nil, false
|
|
}
|
|
|
|
func scalarValue(s jsonScalar) any {
|
|
var v any
|
|
if err := json.Unmarshal(s.raw, &v); err != nil {
|
|
return string(s.raw)
|
|
}
|
|
return v
|
|
}
|
|
|
|
func parseDateRange(raw string) (time.Time, time.Time, error) {
|
|
startText, endText, ok := strings.Cut(raw, "..")
|
|
if !ok || startText == "" || endText == "" {
|
|
return time.Time{}, time.Time{}, errors.New("range")
|
|
}
|
|
start, err := time.ParseInLocation("2006-01-02", startText, time.UTC)
|
|
if err != nil {
|
|
return time.Time{}, time.Time{}, err
|
|
}
|
|
end, err := time.ParseInLocation("2006-01-02", endText, time.UTC)
|
|
if err != nil {
|
|
return time.Time{}, time.Time{}, err
|
|
}
|
|
if end.Before(start) {
|
|
return time.Time{}, time.Time{}, errors.New("range order")
|
|
}
|
|
return start, end.AddDate(0, 0, 1), nil
|
|
}
|
|
|
|
func scopeArgument(raw string) (any, error) {
|
|
if raw == "" || strings.ContainsAny(raw, " \t;'\"\\") || strings.Contains(raw, "--") {
|
|
return nil, errors.New("value")
|
|
}
|
|
if n, err := strconv.ParseInt(raw, 10, 64); err == nil {
|
|
return n, nil
|
|
}
|
|
if !identifier(raw) {
|
|
return nil, errors.New("value")
|
|
}
|
|
return raw, nil
|
|
}
|
|
|
|
func escapeLike(s string) string {
|
|
s = strings.ReplaceAll(s, `\`, `\\`)
|
|
s = strings.ReplaceAll(s, `%`, `\%`)
|
|
s = strings.ReplaceAll(s, `_`, `\_`)
|
|
return s
|
|
}
|
|
|
|
func qualifiedColumn(db *gorm.DB, table, column string) string {
|
|
return quotedIdent(db, table) + "." + quotedIdent(db, column)
|
|
}
|
|
|
|
func quotedIdent(db *gorm.DB, name string) string {
|
|
if db == nil || db.Dialector == nil || !identifier(name) && !strings.Contains(name, "_") {
|
|
return `"` + strings.ReplaceAll(name, `"`, `""`) + `"`
|
|
}
|
|
var b strings.Builder
|
|
db.Dialector.QuoteTo(&b, name)
|
|
return b.String()
|
|
}
|
|
|
|
func tableName(model any) string {
|
|
if model == nil {
|
|
return ""
|
|
}
|
|
if namer, ok := model.(interface{ TableName() string }); ok {
|
|
return namer.TableName()
|
|
}
|
|
t := reflect.TypeOf(model)
|
|
for t != nil && t.Kind() == reflect.Pointer {
|
|
t = t.Elem()
|
|
}
|
|
if t == nil {
|
|
return ""
|
|
}
|
|
ptr := reflect.New(t)
|
|
if namer, ok := ptr.Interface().(interface{ TableName() string }); ok {
|
|
return namer.TableName()
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func relatedTable(model any, relation string) (string, bool) {
|
|
t := reflect.TypeOf(model)
|
|
for t != nil && t.Kind() == reflect.Pointer {
|
|
t = t.Elem()
|
|
}
|
|
if t == nil || t.Kind() != reflect.Struct {
|
|
return "", false
|
|
}
|
|
field, ok := t.FieldByName(relation)
|
|
if !ok {
|
|
return "", false
|
|
}
|
|
rt := field.Type
|
|
for rt.Kind() == reflect.Pointer {
|
|
rt = rt.Elem()
|
|
}
|
|
if rt.Kind() == reflect.Slice {
|
|
rt = rt.Elem()
|
|
for rt.Kind() == reflect.Pointer {
|
|
rt = rt.Elem()
|
|
}
|
|
}
|
|
if rt.Kind() != reflect.Struct {
|
|
return "", false
|
|
}
|
|
name := tableName(reflect.New(rt).Interface())
|
|
return name, name != ""
|
|
}
|
|
|
|
func primaryColumn(model any) string {
|
|
t := reflect.TypeOf(model)
|
|
for t != nil && t.Kind() == reflect.Pointer {
|
|
t = t.Elem()
|
|
}
|
|
if t == nil || t.Kind() != reflect.Struct {
|
|
return "id"
|
|
}
|
|
for i := 0; i < t.NumField(); i++ {
|
|
field := t.Field(i)
|
|
if field.PkgPath != "" {
|
|
continue
|
|
}
|
|
if strings.Contains(field.Tag.Get("gorm"), "primaryKey") {
|
|
if name := gormColumn(field); name != "" {
|
|
return name
|
|
}
|
|
return field.Name
|
|
}
|
|
}
|
|
return "id"
|
|
}
|
|
|
|
func relatedSelect(row reflect.Value, col ListColumn) (any, bool) {
|
|
for row.Kind() == reflect.Pointer {
|
|
if row.IsNil() {
|
|
return nil, false
|
|
}
|
|
row = row.Elem()
|
|
}
|
|
if row.Kind() != reflect.Struct || col.Relation == "" {
|
|
return nil, false
|
|
}
|
|
field := row.FieldByName(col.Relation)
|
|
if !field.IsValid() {
|
|
return nil, false
|
|
}
|
|
for field.Kind() == reflect.Pointer {
|
|
if field.IsNil() {
|
|
return nil, false
|
|
}
|
|
field = field.Elem()
|
|
}
|
|
if field.Kind() != reflect.Struct {
|
|
return nil, false
|
|
}
|
|
selected := fieldByColumn(field, col.Select)
|
|
if !selected.IsValid() || !selected.CanInterface() {
|
|
return nil, false
|
|
}
|
|
return selected.Interface(), true
|
|
}
|