mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 11:31:57 +00:00
feat(pgsql): add WhereGroup and a podman/docker hardening test
- PgSQLSelectQuery implements common.WhereGrouper so x-custom-sql-or is grouped with the client's own conditions on the pgx adapter too. - Add a container test (opt-in via RESOLVESPEC_TEST_CONTAINERS=1) that starts PostgreSQL with podman or docker and checks the hardening against a real database: parenthesis escape, pg_sleep, catalog subquery, stacked statements, x-custom-sql-or grouping and the legacy behaviour when hardening is switched off.
This commit is contained in:
@@ -0,0 +1,230 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/jackc/pgx/v5/stdlib"
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/pgdialect"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||
)
|
||||
|
||||
// These tests start a throwaway PostgreSQL server with podman or docker (whichever is
|
||||
// installed, podman first) and run the client-SQL hardening against a real database. They
|
||||
// pull an image, so they only run when RESOLVESPEC_TEST_CONTAINERS=1 and not with -short.
|
||||
// The container is removed when the test ends.
|
||||
|
||||
const hardeningPGPassword = "Resolve_Spec_1"
|
||||
|
||||
type hardeningItem struct {
|
||||
bun.BaseModel `bun:"table:items,alias:items"`
|
||||
ID int `bun:"id"`
|
||||
Tenant int `bun:"tenant"`
|
||||
Name string `bun:"name"`
|
||||
}
|
||||
|
||||
func hardeningRuntime(t *testing.T) string {
|
||||
t.Helper()
|
||||
if testing.Short() {
|
||||
t.Skip("container tests are skipped with -short")
|
||||
}
|
||||
if os.Getenv("RESOLVESPEC_TEST_CONTAINERS") != "1" {
|
||||
t.Skip("set RESOLVESPEC_TEST_CONTAINERS=1 to run tests that start a podman/docker container")
|
||||
}
|
||||
for _, rt := range []string{"podman", "docker"} {
|
||||
if p, err := exec.LookPath(rt); err == nil {
|
||||
return p
|
||||
}
|
||||
}
|
||||
t.Skip("neither podman nor docker found in PATH")
|
||||
return ""
|
||||
}
|
||||
|
||||
func hardeningRun(t *testing.T, timeout time.Duration, name string, args ...string) string {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
var out, errb bytes.Buffer
|
||||
cmd := exec.CommandContext(ctx, name, args...)
|
||||
cmd.Stdout, cmd.Stderr = &out, &errb
|
||||
if err := cmd.Run(); err != nil {
|
||||
t.Fatalf("%s %s: %v\n%s", name, strings.Join(args, " "), err, errb.String())
|
||||
}
|
||||
return strings.TrimSpace(out.String())
|
||||
}
|
||||
|
||||
// startHardeningPostgres runs postgres on a random localhost port and returns a ready *sql.DB.
|
||||
func startHardeningPostgres(t *testing.T, rt string) *sql.DB {
|
||||
t.Helper()
|
||||
id := hardeningRun(t, 10*time.Minute, rt, "run", "-d", "--rm", "-p", "127.0.0.1::5432",
|
||||
"-e", "POSTGRES_PASSWORD="+hardeningPGPassword, "docker.io/library/postgres:16-alpine") // first run may pull
|
||||
t.Cleanup(func() { _ = exec.Command(rt, "rm", "-f", id).Run() })
|
||||
|
||||
out := hardeningRun(t, 30*time.Second, rt, "port", id, "5432")
|
||||
line := strings.Fields(out)[len(strings.Fields(out))-1]
|
||||
for _, l := range strings.Split(out, "\n") {
|
||||
if strings.Contains(l, "127.0.0.1:") {
|
||||
line = l[strings.LastIndex(l, " ")+1:]
|
||||
break
|
||||
}
|
||||
}
|
||||
_, port, err := net.SplitHostPort(line)
|
||||
if err != nil {
|
||||
t.Fatalf("cannot parse published port %q: %v", out, err)
|
||||
}
|
||||
dsn := fmt.Sprintf("postgres://postgres:%s@127.0.0.1:%s/postgres?sslmode=disable", hardeningPGPassword, port)
|
||||
|
||||
// The official image restarts once during init: wait, pause, wait again.
|
||||
wait := func(d time.Duration) *sql.DB {
|
||||
deadline := time.Now().Add(d)
|
||||
var last error
|
||||
for time.Now().Before(deadline) {
|
||||
db, err := sql.Open("pgx", dsn)
|
||||
if err == nil {
|
||||
if last = db.Ping(); last == nil {
|
||||
return db
|
||||
}
|
||||
_ = db.Close()
|
||||
} else {
|
||||
last = err
|
||||
}
|
||||
time.Sleep(time.Second)
|
||||
}
|
||||
t.Fatalf("postgres not ready within %s: %v", d, last)
|
||||
return nil
|
||||
}
|
||||
_ = wait(90 * time.Second).Close()
|
||||
time.Sleep(2 * time.Second)
|
||||
db := wait(60 * time.Second)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
return db
|
||||
}
|
||||
|
||||
// setHardeningConfig overrides the global hardening switches for the test.
|
||||
func setHardeningConfig(t *testing.T, h config.HardeningConfig) {
|
||||
t.Helper()
|
||||
m := config.GetConfigManager()
|
||||
cfg, err := m.GetConfig()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
old := cfg.Hardening
|
||||
cfg.Hardening = h
|
||||
if err := m.SetConfig(cfg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
cfg.Hardening = old
|
||||
_ = m.SetConfig(cfg)
|
||||
})
|
||||
}
|
||||
|
||||
// clientWhere mirrors the restheadspec handler pipeline for x-custom-sql-w.
|
||||
func clientWhere(raw string) string {
|
||||
w := common.AddTablePrefixToColumns(raw, "items")
|
||||
w = common.SanitizeWhereClause(w, "items")
|
||||
return common.EnsureOuterParentheses(w)
|
||||
}
|
||||
|
||||
func TestHardeningAgainstPostgresContainer(t *testing.T) {
|
||||
rt := hardeningRuntime(t)
|
||||
sqldb := startHardeningPostgres(t, rt)
|
||||
if _, err := sqldb.Exec(`
|
||||
CREATE TABLE items (id int PRIMARY KEY, tenant int, name text);
|
||||
INSERT INTO items VALUES (1,5,'mine-a'),(2,5,'mine-b'),(3,6,'other-a'),(4,6,'awaiting update approval');`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bdb := bun.NewDB(sqldb, pgdialect.New())
|
||||
adapter := NewBunAdapter(bdb)
|
||||
ctx := context.Background()
|
||||
|
||||
// list runs the handler-shaped query: client x-custom-sql-w, then the server tenant filter.
|
||||
list := func(t *testing.T, where string) []hardeningItem {
|
||||
t.Helper()
|
||||
var rows []hardeningItem
|
||||
q := adapter.NewSelect().Model(&rows)
|
||||
if w := clientWhere(where); w != "" {
|
||||
q = q.Where(w)
|
||||
}
|
||||
q = q.Where("items.tenant = ?", 5)
|
||||
if err := q.Scan(ctx, &rows); err != nil {
|
||||
t.Fatalf("query failed for %q: %v", where, err)
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
t.Run("strict", func(t *testing.T) {
|
||||
setHardeningConfig(t, config.HardeningConfig{CORSStrictOrigins: true, SortStrict: true, SQLStrict: true})
|
||||
|
||||
t.Run("legitimate filters keep working", func(t *testing.T) {
|
||||
if got := list(t, "name = 'mine-a'"); len(got) != 1 || got[0].ID != 1 {
|
||||
t.Errorf("simple filter: %v", got)
|
||||
}
|
||||
if got := list(t, "id in (select id from items where tenant = 5)"); len(got) != 2 {
|
||||
t.Errorf("subquery filter: %v", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("parenthesis escape cannot leave the tenant", func(t *testing.T) {
|
||||
if got := list(t, "1=1)) OR ((1=1"); len(got) != 0 {
|
||||
t.Errorf("escape not rejected closed, got rows: %v", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("hostile fragments fail closed", func(t *testing.T) {
|
||||
for _, w := range []string{
|
||||
"id = 1 and pg_sleep(10) is not null",
|
||||
"id = 1 or (select count(*) from pg_shadow) > 0",
|
||||
"id = 1; delete/**/from items",
|
||||
} {
|
||||
start := time.Now()
|
||||
if got := list(t, w); len(got) != 0 {
|
||||
t.Errorf("%q returned rows: %v", w, got)
|
||||
}
|
||||
if time.Since(start) > 5*time.Second {
|
||||
t.Errorf("%q was executed (took %s)", w, time.Since(start))
|
||||
}
|
||||
}
|
||||
var n int
|
||||
if err := sqldb.QueryRow("SELECT count(*) FROM items").Scan(&n); err != nil || n != 4 {
|
||||
t.Errorf("items table modified: count=%d err=%v", n, err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("x-custom-sql-or stays inside the tenant", func(t *testing.T) {
|
||||
var rows []hardeningItem
|
||||
q := adapter.NewSelect().Model(&rows)
|
||||
orClause := common.EnsureOuterParentheses(common.SanitizeWhereClause("items.name = 'other-a'", "items"))
|
||||
q = q.(common.WhereGrouper).WhereGroup(func(g common.SelectQuery) common.SelectQuery {
|
||||
return g.Where("items.name = ?", "mine-a").WhereOr(orClause)
|
||||
})
|
||||
q = q.Where("items.tenant = ?", 5)
|
||||
if err := q.Scan(ctx, &rows); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(rows) != 1 || rows[0].ID != 1 {
|
||||
t.Errorf("OR clause leaked outside tenant filter: %v", rows)
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("switch off restores legacy behaviour", func(t *testing.T) {
|
||||
setHardeningConfig(t, config.HardeningConfig{})
|
||||
// Proves the strict assertions above are meaningful: without hardening the same
|
||||
// escape returns the other tenant's rows.
|
||||
if got := list(t, "1=1)) OR ((1=1"); len(got) < 3 {
|
||||
t.Errorf("expected the legacy escape to leak rows, got %v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -325,6 +325,32 @@ func (p *PgSQLSelectQuery) WhereOr(query string, args ...interface{}) common.Sel
|
||||
return p
|
||||
}
|
||||
|
||||
// WhereGroup wraps the conditions added by fn (Where = AND, WhereOr = OR) in one
|
||||
// parenthesised group ANDed with the rest of the query. Inside the group the
|
||||
// semantics match Bun: `w1 AND w2 OR o1 OR o2`.
|
||||
func (p *PgSQLSelectQuery) WhereGroup(fn func(common.SelectQuery) common.SelectQuery) common.SelectQuery {
|
||||
sub := &PgSQLSelectQuery{driverName: p.driverName, paramCounter: p.paramCounter, args: make([]interface{}, 0)}
|
||||
res, ok := fn(sub).(*PgSQLSelectQuery)
|
||||
if !ok {
|
||||
res = sub
|
||||
}
|
||||
var group string
|
||||
switch {
|
||||
case len(res.whereClauses) > 0 && len(res.orClauses) > 0:
|
||||
group = "(" + strings.Join(res.whereClauses, " AND ") + ") OR " + strings.Join(res.orClauses, " OR ")
|
||||
case len(res.whereClauses) > 0:
|
||||
group = strings.Join(res.whereClauses, " AND ")
|
||||
case len(res.orClauses) > 0:
|
||||
group = strings.Join(res.orClauses, " OR ")
|
||||
default:
|
||||
return p
|
||||
}
|
||||
p.whereClauses = append(p.whereClauses, "("+group+")")
|
||||
p.args = append(p.args, res.args...)
|
||||
p.paramCounter = res.paramCounter
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLSelectQuery) Join(query string, args ...interface{}) common.SelectQuery {
|
||||
query = p.replacePlaceholders(query, len(args))
|
||||
p.joins = append(p.joins, "JOIN "+query)
|
||||
|
||||
@@ -29,3 +29,22 @@ func TestBunWhereGroupConfinesOr(t *testing.T) {
|
||||
t.Fatalf("unexpected SQL:\n got: %s\nwant to contain: %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPgSQLWhereGroupConfinesOr(t *testing.T) {
|
||||
var q common.SelectQuery = &PgSQLSelectQuery{driverName: "postgres", tableName: "items", columns: []string{"*"}, args: []interface{}{}}
|
||||
q = q.Where("a = ?", 1)
|
||||
q = q.(common.WhereGrouper).WhereGroup(func(g common.SelectQuery) common.SelectQuery {
|
||||
return g.Where("b = ?", 2).WhereOr("(c = 3)")
|
||||
})
|
||||
q = q.Where("tenant = ?", 5)
|
||||
|
||||
pq := q.(*PgSQLSelectQuery)
|
||||
got := pq.buildSQL()
|
||||
want := `WHERE (a = $1 AND ((b = $2) OR (c = 3)) AND tenant = $3)`
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("unexpected SQL:\n got: %s\nwant to contain: %s", got, want)
|
||||
}
|
||||
if len(pq.args) != 3 || pq.args[0] != 1 || pq.args[1] != 2 || pq.args[2] != 5 {
|
||||
t.Fatalf("args out of order: %v", pq.args)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user