From 6a57f63e812858707f017841e5170b5c26279d63 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Thu, 24 Sep 2026 19:03:44 +0200 Subject: [PATCH] feat(09-04): execute allowlisted deterministic list queries - 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 --- cabana/contracts.go | 10 +- cabana/http.go | 71 ++---- cabana/query.go | 511 ++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 537 insertions(+), 55 deletions(-) create mode 100644 cabana/query.go diff --git a/cabana/contracts.go b/cabana/contracts.go index 785e970..edd6c9c 100644 --- a/cabana/contracts.go +++ b/cabana/contracts.go @@ -129,11 +129,19 @@ func WriteData(w http.ResponseWriter, status int, data, meta any) { // WriteError writes a D-10 error envelope. details is always an object. func WriteError(w http.ResponseWriter, status int, code, message string) { + WriteErrorDetails(w, status, code, message, nil) +} + +// WriteErrorDetails writes a D-10 error envelope with field messages. +func WriteErrorDetails(w http.ResponseWriter, status int, code, message string, details map[string]any) { + if details == nil { + details = map[string]any{} + } writeJSON(w, status, map[string]any{ "error": map[string]any{ "code": code, "message": message, - "details": map[string]any{}, + "details": details, }, }) } diff --git a/cabana/http.go b/cabana/http.go index 370424a..8a08742 100644 --- a/cabana/http.go +++ b/cabana/http.go @@ -172,28 +172,21 @@ func (s *service) list(w http.ResponseWriter, r *http.Request) { WriteError(w, http.StatusInternalServerError, "error", msgServerError) return } - rows, total, err := queryList(r.Context(), db, cc) + result, err := ExecuteList(r.Context(), db, cc, listQueryFromRequest(r)) + var invalid *ListValidationError + if errors.As(err, &invalid) { + WriteErrorDetails(w, http.StatusUnprocessableEntity, "validation_failed", "Validation failed", invalid.Details) + return + } if err != nil { WriteError(w, http.StatusInternalServerError, "error", msgServerError) return } - data := make([]map[string]any, 0, len(rows)) - for _, row := range rows { - data = append(data, projectRow(row, cc.List.Columns)) - } - per := cc.List.RecordsPerPage - if per < 1 { - per = 20 - } - last := int((total + int64(per) - 1) / int64(per)) - if last < 1 { - last = 1 - } - WriteData(w, http.StatusOK, data, map[string]any{ - "page": 1, - "per_page": per, - "total": total, - "last_page": last, + WriteData(w, http.StatusOK, result.Data, map[string]any{ + "page": result.Meta.Page, + "per_page": result.Meta.PerPage, + "total": result.Meta.Total, + "last_page": result.Meta.LastPage, }) }) } @@ -220,42 +213,6 @@ func (s *service) protect(w http.ResponseWriter, r *http.Request, fn func(*Compi fn(cc) } -func queryList(ctx context.Context, db *gorm.DB, cc *CompiledController) ([]any, int64, error) { - src, ok := cc.Controller.(pact.AdminRecordSource) - if !ok || src == nil { - return nil, 0, errors.New("cabana: admin controller has no record source") - } - model := src.NewRecord() - mt := reflect.TypeOf(model) - if mt == nil || mt.Kind() != reflect.Pointer { - return nil, 0, errors.New("cabana: admin model must be a pointer") - } - 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 - } - } - var total int64 - if err := q.Session(&gorm.Session{}).Count(&total).Error; err != nil { - return nil, 0, err - } - per := 20 - if cc.List != nil && cc.List.RecordsPerPage > 0 { - per = cc.List.RecordsPerPage - } - slice := reflect.New(reflect.SliceOf(mt.Elem())) - if err := q.Session(&gorm.Session{}).Limit(per).Find(slice.Interface()).Error; err != nil { - return nil, 0, err - } - values := slice.Elem() - out := make([]any, values.Len()) - for i := 0; i < values.Len(); i++ { - out[i] = values.Index(i).Addr().Interface() - } - return out, total, nil -} - func projectRow(row any, cols []ListColumn) map[string]any { v := reflect.ValueOf(row) for v.Kind() == reflect.Pointer { @@ -269,6 +226,12 @@ func projectRow(row any, cols []ListColumn) map[string]any { out["id"] = id.Interface() } for _, col := range cols { + if col.Relation != "" { + if value, ok := relatedSelect(v, col); ok { + out[col.Key] = value + } + continue + } field := fieldByColumn(v, col.Key) if !field.IsValid() || !field.CanInterface() { continue diff --git a/cabana/query.go b/cabana/query.go new file mode 100644 index 0000000..1852e02 --- /dev/null +++ b/cabana/query.go @@ -0,0 +1,511 @@ +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 +}