package lagoon import ( "fmt" "strings" "gorm.io/gorm" ) // maxCollationName is PostgreSQL's identifier limit (NAMEDATALEN-1). const maxCollationName = 63 // OrderOption changes how OrderBy builds its ORDER BY clause. type OrderOption func(*orderOptions) type orderOptions struct { collation string hasCollation bool } // Collate sorts the column with the named PostgreSQL collation, for example // the ICU collation pl-x-icu for Polish alphabetical order. ICU collations // exist in any ICU-enabled PostgreSQL server, so they need no database // setup. The name is validated and emitted as a double-quoted identifier; // an invalid name makes OrderBy return an error. func Collate(name string) OrderOption { return func(o *orderOptions) { o.collation = name o.hasCollation = true } } // OrderBy appends an ORDER BY for an allow-listed qualified column. // Identifiers and directions are never taken from untrusted input: column // must match allowed exactly, and dir must be asc or desc. No COLLATE is // emitted unless Collate is passed; without it, text sorts by the // database's default collation. func OrderBy(db *gorm.DB, column, dir string, allowed []string, opts ...OrderOption) (*gorm.DB, error) { if db == nil { return nil, fmt.Errorf("lagoon: gorm db is nil") } clause, err := orderClause(column, dir, allowed, resolveOrderOptions(opts)) if err != nil { return nil, err } return db.Order(clause), nil } func resolveOrderOptions(opts []OrderOption) orderOptions { var o orderOptions for _, opt := range opts { if opt != nil { opt(&o) } } return o } func orderClause(column, dir string, allowed []string, opts orderOptions) (string, error) { if !allowListed(column, allowed) { return "", fmt.Errorf("lagoon: order column %q is not allow-listed", column) } var direction string switch strings.ToLower(strings.TrimSpace(dir)) { case "asc": direction = "ASC" case "desc": direction = "DESC" default: return "", fmt.Errorf("lagoon: order direction %q is not allow-listed", dir) } if !opts.hasCollation { return column + " " + direction, nil } if !validCollationName(opts.collation) { return "", fmt.Errorf("lagoon: order collation %q is not a valid collation name", opts.collation) } return column + " COLLATE " + quoteCollation(opts.collation) + " " + direction, nil } // validCollationName accepts 1-63 ASCII bytes: letters, digits and // underscore first, then letters, digits, '_', '-', '.' and '@'. The set // excludes the double quote, so a valid name cannot leave its identifier. func validCollationName(name string) bool { if len(name) == 0 || len(name) > maxCollationName { return false } for i := 0; i < len(name); i++ { c := name[i] switch { case c >= 'a' && c <= 'z', c >= 'A' && c <= 'Z', c >= '0' && c <= '9', c == '_': case i > 0 && (c == '-' || c == '.' || c == '@'): default: return false } } return true } // quoteCollation wraps a validated collation name as a PostgreSQL quoted // identifier. func quoteCollation(name string) string { return `"` + strings.ReplaceAll(name, `"`, `""`) + `"` } func allowListed(column string, allowed []string) bool { for _, a := range allowed { if a == column { return true } } return false }