Skip to content

Commit a17f20f

Browse files
committed
Sanitize sort column to prevent SQL injection vulnerabilities
1 parent 3c6bc48 commit a17f20f

2 files changed

Lines changed: 10 additions & 16 deletions

File tree

page.go

Lines changed: 9 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ import (
66
"strings"
77

88
sq "github.com/Masterminds/squirrel"
9+
"github.com/jackc/pgx/v5"
910
)
1011

1112
const (
@@ -37,20 +38,16 @@ func (s Sort) String() string {
3738
return fmt.Sprintf("%s %s", s.Column, s.Order)
3839
}
3940

40-
func (s Sort) IsValid() bool {
41-
return s.Column != "" && _MatcherOrderBy.MatchString(s.Column)
42-
}
43-
44-
var _MatcherOrderBy = regexp.MustCompile(`^-?([a-zA-Z_][a-zA-Z0-9_]*)$`)
41+
var _MatcherOrderBy = regexp.MustCompile(`-?([a-zA-Z0-9]+)`)
4542

4643
func NewSort(s string) (Sort, bool) {
44+
if s == "" || !_MatcherOrderBy.MatchString(s) {
45+
return Sort{}, false
46+
}
4747
sort := Sort{
4848
Column: s,
4949
Order: Asc,
5050
}
51-
if !sort.IsValid() {
52-
return Sort{}, false
53-
}
5451
if strings.HasPrefix(s, "-") {
5552
sort.Column = s[1:]
5653
sort.Order = Desc
@@ -83,13 +80,11 @@ func NewPage(size, page uint32, sort ...Sort) *Page {
8380
func (p *Page) GetOrder(defaultSort ...string) []Sort {
8481
// if page has sort, use it
8582
if p != nil && len(p.Sort) != 0 {
86-
sort := make([]Sort, 0, len(p.Sort))
87-
for _, s := range p.Sort {
88-
if s.IsValid() {
89-
sort = append(sort, s)
90-
}
83+
for i, s := range p.Sort {
84+
s.Column = pgx.Identifier(strings.Split(s.Column, ".")).Sanitize()
85+
p.Sort[i] = s
9186
}
92-
return sort
87+
return p.Sort
9388
}
9489
// if page has column, use default sort
9590
if p == nil || p.Column == "" {

page_test.go

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,6 @@ func TestInvalidSort(t *testing.T) {
5858

5959
sql, args, err := query.ToSql()
6060
require.NoError(t, err)
61-
require.Equal(t, "SELECT * FROM t ORDER BY name DESC LIMIT 11 OFFSET 0", sql)
61+
require.Equal(t, "SELECT * FROM t ORDER BY \"ID; DROP TABLE users;\" ASC, \"name\" DESC LIMIT 11 OFFSET 0", sql)
6262
require.Empty(t, args)
63-
6463
}

0 commit comments

Comments
 (0)