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, cc.List, model, sortCol, sortDesc) for _, col := range cc.List.Columns { if col.Relation == "" { continue } if field, ok := relationFieldName(model, col.Relation); ok { q = q.Preload(field) } } 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 { 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, schema *ListSchema, model any, column string, desc bool) *gorm.DB { table := tableName(model) pk := primaryColumn(model) if column != "" { orderTable := table orderColumn := column for _, col := range schema.Columns { if col.Key != column || col.Relation == "" { continue } field, ok := relationFieldName(model, col.Relation) if !ok { continue } orderTable = field orderColumn = col.Select db = db.Joins(field) break } db = db.Order(clause.OrderByColumn{ Column: clause.Column{Table: orderTable, Name: orderColumn}, 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 != "" { field, ok := relationFieldName(model, col.Relation) if !ok || !identifier(col.Select) { return nil, listInvalid("search", "is not searchable") } // GORM aliases the joined association with the Go field name. table = field column = col.Select if _, done := joined[field]; !done { db = db.Joins(field) joined[field] = 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 "" } // relationFieldName resolves a Winter relation key onto the exported Go field. // Exact names win; otherwise the match is case-insensitive, so relation: genre // selects field Genre. func relationFieldName(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 || relation == "" { return "", false } if field, ok := t.FieldByName(relation); ok && isListRelation(field.Type) { return field.Name, true } for i := 0; i < t.NumField(); i++ { field := t.Field(i) if field.PkgPath != "" || !isListRelation(field.Type) { continue } if strings.EqualFold(field.Name, relation) { return field.Name, true } } return "", false } 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 } name, ok := relationFieldName(row.Interface(), col.Relation) if !ok { return nil, false } field := row.FieldByName(name) 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 }