fix: disable PostgreSQL TLS for refusing RDS endpoint

This commit is contained in:
brother7 committed 2026-08-16 23:26:03 +08:00
1 parent acf368b6fe
commit ed978142eb
17 files changed
+340 -180

No files matched your search

+17 -31
View File
@@ -2,7 +2,6 @@ package postgres
import (
"crypto/tls"
"crypto/x509"
"fmt"
"net/url"
"strconv"
@@ -39,7 +38,7 @@ type Config struct {
ApplicationName string
}
func ParseConfig(getenv Getenv, readFile ReadFile) (Config, error) {
func ParseConfig(getenv Getenv, _ ReadFile) (Config, error) {
if getenv == nil {
return Config{}, fmt.Errorf("environment getter is required")
}
@@ -49,10 +48,9 @@ func ParseConfig(getenv Getenv, readFile ReadFile) (Config, error) {
ConnectionTimeout: 10 * time.Second,
StatementTimeout: 30 * time.Second,
ApplicationName: "zhinian-go",
// DATABASE_URL is the only database setting. Production defaults to
// full TLS verification; sslrootcert may be supplied in the URL when
// the RDS CA is not part of the container's system trust store.
SSLMode: SSLVerifyFull,
// Production RDS currently refuses TLS. Keep the transport choice in
// code so an older Secret cannot silently re-enable negotiation.
SSLMode: SSLDisable,
}
backend := strings.ToLower(strings.TrimSpace(getenv("ZHINIAN_DATA_BACKEND")))
@@ -97,35 +95,23 @@ func ParseConfig(getenv Getenv, readFile ReadFile) (Config, error) {
return Config{}, fmt.Errorf("DATABASE_URL must use the postgres:// or postgresql:// scheme")
}
query := parsed.Query()
for key := range query {
for key, values := range query {
lower := strings.ToLower(key)
if strings.HasPrefix(lower, "ssl") && lower != "sslmode" && lower != "sslrootcert" {
return Config{}, fmt.Errorf("DATABASE_URL contains unsupported SSL query parameter %s; use sslmode and optional sslrootcert", key)
return Config{}, fmt.Errorf("DATABASE_URL contains unsupported SSL query parameter %s; PostgreSQL transport is forced to sslmode=disable", key)
}
if lower == "sslmode" {
for _, value := range values {
mode := strings.ToLower(strings.TrimSpace(value))
if mode != "" && mode != string(SSLDisable) && mode != string(SSLVerifyFull) {
return Config{}, fmt.Errorf("DATABASE_URL sslmode must be 'disable' or 'verify-full'")
}
}
}
}
if mode := strings.ToLower(strings.TrimSpace(query.Get("sslmode"))); mode != "" {
cfg.SSLMode = SSLMode(mode)
}
switch cfg.SSLMode {
case SSLDisable:
case SSLVerifyFull:
cfg.TLSConfig = &tls.Config{MinVersion: tls.VersionTLS12}
if path := strings.TrimSpace(query.Get("sslrootcert")); path != "" {
if readFile == nil {
return Config{}, fmt.Errorf("a certificate reader is required when DATABASE_URL contains sslrootcert")
}
pem, readErr := readFile(path)
if readErr != nil {
return Config{}, fmt.Errorf("read DATABASE_URL sslrootcert: %w", readErr)
}
roots := x509.NewCertPool()
if !roots.AppendCertsFromPEM(pem) {
return Config{}, fmt.Errorf("DATABASE_URL sslrootcert does not contain a valid CA certificate")
}
cfg.TLSConfig.RootCAs = roots
}
default:
return Config{}, fmt.Errorf("DATABASE_URL sslmode must be 'disable' or 'verify-full'")
cfg.DatabaseURL, err = forcePlaintextDatabaseURL(cfg.DatabaseURL)
if err != nil {
return Config{}, fmt.Errorf("normalize DATABASE_URL: %w", err)
}
return cfg, nil
}