fix(sec): escaping datastore user and schema in sql statements (#1239)

Signed-off-by: Dario Tranchitella <dario@tranchitella.eu>
This commit is contained in:
Dario Tranchitella
2026-07-15 17:56:36 +02:00
committed by GitHub
parent 4232a9df7c
commit 6a9f3e10ae
5 changed files with 291 additions and 42 deletions
+6 -4
View File
@@ -25,10 +25,12 @@ func NewStorageConnection(ctx context.Context, client client.Client, ds kamajiv1
cc.TLSConfig.ServerName = cc.Endpoints[0].Host
}
cc.Parameters = map[string][]string{
"multiStatements": {"true"},
}
// NOTE: multiStatements is intentionally NOT enabled here. Only the dump
// import performed during Migrate needs to execute a batch of statements
// in a single call, and it opens its own scoped connection for that.
// Keeping the primary connection single-statement prevents any SQL
// injection on the interpolated DDL statements from escalating into
// stacked queries.
return NewMySQLConnection(*cc)
case kamajiv1alpha1.KinePostgreSQLDriver:
if ds.Spec.TLSConfig != nil {
+107 -21
View File
@@ -9,6 +9,7 @@ import (
"fmt"
"net/url"
"os"
"strings"
"time"
"github.com/JamesStewy/go-mysqldump"
@@ -24,20 +25,24 @@ const (
)
const (
// Identifiers (database and user names) cannot be passed as bind parameters,
// so the `%s` verbs below must only ever be fed values run through
// quoteMySQLIdentifier; the password literal must be fed escapeMySQLString.
mysqlFetchUserStatement = "SELECT User FROM mysql.user WHERE User= ? LIMIT 1"
mysqlFetchDBStatement = "SELECT SCHEMA_NAME FROM INFORMATION_SCHEMA.SCHEMATA WHERE SCHEMA_NAME=? LIMIT 1"
mysqlShowGrantsStatement = "SHOW GRANTS FOR `%s`@`%%`"
mysqlCreateDBStatement = "CREATE DATABASE IF NOT EXISTS `%s`"
mysqlCreateUserStatement = "CREATE USER `%s`@`%%` IDENTIFIED BY '%s'"
mysqlUpdateUserStatement = "ALTER USER `%s`@`%%` IDENTIFIED BY '%s'"
mysqlGrantPrivilegesStatement = "GRANT SELECT, INSERT, UPDATE, DELETE, CREATE, ALTER, INDEX ON `%s`.* TO `%s`@`%%`"
mysqlDropDBStatement = "DROP DATABASE IF EXISTS `%s`"
mysqlDropUserStatement = "DROP USER IF EXISTS `%s`"
mysqlRevokePrivilegesStatement = "REVOKE ALL PRIVILEGES ON `%s`.* FROM `%s`"
mysqlShowGrantsStatement = "SHOW GRANTS FOR %s@`%%`"
mysqlCreateDBStatement = "CREATE DATABASE IF NOT EXISTS %s"
mysqlCreateUserStatement = "CREATE USER %s@`%%` IDENTIFIED BY '%s'"
mysqlUpdateUserStatement = "ALTER USER %s@`%%` IDENTIFIED BY '%s'"
mysqlGrantPrivilegesStatement = "GRANT SELECT, INSERT, UPDATE, DELETE, CREATE, ALTER, INDEX ON %s.* TO %s@`%%`"
mysqlDropDBStatement = "DROP DATABASE IF EXISTS %s"
mysqlDropUserStatement = "DROP USER IF EXISTS %s"
mysqlRevokePrivilegesStatement = "REVOKE ALL PRIVILEGES ON %s.* FROM %s"
)
type MySQLConnection struct {
db *sql.DB
config *mysql.Config
connector ConnectionEndpoint
}
@@ -59,7 +64,7 @@ func (c *MySQLConnection) Migrate(ctx context.Context, tcp kamajiv1alpha1.Tenant
}
defer os.RemoveAll(dir)
if _, err = c.db.ExecContext(ctx, fmt.Sprintf("USE %s", tcp.Status.Storage.Setup.Schema)); err != nil {
if _, err = c.db.ExecContext(ctx, fmt.Sprintf("USE %s", quoteMySQLIdentifier(tcp.Status.Storage.Setup.Schema))); err != nil {
return fmt.Errorf("unable to switch DB for MySQL migration: %w", err)
}
@@ -81,11 +86,21 @@ func (c *MySQLConnection) Migrate(ctx context.Context, tcp kamajiv1alpha1.Tenant
// Executing the import to the target datastore
targetClient := target.(*MySQLConnection) //nolint:forcetypeassert
if _, err = targetClient.db.ExecContext(ctx, fmt.Sprintf("USE %s_%s", tcp.GetNamespace(), tcp.GetName())); err != nil {
// The dump is a batch of semicolon-separated statements, so it must run over
// a connection with multiStatements enabled. That connection is scoped to
// this import alone and closed right after, keeping the primary connection
// single-statement.
importDB, err := targetClient.multiStatementConn()
if err != nil {
return fmt.Errorf("unable to open MySQL multi-statement connection for migration: %w", err)
}
defer importDB.Close()
if _, err = importDB.ExecContext(ctx, fmt.Sprintf("USE %s", quoteMySQLIdentifier(tcp.Status.Storage.Setup.Schema))); err != nil {
return fmt.Errorf("unable to switch DB for MySQL migration: %w", err)
}
if _, err = targetClient.db.ExecContext(ctx, string(statements)); err != nil {
if _, err = importDB.ExecContext(ctx, string(statements)); err != nil {
return fmt.Errorf("cannot execute dump statements for MySQL: %w", err)
}
@@ -128,7 +143,31 @@ func NewMySQLConnection(config ConnectionConfig) (Connection, error) {
return nil, err
}
return &MySQLConnection{db: db, connector: config.Endpoints[0]}, nil
return &MySQLConnection{db: db, config: mysqlConfig, connector: config.Endpoints[0]}, nil
}
// multiStatementConn opens a dedicated connection whose DSN carries the
// multiStatements driver parameter, required to execute a mysqldump batch
// (many semicolon-separated statements) in a single ExecContext call. It is
// deliberately kept out of the primary connection so that statements
// interpolating tenant-controlled identifiers never run over a connection that
// permits stacked queries. The pool is capped at a single connection so the
// USE statement and the subsequent import share the same session.
func (c *MySQLConnection) multiStatementConn() (*sql.DB, error) {
cfg := c.config.Clone()
if cfg.Params == nil {
cfg.Params = map[string]string{}
}
cfg.Params["multiStatements"] = "true"
db, err := sql.Open("mysql", cfg.FormatDSN())
if err != nil {
return nil, err
}
db.SetMaxOpenConns(1)
return db, nil
}
func (c *MySQLConnection) GetConnectionString() string {
@@ -152,7 +191,7 @@ func (c *MySQLConnection) Check(ctx context.Context) error {
}
func (c *MySQLConnection) CreateUser(ctx context.Context, user, password string) error {
if err := c.mutate(ctx, mysqlCreateUserStatement, user, password); err != nil {
if err := c.mutate(ctx, mysqlCreateUserStatement, quoteMySQLIdentifier(user), escapeMySQLString(password)); err != nil {
return errors.NewCreateUserError(err)
}
@@ -160,7 +199,7 @@ func (c *MySQLConnection) CreateUser(ctx context.Context, user, password string)
}
func (c *MySQLConnection) UpdateUser(ctx context.Context, user, password string) error {
if err := c.mutate(ctx, mysqlUpdateUserStatement, user, password); err != nil {
if err := c.mutate(ctx, mysqlUpdateUserStatement, quoteMySQLIdentifier(user), escapeMySQLString(password)); err != nil {
return errors.NewUpdateUserError(err)
}
@@ -168,7 +207,7 @@ func (c *MySQLConnection) UpdateUser(ctx context.Context, user, password string)
}
func (c *MySQLConnection) CreateDB(ctx context.Context, dbName string) error {
if err := c.mutate(ctx, mysqlCreateDBStatement, dbName); err != nil {
if err := c.mutate(ctx, mysqlCreateDBStatement, quoteMySQLIdentifier(dbName)); err != nil {
return errors.NewCreateDBError(err)
}
@@ -176,7 +215,7 @@ func (c *MySQLConnection) CreateDB(ctx context.Context, dbName string) error {
}
func (c *MySQLConnection) GrantPrivileges(ctx context.Context, user, dbName string) error {
if err := c.mutate(ctx, mysqlGrantPrivilegesStatement, dbName, user); err != nil {
if err := c.mutate(ctx, mysqlGrantPrivilegesStatement, quoteMySQLIdentifier(dbName), quoteMySQLIdentifier(user)); err != nil {
return errors.NewGrantPrivilegesError(err)
}
@@ -228,7 +267,7 @@ func (c *MySQLConnection) DBExists(ctx context.Context, dbName string) (bool, er
}
func (c *MySQLConnection) GrantPrivilegesExists(_ context.Context, user, dbName string) (bool, error) {
statementShowGrantsStatement := fmt.Sprintf(mysqlShowGrantsStatement, user)
statementShowGrantsStatement := fmt.Sprintf(mysqlShowGrantsStatement, quoteMySQLIdentifier(user))
rows, err := c.db.Query(statementShowGrantsStatement) //nolint:sqlclosecheck
if err != nil {
return false, errors.NewGrantPrivilegesError(err)
@@ -238,7 +277,7 @@ func (c *MySQLConnection) GrantPrivilegesExists(_ context.Context, user, dbName
return false, errors.NewGrantPrivilegesError(err)
}
expected := fmt.Sprintf(mysqlGrantPrivilegesStatement, dbName, user)
expected := fmt.Sprintf(mysqlGrantPrivilegesStatement, quoteMySQLIdentifier(dbName), quoteMySQLIdentifier(user))
var grant string
for rows.Next() {
@@ -255,7 +294,7 @@ func (c *MySQLConnection) GrantPrivilegesExists(_ context.Context, user, dbName
}
func (c *MySQLConnection) DeleteUser(ctx context.Context, user string) error {
if err := c.mutate(ctx, mysqlDropUserStatement, user); err != nil {
if err := c.mutate(ctx, mysqlDropUserStatement, quoteMySQLIdentifier(user)); err != nil {
return errors.NewDeleteUserError(err)
}
@@ -263,7 +302,7 @@ func (c *MySQLConnection) DeleteUser(ctx context.Context, user string) error {
}
func (c *MySQLConnection) DeleteDB(ctx context.Context, dbName string) error {
if err := c.mutate(ctx, mysqlDropDBStatement, dbName); err != nil {
if err := c.mutate(ctx, mysqlDropDBStatement, quoteMySQLIdentifier(dbName)); err != nil {
return errors.NewCannotDeleteDatabaseError(err)
}
@@ -271,7 +310,7 @@ func (c *MySQLConnection) DeleteDB(ctx context.Context, dbName string) error {
}
func (c *MySQLConnection) RevokePrivileges(ctx context.Context, user, dbName string) error {
if err := c.mutate(ctx, mysqlRevokePrivilegesStatement, dbName, user); err != nil {
if err := c.mutate(ctx, mysqlRevokePrivilegesStatement, quoteMySQLIdentifier(dbName), quoteMySQLIdentifier(user)); err != nil {
return errors.NewRevokePrivilegesError(err)
}
@@ -302,3 +341,50 @@ func (c *MySQLConnection) mutate(ctx context.Context, nonFilledStatement string,
func (c *MySQLConnection) checkEmptyQueryResult(err error) bool {
return err.Error() == sqlErrorNoRows
}
// quoteMySQLIdentifier safely quotes a MySQL identifier (such as a database or
// user name) so it can be interpolated into a statement: it wraps the value in
// backticks and doubles any embedded backtick, neutralising attempts to break
// out of the identifier. NUL bytes, which are illegal in identifiers, are
// stripped. Identifiers cannot be supplied as bind parameters, hence the manual
// quoting.
func quoteMySQLIdentifier(identifier string) string {
identifier = strings.ReplaceAll(identifier, "\x00", "")
return "`" + strings.ReplaceAll(identifier, "`", "``") + "`"
}
// escapeMySQLString escapes a value for safe embedding inside a single-quoted
// MySQL string literal. Both ways of breaking out of such a literal are closed
// in a manner that holds under every sql_mode: single quotes are doubled (” is
// a literal quote regardless of NO_BACKSLASH_ESCAPES) and backslashes are
// doubled (so a trailing backslash cannot escape the closing quote when
// backslash escaping is enabled). The remaining control-character escapes are
// conveniences for the default sql_mode. DDL statements such as CREATE USER
// cannot bind the password as a parameter, hence the manual escaping.
func escapeMySQLString(value string) string {
var b strings.Builder
for _, r := range value {
switch r {
case 0:
b.WriteString(`\0`)
case '\'':
b.WriteString(`''`)
case '"':
b.WriteString(`\"`)
case '\\':
b.WriteString(`\\`)
case '\n':
b.WriteString(`\n`)
case '\r':
b.WriteString(`\r`)
case 26: // Ctrl+Z
b.WriteString(`\Z`)
default:
b.WriteRune(r)
}
}
return b.String()
}
@@ -0,0 +1,87 @@
// Copyright 2022 Clastix Labs
// SPDX-License-Identifier: Apache-2.0
package datastore
import (
"strings"
"testing"
)
// TestQuoteMySQLIdentifierNeutralisesInjection ensures a malicious identifier
// cannot break out of the backtick quoting and inject additional SQL, which is
// the regression behind the dataStoreSchema/dataStoreUsername SQL injection.
func TestQuoteMySQLIdentifierNeutralisesInjection(t *testing.T) {
const bt = "`"
// A tenant-controlled value attempting to stack a DROP DATABASE statement by
// closing the identifier with a backtick.
payload := "x" + bt + "; DROP DATABASE victim; -- "
// Every embedded backtick is doubled, so the whole payload stays a single
// quoted identifier instead of terminating it early.
want := bt + "x" + bt + bt + "; DROP DATABASE victim; -- " + bt
if got := quoteMySQLIdentifier(payload); got != want {
t.Fatalf("quoteMySQLIdentifier(%q) = %q, want %q", payload, got, want)
}
// Structural safety property: stripping the outer backticks and collapsing
// the doubled backticks must leave no lone backtick that could close the
// identifier prematurely.
inner := want[1 : len(want)-1]
if strings.Contains(strings.ReplaceAll(inner, bt+bt, ""), bt) {
t.Fatalf("a lone backtick survived quoting: %q", want)
}
}
func TestEscapeMySQLStringNeutralisesInjection(t *testing.T) {
payload := "p'; DROP DATABASE victim; -- "
// The single quote that would terminate the literal is doubled ('') rather
// than backslash-escaped, so the payload cannot escape the surrounding '...'
// literal under any sql_mode, including NO_BACKSLASH_ESCAPES.
want := `p''; DROP DATABASE victim; -- `
if got := escapeMySQLString(payload); got != want {
t.Fatalf("escapeMySQLString(%q) = %q, want %q", payload, got, want)
}
}
// TestEscapeMySQLStringQuoteEscapingIsModeIndependent guards against a
// regression back to backslash-escaping the quote (\'), which is unsafe when the
// server runs with NO_BACKSLASH_ESCAPES. Quotes must be doubled and backslashes
// must be doubled independently.
func TestEscapeMySQLStringQuoteEscapingIsModeIndependent(t *testing.T) {
// Quote-only input: the sole transformation is '' doubling, so an exact match
// unambiguously proves the quote is not backslash-escaped.
if got, want := escapeMySQLString("a'b'c"), "a''b''c"; got != want {
t.Fatalf("escapeMySQLString(%q) = %q, want %q", "a'b'c", got, want)
}
// Combined quote + backslash: both are doubled, still without producing a
// backslash-escaped quote.
if got, want := escapeMySQLString(`'\`), `''\\`; got != want {
t.Fatalf("escapeMySQLString(%q) = %q, want %q", `'\`, got, want)
}
}
func TestEscapeMySQLStringEscapesBackslash(t *testing.T) {
// A backslash must be doubled, otherwise `\'` could be smuggled in.
if got, want := escapeMySQLString(`a\`), `a\\`; got != want {
t.Fatalf("escapeMySQLString(%q) = %q, want %q", `a\`, got, want)
}
}
func TestQuoteMySQLIdentifierStripsNUL(t *testing.T) {
if got := quoteMySQLIdentifier("a\x00b"); strings.ContainsRune(got, 0) {
t.Fatalf("NUL byte survived quoting: %q", got)
}
}
// TestQuoteMySQLIdentifierPreservesValidIdentifiers guards the SHOW GRANTS
// comparison in GrantPrivilegesExists: well-formed identifiers must keep the
// exact backtick-wrapped shape MySQL itself emits.
func TestQuoteMySQLIdentifierPreservesValidIdentifiers(t *testing.T) {
if got, want := quoteMySQLIdentifier("tenant_namespace_cp"), "`tenant_namespace_cp`"; got != want {
t.Fatalf("valid identifier altered: got %q want %q", got, want)
}
}
+35 -17
View File
@@ -16,20 +16,24 @@ import (
)
const (
// Identifiers (database and role names) cannot be passed as bind parameters,
// so the `%s` verbs below must only ever be fed values run through
// quotePostgreSQLIdentifier. Plain values keep using the `?` placeholder as
// regular bind parameters.
postgresqlFetchDBStatement = "SELECT FROM pg_database WHERE datname = ?"
postgresqlCreateDBStatement = `CREATE DATABASE "%s"`
postgresqlCreateDBStatement = `CREATE DATABASE %s`
postgresqlUserExists = "SELECT 1 FROM pg_roles WHERE rolname = ?"
postgresqlCreateUserStatement = `CREATE ROLE "%s" LOGIN PASSWORD ?`
postgresqlUpdateUserStatement = `ALTER ROLE "%s" WITH PASSWORD ?`
postgresqlCreateUserStatement = `CREATE ROLE %s LOGIN PASSWORD ?`
postgresqlUpdateUserStatement = `ALTER ROLE %s WITH PASSWORD ?`
postgresqlShowGrantsStatement = "SELECT has_database_privilege(rolname, ?, 'create') from pg_roles where rolcanlogin and rolname = ?"
postgresqlShowOwnershipStatement = "SELECT 't' FROM pg_catalog.pg_database AS d WHERE d.datname = ? AND pg_catalog.pg_get_userbyid(d.datdba) = ?"
postgresqlShowTableOwnershipStatement = "SELECT 't' from pg_tables where tableowner = ? AND tablename = ?"
postgresqlKineTableExistsStatement = "SELECT 't' FROM pg_tables WHERE schemaname = ? AND tablename = ?"
postgresqlGrantPrivilegesStatement = `GRANT CONNECT, CREATE ON DATABASE "%s" TO "%s"`
postgresqlChangeOwnerStatement = `ALTER DATABASE "%s" OWNER TO "%s"`
postgresqlRevokePrivilegesStatement = `REVOKE ALL PRIVILEGES ON DATABASE "%s" FROM "%s"`
postgresqlDropRoleStatement = `DROP ROLE "%s"`
postgresqlDropDBStatement = `DROP DATABASE "%s" WITH (FORCE)`
postgresqlGrantPrivilegesStatement = `GRANT CONNECT, CREATE ON DATABASE %s TO %s`
postgresqlChangeOwnerStatement = `ALTER DATABASE %s OWNER TO %s`
postgresqlRevokePrivilegesStatement = `REVOKE ALL PRIVILEGES ON DATABASE %s FROM %s`
postgresqlDropRoleStatement = `DROP ROLE %s`
postgresqlDropDBStatement = `DROP DATABASE %s WITH (FORCE)`
)
type PostgreSQLConnection struct {
@@ -135,7 +139,7 @@ func (r *PostgreSQLConnection) UserExists(ctx context.Context, user string) (boo
}
func (r *PostgreSQLConnection) CreateUser(ctx context.Context, user, password string) error {
_, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlCreateUserStatement, user), password)
_, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlCreateUserStatement, quotePostgreSQLIdentifier(user)), password)
if err != nil {
return errors.NewCreateUserError(err)
}
@@ -144,7 +148,7 @@ func (r *PostgreSQLConnection) CreateUser(ctx context.Context, user, password st
}
func (r *PostgreSQLConnection) UpdateUser(ctx context.Context, user, password string) error {
_, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlUpdateUserStatement, user), password)
_, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlUpdateUserStatement, quotePostgreSQLIdentifier(user)), password)
if err != nil {
return errors.NewUpdateUserError(err)
}
@@ -162,7 +166,7 @@ func (r *PostgreSQLConnection) DBExists(ctx context.Context, dbName string) (boo
}
func (r *PostgreSQLConnection) CreateDB(ctx context.Context, dbName string) error {
_, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlCreateDBStatement, dbName))
_, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlCreateDBStatement, quotePostgreSQLIdentifier(dbName)))
if err != nil {
return errors.NewCreateDBError(err)
}
@@ -210,14 +214,14 @@ func (r *PostgreSQLConnection) GrantPrivilegesExists(ctx context.Context, user,
}
func (r *PostgreSQLConnection) GrantPrivileges(ctx context.Context, user, dbName string) error {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlGrantPrivilegesStatement, dbName, user)); err != nil {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlGrantPrivilegesStatement, quotePostgreSQLIdentifier(dbName), quotePostgreSQLIdentifier(user))); err != nil {
return errors.NewGrantPrivilegesError(err)
}
dbConn := r.switchDatabaseFn(dbName)
defer dbConn.Close()
if _, err := dbConn.ExecContext(ctx, fmt.Sprintf(postgresqlChangeOwnerStatement, dbName, user)); err != nil {
if _, err := dbConn.ExecContext(ctx, fmt.Sprintf(postgresqlChangeOwnerStatement, quotePostgreSQLIdentifier(dbName), quotePostgreSQLIdentifier(user))); err != nil {
return errors.NewGrantPrivilegesError(err)
}
@@ -227,7 +231,7 @@ func (r *PostgreSQLConnection) GrantPrivileges(ctx context.Context, user, dbName
}
if tableExists {
if _, err = dbConn.ExecContext(ctx, fmt.Sprintf(`ALTER TABLE kine OWNER TO "%s"`, user)); err != nil {
if _, err = dbConn.ExecContext(ctx, fmt.Sprintf("ALTER TABLE kine OWNER TO %s", quotePostgreSQLIdentifier(user))); err != nil {
return errors.NewGrantPrivilegesError(err)
}
}
@@ -236,7 +240,7 @@ func (r *PostgreSQLConnection) GrantPrivileges(ctx context.Context, user, dbName
}
func (r *PostgreSQLConnection) DeleteUser(ctx context.Context, user string) error {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlDropRoleStatement, user)); err != nil {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlDropRoleStatement, quotePostgreSQLIdentifier(user))); err != nil {
return errors.NewDeleteUserError(err)
}
@@ -248,7 +252,7 @@ func (r *PostgreSQLConnection) DeleteDB(ctx context.Context, dbName string) erro
return errors.NewCannotDeleteDatabaseError(fmt.Errorf("cannot grant privileges to root user: %w", err))
}
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlDropDBStatement, dbName)); err != nil {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlDropDBStatement, quotePostgreSQLIdentifier(dbName))); err != nil {
return errors.NewCannotDeleteDatabaseError(err)
}
@@ -256,7 +260,7 @@ func (r *PostgreSQLConnection) DeleteDB(ctx context.Context, dbName string) erro
}
func (r *PostgreSQLConnection) RevokePrivileges(ctx context.Context, user, dbName string) error {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlRevokePrivilegesStatement, dbName, user)); err != nil {
if _, err := r.db.ExecContext(ctx, fmt.Sprintf(postgresqlRevokePrivilegesStatement, quotePostgreSQLIdentifier(dbName), quotePostgreSQLIdentifier(user))); err != nil {
return errors.NewRevokePrivilegesError(err)
}
@@ -292,3 +296,17 @@ func (r *PostgreSQLConnection) kineTableExists(ctx context.Context, db *pg.DB) (
return tableExists == "t", nil
}
// quotePostgreSQLIdentifier safely quotes a PostgreSQL identifier (such as a
// database or role name) so it can be interpolated into a statement: it wraps
// the value in double quotes and doubles any embedded double quote, preventing
// a malicious value from breaking out of the identifier and injecting SQL. NUL
// bytes, which are illegal in identifiers, are stripped. Identifiers cannot be
// supplied as bind parameters, hence the manual quoting. The whole value is
// treated as a single identifier (dots are not treated as schema separators),
// preserving the historical behaviour for dotted database/role names.
func quotePostgreSQLIdentifier(identifier string) string {
identifier = strings.ReplaceAll(identifier, "\x00", "")
return `"` + strings.ReplaceAll(identifier, `"`, `""`) + `"`
}
@@ -0,0 +1,56 @@
// Copyright 2022 Clastix Labs
// SPDX-License-Identifier: Apache-2.0
package datastore
import (
"fmt"
"strings"
"testing"
)
// TestQuotePostgreSQLIdentifierNeutralisesInjection ensures a malicious
// identifier cannot break out of the double-quote quoting and stack additional
// statements, which is the regression behind the dataStoreSchema/
// dataStoreUsername SQL injection.
func TestQuotePostgreSQLIdentifierNeutralisesInjection(t *testing.T) {
// A tenant-controlled value attempting to stack a DROP DATABASE statement by
// closing the identifier with a double quote (PostgreSQL executes stacked
// statements over the simple query protocol used for these DDL calls).
payload := `x"; DROP DATABASE victim; -- `
// Every embedded double quote is doubled, so the whole payload stays a single
// quoted identifier instead of terminating it early.
want := `"x""; DROP DATABASE victim; -- "`
if got := quotePostgreSQLIdentifier(payload); got != want {
t.Fatalf("quotePostgreSQLIdentifier(%q) = %q, want %q", payload, got, want)
}
// Structural safety property: stripping the outer quotes and collapsing the
// doubled quotes must leave no lone double quote that could close the
// identifier prematurely.
inner := want[1 : len(want)-1]
if strings.Contains(strings.ReplaceAll(inner, `""`, ""), `"`) {
t.Fatalf("a lone double quote survived quoting: %q", want)
}
}
// TestQuotePostgreSQLIdentifierPreservesDottedNames guards the defaulter, which
// can produce dotted database/role names from a TenantControlPlane whose name
// contains a dot: the value must remain a single identifier.
func TestQuotePostgreSQLIdentifierPreservesDottedNames(t *testing.T) {
if got, want := quotePostgreSQLIdentifier("ns_foo.bar"), `"ns_foo.bar"`; got != want {
t.Fatalf("dotted identifier altered: got %q want %q", got, want)
}
// Confirm it interpolates into a single-identifier CREATE DATABASE.
if got, want := fmt.Sprintf(postgresqlCreateDBStatement, quotePostgreSQLIdentifier("ns_foo.bar")), `CREATE DATABASE "ns_foo.bar"`; got != want {
t.Fatalf("statement = %q, want %q", got, want)
}
}
func TestQuotePostgreSQLIdentifierStripsNUL(t *testing.T) {
if got := quotePostgreSQLIdentifier("a\x00b"); strings.ContainsRune(got, 0) {
t.Fatalf("NUL byte survived quoting: %q", got)
}
}