fix: disable PostgreSQL TLS for refusing RDS endpoint

This commit is contained in:
2026-08-16 23:26:03 +08:00
parent acf368b6fe
commit ed978142eb
17 changed files with 340 additions and 180 deletions

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
}

View File

@@ -1,15 +1,8 @@
package postgres
import (
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"net/url"
"os"
"path/filepath"
"strings"
"testing"
"time"
@@ -60,52 +53,36 @@ func TestParseConfigAcceptsOnlyPostgresSchemesAndRejectsUnsupportedSSLQueryParam
}
}
func TestParseConfigBuildsVerifyFullTLSFromURLCA(t *testing.T) {
dir := t.TempDir()
caPath := filepath.Join(dir, "ca.pem")
const ca = "-----BEGIN CERTIFICATE-----\nMIIB\n-----END CERTIFICATE-----\n"
read := func(path string) ([]byte, error) {
if path != caPath {
t.Fatalf("read path = %q, want %q", path, caPath)
}
return []byte(ca), nil
}
_, err := ParseConfig(env(map[string]string{
"ZHINIAN_DATA_BACKEND": "postgres",
"DATABASE_URL": "postgresql://db.example/app?sslmode=verify-full&sslrootcert=" + url.QueryEscape(caPath),
}), read)
if err == nil || !strings.Contains(err.Error(), "CA certificate") {
t.Fatalf("ParseConfig() error = %v, want invalid CA certificate error", err)
}
}
func TestParseConfigVerifyFullBuildsRootsWithoutDisablingVerification(t *testing.T) {
key, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "test CA"},
NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour),
IsCA: true, BasicConstraintsValid: true, KeyUsage: x509.KeyUsageCertSign,
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
if err != nil {
t.Fatal(err)
}
ca := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
func TestParseConfigForcesPlaintextAndDoesNotReadURLCA(t *testing.T) {
readCalled := false
cfg, err := ParseConfig(env(map[string]string{
"ZHINIAN_DATA_BACKEND": "postgres",
"DATABASE_URL": "postgresql://db.example/app?sslmode=verify-full&sslrootcert=%2Fca.pem",
}), func(string) ([]byte, error) { return ca, nil })
"DATABASE_URL": "postgresql://db.example/app?connect_timeout=5&sslmode=verify-full&sslrootcert=%2Fetc%2Fzhinian%2Frds%2Fca.pem",
}), func(string) ([]byte, error) {
readCalled = true
return nil, os.ErrNotExist
})
if err != nil {
t.Fatalf("ParseConfig() error = %v", err)
}
if cfg.TLSConfig == nil || cfg.TLSConfig.RootCAs == nil {
t.Fatal("TLSConfig.RootCAs is nil")
if readCalled {
t.Fatal("ParseConfig() read sslrootcert, want plaintext configuration without CA access")
}
if cfg.TLSConfig.InsecureSkipVerify {
t.Fatal("TLSConfig.InsecureSkipVerify = true, want full certificate and hostname verification")
if cfg.SSLMode != SSLDisable || cfg.TLSConfig != nil {
t.Fatalf("TLS configuration = mode %q config %v, want disabled and nil", cfg.SSLMode, cfg.TLSConfig)
}
parsed, err := url.Parse(cfg.DatabaseURL)
if err != nil {
t.Fatalf("parse normalized DATABASE_URL: %v", err)
}
if got := parsed.Query().Get("sslmode"); got != "disable" {
t.Fatalf("normalized sslmode = %q, want disable", got)
}
if parsed.Query().Has("sslrootcert") {
t.Fatal("normalized DATABASE_URL retains sslrootcert")
}
if got := parsed.Query().Get("connect_timeout"); got != "5" {
t.Fatalf("normalized connect_timeout = %q, want 5", got)
}
}
@@ -120,11 +97,11 @@ func TestParseConfigDefaultsAndNumericValidation(t *testing.T) {
if cfg.PoolMax != 10 || cfg.IdleTimeout != 30*time.Second || cfg.ConnectionTimeout != 10*time.Second || cfg.StatementTimeout != 30*time.Second {
t.Fatalf("unexpected defaults: %+v", cfg)
}
if cfg.ApplicationName != "zhinian-go" || cfg.SSLMode != SSLVerifyFull {
if cfg.ApplicationName != "zhinian-go" || cfg.SSLMode != SSLDisable {
t.Fatalf("unexpected identity/TLS defaults: %+v", cfg)
}
if cfg.TLSConfig == nil || cfg.TLSConfig.InsecureSkipVerify {
t.Fatal("default PostgreSQL configuration must verify the server certificate")
if cfg.TLSConfig != nil {
t.Fatal("default PostgreSQL configuration must not negotiate TLS")
}
for name, value := range map[string]string{

View File

@@ -3,7 +3,9 @@ package postgres
import (
"context"
"fmt"
"net/url"
"strconv"
"strings"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
@@ -27,7 +29,11 @@ func Open(ctx context.Context, config Config) (*Module, error) {
if config.Backend != BackendPostgres {
return nil, fmt.Errorf("unsupported data backend %q", config.Backend)
}
poolConfig, err := pgxpool.ParseConfig(config.DatabaseURL)
databaseURL, err := forcePlaintextDatabaseURL(config.DatabaseURL)
if err != nil {
return nil, fmt.Errorf("parse DATABASE_URL: %w", err)
}
poolConfig, err := pgxpool.ParseConfig(databaseURL)
if err != nil {
return nil, fmt.Errorf("parse DATABASE_URL: %w", err)
}
@@ -36,23 +42,8 @@ func Open(ctx context.Context, config Config) (*Module, error) {
poolConfig.ConnConfig.ConnectTimeout = config.ConnectionTimeout
poolConfig.ConnConfig.RuntimeParams["statement_timeout"] = strconv.FormatInt(config.StatementTimeout.Milliseconds(), 10)
poolConfig.ConnConfig.RuntimeParams["application_name"] = config.ApplicationName
switch config.SSLMode {
case SSLDisable:
poolConfig.ConnConfig.TLSConfig = nil
poolConfig.ConnConfig.Fallbacks = nil
case SSLVerifyFull:
if config.TLSConfig == nil {
return nil, fmt.Errorf("TLS configuration is required when DATABASE_URL sslmode=verify-full")
}
tlsConfig := config.TLSConfig.Clone()
if tlsConfig.ServerName == "" {
tlsConfig.ServerName = poolConfig.ConnConfig.Host
}
poolConfig.ConnConfig.TLSConfig = tlsConfig
poolConfig.ConnConfig.Fallbacks = nil
default:
return nil, fmt.Errorf("unsupported DATABASE_URL sslmode %q", config.SSLMode)
}
poolConfig.ConnConfig.TLSConfig = nil
poolConfig.ConnConfig.Fallbacks = nil
pool, err := pgxpool.NewWithConfig(ctx, poolConfig)
if err != nil {
return nil, fmt.Errorf("open PostgreSQL pool: %w", err)
@@ -61,6 +52,23 @@ func Open(ctx context.Context, config Config) (*Module, error) {
return &Module{pool: adapter, Store: NewDatabase(config, adapter)}, nil
}
func forcePlaintextDatabaseURL(databaseURL string) (string, error) {
parsed, err := url.Parse(databaseURL)
if err != nil {
return "", err
}
query := parsed.Query()
for key := range query {
switch strings.ToLower(key) {
case "sslmode", "sslrootcert":
query.Del(key)
}
}
query.Set("sslmode", string(SSLDisable))
parsed.RawQuery = query.Encode()
return parsed.String(), nil
}
type pgxPoolAdapter struct {
pool *pgxpool.Pool
}

View File

@@ -2,8 +2,12 @@ package postgres
import (
"context"
"strings"
"crypto/tls"
"encoding/binary"
"io"
"net"
"testing"
"time"
)
func TestOpenLocalReturnsStoreWithoutPool(t *testing.T) {
@@ -48,13 +52,87 @@ func TestParseConfigIgnoresPostgresTLSSettingsForLocalBackend(t *testing.T) {
}
}
func TestOpenPostgresRejectsInvalidConfiguredTLSMode(t *testing.T) {
_, err := Open(context.Background(), Config{
Backend: BackendPostgres,
DatabaseURL: "postgresql://app:secret@db.example/app",
SSLMode: SSLMode("prefer"),
func TestOpenPostgresForcesPlaintextDespiteTLSBearingConfig(t *testing.T) {
module, err := Open(context.Background(), Config{
Backend: BackendPostgres,
DatabaseURL: "postgresql://app:secret@127.0.0.1:1/app?sslmode=verify-full&sslrootcert=/missing/ca.pem",
SSLMode: SSLVerifyFull,
TLSConfig: &tls.Config{MinVersion: tls.VersionTLS13},
PoolMax: 1,
ApplicationName: "zhinian-go-test",
})
if err == nil || !strings.Contains(err.Error(), "unsupported DATABASE_URL sslmode") {
t.Fatalf("Open() error = %v, want DATABASE_URL sslmode error", err)
if err != nil {
t.Fatalf("Open() error = %v, want TLS settings ignored", err)
}
module.Close()
}
func TestOpenPostgresSendsPlaintextStartupMessageWithoutTLSFallback(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("Listen() error = %v", err)
}
defer listener.Close()
firstPacket := make(chan [8]byte, 1)
serverError := make(chan error, 1)
go func() {
connection, acceptErr := listener.Accept()
if acceptErr != nil {
serverError <- acceptErr
return
}
defer connection.Close()
var packet [8]byte
if _, readErr := io.ReadFull(connection, packet[:]); readErr != nil {
serverError <- readErr
return
}
firstPacket <- packet
}()
module, err := Open(context.Background(), Config{
Backend: BackendPostgres,
DatabaseURL: "postgresql://app:secret@" + listener.Addr().String() + "/app?sslmode=verify-full&sslrootcert=/missing/ca.pem",
SSLMode: SSLVerifyFull,
TLSConfig: &tls.Config{MinVersion: tls.VersionTLS13},
PoolMax: 1,
ConnectionTimeout: time.Second,
StatementTimeout: time.Second,
ApplicationName: "zhinian-go-test",
})
if err != nil {
t.Fatalf("Open() error = %v", err)
}
defer module.Close()
adapter, ok := module.pool.(*pgxPoolAdapter)
if !ok {
t.Fatalf("pool type = %T, want *pgxPoolAdapter", module.pool)
}
connectionConfig := adapter.pool.Config().ConnConfig
if connectionConfig.TLSConfig != nil {
t.Fatal("pgx TLSConfig is not nil")
}
if len(connectionConfig.Fallbacks) != 0 {
t.Fatalf("pgx Fallbacks = %v, want empty", connectionConfig.Fallbacks)
}
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
if readinessErr := module.Store.Readiness(ctx); readinessErr == nil {
t.Fatal("Readiness() error = nil, want fake server disconnect")
}
select {
case packet := <-firstPacket:
protocolCode := binary.BigEndian.Uint32(packet[4:8])
if protocolCode != 196608 {
t.Fatalf("first PostgreSQL protocol code = %d, want plaintext StartupMessage 196608 (SSLRequest is 80877103)", protocolCode)
}
case serverErr := <-serverError:
t.Fatalf("fake PostgreSQL server error = %v", serverErr)
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for PostgreSQL startup packet")
}
}