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 1fe5e54bfc
9 changed files with 686 additions and 16 deletions
+77
View File
@@ -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 "<prefix>.<column>" into "<column>" 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 "<main table or alias>.<column>" 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))
}
+55
View File
@@ -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)
}
}
+92
View File
@@ -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)
}
}
+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)
+180
View File
@@ -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)
}
}
+3
View File
@@ -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 "<main table or alias>.<column>" 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)