139 lines
4.0 KiB
Go
139 lines
4.0 KiB
Go
package postgres
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"encoding/binary"
|
|
"io"
|
|
"net"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestOpenLocalReturnsStoreWithoutPool(t *testing.T) {
|
|
module, err := Open(context.Background(), Config{Backend: BackendLocal})
|
|
if err != nil {
|
|
t.Fatalf("Open() error = %v", err)
|
|
}
|
|
if module.Store == nil {
|
|
t.Fatal("Store = nil")
|
|
}
|
|
module.Close()
|
|
if err := module.Store.Readiness(context.Background()); err != nil {
|
|
t.Fatalf("Readiness() error = %v", err)
|
|
}
|
|
}
|
|
|
|
func TestOpenPostgresDoesNotProbeBeforeReadiness(t *testing.T) {
|
|
module, err := Open(context.Background(), Config{
|
|
Backend: BackendPostgres,
|
|
DatabaseURL: "postgresql://app:secret@127.0.0.1:1/app",
|
|
SSLMode: SSLDisable,
|
|
PoolMax: 1,
|
|
ApplicationName: "zhinian-go-test",
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Open() error = %v; connection failures belong to readiness", err)
|
|
}
|
|
module.Close()
|
|
}
|
|
|
|
func TestParseConfigIgnoresPostgresTLSSettingsForLocalBackend(t *testing.T) {
|
|
cfg, err := ParseConfig(env(map[string]string{
|
|
"ZHINIAN_DATA_BACKEND": "local",
|
|
"DATABASE_POOL_MAX": "0",
|
|
"DATABASE_CONNECTION_TIMEOUT_MS": "invalid",
|
|
}), nil)
|
|
if err != nil {
|
|
t.Fatalf("ParseConfig() error = %v", err)
|
|
}
|
|
if cfg.Backend != BackendLocal || cfg.TLSConfig != nil {
|
|
t.Fatalf("config = %+v", cfg)
|
|
}
|
|
}
|
|
|
|
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 {
|
|
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")
|
|
}
|
|
}
|