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
<table>.<column> from clients in resolvespec and restheadspec instead
of silently dropping the filter. Add join/preload tests for
resolvespec, restheadspec and funcspec.
This commit is contained in:
2026-10-07 20:58:44 +02:00
parent 3efd539e0f
commit e4c4315f4b
9 changed files with 686 additions and 16 deletions
+183
View File
@@ -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)
}
}
+65
View File
@@ -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)
}
}
+30 -15
View File
@@ -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 "<main table or alias>.<column>" 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 {
+1 -1
View File
@@ -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)