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