Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 31 additions & 8 deletions storage/sql.go
Original file line number Diff line number Diff line change
Expand Up @@ -285,7 +285,10 @@ func (s *SQLAdapter) GetContext(ctx context.Context, dest any, filter map[string
if len(filter) == 0 {
return errors.New("filtering is required when getting a resource")
}
query, bindings := s.buildQuery(filter)
query, bindings, err := s.buildQuery(filter)
if err != nil {
return err
}
result := s.dbWithCtx(ctx).Where(query, bindings).Find(dest)
if result.RowsAffected == 0 {
return ErrNotFound
Expand All @@ -301,7 +304,10 @@ func (s *SQLAdapter) UpdateContext(ctx context.Context, item any, filter map[str
if len(filter) == 0 {
return errors.New("filtering is required when updating a resource")
}
query, bindings := s.buildQuery(filter)
query, bindings, err := s.buildQuery(filter)
if err != nil {
return err
}
result := s.dbWithCtx(ctx).Where(query, bindings).Save(item)
return result.Error
}
Expand All @@ -314,7 +320,10 @@ func (s *SQLAdapter) DeleteContext(ctx context.Context, item any, filter map[str
if len(filter) == 0 {
return errors.New("filtering is required when deleting a resource")
}
query, bindings := s.buildQuery(filter)
query, bindings, err := s.buildQuery(filter)
if err != nil {
return err
}
result := s.dbWithCtx(ctx).Where(query, bindings).Delete(item)
return result.Error
}
Expand Down Expand Up @@ -411,9 +420,15 @@ func (s *SQLAdapter) ListContext(ctx context.Context, dest any, sortKey string,
if err != nil {
return "", fmt.Errorf("failed to list: %w", err)
}
var query string
var bindings map[string]any
if len(filter) > 0 {
if query, bindings, err = s.buildQuery(filter); err != nil {
return "", err
}
}
return s.executePaginatedQuery(ctx, dest, sortKey, sortDirection, limit, cursor, func(q *gorm.DB) *gorm.DB {
if len(filter) > 0 {
query, bindings := s.buildQuery(filter)
if query != "" {
return q.Where(query, bindings)
}
return q
Expand Down Expand Up @@ -473,7 +488,10 @@ func (s *SQLAdapter) CountContext(ctx context.Context, dest any, filter map[stri
q := s.dbWithCtx(ctx).Model(dest)

if len(filter) > 0 {
query, bindings := s.buildQuery(filter)
query, bindings, err := s.buildQuery(filter)
if err != nil {
return 0, err
}
q = q.Where(query, bindings)
}

Expand All @@ -494,11 +512,16 @@ func (s *SQLAdapter) QueryContext(ctx context.Context, dest any, statement strin
}


func (s *SQLAdapter) buildQuery(filter map[string]any) (string, map[string]any) {
// buildQuery rejects keys that aren't plain column names: keys are written
// into the SQL as-is, only values are bound.
func (s *SQLAdapter) buildQuery(filter map[string]any) (string, map[string]any, error) {
clauses := []string{}
bindings := make(map[string]any)

for key, value := range filter {
if !validColumnName.MatchString(key) {
return "", nil, fmt.Errorf("invalid filter key %q: must match [a-zA-Z_][a-zA-Z0-9_]*", key)
}
if value == nil {
// For nil values, use IS NULL instead of = @key
clauses = append(clauses, fmt.Sprintf("%s IS NULL", key))
Expand All @@ -508,5 +531,5 @@ func (s *SQLAdapter) buildQuery(filter map[string]any) (string, map[string]any)
bindings[key] = value
}
}
return strings.Join(clauses, " AND "), bindings
return strings.Join(clauses, " AND "), bindings, nil
}
43 changes: 43 additions & 0 deletions storage/sql_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -210,6 +210,49 @@ func TestSQLAdapterListRejectsInvalidSortKey(t *testing.T) {
}
}

func TestSQLAdapterRejectsUnsafeFilterKeys(t *testing.T) {
_, sql := setupSQLCoverage(t)
for _, id := range []string{"a", "b", "c"} {
if err := sql.Create(&sqlCoverageItem{Id: id, Name: "orig"}); err != nil {
t.Fatalf("Create %s: %v", id, err)
}
}
// Without validation this key closes the filter's parentheses and turns
// the WHERE into "(1=1) OR (...)", which matches every row.
filter := map[string]any{"1=1) OR (1=1": nil, "name": "orig"}

var got sqlCoverageItem
if err := sql.Get(&got, filter); err == nil {
t.Fatal("Get = nil error; want the filter key rejected")
}
var page []sqlCoverageItem
if _, err := sql.List(&page, "id", filter, 10, ""); err == nil {
t.Fatal("List = nil error; want the filter key rejected")
}
if _, err := sql.Count(&[]sqlCoverageItem{}, filter); err == nil {
t.Fatal("Count = nil error; want the filter key rejected")
}
if err := sql.Update(&sqlCoverageItem{Id: "a", Name: "changed"}, filter); err == nil {
t.Fatal("Update = nil error; want the filter key rejected")
}
if err := sql.Delete(&sqlCoverageItem{}, filter); err == nil {
t.Fatal("Delete = nil error; want the filter key rejected")
}

var rows []sqlCoverageItem
if _, err := sql.List(&rows, "id", nil, 10, ""); err != nil {
t.Fatalf("List: %v", err)
}
if len(rows) != 3 {
t.Fatalf("rows = %v; want all 3 left", rows)
}
for _, row := range rows {
if row.Name != "orig" {
t.Fatalf("rows = %v; want none changed", rows)
}
}
}

func TestSQLAdapterListRejectsInvalidSortDirection(t *testing.T) {
_, sql := setupSQLCoverage(t)
var page []sqlCoverageItem
Expand Down
Loading