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"
|
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
|
// 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
|
// (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
|
// 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
|
// 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.
|
// happen to run into a case where this matters we can address it then.
|
||||||
var (
|
var (
|
||||||
_ detectors.Detector = (*Scanner)(nil)
|
_ detectors.Detector = (*Scanner)(nil)
|
||||||
uriPattern = regexp.MustCompile(`\b(?i)(postgres(?:ql)?)://\S+\b`)
|
uriPattern = regexp.MustCompile(`\b(?i)(postgres(?:ql)?)://\S+\b`)
|
||||||
connStrPartPattern = regexp.MustCompile(`([[:alpha:]]+)='(.+?)' ?`)
|
connStrPartPattern = regexp.MustCompile(`([[:alpha:]_]+)='(.+?)' ?`)
|
||||||
)
|
)
|
||||||
|
|
||||||
type Scanner struct {
|
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 {
|
func pgxConnString(params map[string]string) string {
|
||||||
var connStr strings.Builder
|
var connStr strings.Builder
|
||||||
for key, value := range params {
|
for key, value := range params {
|
||||||
if key == pgDbType || key == pgRequiressl {
|
if key == pgDbType || key == pgRequiressl || isNonConnectionParam(key) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
fmt.Fprintf(&connStr, "%s='%s'", key, value)
|
fmt.Fprintf(&connStr, "%s='%s'", key, value)
|
||||||
@@ -370,6 +388,9 @@ func verifyPostgresPq(params map[string]string) (bool, error) {
|
|||||||
|
|
||||||
var connStr string
|
var connStr string
|
||||||
for key, value := range params {
|
for key, value := range params {
|
||||||
|
if isNonConnectionParam(key) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
connStr += fmt.Sprintf("%s='%s'", key, value)
|
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, expectedRaw, string(res.RawV2))
|
||||||
assert.Equal(t, input, res.GetPrimarySecretValue())
|
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