mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 19:41:57 +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
|
||||
}
|
||||
|
||||
// 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 {
|
||||
// Extract optional prefix from args
|
||||
// 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
|
||||
}
|
||||
|
||||
// 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 {
|
||||
// Extract optional prefix from args
|
||||
// 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()
|
||||
cfg, _ := configManager.GetConfig()
|
||||
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()
|
||||
|
||||
@@ -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) {
|
||||
// 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")
|
||||
if origin == "" {
|
||||
origin = "*"
|
||||
// Not a cross-origin browser request; nothing to protect.
|
||||
w.SetHeader("Access-Control-Allow-Origin", "*")
|
||||
} else {
|
||||
// Vary must be set so caches don't serve one origin's response to another
|
||||
httpW := w.UnderlyingResponseWriter()
|
||||
httpW.Header().Set("Vary", "Origin")
|
||||
w.UnderlyingResponseWriter().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
|
||||
if len(config.AllowedMethods) > 0 {
|
||||
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")
|
||||
if 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))
|
||||
}
|
||||
|
||||
// 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 != "*" {
|
||||
w.SetHeader("Access-Control-Allow-Credentials", "true")
|
||||
}
|
||||
|
||||
// Expose headers that clients can read
|
||||
exposeHeaders := config.AllowedHeaders
|
||||
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, ", "))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
// 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
|
||||
}
|
||||
|
||||
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
|
||||
// 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)
|
||||
|
||||
// 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)
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -102,25 +102,25 @@ func TestSanitizeWhereClause(t *testing.T) {
|
||||
name: "dangerous DELETE keyword - blocked",
|
||||
where: "status = 'active'; DELETE FROM users",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "dangerous UPDATE keyword - blocked",
|
||||
where: "1=1; UPDATE users SET admin = true",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "dangerous TRUNCATE keyword - blocked",
|
||||
where: "status = 'active' OR TRUNCATE TABLE users",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "dangerous DROP keyword - blocked",
|
||||
where: "status = 'active'; DROP TABLE users",
|
||||
tableName: "users",
|
||||
expected: "",
|
||||
expected: "(1=0)", // fail closed,
|
||||
},
|
||||
{
|
||||
name: "subquery with table alias should not be modified",
|
||||
|
||||
@@ -3,6 +3,7 @@ package common
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
@@ -105,7 +106,7 @@ func (v *ColumnValidator) ValidateColumn(column string) error {
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
@@ -275,8 +276,14 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
validSorts = append(validSorts, sort)
|
||||
} else {
|
||||
foundJoin := false
|
||||
strictSort := Hardening().SortStrict
|
||||
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
|
||||
break
|
||||
}
|
||||
@@ -287,7 +294,7 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
}
|
||||
if strings.HasPrefix(sort.Column, "(") && strings.HasSuffix(sort.Column, ")") {
|
||||
// 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)
|
||||
} else {
|
||||
logger.Warn("Unsafe sort expression '%s' removed", sort.Column)
|
||||
@@ -376,6 +383,56 @@ func (v *ColumnValidator) FilterRequestOptions(options RequestOptions) RequestOp
|
||||
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
|
||||
// and doesn't contain SQL injection attempts or dangerous commands
|
||||
func IsSafeSortExpression(expr string) bool {
|
||||
|
||||
@@ -435,19 +435,19 @@ func TestFilterRequestOptions_WithSortExpressions(t *testing.T) {
|
||||
|
||||
options := RequestOptions{
|
||||
Sort: []SortOption{
|
||||
{Column: "id", Direction: "ASC"}, // Valid column
|
||||
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
|
||||
{Column: "name", Direction: "ASC"}, // Valid column
|
||||
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
|
||||
{Column: "invalid_col", Direction: "ASC"}, // Invalid column
|
||||
{Column: "id", Direction: "ASC"}, // Valid column
|
||||
{Column: "(SELECT MAX(age) FROM users)", Direction: "DESC"}, // Safe expression
|
||||
{Column: "name", Direction: "ASC"}, // Valid column
|
||||
{Column: "(id); DROP TABLE users; --", Direction: "DESC"}, // Dangerous expression
|
||||
{Column: "invalid_col", Direction: "ASC"}, // Invalid column
|
||||
{Column: "(CASE WHEN age > 18 THEN 1 ELSE 0 END)", Direction: "ASC"}, // Safe expression
|
||||
},
|
||||
}
|
||||
|
||||
filtered := validator.FilterRequestOptions(options)
|
||||
|
||||
// Should keep: id, safe expression, name, another safe expression
|
||||
// Should remove: dangerous expression, invalid column
|
||||
// Keeps: id, subquery expression, name, CASE expression
|
||||
// Removes: dangerous expression, invalid column
|
||||
expectedCount := 4
|
||||
if len(filtered.Sort) != expectedCount {
|
||||
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
|
||||
// the relation field is the name used in x-preload headers.
|
||||
type PreloadParentModel struct {
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
RELATED *RelatedModel `json:"RELATED" bun:"rel:has-one,join:id=related_id"`
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user