From 389aac851435626fbad5095acbc512e4f71e4fef Mon Sep 17 00:00:00 2001 From: Martin Yankovs Date: Fri, 18 Sep 2026 15:01:17 +0300 Subject: [PATCH] fix(storage/sql): reject filter keys that aren't plain column names buildQuery wrote filter keys into the WHERE clause unchanged. Only values were bound, so a key like "1=1) OR (1=1" widened the clause to every row, and Update and Delete then wrote or removed all of them. Check each key against the same pattern sort keys use, and return an error from Get, List, Count, Update and Delete when one doesn't match. --- storage/sql.go | 39 +++++++++++++++++++++++++++++++-------- storage/sql_test.go | 43 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 74 insertions(+), 8 deletions(-) diff --git a/storage/sql.go b/storage/sql.go index eb19d35..2fd5f97 100644 --- a/storage/sql.go +++ b/storage/sql.go @@ -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 @@ -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 } @@ -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 } @@ -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 @@ -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) } @@ -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)) @@ -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 } diff --git a/storage/sql_test.go b/storage/sql_test.go index d266e2e..141869f 100644 --- a/storage/sql_test.go +++ b/storage/sql_test.go @@ -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