From ca89cb8a73870db1d4066b1963d638cad9f08ee8 Mon Sep 17 00:00:00 2001 From: Hein Date: Thu, 1 Oct 2026 13:46:01 +0200 Subject: [PATCH] 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. --- pkg/common/adapters/database/bun.go | 13 ++ pkg/common/adapters/database/gorm.go | 10 ++ .../adapters/database/where_group_test.go | 31 +++++ pkg/common/cors.go | 97 +++++++++++-- pkg/common/hardening.go | 16 +++ pkg/common/hardening_test.go | 121 ++++++++++++++++ pkg/common/interfaces.go | 8 ++ pkg/common/sql_helpers.go | 96 ++++++++++++- pkg/common/sql_helpers_test.go | 8 +- pkg/common/validation.go | 63 ++++++++- pkg/common/validation_test.go | 18 +-- pkg/config/config.go | 20 +++ pkg/config/manager.go | 7 + pkg/restheadspec/handler.go | 131 ++++++++++-------- 14 files changed, 552 insertions(+), 87 deletions(-) create mode 100644 pkg/common/adapters/database/where_group_test.go create mode 100644 pkg/common/hardening.go create mode 100644 pkg/common/hardening_test.go diff --git a/pkg/common/adapters/database/bun.go b/pkg/common/adapters/database/bun.go index 7d87399..043a81b 100644 --- a/pkg/common/adapters/database/bun.go +++ b/pkg/common/adapters/database/bun.go @@ -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 diff --git a/pkg/common/adapters/database/gorm.go b/pkg/common/adapters/database/gorm.go index 46bcb35..2fd2975 100644 --- a/pkg/common/adapters/database/gorm.go +++ b/pkg/common/adapters/database/gorm.go @@ -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 diff --git a/pkg/common/adapters/database/where_group_test.go b/pkg/common/adapters/database/where_group_test.go new file mode 100644 index 0000000..5682ef9 --- /dev/null +++ b/pkg/common/adapters/database/where_group_test.go @@ -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) + } +} diff --git a/pkg/common/cors.go b/pkg/common/cors.go index 58b7e1b..5bbc24a 100644 --- a/pkg/common/cors.go +++ b/pkg/common/cors.go @@ -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, ", ")) } diff --git a/pkg/common/hardening.go b/pkg/common/hardening.go new file mode 100644 index 0000000..62e2488 --- /dev/null +++ b/pkg/common/hardening.go @@ -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() } diff --git a/pkg/common/hardening_test.go b/pkg/common/hardening_test.go new file mode 100644 index 0000000..d9f67e9 --- /dev/null +++ b/pkg/common/hardening_test.go @@ -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") + } +} diff --git a/pkg/common/interfaces.go b/pkg/common/interfaces.go index 00e67e7..9729134 100644 --- a/pkg/common/interfaces.go +++ b/pkg/common/interfaces.go @@ -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 +} diff --git a/pkg/common/sql_helpers.go b/pkg/common/sql_helpers.go index 906085c..ea2d90d 100644 --- a/pkg/common/sql_helpers.go +++ b/pkg/common/sql_helpers.go @@ -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 "" } diff --git a/pkg/common/sql_helpers_test.go b/pkg/common/sql_helpers_test.go index 8c73649..e1f731d 100644 --- a/pkg/common/sql_helpers_test.go +++ b/pkg/common/sql_helpers_test.go @@ -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", diff --git a/pkg/common/validation.go b/pkg/common/validation.go index 4ef9d26..ca03e77 100644 --- a/pkg/common/validation.go +++ b/pkg/common/validation.go @@ -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 ".". +// 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 { diff --git a/pkg/common/validation_test.go b/pkg/common/validation_test.go index c07b71b..b4e8446 100644 --- a/pkg/common/validation_test.go +++ b/pkg/common/validation_test.go @@ -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"` } diff --git a/pkg/config/config.go b/pkg/config/config.go index 9bf5338..6856c4f 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -17,6 +17,7 @@ type Config struct { EventBroker EventBrokerConfig `mapstructure:"event_broker"` DBManager DBManagerConfig `mapstructure:"dbmanager"` DBTrace DBTraceConfig `mapstructure:"db_trace"` + Hardening HardeningConfig `mapstructure:"hardening"` Paths PathsConfig `mapstructure:"paths"` Extensions map[string]interface{} `mapstructure:"extensions"` } @@ -143,6 +144,25 @@ type CORSConfig struct { 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). // Env: RESOLVESPEC_DB_TRACE_ENABLED, _MIN_CALLS, _MIN_DURATION, _POOL_LOG. type DBTraceConfig struct { diff --git a/pkg/config/manager.go b/pkg/config/manager.go index d66fedc..b3f6d65 100644 --- a/pkg/config/manager.go +++ b/pkg/config/manager.go @@ -168,6 +168,7 @@ func (m *Manager) SetConfig(cfg *Config) error { m.v.Set("event_broker", cfg.EventBroker) m.v.Set("dbmanager", cfg.DBManager) m.v.Set("db_trace", cfg.DBTrace) + m.v.Set("hardening", cfg.Hardening) m.v.Set("paths", cfg.Paths) m.v.Set("extensions", cfg.Extensions) @@ -279,6 +280,12 @@ func setDefaults(v *viper.Viper) { v.SetDefault("cors.allowed_headers", []string{"*"}) 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 v.SetDefault("database.url", "") diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 3600e3d..a67dbb8 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -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 } - // 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] + // Client-controlled conditions (filters + x-custom-sql-w) are built by a closure so they can + // be wrapped in a single group together with x-custom-sql-or: the OR then only widens the + // client's own conditions and can never escape the server-side filters ANDed around it. + 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 - castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model) + // Validate and adjust filter based on column type + castInfo := h.ValidateAndAdjustFilterForColumnType(filter, model) - // Default to AND if LogicOperator is not set - logicOp := filter.LogicOperator - if logicOp == "" { - 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 - } + // Default to AND if LogicOperator is not set + logicOp := filter.LogicOperator + if logicOp == "" { + logicOp = "AND" } - // 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++ + // 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 + 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) - 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) + sanitizedOr := "" if options.CustomSQLOr != "" { logger.Debug("Applying custom SQL OR: %s", options.CustomSQLOr) customOr := common.AddTablePrefixToColumns(options.CustomSQLOr, reflection.ExtractTableNameOnly(tableName)) // 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 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 != "" { query = query.WhereOr(sanitizedOr) }