From e4c4315f4b8ed3ebe49841fee8505abe41dd7371 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 7 Oct 2026 20:58:44 +0200 Subject: [PATCH] fix(filters): qualify model columns to avoid ambiguity with joins Qualify plain model columns with the main table alias in resolvespec filters so preload joins no longer cause SQLSTATE 42702, and accept . from clients in resolvespec and restheadspec instead of silently dropping the filter. Add join/preload tests for resolvespec, restheadspec and funcspec. --- pkg/common/qualify_column.go | 77 ++++++++ pkg/common/qualify_column_test.go | 55 ++++++ pkg/funcspec/filter_join_test.go | 92 ++++++++++ pkg/resolvespec/filter_join_sqlite_test.go | 183 ++++++++++++++++++++ pkg/resolvespec/filter_qualify_test.go | 65 +++++++ pkg/resolvespec/handler.go | 45 +++-- pkg/resolvespec/json_column_test.go | 2 +- pkg/restheadspec/filter_join_sqlite_test.go | 180 +++++++++++++++++++ pkg/restheadspec/handler.go | 3 + 9 files changed, 686 insertions(+), 16 deletions(-) create mode 100644 pkg/common/qualify_column.go create mode 100644 pkg/common/qualify_column_test.go create mode 100644 pkg/funcspec/filter_join_test.go create mode 100644 pkg/resolvespec/filter_join_sqlite_test.go create mode 100644 pkg/resolvespec/filter_qualify_test.go create mode 100644 pkg/restheadspec/filter_join_sqlite_test.go diff --git a/pkg/common/qualify_column.go b/pkg/common/qualify_column.go new file mode 100644 index 0000000..dffa51e --- /dev/null +++ b/pkg/common/qualify_column.go @@ -0,0 +1,77 @@ +package common + +import ( + "regexp" + "strings" + + "github.com/bitechdev/ResolveSpec/pkg/reflection" +) + +var rePlainIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) + +// MainTableAlias returns the alias the main table gets in the SELECT: +// the model's TableAlias() when provided, otherwise the bare table name. +func MainTableAlias(model interface{}, tableName string) string { + if p, ok := model.(TableAliasProvider); ok { + if a := p.TableAlias(); a != "" { + return a + } + } + return reflection.ExtractTableNameOnly(tableName) +} + +func isModelSQLColumn(model interface{}, column string) bool { + for _, c := range reflection.GetSQLModelColumns(model) { + if strings.EqualFold(c, column) { + return true + } + } + return false +} + +// QualifyModelColumn returns "alias"."column" when column is a plain identifier +// that exists on the model, so it stays unambiguous once joins are added. +// Anything else (expressions, JSON paths, other tables' columns) is returned unchanged. +func QualifyModelColumn(model interface{}, alias, column string) string { + if alias == "" || model == nil || !rePlainIdent.MatchString(column) || !isModelSQLColumn(model, column) { + return column + } + return QuoteIdent(alias) + "." + QuoteIdent(column) +} + +// StripMainTablePrefix turns "." into "" when prefix is +// one of the main table's names/aliases and column exists on the model. +// Any other input is returned unchanged. +func StripMainTablePrefix(model interface{}, column string, prefixes ...string) string { + idx := strings.Index(column, ".") + if idx <= 0 || model == nil { + return column + } + prefix := strings.Trim(column[:idx], `"`) + col := strings.Trim(column[idx+1:], `"`) + if !rePlainIdent.MatchString(prefix) || !rePlainIdent.MatchString(col) { + return column + } + for _, p := range prefixes { + if p != "" && strings.EqualFold(p, prefix) && isModelSQLColumn(model, col) { + return col + } + } + return column +} + +// StripMainTablePrefixFromFilters applies StripMainTablePrefix to every filter column. +func StripMainTablePrefixFromFilters(model interface{}, filters []FilterOption, prefixes ...string) { + for i := range filters { + filters[i].Column = StripMainTablePrefix(model, filters[i].Column, prefixes...) + } +} + +// NormalizeMainTableFilters rewrites "
." filter +// columns on opts to the bare model column (on a copy, never the caller's +// slice) so the column validator keeps them. +func NormalizeMainTableFilters(model interface{}, tableName string, opts *RequestOptions) { + opts.Filters = append([]FilterOption(nil), opts.Filters...) + StripMainTablePrefixFromFilters(model, opts.Filters, + MainTableAlias(model, tableName), reflection.ExtractTableNameOnly(tableName)) +} diff --git a/pkg/common/qualify_column_test.go b/pkg/common/qualify_column_test.go new file mode 100644 index 0000000..2203bbe --- /dev/null +++ b/pkg/common/qualify_column_test.go @@ -0,0 +1,55 @@ +package common + +import "testing" + +type qualifyModel struct { + ID int64 `bun:"id,pk"` + Name string `bun:"name"` +} + +func TestQualifyModelColumn(t *testing.T) { + m := qualifyModel{} + cases := []struct{ alias, col, want string }{ + {"t", "name", `"t"."name"`}, + {"t", "NAME", `"t"."NAME"`}, + {"t", "missing", "missing"}, + {"t", "data->>'x'", "data->>'x'"}, + {"t", "rel.name", "rel.name"}, + {"", "name", "name"}, + } + for _, c := range cases { + if got := QualifyModelColumn(m, c.alias, c.col); got != c.want { + t.Errorf("Qualify(%q,%q) = %q, want %q", c.alias, c.col, got, c.want) + } + } + if got := QualifyModelColumn(nil, "t", "name"); got != "name" { + t.Errorf("nil model: %q", got) + } +} + +func TestStripMainTablePrefix(t *testing.T) { + m := qualifyModel{} + cases := []struct{ col, want string }{ + {"province_state.name", "name"}, + {`"province_state"."name"`, "name"}, + {"PROVINCE_STATE.name", "name"}, + {"rel_rid_country.name", "rel_rid_country.name"}, + {"province_state.missing", "province_state.missing"}, + {"name", "name"}, + } + for _, c := range cases { + if got := StripMainTablePrefix(m, c.col, "province_state"); got != c.want { + t.Errorf("Strip(%q) = %q, want %q", c.col, got, c.want) + } + } +} + +func TestStripMainTablePrefix_ValidatorKeepsFilter(t *testing.T) { + m := qualifyModel{} + filters := []FilterOption{{Column: "province_state.name", Operator: "eq", Value: "x"}} + StripMainTablePrefixFromFilters(m, filters, "province_state") + out := NewColumnValidator(m).FilterRequestOptions(RequestOptions{Filters: filters}) + if len(out.Filters) != 1 || out.Filters[0].Column != "name" { + t.Fatalf("filters = %+v", out.Filters) + } +} diff --git a/pkg/funcspec/filter_join_test.go b/pkg/funcspec/filter_join_test.go new file mode 100644 index 0000000..2e8a407 --- /dev/null +++ b/pkg/funcspec/filter_join_test.go @@ -0,0 +1,92 @@ +package funcspec + +import ( + "database/sql" + "strings" + "testing" + + "github.com/uptrace/bun/driver/sqliteshim" +) + +const joinedBaseSQL = "SELECT p.id, p.name, c.name AS country_name FROM province_state p LEFT JOIN country c ON c.id = p.rid_country" + +// funcspec filters are appended to author-written SQL with no model, so columns +// are used verbatim: clients disambiguate by sending "alias.column", and the +// dot must survive ValidSQL and every filter path. +func TestApplyFilters_QualifiedColumnsPreservedWithJoin(t *testing.T) { + h := NewHandler(&MockDatabase{}) + + cases := []struct { + name string + params *RequestParameters + want string + }{ + {"field filter", &RequestParameters{FieldFilters: map[string]string{"p.name": "abc"}}, "p.name = abc"}, + {"search filter", &RequestParameters{SearchFilters: map[string]string{"p.name": "abc"}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"}, + {"search op eq", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "eq", Value: "abc", Logic: "AND"}}}, "p.name = 'abc'"}, + {"search op contains", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "contains", Value: "abc", Logic: "AND"}}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got := h.ApplyFilters(joinedBaseSQL, c.params) + if !strings.Contains(got, c.want) { + t.Fatalf("SQL %q does not contain %q", got, c.want) + } + if !strings.Contains(got, "LEFT JOIN country") || !strings.Contains(got, " WHERE ") { + t.Fatalf("join/where lost: %q", got) + } + }) + } +} + +func TestApplyFilters_WithAndWithoutJoin_RealQuery(t *testing.T) { + sqldb, err := sql.Open(sqliteshim.ShimName, "file:funcspecjoin?mode=memory&cache=private") + if err != nil { + t.Fatal(err) + } + defer sqldb.Close() + for _, stmt := range []string{ + "CREATE TABLE country (id INTEGER PRIMARY KEY, name TEXT)", + "CREATE TABLE province_state (id INTEGER PRIMARY KEY, name TEXT, rid_country INTEGER)", + "INSERT INTO country VALUES (1,'Abcland'),(2,'Other')", + "INSERT INTO province_state VALUES (1,'abc one',1),(2,'xyz two',1),(3,'nope',2)", + } { + if _, err := sqldb.Exec(stmt); err != nil { + t.Fatal(err) + } + } + + h := NewHandler(&MockDatabase{}) + count := func(base string, params *RequestParameters) (int, error) { + rows, err := sqldb.Query(h.ApplyFilters(base, params)) + if err != nil { + return 0, err + } + defer rows.Close() + n := 0 + for rows.Next() { + n++ + } + return n, rows.Err() + } + + const noJoin = "SELECT p.id, p.name FROM province_state p" + params := func(col string) *RequestParameters { + return &RequestParameters{SearchOps: map[string]FilterOperator{col: {Operator: "eq", Value: "nope", Logic: "AND"}}} + } + + if n, err := count(noJoin, params("name")); err != nil || n != 1 { + t.Fatalf("no join, unqualified: n=%d err=%v", n, err) + } + if n, err := count(joinedBaseSQL, params("p.name")); err != nil || n != 1 { + t.Fatalf("join, qualified: n=%d err=%v", n, err) + } + if n, err := count(joinedBaseSQL, params("c.name")); err != nil || n != 0 { + t.Fatalf("join, joined-table column: n=%d err=%v", n, err) + } + // Documented limit: with a join in the author's SQL, an unqualified shared + // column is ambiguous and the client must send "alias.column". + if _, err := count(joinedBaseSQL, params("name")); err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") { + t.Fatalf("expected ambiguous error for unqualified shared column, got %v", err) + } +} diff --git a/pkg/resolvespec/filter_join_sqlite_test.go b/pkg/resolvespec/filter_join_sqlite_test.go new file mode 100644 index 0000000..1616489 --- /dev/null +++ b/pkg/resolvespec/filter_join_sqlite_test.go @@ -0,0 +1,183 @@ +package resolvespec + +import ( + "context" + "database/sql" + "strings" + "testing" + + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/sqlitedialect" + "github.com/uptrace/bun/driver/sqliteshim" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" +) + +type joinCountry struct { + bun.BaseModel `bun:"table:country,alias:country"` + ID int64 `bun:"id,pk"` + Name string `bun:"name"` +} + +type joinProvince struct { + bun.BaseModel `bun:"table:province_state,alias:province_state"` + ID int64 `bun:"id,pk"` + Name string `bun:"name"` + Abbreviation string `bun:"abbreviation"` + RidCountry int64 `bun:"rid_country"` +} + +// SQLite has no ILIKE; the ilike operator is covered by SQL-string tests, and +// LIKE here proves the CAST(... AS TEXT) wrapping executes with qualified columns. +func setupJoinDB(t *testing.T) *bun.DB { + t.Helper() + sqldb, err := sql.Open(sqliteshim.ShimName, "file:filterjoin?mode=memory&cache=private") + if err != nil { + t.Fatal(err) + } + db := bun.NewDB(sqldb, sqlitedialect.New()) + t.Cleanup(func() { _ = db.Close() }) + ctx := context.Background() + for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} { + if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil { + t.Fatal(err) + } + } + if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil { + t.Fatal(err) + } + if _, err := db.NewInsert().Model(&[]joinProvince{ + {ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1}, + {ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1}, + {ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2}, + }).Exec(ctx); err != nil { + t.Fatal(err) + } + return db +} + +const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country" + +func countWith(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption, alias string) (int, error) { + t.Helper() + h := &Handler{} + var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{}) + if withJoin { + q = q.Join(joinSQL) + } + q = h.applyFilters(q, filters, &joinProvince{}, alias) + return q.Count(context.Background()) +} + +func TestFilters_AmbiguousColumnWithJoin_RealQuery(t *testing.T) { + db := setupJoinDB(t) + + filters := []common.FilterOption{ + {Column: "name", Operator: "like", Value: "%abc%"}, + {Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"}, + } + + t.Run("unqualified with join is ambiguous (the bug)", func(t *testing.T) { + _, err := countWith(t, db, true, filters, "") + if err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") { + t.Fatalf("expected ambiguous column error, got %v", err) + } + }) + + t.Run("qualified with join", func(t *testing.T) { + n, err := countWith(t, db, true, filters, "province_state") + if err != nil { + t.Fatal(err) + } + if n != 2 { + t.Fatalf("count = %d, want 2", n) + } + }) + + t.Run("qualified without join", func(t *testing.T) { + n, err := countWith(t, db, false, filters, "province_state") + if err != nil { + t.Fatal(err) + } + if n != 2 { + t.Fatalf("count = %d, want 2", n) + } + }) + + t.Run("unqualified without join still works", func(t *testing.T) { + n, err := countWith(t, db, false, filters, "") + if err != nil { + t.Fatal(err) + } + if n != 2 { + t.Fatalf("count = %d, want 2", n) + } + }) +} + +func TestFilters_AllOperatorsWithJoin_RealQuery(t *testing.T) { + db := setupJoinDB(t) + cases := []struct { + name string + filter common.FilterOption + want int + }{ + {"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "nope"}, 1}, + {"neq", common.FilterOption{Column: "name", Operator: "neq", Value: "nope"}, 2}, + {"like", common.FilterOption{Column: "name", Operator: "like", Value: "abc%"}, 1}, + {"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}, 2}, + {"gt", common.FilterOption{Column: "id", Operator: "gt", Value: 1}, 2}, + {"gte", common.FilterOption{Column: "id", Operator: "gte", Value: 2}, 2}, + {"lt", common.FilterOption{Column: "id", Operator: "lt", Value: 3}, 2}, + {"lte", common.FilterOption{Column: "id", Operator: "lte", Value: 1}, 1}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + n, err := countWith(t, db, true, []common.FilterOption{c.filter}, "province_state") + if err != nil { + t.Fatal(err) + } + if n != c.want { + t.Fatalf("count = %d, want %d", n, c.want) + } + }) + } +} + +// A filter on the joined table's own column, sent already qualified, must pass through untouched. +func TestFilters_JoinedTableColumnPassesThrough_RealQuery(t *testing.T) { + db := setupJoinDB(t) + n, err := countWith(t, db, true, []common.FilterOption{ + {Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"}, + }, "province_state") + if err != nil { + t.Fatal(err) + } + if n != 2 { + t.Fatalf("count = %d, want 2", n) + } +} + +// Client sends "province_state.name": it must survive validation and be applied. +func TestFilters_ClientQualifiedColumn_NotDropped(t *testing.T) { + db := setupJoinDB(t) + model := &joinProvince{} + + opts := common.RequestOptions{Filters: []common.FilterOption{ + {Column: "province_state.name", Operator: "like", Value: "%abc%"}, + }} + common.NormalizeMainTableFilters(model, "public.province_state", &opts) + opts = common.NewColumnValidator(model).FilterRequestOptions(opts) + if len(opts.Filters) != 1 { + t.Fatalf("filter was dropped: %+v", opts.Filters) + } + + n, err := countWith(t, db, true, opts.Filters, "province_state") + if err != nil { + t.Fatal(err) + } + if n != 1 { + t.Fatalf("count = %d, want 1 (filter must narrow the result)", n) + } +} diff --git a/pkg/resolvespec/filter_qualify_test.go b/pkg/resolvespec/filter_qualify_test.go new file mode 100644 index 0000000..da8cc57 --- /dev/null +++ b/pkg/resolvespec/filter_qualify_test.go @@ -0,0 +1,65 @@ +package resolvespec + +import ( + "testing" + + "github.com/bitechdev/ResolveSpec/pkg/common" +) + +func TestBuildFilterConditionAlias_QualifiesModelColumns(t *testing.T) { + h := &Handler{} + model := jsonColModel{} + + tests := []struct { + name string + filter common.FilterOption + want string + }{ + {"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "x"}, `"province_state"."name" = ?`}, + {"ilike", common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) ILIKE ?`}, + {"like", common.FilterOption{Column: "name", Operator: "like", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) LIKE ?`}, + {"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"a", "b"}}, `"province_state"."name" IN (?,?)`}, + {"non-model column untouched", common.FilterOption{Column: "other", Operator: "eq", Value: 1}, `other = ?`}, + {"already qualified untouched", common.FilterOption{Column: "rel.name", Operator: "eq", Value: 1}, `rel.name = ?`}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, _ := h.buildFilterConditionAlias(tt.filter, model, "province_state") + if got != tt.want { + t.Fatalf("got %q, want %q", got, tt.want) + } + }) + } +} + +func TestBuildFilterConditionAlias_NoAliasUnchanged(t *testing.T) { + h := &Handler{} + got, _ := h.buildFilterConditionAlias(common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "") + if got != "CAST(name AS TEXT) ILIKE ?" { + t.Fatalf("got %q", got) + } +} + +func TestApplyFilter_QualifiesWithAlias(t *testing.T) { + h := &Handler{} + q := &jsonCapQuery{} + h.applyFilter(q, common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "province_state") + c := q.only(t) + if c.query != `CAST("province_state"."name" AS TEXT) ILIKE ?` { + t.Fatalf("query = %q", c.query) + } +} + +func TestApplyFilters_OrGroupQualified(t *testing.T) { + h := &Handler{} + q := &jsonCapQuery{} + h.applyFilters(q, []common.FilterOption{ + {Column: "name", Operator: "ilike", Value: "%a%"}, + {Column: "id", Operator: "eq", Value: 1, LogicOperator: "OR"}, + }, jsonColModel{}, "t") + c := q.only(t) + want := `(CAST("t"."name" AS TEXT) ILIKE ? OR "t"."id" = ?)` + if c.query != want { + t.Fatalf("query = %q, want %q", c.query, want) + } +} diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index c1b397c..0eae678 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -172,6 +172,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s // Add request-scoped data to context ctx = WithRequestData(ctx, schema, entity, tableName, model, modelPtr) + // Accept "
." for model columns before validation drops them + common.NormalizeMainTableFilters(model, tableName, &req.Options) + // Validate and filter columns in options (log warnings for invalid columns) validator := common.NewColumnValidator(model) req.Options = validator.FilterRequestOptions(req.Options) @@ -411,7 +414,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st } // Apply filters with proper grouping for OR logic - query = h.applyFilters(query, options.Filters, model) + query = h.applyFilters(query, options.Filters, model, common.MainTableAlias(model, tableName)) // Apply custom operators for _, customOp := range options.CustomOperators { @@ -558,7 +561,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st // Apply the same filters as the main query for _, filter := range options.Filters { - rowNumQuery = h.applyFilter(rowNumQuery, filter, model) + rowNumQuery = h.applyFilter(rowNumQuery, filter, model, common.MainTableAlias(model, tableName)) } // Apply custom operators @@ -1932,7 +1935,7 @@ func (h *Handler) executeDelete(ctx context.Context, tx common.Database, hookCtx // applyFilters applies all filters with proper grouping for OR logic // Groups consecutive OR filters together to ensure proper query precedence // Example: [A, B(OR), C(OR), D(AND)] => WHERE (A OR B OR C) AND D -func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery { +func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery { if len(filters) == 0 { return query } @@ -1952,11 +1955,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter } // Apply the OR group as a single grouped WHERE clause - query = h.applyFilterGroup(query, orGroup, model) + query = h.applyFilterGroup(query, orGroup, model, alias) i = j } else { // Single filter with AND logic (or first filter) - condition, args := h.buildFilterCondition(filters[i], model) + condition, args := h.buildFilterConditionAlias(filters[i], model, alias) if condition != "" { query = query.Where(condition, args...) } @@ -1969,7 +1972,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter // applyFilterGroup applies a group of filters that should be OR'd together // Always wraps them in parentheses and applies as a single WHERE clause -func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery { +func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery { if len(filters) == 0 { return query } @@ -1979,7 +1982,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi var args []interface{} for _, filter := range filters { - condition, filterArgs := h.buildFilterCondition(filter, model) + condition, filterArgs := h.buildFilterConditionAlias(filter, model, alias) if condition != "" { conditions = append(conditions, condition) args = append(args, filterArgs...) @@ -2005,6 +2008,12 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi // or the dotted data.x shorthand for a JSON column) resolve to a safe, // parameterised expression before the ordinary operator handling below. func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) { + return h.buildFilterConditionAlias(filter, model, "") +} + +// buildFilterConditionAlias is buildFilterCondition with plain model columns +// qualified by the main table alias, so joins from preloads can't make them ambiguous. +func (h *Handler) buildFilterConditionAlias(filter common.FilterOption, model interface{}, alias string) (conditionString string, conditionArgs []interface{}) { var condition string var args []interface{} @@ -2012,6 +2021,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa return cond, jargs } + rawColumn := filter.Column + filter.Column = common.QualifyModelColumn(model, alias, filter.Column) + switch filter.Operator { case "eq", "=": condition = fmt.Sprintf("%s = ?", filter.Column) @@ -2032,10 +2044,10 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa condition = fmt.Sprintf("%s <= ?", filter.Column) args = []interface{}{filter.Value} case "like": - condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model)) + condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model)) args = []interface{}{filter.Value} case "ilike": - condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model)) + condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model)) args = []interface{}{filter.Value} case "in": condition, args = common.BuildInCondition(filter.Column, filter.Value) @@ -2073,14 +2085,14 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa // CAST(... AS TEXT) would switch to case-sensitive matching and defeat a // citext index. Every other column is cast to TEXT so LIKE/ILIKE also works // against date/time/timestamp and numeric columns. -func likeColumn(column string, model interface{}) string { - if reflection.IsCitextColumn(model, column) { +func likeColumn(column, rawColumn string, model interface{}) string { + if reflection.IsCitextColumn(model, rawColumn) { return column } return fmt.Sprintf("CAST(%s AS TEXT)", column) } -func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}) common.SelectQuery { +func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}, alias string) common.SelectQuery { // Determine which method to use based on LogicOperator useOrLogic := strings.EqualFold(filter.LogicOperator, "OR") @@ -2094,6 +2106,9 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti return query.Where(cond, jargs...) } + rawColumn := filter.Column + filter.Column = common.QualifyModelColumn(model, alias, filter.Column) + switch filter.Operator { case "eq", "=": condition = fmt.Sprintf("%s = ?", filter.Column) @@ -2114,10 +2129,10 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti condition = fmt.Sprintf("%s <= ?", filter.Column) args = []interface{}{filter.Value} case "like": - condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model)) + condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model)) args = []interface{}{filter.Value} case "ilike": - condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model)) + condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model)) args = []interface{}{filter.Value} case "in": condition, args = common.BuildInCondition(filter.Column, filter.Value) @@ -2523,7 +2538,7 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre if len(preload.Filters) > 0 { for _, filter := range preload.Filters { - sq = h.applyFilter(sq, filter, nil) + sq = h.applyFilter(sq, filter, nil, "") } } if len(preload.Sort) > 0 { diff --git a/pkg/resolvespec/json_column_test.go b/pkg/resolvespec/json_column_test.go index cee1368..dcd3f69 100644 --- a/pkg/resolvespec/json_column_test.go +++ b/pkg/resolvespec/json_column_test.go @@ -142,7 +142,7 @@ func TestApplyFilter_JSONColumn(t *testing.T) { q := &jsonCapQuery{} h.applyFilter(q, common.FilterOption{ Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR", - }, model) + }, model, "") c := q.only(t) if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` { t.Fatalf("call = %+v", c) diff --git a/pkg/restheadspec/filter_join_sqlite_test.go b/pkg/restheadspec/filter_join_sqlite_test.go new file mode 100644 index 0000000..0a154be --- /dev/null +++ b/pkg/restheadspec/filter_join_sqlite_test.go @@ -0,0 +1,180 @@ +package restheadspec + +import ( + "context" + "database/sql" + "testing" + + "github.com/uptrace/bun" + "github.com/uptrace/bun/dialect/sqlitedialect" + "github.com/uptrace/bun/driver/sqliteshim" + + "github.com/bitechdev/ResolveSpec/pkg/common" + "github.com/bitechdev/ResolveSpec/pkg/common/adapters/database" +) + +type joinCountry struct { + bun.BaseModel `bun:"table:country,alias:country"` + ID int64 `bun:"id,pk"` + Name string `bun:"name"` +} + +type joinProvince struct { + bun.BaseModel `bun:"table:province_state,alias:province_state"` + ID int64 `bun:"id,pk"` + Name string `bun:"name"` + Abbreviation string `bun:"abbreviation"` + RidCountry int64 `bun:"rid_country"` +} + +func setupJoinDB(t *testing.T) *bun.DB { + t.Helper() + sqldb, err := sql.Open(sqliteshim.ShimName, "file:rhsfilterjoin?mode=memory&cache=private") + if err != nil { + t.Fatal(err) + } + db := bun.NewDB(sqldb, sqlitedialect.New()) + t.Cleanup(func() { _ = db.Close() }) + ctx := context.Background() + for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} { + if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil { + t.Fatal(err) + } + } + if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil { + t.Fatal(err) + } + if _, err := db.NewInsert().Model(&[]joinProvince{ + {ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1}, + {ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1}, + {ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2}, + }).Exec(ctx); err != nil { + t.Fatal(err) + } + return db +} + +const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country" + +// countFilters applies filters the way handleRead does (single AND filters via +// applyFilter, consecutive OR filters via applyOrFilterGroup) and counts. +func countFilters(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption) (int, error) { + t.Helper() + h := &Handler{} + model := &joinProvince{} + var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{}) + if withJoin { + q = q.Join(joinSQL) + } + for i := 0; i < len(filters); { + f := filters[i] + castInfo := h.ValidateAndAdjustFilterForColumnType(&f, model) + if f.LogicOperator == "OR" { + group := []*common.FilterOption{&f} + info := []ColumnCastInfo{castInfo} + j := i + 1 + for j < len(filters) && filters[j].LogicOperator == "OR" { + g := filters[j] + info = append(info, h.ValidateAndAdjustFilterForColumnType(&g, model)) + group = append(group, &g) + j++ + } + q = h.applyOrFilterGroup(q, group, info, "public.province_state", model) + i = j + continue + } + q = h.applyFilter(q, f, "public.province_state", castInfo.NeedsCast, "AND", model) + i++ + } + return q.Count(context.Background()) +} + +func TestRHSFilters_WithAndWithoutJoin_RealQuery(t *testing.T) { + db := setupJoinDB(t) + cases := []struct { + name string + filters []common.FilterOption + want int + }{ + {"eq", []common.FilterOption{{Column: "name", Operator: "eq", Value: "nope"}}, 1}, + {"neq", []common.FilterOption{{Column: "name", Operator: "neq", Value: "nope"}}, 2}, + {"like", []common.FilterOption{{Column: "name", Operator: "like", Value: "%abc%"}}, 1}, + {"in", []common.FilterOption{{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}}, 2}, + {"gt", []common.FilterOption{{Column: "id", Operator: "gt", Value: 1}}, 2}, + {"between", []common.FilterOption{{Column: "id", Operator: "between", Value: []interface{}{0, 3}}}, 2}, + {"between_inclusive", []common.FilterOption{{Column: "id", Operator: "between_inclusive", Value: []interface{}{1, 3}}}, 3}, + {"is_not_null", []common.FilterOption{{Column: "name", Operator: "is_not_null"}}, 3}, + {"or group", []common.FilterOption{ + {Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"}, + {Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"}, + }, 2}, + {"or group then and", []common.FilterOption{ + {Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"}, + {Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"}, + {Column: "id", Operator: "eq", Value: 2}, + }, 1}, + } + for _, c := range cases { + for _, withJoin := range []bool{false, true} { + name := c.name + "/no_join" + if withJoin { + name = c.name + "/join" + } + t.Run(name, func(t *testing.T) { + n, err := countFilters(t, db, withJoin, c.filters) + if err != nil { + t.Fatalf("query failed (ambiguous column?): %v", err) + } + if n != c.want { + t.Fatalf("count = %d, want %d", n, c.want) + } + }) + } + } +} + +func TestRHSFilters_JoinedTableColumnPassesThrough(t *testing.T) { + db := setupJoinDB(t) + n, err := countFilters(t, db, true, []common.FilterOption{ + {Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"}, + }) + if err != nil { + t.Fatal(err) + } + if n != 2 { + t.Fatalf("count = %d, want 2", n) + } +} + +// Client sends "province_state.name": it must survive validation and be applied. +func TestRHSFilters_ClientQualifiedColumn_NotDropped(t *testing.T) { + db := setupJoinDB(t) + h := &Handler{} + model := &joinProvince{} + + opts := ExtendedRequestOptions{} + opts.Filters = []common.FilterOption{{Column: "province_state.name", Operator: "like", Value: "%abc%"}} + common.NormalizeMainTableFilters(model, "public.province_state", &opts.RequestOptions) + opts = h.filterExtendedOptions(common.NewColumnValidator(model), opts, model) + if len(opts.Filters) != 1 || opts.Filters[0].Column != "name" { + t.Fatalf("filter was dropped or not normalised: %+v", opts.Filters) + } + + n, err := countFilters(t, db, true, opts.Filters) + if err != nil { + t.Fatal(err) + } + if n != 1 { + t.Fatalf("count = %d, want 1", n) + } +} + +func TestRHSFilters_ILikeSQLQualified(t *testing.T) { + h := &Handler{} + q := &jsonCapQuery{} + f := common.FilterOption{Column: "name", Operator: "ilike", Value: "%abc%"} + h.applyFilter(q, f, "info.province_state", false, "AND", jsonColModel{}) + if got := q.only(t).query; got != "CAST(province_state.name AS TEXT) ILIKE ?" { + t.Fatalf("query = %q", got) + } +} diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index a89dd9f..b357c83 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -164,6 +164,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s // Parse options from headers - this now includes relation name resolution options := h.parseOptionsFromHeaders(r, model) + // Accept "
." for model columns before validation drops them + common.NormalizeMainTableFilters(model, tableName, &options.RequestOptions) + // Validate and filter columns in options (log warnings for invalid columns) validator := common.NewColumnValidator(model) options = h.filterExtendedOptions(validator, options, model)