Files
pawsql/internal/server/server_test.go
Shaun Campbell e67698f3c2
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 5m3s
Test and Release PawSQL / test (push) Successful in 43s
Test and Release PawSQL / release (push) Successful in 7s
feat: serve plaintext PostgreSQL when TLS is not configured
2026-09-15 22:03:58 -04:00

324 lines
8.9 KiB
Go

package server
import (
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/binary"
"encoding/pem"
"io"
"math/big"
"net"
"testing"
"time"
config "cloud.campbellwireless.net/git/barkstack/barkfile-parser"
"github.com/barkstack/pawsql/internal/pgwire"
"github.com/barkstack/pawsql/internal/router"
)
func TestServerHandlesPostgreSQLSSLRequestBeforeTLS(t *testing.T) {
backendListener := listen(t)
defer backendListener.Close()
backendReceived := make(chan []byte, 1)
go func() {
connection, err := backendListener.Accept()
if err != nil {
return
}
defer connection.Close()
payload := make([]byte, 4)
if _, err := io.ReadFull(connection, payload); err == nil {
backendReceived <- payload
_, _ = connection.Write([]byte("backend"))
}
}()
staticResolver, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "foo", Hostname: "foo.pawsql.test", Upstream: backendListener.Addr().String()}})
if err != nil {
t.Fatal(err)
}
resolver := &releaseTrackingResolver{BackendResolver: staticResolver, released: make(chan string, 1), traffic: make(chan trafficEvent, 2)}
routingServer, err := New(testCertificate(t), resolver, nil)
if err != nil {
t.Fatal(err)
}
listener := listen(t)
defer func() {
_ = routingServer.Shutdown()
routingServer.Wait()
_ = listener.Close()
}()
go func() { _ = routingServer.Serve(listener) }()
connection, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer connection.Close()
request := pgwire.SSLRequest()
if _, err := connection.Write(request[:]); err != nil {
t.Fatal(err)
}
response := make([]byte, 1)
if _, err := io.ReadFull(connection, response); err != nil {
t.Fatal(err)
}
if response[0] != 'S' {
t.Fatalf("SSL response = %q, want S", response)
}
client := tls.Client(connection, &tls.Config{ServerName: "FOO.pawsql.test", InsecureSkipVerify: true}) // test certificate is ephemeral
if err := client.Handshake(); err != nil {
t.Fatalf("TLS handshake: %v", err)
}
if _, err := client.Write([]byte("ping")); err != nil {
t.Fatal(err)
}
response = make([]byte, len("backend"))
if _, err := io.ReadFull(client, response); err != nil {
t.Fatal(err)
}
if string(response) != "backend" {
t.Fatalf("proxied response = %q", response)
}
if err := client.Close(); err != nil {
t.Fatal(err)
}
select {
case payload := <-backendReceived:
if string(payload) != "ping" {
t.Errorf("backend payload = %q", payload)
}
case <-time.After(time.Second):
t.Fatal("backend did not receive decrypted PostgreSQL bytes")
}
select {
case database := <-resolver.released:
if database != "foo" {
t.Errorf("released database = %q, want foo", database)
}
case <-time.After(time.Second):
t.Fatal("server did not release backend lifecycle lease")
}
var clientBytes, backendBytes int64
for range 2 {
select {
case event := <-resolver.traffic:
if event.clientToBackend {
clientBytes += event.bytes
} else {
backendBytes += event.bytes
}
case <-time.After(time.Second):
t.Fatal("server did not report proxied traffic")
}
}
if clientBytes != int64(len("ping")) || backendBytes != int64(len("backend")) {
t.Errorf("proxied traffic = client %d, backend %d", clientBytes, backendBytes)
}
}
func TestServerRoutesNoSNIByConfiguredDatabase(t *testing.T) {
backendListener := listen(t)
defer backendListener.Close()
backendDatabase := make(chan string, 1)
go func() {
connection, err := backendListener.Accept()
if err != nil {
return
}
defer connection.Close()
startup, err := pgwire.ReadStartupMessage(connection)
if err != nil {
return
}
backendDatabase <- startup.Database
_, _ = connection.Write([]byte("backend"))
}()
resolver, err := router.NewStaticResolver([]config.DatabaseConfig{{
Name: "analytics",
Upstream: backendListener.Addr().String(),
}})
if err != nil {
t.Fatal(err)
}
routingServer, err := New(testCertificate(t), resolver, nil)
if err != nil {
t.Fatal(err)
}
listener := listen(t)
defer func() {
_ = routingServer.Shutdown()
routingServer.Wait()
_ = listener.Close()
}()
go func() { _ = routingServer.Serve(listener) }()
connection, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer connection.Close()
request := pgwire.SSLRequest()
if _, err := connection.Write(request[:]); err != nil {
t.Fatal(err)
}
response := make([]byte, 1)
if _, err := io.ReadFull(connection, response); err != nil || response[0] != 'S' {
t.Fatalf("SSL response = %q, error = %v", response, err)
}
client := tls.Client(connection, &tls.Config{InsecureSkipVerify: true}) // no ServerName intentionally omits SNI
if err := client.Handshake(); err != nil {
t.Fatal(err)
}
if _, err := client.Write(startupMessage("analytics")); err != nil {
t.Fatal(err)
}
response = make([]byte, len("backend"))
if _, err := io.ReadFull(client, response); err != nil {
t.Fatal(err)
}
if string(response) != "backend" {
t.Fatalf("proxied response = %q", response)
}
select {
case database := <-backendDatabase:
if database != "analytics" {
t.Errorf("backend database = %q, want analytics", database)
}
case <-time.After(time.Second):
t.Fatal("backend did not receive original StartupMessage")
}
}
func TestServerRoutesPlaintextConnectionsByDatabase(t *testing.T) {
backendListener := listen(t)
defer backendListener.Close()
backendDatabase := make(chan string, 1)
go func() {
connection, err := backendListener.Accept()
if err != nil {
return
}
defer connection.Close()
startup, err := pgwire.ReadStartupMessage(connection)
if err != nil {
return
}
backendDatabase <- startup.Database
_, _ = connection.Write([]byte("backend"))
}()
staticResolver, err := router.NewStaticResolver([]config.DatabaseConfig{{
Name: "analytics",
Upstream: backendListener.Addr().String(),
}})
if err != nil {
t.Fatal(err)
}
routingServer, err := New(nil, staticResolver, nil)
if err != nil {
t.Fatal(err)
}
listener := listen(t)
defer func() {
_ = routingServer.Shutdown()
routingServer.Wait()
_ = listener.Close()
}()
go func() { _ = routingServer.Serve(listener) }()
connection, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer connection.Close()
if _, err := connection.Write(startupMessage("analytics")); err != nil {
t.Fatal(err)
}
response := make([]byte, len("backend"))
if _, err := io.ReadFull(connection, response); err != nil {
t.Fatal(err)
}
if string(response) != "backend" {
t.Fatalf("proxied response = %q", response)
}
select {
case database := <-backendDatabase:
if database != "analytics" {
t.Errorf("backend database = %q, want analytics", database)
}
case <-time.After(time.Second):
t.Fatal("backend did not receive the plaintext StartupMessage")
}
}
func startupMessage(database string) []byte {
body := make([]byte, 4)
binary.BigEndian.PutUint32(body, 196608)
body = append(body, "user"...)
body = append(body, 0)
body = append(body, "ruckstack"...)
body = append(body, 0)
body = append(body, "database"...)
body = append(body, 0)
body = append(body, database...)
body = append(body, 0, 0)
message := make([]byte, 4, 4+len(body))
message = append(message, body...)
binary.BigEndian.PutUint32(message, uint32(len(message)))
return message
}
type trafficEvent struct {
clientToBackend bool
bytes int64
}
type releaseTrackingResolver struct {
router.BackendResolver
released chan string
traffic chan trafficEvent
}
func (r *releaseTrackingResolver) RecordTraffic(_ string, clientToBackend bool, bytes int64) {
r.traffic <- trafficEvent{clientToBackend: clientToBackend, bytes: bytes}
}
func (r *releaseTrackingResolver) ReleaseConnection(database string) {
r.released <- database
}
func listen(t *testing.T) net.Listener {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
return listener
}
func testCertificate(t *testing.T) *tls.Config {
t.Helper()
key, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
certificateTemplate := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "pawsql test"}, NotBefore: time.Now().Add(-time.Minute), NotAfter: time.Now().Add(time.Hour), DNSNames: []string{"*.pawsql.test"}, KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}}
certificateDER, err := x509.CreateCertificate(rand.Reader, certificateTemplate, certificateTemplate, &key.PublicKey, key)
if err != nil {
t.Fatal(err)
}
certificatePEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificateDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
certificate, err := tls.X509KeyPair(certificatePEM, keyPEM)
if err != nil {
t.Fatal(err)
}
return &tls.Config{Certificates: []tls.Certificate{certificate}}
}