Files
NianAIGC/backend/internal/postgres/open_test.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")
}
}