Postgres: drop non-connection URI params before verifying (#5296)
* Postgres: drop non-connection URI params before verifying * Parameter names may contain underscores (e.g. connect_timeout, application_name, and ORM keys like connection_limit). Matching only [[:alpha:]] truncates these to their last segment (connection_limit to limit), which both breaks legitimate keys and lets ORM keys evade the non-connection-param filter.
This commit is contained in:
@@ -38,6 +38,24 @@ const (
|
||||
pgDbType = "db_type"
|
||||
)
|
||||
|
||||
// nonConnectionParams are query-string arguments that ORMs append to
|
||||
// Postgres connection URIs but that are not libpq connection keywords. lib/pq and pgx
|
||||
// forward any key they don't recognize to the server as a startup runtime parameter,
|
||||
// which the server rejects with 42704. Exluding them prevents this
|
||||
var nonConnectionParams = map[string]struct{}{
|
||||
"schema": {}, // Prisma: search_path selector
|
||||
"connection_limit": {}, // Prisma: client-side pool size
|
||||
"pool_timeout": {}, // Prisma: pool-acquisition wait
|
||||
"socket_timeout": {}, // Prisma: per-query timeout
|
||||
"pgbouncer": {}, // Prisma: PgBouncer compatibility mode
|
||||
"sslidentity": {}, // Prisma: PKCS12 certificate path
|
||||
}
|
||||
|
||||
func isNonConnectionParam(key string) bool {
|
||||
_, ok := nonConnectionParams[key]
|
||||
return ok
|
||||
}
|
||||
|
||||
// This detector currently only finds Postgres connection string URIs
|
||||
// (https://www.postgresql.org/docs/current/libpq-connect.html#LIBPQ-CONNSTRING-URIS) When it finds one, it uses
|
||||
// pq.ParseURI to normalize this into space-separated key-value pair Postgres connection string, and then uses a regular
|
||||
@@ -49,9 +67,9 @@ const (
|
||||
// Multi-host connection string URIs are currently not supported because pq.ParseURI doesn't parse them correctly. If we
|
||||
// happen to run into a case where this matters we can address it then.
|
||||
var (
|
||||
_ detectors.Detector = (*Scanner)(nil)
|
||||
uriPattern = regexp.MustCompile(`\b(?i)(postgres(?:ql)?)://\S+\b`)
|
||||
connStrPartPattern = regexp.MustCompile(`([[:alpha:]]+)='(.+?)' ?`)
|
||||
_ detectors.Detector = (*Scanner)(nil)
|
||||
uriPattern = regexp.MustCompile(`\b(?i)(postgres(?:ql)?)://\S+\b`)
|
||||
connStrPartPattern = regexp.MustCompile(`([[:alpha:]_]+)='(.+?)' ?`)
|
||||
)
|
||||
|
||||
type Scanner struct {
|
||||
@@ -320,7 +338,7 @@ func verifyPostgresPgx(ctx context.Context, params map[string]string) (bool, err
|
||||
func pgxConnString(params map[string]string) string {
|
||||
var connStr strings.Builder
|
||||
for key, value := range params {
|
||||
if key == pgDbType || key == pgRequiressl {
|
||||
if key == pgDbType || key == pgRequiressl || isNonConnectionParam(key) {
|
||||
continue
|
||||
}
|
||||
fmt.Fprintf(&connStr, "%s='%s'", key, value)
|
||||
@@ -370,6 +388,9 @@ func verifyPostgresPq(params map[string]string) (bool, error) {
|
||||
|
||||
var connStr string
|
||||
for key, value := range params {
|
||||
if isNonConnectionParam(key) {
|
||||
continue
|
||||
}
|
||||
connStr += fmt.Sprintf("%s='%s'", key, value)
|
||||
}
|
||||
|
||||
|
||||
@@ -253,3 +253,36 @@ func TestPostgres_RawVsPrimarySecret(t *testing.T) {
|
||||
assert.Equal(t, expectedRaw, string(res.RawV2))
|
||||
assert.Equal(t, input, res.GetPrimarySecretValue())
|
||||
}
|
||||
|
||||
// TestVerifyConnString_FiltersNonConnectionParams verifies that ORM query-string
|
||||
// arguments which are not libpq connection keywords are dropped from the verification
|
||||
// connection string.
|
||||
func TestVerifyConnString_FiltersNonConnectionParams(t *testing.T) {
|
||||
uri := `postgresql://u:p@h:5432/db?sslmode=require&schema=public&connection_limit=5&pool_timeout=10&socket_timeout=30&pgbouncer=true&sslidentity=/tmp/i.p12&connect_timeout=7`
|
||||
|
||||
matches := findUriMatches([]byte(uri), nil)
|
||||
require.Len(t, matches, 1)
|
||||
got := pgxConnString(matches[0].params)
|
||||
|
||||
|
||||
for _, key := range []string{pgUser, pgPassword, pgHost, pgPort, pgDbname, pgSslmode, pgConnectTimeout} {
|
||||
assert.Containsf(t, got, key+"=", "expected connection param %q to be preserved", key)
|
||||
}
|
||||
for _, key := range []string{pgDbType, "schema", "connection_limit", "pool_timeout", "socket_timeout", "pgbouncer", "sslidentity"} {
|
||||
assert.NotContainsf(t, got, key+"=", "expected non-connection param %q to be filtered out", key)
|
||||
}
|
||||
for _, mangled := range []string{"limit=", "identity="} {
|
||||
assert.NotContainsf(t, got, mangled, "found truncated key fragment %q — connStrPartPattern lost the underscore", mangled)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNonConnectionParam(t *testing.T) {
|
||||
// ORM arguments that are not libpq keywords.
|
||||
for _, key := range []string{"schema", "connection_limit", "pool_timeout", "socket_timeout", "pgbouncer", "sslidentity"} {
|
||||
assert.Truef(t, isNonConnectionParam(key), "%q should be treated as a non-connection param", key)
|
||||
}
|
||||
// Real libpq keywords that ORMs also use — must not be filtered.
|
||||
for _, key := range []string{pgConnectTimeout, pgSslmode, "sslcert", pgUser, pgPassword, pgHost, pgPort, pgDbname} {
|
||||
assert.Falsef(t, isNonConnectionParam(key), "%q is a real connection parameter and must not be filtered", key)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user