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

+8 -5
View File
@@ -20,9 +20,10 @@ Implemented Modules:
tenant-scoped reports.
- `templates`, `prompt`, `settings`, and `logging`: the remaining compatibility
modules used by the HTTP surface.
- `postgres`: fail-closed configuration, verified-CA TLS, readiness, atomic
account mutations, and calls to the existing claim and wallet PostgreSQL
functions. PostgreSQL is the production relational source of truth.
- `postgres`: fail-closed configuration, code-enforced plaintext
`sslmode=disable`, readiness, atomic account mutations, and calls to the
existing claim and wallet PostgreSQL functions. PostgreSQL is the production
relational source of truth.
- `localstore`: a mutex-protected, non-durable, single-process development
store covering the same business Module ports.
- `httpapi`: the complete checked-in route compatibility surface.
@@ -107,5 +108,7 @@ bootstrap-creates accounts.
Production routing, Secret ownership, probes, and rollout commands are defined
in [`../docs/DEPLOYMENT.md`](../docs/DEPLOYMENT.md) and `../deploy/ack/`. Go owns
the session signing Secret and backend runtime configuration; the static Web
workload receives neither. Validate RDS/CA, OSS, providers, Webhooks, embedded
Worker recovery, and rollback behavior for each production release.
workload receives neither. PostgreSQL does not use TLS, so production must use
the RDS internal endpoint and restrict access with VPC boundaries, security
groups, and allowlists. Validate RDS connectivity, OSS, providers, Webhooks,
embedded Worker recovery, and rollback behavior for each production release.
+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
}
+27 -50
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{
+26 -18
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
}
+86 -8
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")
}
}