mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-03 03:51:59 +00:00
fix(common): harden CORS, sort, raw-SQL WHERE and x-custom-sql-or
Add a `hardening` config section (RESOLVESPEC_HARDENING_*) so each
fix can be switched off to restore the previous behaviour:
- cors_strict_origins: only reflect origins listed in
cors.allowed_origins / server URLs, with credentials; `*` never
sends credentials; fix shared-slice append of expose headers.
- sort_strict: join aliases must match `alias.identifier` (empty alias
no longer matches everything); sort expressions reject dangerous
functions/catalogs; cql* columns must be identifier-safe.
- sql_strict: client raw-SQL fragments must have balanced parens and
quotes, no comments/`;`/`$$`, DML keywords, dangerous functions or
system catalogs; a rejected fragment now fails closed ("(1=0)")
instead of dropping the filter. Subqueries stay allowed unless
sql_block_subqueries is set.
- x-custom-sql-or is grouped together with the client's own
conditions (new optional WhereGrouper, implemented for bun and gorm)
so it can no longer OR past server-side filters.
This commit is contained in:
@@ -537,6 +537,19 @@ func (b *BunSelectQuery) WhereOr(query string, args ...interface{}) common.Selec
|
|||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WhereGroup wraps the conditions added by fn in one parenthesised group ANDed with the rest.
|
||||||
|
func (b *BunSelectQuery) WhereGroup(fn func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||||
|
b.query = b.query.WhereGroup(" AND ", func(q *bun.SelectQuery) *bun.SelectQuery {
|
||||||
|
inner := *b
|
||||||
|
inner.query = q
|
||||||
|
if res, ok := fn(&inner).(*BunSelectQuery); ok {
|
||||||
|
return res.query
|
||||||
|
}
|
||||||
|
return q
|
||||||
|
})
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
func (b *BunSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
|
func (b *BunSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
|
||||||
// Extract optional prefix from args
|
// Extract optional prefix from args
|
||||||
// If the last arg is a string that looks like a table prefix, use it
|
// If the last arg is a string that looks like a table prefix, use it
|
||||||
|
|||||||
@@ -377,6 +377,16 @@ func (g *GormSelectQuery) WhereOr(query string, args ...interface{}) common.Sele
|
|||||||
return g
|
return g
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WhereGroup wraps the conditions added by fn in one parenthesised group ANDed with the rest.
|
||||||
|
func (g *GormSelectQuery) WhereGroup(fn func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||||
|
inner := *g
|
||||||
|
inner.db = g.db.Session(&gorm.Session{NewDB: true})
|
||||||
|
if res, ok := fn(&inner).(*GormSelectQuery); ok {
|
||||||
|
g.db = g.db.Where(res.db)
|
||||||
|
}
|
||||||
|
return g
|
||||||
|
}
|
||||||
|
|
||||||
func (g *GormSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
|
func (g *GormSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
|
||||||
// Extract optional prefix from args
|
// Extract optional prefix from args
|
||||||
// If the last arg is a string that looks like a table prefix, use it
|
// If the last arg is a string that looks like a table prefix, use it
|
||||||
|
|||||||
@@ -0,0 +1,31 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bun/dialect/pgdialect"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The client's OR must widen only the client's own conditions, never the
|
||||||
|
// server-side condition ANDed after it.
|
||||||
|
func TestBunWhereGroupConfinesOr(t *testing.T) {
|
||||||
|
db := bun.NewDB(&sql.DB{}, pgdialect.New())
|
||||||
|
q := &BunSelectQuery{query: db.NewSelect().TableExpr("items"), db: db, driverName: "postgres"}
|
||||||
|
|
||||||
|
var sq common.SelectQuery = q.Where("a = 1")
|
||||||
|
sq = sq.(common.WhereGrouper).WhereGroup(func(g common.SelectQuery) common.SelectQuery {
|
||||||
|
return g.Where("b = 2").WhereOr("(c = 3)")
|
||||||
|
})
|
||||||
|
sq = sq.Where("tenant = 5")
|
||||||
|
|
||||||
|
got := sq.(*BunSelectQuery).query.String()
|
||||||
|
want := `WHERE (a = 1) AND ((b = 2) OR ((c = 3))) AND (tenant = 5)`
|
||||||
|
if !strings.Contains(got, want) {
|
||||||
|
t.Fatalf("unexpected SQL:\n got: %s\nwant to contain: %s", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
+85
-12
@@ -20,7 +20,15 @@ func DefaultCORSConfig() CORSConfig {
|
|||||||
configManager := config.GetConfigManager()
|
configManager := config.GetConfigManager()
|
||||||
cfg, _ := configManager.GetConfig()
|
cfg, _ := configManager.GetConfig()
|
||||||
hosts := make([]string, 0)
|
hosts := make([]string, 0)
|
||||||
// hosts = append(hosts, "*")
|
if cfg == nil {
|
||||||
|
return CORSConfig{
|
||||||
|
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
|
||||||
|
AllowedHeaders: GetHeadSpecHeaders(),
|
||||||
|
MaxAge: 86400,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Explicitly configured origins (cors.allowed_origins); "*" allows any origin without credentials
|
||||||
|
hosts = append(hosts, cfg.CORS.AllowedOrigins...)
|
||||||
|
|
||||||
_, _, ipsList := config.GetIPs()
|
_, _, ipsList := config.GetIPs()
|
||||||
|
|
||||||
@@ -113,25 +121,63 @@ func GetHeadSpecHeaders() []string {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetCORSHeaders sets CORS headers on a response writer
|
// originAllowed reports whether origin matches config.AllowedOrigins exactly
|
||||||
|
// (case-insensitive, trailing slash ignored). wildcard is true when the list
|
||||||
|
// contains "*".
|
||||||
|
func originAllowed(origin string, allowed []string) (ok bool, wildcard bool) {
|
||||||
|
norm := strings.ToLower(strings.TrimRight(strings.TrimSpace(origin), "/"))
|
||||||
|
for _, a := range allowed {
|
||||||
|
a = strings.ToLower(strings.TrimRight(strings.TrimSpace(a), "/"))
|
||||||
|
if a == "*" {
|
||||||
|
wildcard = true
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if a != "" && a == norm {
|
||||||
|
return true, wildcard
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false, wildcard
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCORSHeaders sets CORS headers on a response writer.
|
||||||
|
//
|
||||||
|
// The request Origin is only reflected (and credentials only allowed) when it
|
||||||
|
// is listed in config.AllowedOrigins. A "*" entry allows any origin but never
|
||||||
|
// with credentials. Unlisted origins get no CORS headers, so browsers block
|
||||||
|
// the cross-origin read.
|
||||||
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
||||||
// Reflect the request origin; fall back to wildcard only when no origin is present
|
if !Hardening().CORSStrictOrigins {
|
||||||
|
setCORSHeadersLegacy(w, r, config)
|
||||||
|
return
|
||||||
|
}
|
||||||
origin := r.Header("Origin")
|
origin := r.Header("Origin")
|
||||||
if origin == "" {
|
if origin == "" {
|
||||||
origin = "*"
|
// Not a cross-origin browser request; nothing to protect.
|
||||||
|
w.SetHeader("Access-Control-Allow-Origin", "*")
|
||||||
} else {
|
} else {
|
||||||
// Vary must be set so caches don't serve one origin's response to another
|
// Vary must be set so caches don't serve one origin's response to another
|
||||||
httpW := w.UnderlyingResponseWriter()
|
w.UnderlyingResponseWriter().Header().Set("Vary", "Origin")
|
||||||
httpW.Header().Set("Vary", "Origin")
|
|
||||||
|
ok, wildcard := originAllowed(origin, config.AllowedOrigins)
|
||||||
|
switch {
|
||||||
|
case ok:
|
||||||
|
w.SetHeader("Access-Control-Allow-Origin", origin)
|
||||||
|
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||||
|
case wildcard:
|
||||||
|
w.SetHeader("Access-Control-Allow-Origin", "*")
|
||||||
|
default:
|
||||||
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
w.SetHeader("Access-Control-Allow-Origin", origin)
|
|
||||||
|
|
||||||
// Set allowed methods
|
// Set allowed methods
|
||||||
if len(config.AllowedMethods) > 0 {
|
if len(config.AllowedMethods) > 0 {
|
||||||
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reflect the preflight request headers when present; otherwise use the explicit config list
|
// The origin is trusted at this point, so reflecting the preflight request
|
||||||
|
// headers is safe (the config list contains "X-Foo-*" patterns that browsers
|
||||||
|
// cannot match literally).
|
||||||
requestedHeaders := r.Header("Access-Control-Request-Headers")
|
requestedHeaders := r.Header("Access-Control-Request-Headers")
|
||||||
if requestedHeaders != "" {
|
if requestedHeaders != "" {
|
||||||
w.SetHeader("Access-Control-Allow-Headers", requestedHeaders)
|
w.SetHeader("Access-Control-Allow-Headers", requestedHeaders)
|
||||||
@@ -144,13 +190,40 @@ func SetCORSHeaders(w ResponseWriter, r Request, config CORSConfig) {
|
|||||||
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Allow credentials only when a specific origin is reflected (not wildcard)
|
// Expose headers that clients can read (fresh slice: avoid appending into
|
||||||
|
// config.AllowedHeaders' backing array)
|
||||||
|
exposeHeaders := make([]string, 0, len(config.AllowedHeaders)+3)
|
||||||
|
exposeHeaders = append(exposeHeaders, config.AllowedHeaders...)
|
||||||
|
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
|
||||||
|
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
|
||||||
|
}
|
||||||
|
|
||||||
|
// setCORSHeadersLegacy is the pre-hardening behaviour (reflect any origin with
|
||||||
|
// credentials). Used only when hardening.cors_strict_origins is false.
|
||||||
|
func setCORSHeadersLegacy(w ResponseWriter, r Request, config CORSConfig) {
|
||||||
|
origin := r.Header("Origin")
|
||||||
|
if origin == "" {
|
||||||
|
origin = "*"
|
||||||
|
} else {
|
||||||
|
w.UnderlyingResponseWriter().Header().Set("Vary", "Origin")
|
||||||
|
}
|
||||||
|
w.SetHeader("Access-Control-Allow-Origin", origin)
|
||||||
|
if len(config.AllowedMethods) > 0 {
|
||||||
|
w.SetHeader("Access-Control-Allow-Methods", strings.Join(config.AllowedMethods, ", "))
|
||||||
|
}
|
||||||
|
if requested := r.Header("Access-Control-Request-Headers"); requested != "" {
|
||||||
|
w.SetHeader("Access-Control-Allow-Headers", requested)
|
||||||
|
} else if len(config.AllowedHeaders) > 0 {
|
||||||
|
w.SetHeader("Access-Control-Allow-Headers", strings.Join(config.AllowedHeaders, ", "))
|
||||||
|
}
|
||||||
|
if config.MaxAge > 0 {
|
||||||
|
w.SetHeader("Access-Control-Max-Age", fmt.Sprintf("%d", config.MaxAge))
|
||||||
|
}
|
||||||
if origin != "*" {
|
if origin != "*" {
|
||||||
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||||
}
|
}
|
||||||
|
exposeHeaders := make([]string, 0, len(config.AllowedHeaders)+3)
|
||||||
// Expose headers that clients can read
|
exposeHeaders = append(exposeHeaders, config.AllowedHeaders...)
|
||||||
exposeHeaders := config.AllowedHeaders
|
|
||||||
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
|
exposeHeaders = append(exposeHeaders, "Content-Range", "X-Api-Range-Total", "X-Api-Range-Size")
|
||||||
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
|
w.SetHeader("Access-Control-Expose-Headers", strings.Join(exposeHeaders, ", "))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import "github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
|
|
||||||
|
// hardeningProvider returns the active hardening switches. Tests may replace it.
|
||||||
|
var hardeningProvider = func() config.HardeningConfig {
|
||||||
|
cfg, err := config.GetConfigManager().GetConfig()
|
||||||
|
if err != nil || cfg == nil {
|
||||||
|
// Fail secure: hardening on when config is unavailable.
|
||||||
|
return config.HardeningConfig{CORSStrictOrigins: true, SortStrict: true, SQLStrict: true}
|
||||||
|
}
|
||||||
|
return cfg.Hardening
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hardening returns the security hardening toggles (config section "hardening").
|
||||||
|
func Hardening() config.HardeningConfig { return hardeningProvider() }
|
||||||
@@ -0,0 +1,121 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setHardening(t *testing.T, h config.HardeningConfig) {
|
||||||
|
t.Helper()
|
||||||
|
old := hardeningProvider
|
||||||
|
hardeningProvider = func() config.HardeningConfig { return h }
|
||||||
|
t.Cleanup(func() { hardeningProvider = old })
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
hardOn = config.HardeningConfig{CORSStrictOrigins: true, SortStrict: true, SQLStrict: true}
|
||||||
|
hardOff = config.HardeningConfig{}
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSanitizeWhereClause_Strict(t *testing.T) {
|
||||||
|
setHardening(t, hardOn)
|
||||||
|
allowed := []string{
|
||||||
|
"status = 'awaiting update approval'",
|
||||||
|
"last_update > '2020-01-01'",
|
||||||
|
"name = 'it''s; fine'",
|
||||||
|
"(a = 1 or b = 2) and ifblnk(c) = 'x'",
|
||||||
|
}
|
||||||
|
for _, w := range allowed {
|
||||||
|
if got := SanitizeWhereClause(w, ""); got == "" || got == "(1=0)" {
|
||||||
|
t.Errorf("legitimate clause %q rejected: %q", w, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
hostile := []string{
|
||||||
|
"1=1)) OR ((1=1",
|
||||||
|
"id = 1 or (select count(*) from pg_shadow) > 0",
|
||||||
|
"id in (select oid from pg_catalog.pg_class)",
|
||||||
|
"id = 1 and pg_sleep(5) is not null",
|
||||||
|
"id = 1; delete/**/from items",
|
||||||
|
"id = 1 -- x",
|
||||||
|
"name = 'unterminated",
|
||||||
|
}
|
||||||
|
for _, w := range hostile {
|
||||||
|
if got := SanitizeWhereClause(w, "t"); got != "(1=0)" {
|
||||||
|
t.Errorf("hostile clause %q not rejected closed: %q", w, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSanitizeWhereClause_Subqueries(t *testing.T) {
|
||||||
|
q := "id in (select l.id from other l where l.x = 1)"
|
||||||
|
setHardening(t, hardOn)
|
||||||
|
if got := SanitizeWhereClause(q, "t"); got == "(1=0)" || got == "" {
|
||||||
|
t.Errorf("subquery should be allowed by default: %q", got)
|
||||||
|
}
|
||||||
|
h := hardOn
|
||||||
|
h.SQLBlockSubqueries = true
|
||||||
|
setHardening(t, h)
|
||||||
|
if got := SanitizeWhereClause(q, "t"); got != "(1=0)" {
|
||||||
|
t.Errorf("subquery should be blocked with sql_block_subqueries: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSanitizeWhereClause_StrictOff_LegacyBehaviour(t *testing.T) {
|
||||||
|
setHardening(t, hardOff)
|
||||||
|
if got := SanitizeWhereClause("a = 1 and drop table x", "t"); got != "" {
|
||||||
|
t.Errorf("legacy fail-open expected empty, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSortStrict(t *testing.T) {
|
||||||
|
v := NewColumnValidator(TestModel{})
|
||||||
|
opts := RequestOptions{
|
||||||
|
JoinAliases: []string{""},
|
||||||
|
Sort: []SortOption{{Column: "x.id, (select pg_sleep(10))"}, {Column: "(select pg_sleep(1))"}},
|
||||||
|
}
|
||||||
|
setHardening(t, hardOn)
|
||||||
|
if got := v.FilterRequestOptions(opts).Sort; len(got) != 0 {
|
||||||
|
t.Errorf("strict: expected no sorts, got %v", got)
|
||||||
|
}
|
||||||
|
opts.JoinAliases = []string{"j"}
|
||||||
|
opts.Sort = []SortOption{{Column: "j.id"}, {Column: "j.id, (select 1)"}, {Column: "(select max(age) from users)"}}
|
||||||
|
if got := v.FilterRequestOptions(opts).Sort; len(got) != 2 {
|
||||||
|
t.Errorf("strict: join column and plain subquery sort must be allowed, got %v", got)
|
||||||
|
}
|
||||||
|
opts.Sort = []SortOption{{Column: "j.id"}, {Column: "j.id, (select 1)"}}
|
||||||
|
if got := v.FilterRequestOptions(opts).Sort; len(got) != 1 || got[0].Column != "j.id" {
|
||||||
|
t.Errorf("strict: expected only j.id, got %v", got)
|
||||||
|
}
|
||||||
|
setHardening(t, hardOff)
|
||||||
|
opts.JoinAliases = []string{""}
|
||||||
|
opts.Sort = []SortOption{{Column: "x.id, (select pg_sleep(10))"}}
|
||||||
|
if got := v.FilterRequestOptions(opts).Sort; len(got) != 1 {
|
||||||
|
t.Errorf("legacy: expected sort kept, got %v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCQLColumn(t *testing.T) {
|
||||||
|
v := NewColumnValidator(TestModel{})
|
||||||
|
setHardening(t, hardOn)
|
||||||
|
if !v.IsValidColumn("cqlComputed1") || v.IsValidColumn("cql1); drop") {
|
||||||
|
t.Error("strict cql validation wrong")
|
||||||
|
}
|
||||||
|
setHardening(t, hardOff)
|
||||||
|
if !v.IsValidColumn("cql1); drop") {
|
||||||
|
t.Error("legacy cql behaviour should be permissive")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOriginAllowed(t *testing.T) {
|
||||||
|
list := []string{"https://app.example.com/", "http://localhost:8080"}
|
||||||
|
if ok, _ := originAllowed("https://APP.example.com", list); !ok {
|
||||||
|
t.Error("listed origin should match case-insensitively, ignoring trailing slash")
|
||||||
|
}
|
||||||
|
if ok, _ := originAllowed("https://evil.example", list); ok {
|
||||||
|
t.Error("unlisted origin must not match")
|
||||||
|
}
|
||||||
|
if ok, wc := originAllowed("https://evil.example", []string{"*"}); ok || !wc {
|
||||||
|
t.Error("wildcard must be reported separately and never as an exact match")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -319,3 +319,11 @@ type QueryHandler interface {
|
|||||||
SpecHandler
|
SpecHandler
|
||||||
// Methods are defined in funcspec package due to different function signature requirements
|
// Methods are defined in funcspec package due to different function signature requirements
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WhereGrouper is implemented by query builders that can wrap a set of
|
||||||
|
// conditions (including WhereOr) in one parenthesised group that is ANDed with
|
||||||
|
// the rest of the query. It is optional so existing SelectQuery implementations
|
||||||
|
// keep compiling.
|
||||||
|
type WhereGrouper interface {
|
||||||
|
WhereGroup(fn func(SelectQuery) SelectQuery) SelectQuery
|
||||||
|
}
|
||||||
|
|||||||
@@ -148,6 +148,93 @@ func validateWhereClauseSecurity(where string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
reStrictDML = regexp.MustCompile(`(?i)\b(delete|update|truncate|drop|alter|create|insert|grant|revoke|exec|execute|copy|call|do|merge|vacuum|listen|notify|set|returning|into)\b`)
|
||||||
|
// reStrictSubquery is only enforced when hardening.sql_block_subqueries is on.
|
||||||
|
reStrictSubquery = regexp.MustCompile(`(?i)\b(select|union|lateral|with)\b`)
|
||||||
|
// reStrictDangerousFunc matches functions/schemas that enable DoS or data
|
||||||
|
// exfiltration through a WHERE fragment (sleep, file/large-object access,
|
||||||
|
// dblink, config access, catalogs).
|
||||||
|
reStrictDangerousFunc = regexp.MustCompile(`(?i)\b(pg_[a-z0-9_]*|lo_[a-z0-9_]*|dblink[a-z0-9_]*|set_config|current_setting|query_to_xml[a-z_]*|xpath[a-z_]*|generate_series|repeat|crypt|information_schema|sleep|benchmark)\b`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// stripSQLLiterals blanks out single-quoted literals (honouring ”) and
|
||||||
|
// double-quoted identifiers so structural checks only see SQL syntax.
|
||||||
|
// ok is false when a quote is left unterminated.
|
||||||
|
func stripSQLLiterals(s string) (out string, ok bool) {
|
||||||
|
var b strings.Builder
|
||||||
|
for i := 0; i < len(s); i++ {
|
||||||
|
ch := s[i]
|
||||||
|
if ch != '\'' && ch != '"' {
|
||||||
|
b.WriteByte(ch)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
q := ch
|
||||||
|
closed := false
|
||||||
|
for i++; i < len(s); i++ {
|
||||||
|
if s[i] == q {
|
||||||
|
if i+1 < len(s) && s[i+1] == q { // escaped quote
|
||||||
|
i++
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
closed = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !closed {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
b.WriteString("''")
|
||||||
|
}
|
||||||
|
return b.String(), true
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateWhereClauseStrict is the hardened check for client raw-SQL fragments
|
||||||
|
// (hardening.sql_strict). It inspects syntax outside string literals: quotes
|
||||||
|
// and parentheses must be balanced, and comments, statement separators,
|
||||||
|
// dollar-quoting, DML keywords, dangerous functions and system catalogs are
|
||||||
|
// rejected. Subqueries and ordinary functions stay allowed unless
|
||||||
|
// hardening.sql_block_subqueries is set. isJoin (custom joins) still gets all
|
||||||
|
// checks except the subquery block, since joins legitimately use subqueries.
|
||||||
|
func validateWhereClauseStrict(where string, isJoin bool) error {
|
||||||
|
stripped, ok := stripSQLLiterals(where)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("unterminated quote")
|
||||||
|
}
|
||||||
|
for _, bad := range []string{"--", "/*", "*/", ";", "$$", "\\"} {
|
||||||
|
if strings.Contains(stripped, bad) {
|
||||||
|
return fmt.Errorf("forbidden token %q", bad)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
depth := 0
|
||||||
|
for i := 0; i < len(stripped); i++ {
|
||||||
|
switch stripped[i] {
|
||||||
|
case '(':
|
||||||
|
depth++
|
||||||
|
case ')':
|
||||||
|
depth--
|
||||||
|
if depth < 0 {
|
||||||
|
return fmt.Errorf("unbalanced parentheses")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if depth != 0 {
|
||||||
|
return fmt.Errorf("unbalanced parentheses")
|
||||||
|
}
|
||||||
|
if m := reStrictDML.FindString(stripped); m != "" {
|
||||||
|
return fmt.Errorf("forbidden keyword %q", strings.ToLower(m))
|
||||||
|
}
|
||||||
|
if m := reStrictDangerousFunc.FindString(stripped); m != "" {
|
||||||
|
return fmt.Errorf("forbidden function or schema %q", strings.ToLower(m))
|
||||||
|
}
|
||||||
|
if !isJoin && Hardening().SQLBlockSubqueries {
|
||||||
|
if m := reStrictSubquery.FindString(stripped); m != "" {
|
||||||
|
return fmt.Errorf("subqueries not allowed (%q)", strings.ToLower(m))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// SanitizeWhereClause removes trivial conditions and fixes incorrect table prefixes
|
// SanitizeWhereClause removes trivial conditions and fixes incorrect table prefixes
|
||||||
// This function should be used everywhere a WHERE statement is sent to ensure clean, efficient SQL
|
// This function should be used everywhere a WHERE statement is sent to ensure clean, efficient SQL
|
||||||
//
|
//
|
||||||
@@ -174,7 +261,14 @@ func SanitizeWhereClause(where string, tableName string, options ...*RequestOpti
|
|||||||
where = strings.TrimSpace(where)
|
where = strings.TrimSpace(where)
|
||||||
|
|
||||||
// Validate that the WHERE clause doesn't contain dangerous SQL statements
|
// Validate that the WHERE clause doesn't contain dangerous SQL statements
|
||||||
if err := validateWhereClauseSecurity(where); err != nil {
|
if Hardening().SQLStrict {
|
||||||
|
// Strict mode: fail closed. A rejected client fragment must not turn into
|
||||||
|
// "no filter", so substitute a clause that matches no rows.
|
||||||
|
if err := validateWhereClauseStrict(where, tableName == ""); err != nil {
|
||||||
|
logger.Warn("Rejected client SQL fragment (%v): %s", err, where)
|
||||||
|
return "(1=0)"
|
||||||
|
}
|
||||||
|
} else if err := validateWhereClauseSecurity(where); err != nil {
|
||||||
logger.Debug("Security validation failed for WHERE clause: %v", err)
|
logger.Debug("Security validation failed for WHERE clause: %v", err)
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -102,25 +102,25 @@ func TestSanitizeWhereClause(t *testing.T) {
|
|||||||
name: "dangerous DELETE keyword - blocked",
|
name: "dangerous DELETE keyword - blocked",
|
||||||
where: "status = 'active'; DELETE FROM users",
|
where: "status = 'active'; DELETE FROM users",
|
||||||
tableName: "users",
|
tableName: "users",
|
||||||
expected: "",
|
expected: "(1=0)", // fail closed,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "dangerous UPDATE keyword - blocked",
|
name: "dangerous UPDATE keyword - blocked",
|
||||||
where: "1=1; UPDATE users SET admin = true",
|
where: "1=1; UPDATE users SET admin = true",
|
||||||
tableName: "users",
|
tableName: "users",
|
||||||
expected: "",
|
expected: "(1=0)", // fail closed,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "dangerous TRUNCATE keyword - blocked",
|
name: "dangerous TRUNCATE keyword - blocked",
|
||||||
where: "status = 'active' OR TRUNCATE TABLE users",
|
where: "status = 'active' OR TRUNCATE TABLE users",
|
||||||
tableName: "users",
|
tableName: "users",
|
||||||
expected: "",
|
expected: "(1=0)", // fail closed,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "dangerous DROP keyword - blocked",
|
name: "dangerous DROP keyword - blocked",
|
||||||
where: "status = 'active'; DROP TABLE users",
|
where: "status = 'active'; DROP TABLE users",
|
||||||
tableName: "users",
|
tableName: "users",
|
||||||
expected: "",
|
expected: "(1=0)", // fail closed,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "subquery with table alias should not be modified",
|
name: "subquery with table alias should not be modified",
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package common
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"regexp"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -105,7 +106,7 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Allow columns prefixed with "cql" (case insensitive) for computed columns
|
// Allow columns prefixed with "cql" (case insensitive) for computed columns
|
||||||
if strings.HasPrefix(strings.ToLower(column), "cql") {
|
if lc := strings.ToLower(column); strings.HasPrefix(lc, "cql") && (!Hardening().SortStrict || reCQLColumn.MatchString(lc)) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -275,8 +276,14 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
validSorts = append(validSorts, sort)
|
validSorts = append(validSorts, sort)
|
||||||
} else {
|
} else {
|
||||||
foundJoin := false
|
foundJoin := false
|
||||||
|
strictSort := Hardening().SortStrict
|
||||||
for _, j := range options.JoinAliases {
|
for _, j := range options.JoinAliases {
|
||||||
if strings.Contains(sort.Column, j) {
|
if strictSort {
|
||||||
|
if isJoinAliasColumn(sort.Column, j) {
|
||||||
|
foundJoin = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
} else if strings.Contains(sort.Column, j) {
|
||||||
foundJoin = true
|
foundJoin = true
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
@@ -287,7 +294,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
}
|
}
|
||||||
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||||
// Allow sort by expression/subquery, but validate for security
|
// Allow sort by expression/subquery, but validate for security
|
||||||
if IsSafeSortExpression(sort.Column) {
|
if IsSafeSortExpression(sort.Column) && (!strictSort || isSortExpressionRestricted(sort.Column)) {
|
||||||
validSorts = append(validSorts, sort)
|
validSorts = append(validSorts, sort)
|
||||||
} else {
|
} else {
|
||||||
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
||||||
@@ -376,6 +383,56 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
|||||||
return filtered
|
return filtered
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var reCQLColumn = regexp.MustCompile(`^cql[a-z0-9_]*$`)
|
||||||
|
|
||||||
|
var (
|
||||||
|
reJoinColumnIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||||
|
)
|
||||||
|
|
||||||
|
// isJoinAliasColumn reports whether col is exactly "<alias>.<identifier>".
|
||||||
|
// An empty alias never matches.
|
||||||
|
func isJoinAliasColumn(col, alias string) bool {
|
||||||
|
if alias == "" || !strings.HasPrefix(col, alias+".") {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return reJoinColumnIdent.MatchString(col[len(alias)+1:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// isSortExpressionRestricted applies the hardened checks to a client sort
|
||||||
|
// expression. Subqueries and ordinary functions are allowed; dangerous
|
||||||
|
// functions/catalogs (pg_sleep, pg_*, dblink, ...) and unbalanced parentheses are
|
||||||
|
// rejected. Subqueries are blocked only when hardening.sql_block_subqueries is on.
|
||||||
|
func isSortExpressionRestricted(expr string) bool {
|
||||||
|
stripped, ok := stripSQLLiterals(expr)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
depth := 0
|
||||||
|
for i := 0; i < len(stripped); i++ {
|
||||||
|
switch stripped[i] {
|
||||||
|
case '(':
|
||||||
|
depth++
|
||||||
|
case ')':
|
||||||
|
depth--
|
||||||
|
if depth < 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if depth != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if m := reStrictDangerousFunc.FindString(stripped); m != "" {
|
||||||
|
logger.Warn("Forbidden function '%s' in sort expression: %s", m, expr)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if Hardening().SQLBlockSubqueries && reStrictSubquery.MatchString(stripped) {
|
||||||
|
logger.Warn("Subquery in sort expression rejected: %s", expr)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// IsSafeSortExpression validates that a sort expression (enclosed in brackets) is safe
|
// IsSafeSortExpression validates that a sort expression (enclosed in brackets) is safe
|
||||||
// and doesn't contain SQL injection attempts or dangerous commands
|
// and doesn't contain SQL injection attempts or dangerous commands
|
||||||
func IsSafeSortExpression(expr string) bool {
|
func IsSafeSortExpression(expr string) bool {
|
||||||
|
|||||||
@@ -435,19 +435,19 @@ func TestFilterRequestOptions_WithSortExpressions(t *testing.T) {
|
|||||||
|
|
||||||
options := RequestOptions{
|
options := RequestOptions{
|
||||||
Sort: []SortOption{
|
Sort: []SortOption{
|
||||||
{Column: "id", Direction: "ASC"}, // Valid column
|
{Column: "id", Direction: "ASC"}, // Valid column
|
||||||
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
|
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
|
||||||
{Column: "name", Direction: "ASC"}, // Valid column
|
{Column: "name", Direction: "ASC"}, // Valid column
|
||||||
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
|
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
|
||||||
{Column: "invalid_col", Direction: "ASC"}, // Invalid column
|
{Column: "invalid_col", Direction: "ASC"}, // Invalid column
|
||||||
{Column: "(CASE WHEN age > 18 THEN 1 ELSE 0 END)", Direction: "ASC"}, // Safe expression
|
{Column: "(CASE WHEN age > 18 THEN 1 ELSE 0 END)", Direction: "ASC"}, // Safe expression
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
filtered := validator.FilterRequestOptions(options)
|
filtered := validator.FilterRequestOptions(options)
|
||||||
|
|
||||||
// Should keep: id, safe expression, name, another safe expression
|
// Keeps: id, subquery expression, name, CASE expression
|
||||||
// Should remove: dangerous expression, invalid column
|
// Removes: dangerous expression, invalid column
|
||||||
expectedCount := 4
|
expectedCount := 4
|
||||||
if len(filtered.Sort) != expectedCount {
|
if len(filtered.Sort) != expectedCount {
|
||||||
t.Errorf("Expected %d sort options, got %d", expectedCount, len(filtered.Sort))
|
t.Errorf("Expected %d sort options, got %d", expectedCount, len(filtered.Sort))
|
||||||
@@ -474,8 +474,8 @@ type RelatedModel struct {
|
|||||||
// PreloadParentModel has a has-one relation to RelatedModel. The json tag on
|
// PreloadParentModel has a has-one relation to RelatedModel. The json tag on
|
||||||
// the relation field is the name used in x-preload headers.
|
// the relation field is the name used in x-preload headers.
|
||||||
type PreloadParentModel struct {
|
type PreloadParentModel struct {
|
||||||
ID int64 `bun:"id,pk"`
|
ID int64 `bun:"id,pk"`
|
||||||
Name string `bun:"name"`
|
Name string `bun:"name"`
|
||||||
RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"`
|
RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ type Config struct {
|
|||||||
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
EventBroker EventBrokerConfig `mapstructure:"event_broker"`
|
||||||
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
DBManager DBManagerConfig `mapstructure:"dbmanager"`
|
||||||
DBTrace DBTraceConfig `mapstructure:"db_trace"`
|
DBTrace DBTraceConfig `mapstructure:"db_trace"`
|
||||||
|
Hardening HardeningConfig `mapstructure:"hardening"`
|
||||||
Paths PathsConfig `mapstructure:"paths"`
|
Paths PathsConfig `mapstructure:"paths"`
|
||||||
Extensions map[string]interface{} `mapstructure:"extensions"`
|
Extensions map[string]interface{} `mapstructure:"extensions"`
|
||||||
}
|
}
|
||||||
@@ -143,6 +144,25 @@ type CORSConfig struct {
|
|||||||
MaxAge int `mapstructure:"max_age"`
|
MaxAge int `mapstructure:"max_age"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// HardeningConfig toggles security hardening that may reject requests which
|
||||||
|
// older clients relied on. All switches default to true; set one to false to
|
||||||
|
// restore the previous (permissive) behaviour.
|
||||||
|
// Env: RESOLVESPEC_HARDENING_CORS_STRICT_ORIGINS, _SORT_STRICT, _SQL_STRICT, _SQL_BLOCK_SUBQUERIES.
|
||||||
|
type HardeningConfig struct {
|
||||||
|
// CORSStrictOrigins only reflects origins listed in cors.allowed_origins (and the
|
||||||
|
// server URLs); credentials are never sent for unlisted or wildcard origins.
|
||||||
|
CORSStrictOrigins bool `mapstructure:"cors_strict_origins"`
|
||||||
|
// SortStrict stops empty/substring join aliases from admitting arbitrary sort strings.
|
||||||
|
SortStrict bool `mapstructure:"sort_strict"`
|
||||||
|
// SQLStrict hardens client raw-SQL fragments (x-custom-sql-*, preload where, cursor):
|
||||||
|
// balanced parentheses, no subqueries/functions/comments, and rejection instead of
|
||||||
|
// silently dropping the filter.
|
||||||
|
SQLStrict bool `mapstructure:"sql_strict"`
|
||||||
|
// SQLBlockSubqueries additionally rejects subqueries (select/union/with) in client
|
||||||
|
// WHERE fragments (not custom joins). Off by default: existing clients use them.
|
||||||
|
SQLBlockSubqueries bool `mapstructure:"sql_block_subqueries"`
|
||||||
|
}
|
||||||
|
|
||||||
// DBTraceConfig controls database usage logging (off by default).
|
// DBTraceConfig controls database usage logging (off by default).
|
||||||
// Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG.
|
// Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG.
|
||||||
type DBTraceConfig struct {
|
type DBTraceConfig struct {
|
||||||
|
|||||||
@@ -168,6 +168,7 @@ func (m *Manager) SetConfig(cfg *Config) error {
|
|||||||
m.v.Set("event_broker", cfg.EventBroker)
|
m.v.Set("event_broker", cfg.EventBroker)
|
||||||
m.v.Set("dbmanager", cfg.DBManager)
|
m.v.Set("dbmanager", cfg.DBManager)
|
||||||
m.v.Set("db_trace", cfg.DBTrace)
|
m.v.Set("db_trace", cfg.DBTrace)
|
||||||
|
m.v.Set("hardening", cfg.Hardening)
|
||||||
m.v.Set("paths", cfg.Paths)
|
m.v.Set("paths", cfg.Paths)
|
||||||
m.v.Set("extensions", cfg.Extensions)
|
m.v.Set("extensions", cfg.Extensions)
|
||||||
|
|
||||||
@@ -279,6 +280,12 @@ func setDefaults(v *viper.Viper) {
|
|||||||
v.SetDefault("cors.allowed_headers", []string{"*"})
|
v.SetDefault("cors.allowed_headers", []string{"*"})
|
||||||
v.SetDefault("cors.max_age", 3600)
|
v.SetDefault("cors.max_age", 3600)
|
||||||
|
|
||||||
|
// Security hardening defaults (on)
|
||||||
|
v.SetDefault("hardening.cors_strict_origins", true)
|
||||||
|
v.SetDefault("hardening.sort_strict", true)
|
||||||
|
v.SetDefault("hardening.sql_strict", true)
|
||||||
|
v.SetDefault("hardening.sql_block_subqueries", false)
|
||||||
|
|
||||||
// Database defaults
|
// Database defaults
|
||||||
v.SetDefault("database.url", "")
|
v.SetDefault("database.url", "")
|
||||||
|
|
||||||
|
|||||||
+73
-58
@@ -646,77 +646,92 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
// This may need to be handled differently per database adapter
|
// This may need to be handled differently per database adapter
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply filters - validate and adjust for column types first
|
// Client-controlled conditions (filters + x-custom-sql-w) are built by a closure so they can
|
||||||
// Group consecutive OR filters together to prevent OR logic from escaping
|
// be wrapped in a single group together with x-custom-sql-or: the OR then only widens the
|
||||||
for i := 0; i < len(options.Filters); {
|
// client's own conditions and can never escape the server-side filters ANDed around it.
|
||||||
filter := &options.Filters[i]
|
applyUserConds := func(query common.SelectQuery) common.SelectQuery {
|
||||||
|
// Apply filters - validate and adjust for column types first
|
||||||
|
// Group consecutive OR filters together to prevent OR logic from escaping
|
||||||
|
for i := 0; i < len(options.Filters); {
|
||||||
|
filter := &options.Filters[i]
|
||||||
|
|
||||||
// Validate and adjust filter based on column type
|
// Validate and adjust filter based on column type
|
||||||
castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model)
|
castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model)
|
||||||
|
|
||||||
// Default to AND if LogicOperator is not set
|
// Default to AND if LogicOperator is not set
|
||||||
logicOp := filter.LogicOperator
|
logicOp := filter.LogicOperator
|
||||||
if logicOp == "" {
|
if logicOp == "" {
|
||||||
logicOp = "AND"
|
logicOp = "AND"
|
||||||
}
|
|
||||||
|
|
||||||
// Check if this is the start of an OR group
|
|
||||||
if logicOp == "OR" {
|
|
||||||
// Collect all consecutive OR filters
|
|
||||||
orFilters := []*common.FilterOption{filter}
|
|
||||||
orCastInfo := []ColumnCastInfo{castInfo}
|
|
||||||
|
|
||||||
j := i + 1
|
|
||||||
for j < len(options.Filters) {
|
|
||||||
nextFilter := &options.Filters[j]
|
|
||||||
nextLogicOp := nextFilter.LogicOperator
|
|
||||||
if nextLogicOp == "" {
|
|
||||||
nextLogicOp = "AND"
|
|
||||||
}
|
|
||||||
if nextLogicOp == "OR" {
|
|
||||||
nextCastInfo := h.ValidateAndAdjustFilterForColumnType(nextFilter, model)
|
|
||||||
orFilters = append(orFilters, nextFilter)
|
|
||||||
orCastInfo = append(orCastInfo, nextCastInfo)
|
|
||||||
j++
|
|
||||||
} else {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply the OR group as a single grouped condition
|
// Check if this is the start of an OR group
|
||||||
logger.Debug("Applying OR filter group with %d conditions", len(orFilters))
|
if logicOp == "OR" {
|
||||||
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName, model)
|
// Collect all consecutive OR filters
|
||||||
i = j
|
orFilters := []*common.FilterOption{filter}
|
||||||
} else {
|
orCastInfo := []ColumnCastInfo{castInfo}
|
||||||
// Single AND filter - apply normally
|
|
||||||
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp)
|
j := i + 1
|
||||||
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp, model)
|
for j < len(options.Filters) {
|
||||||
i++
|
nextFilter := &options.Filters[j]
|
||||||
|
nextLogicOp := nextFilter.LogicOperator
|
||||||
|
if nextLogicOp == "" {
|
||||||
|
nextLogicOp = "AND"
|
||||||
|
}
|
||||||
|
if nextLogicOp == "OR" {
|
||||||
|
nextCastInfo := h.ValidateAndAdjustFilterForColumnType(nextFilter, model)
|
||||||
|
orFilters = append(orFilters, nextFilter)
|
||||||
|
orCastInfo = append(orCastInfo, nextCastInfo)
|
||||||
|
j++
|
||||||
|
} else {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply the OR group as a single grouped condition
|
||||||
|
logger.Debug("Applying OR filter group with %d conditions", len(orFilters))
|
||||||
|
query = h.applyOrFilterGroup(query, orFilters, orCastInfo, tableName, model)
|
||||||
|
i = j
|
||||||
|
} else {
|
||||||
|
// Single AND filter - apply normally
|
||||||
|
logger.Debug("Applying filter: %s %s %v (needsCast=%v, logic=%s)", filter.Column, filter.Operator, filter.Value, castInfo.NeedsCast, logicOp)
|
||||||
|
query = h.applyFilter(query, *filter, tableName, castInfo.NeedsCast, logicOp, model)
|
||||||
|
i++
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Apply custom SQL WHERE clause (AND condition)
|
||||||
|
if options.CustomSQLWhere != "" {
|
||||||
|
logger.Debug("Applying custom SQL WHERE: %s", options.CustomSQLWhere)
|
||||||
|
// First add table prefixes to unqualified columns (but skip columns inside function calls)
|
||||||
|
prefixedWhere := common.AddTablePrefixToColumns(options.CustomSQLWhere, reflection.ExtractTableNameOnly(tableName))
|
||||||
|
// Then sanitize and allow preload table prefixes since custom SQL may reference multiple tables
|
||||||
|
sanitizedWhere := common.SanitizeWhereClause(prefixedWhere, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
|
||||||
|
// Ensure outer parentheses to prevent OR logic from escaping
|
||||||
|
sanitizedWhere = common.EnsureOuterParentheses(sanitizedWhere)
|
||||||
|
if sanitizedWhere != "" {
|
||||||
|
query = query.Where(sanitizedWhere)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return query
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply custom SQL WHERE clause (AND condition)
|
sanitizedOr := ""
|
||||||
if options.CustomSQLWhere != "" {
|
|
||||||
logger.Debug("Applying custom SQL WHERE: %s", options.CustomSQLWhere)
|
|
||||||
// First add table prefixes to unqualified columns (but skip columns inside function calls)
|
|
||||||
prefixedWhere := common.AddTablePrefixToColumns(options.CustomSQLWhere, reflection.ExtractTableNameOnly(tableName))
|
|
||||||
// Then sanitize and allow preload table prefixes since custom SQL may reference multiple tables
|
|
||||||
sanitizedWhere := common.SanitizeWhereClause(prefixedWhere, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
|
|
||||||
// Ensure outer parentheses to prevent OR logic from escaping
|
|
||||||
sanitizedWhere = common.EnsureOuterParentheses(sanitizedWhere)
|
|
||||||
if sanitizedWhere != "" {
|
|
||||||
query = query.Where(sanitizedWhere)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Apply custom SQL WHERE clause (OR condition)
|
|
||||||
if options.CustomSQLOr != "" {
|
if options.CustomSQLOr != "" {
|
||||||
logger.Debug("Applying custom SQL OR: %s", options.CustomSQLOr)
|
logger.Debug("Applying custom SQL OR: %s", options.CustomSQLOr)
|
||||||
customOr := common.AddTablePrefixToColumns(options.CustomSQLOr, reflection.ExtractTableNameOnly(tableName))
|
customOr := common.AddTablePrefixToColumns(options.CustomSQLOr, reflection.ExtractTableNameOnly(tableName))
|
||||||
// Sanitize and allow preload table prefixes since custom SQL may reference multiple tables
|
// Sanitize and allow preload table prefixes since custom SQL may reference multiple tables
|
||||||
sanitizedOr := common.SanitizeWhereClause(customOr, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
|
sanitizedOr = common.SanitizeWhereClause(customOr, reflection.ExtractTableNameOnly(tableName), &options.RequestOptions)
|
||||||
// Ensure outer parentheses to prevent OR logic from escaping
|
// Ensure outer parentheses to prevent OR logic from escaping
|
||||||
sanitizedOr = common.EnsureOuterParentheses(sanitizedOr)
|
sanitizedOr = common.EnsureOuterParentheses(sanitizedOr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" && common.Hardening().SQLStrict {
|
||||||
|
query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery {
|
||||||
|
return applyUserConds(q).WhereOr(sanitizedOr)
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
query = applyUserConds(query)
|
||||||
if sanitizedOr != "" {
|
if sanitizedOr != "" {
|
||||||
query = query.WhereOr(sanitizedOr)
|
query = query.WhereOr(sanitizedOr)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user