mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-07 22:06:28 +00:00
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:
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user