397 lines
11 KiB
Go
397 lines
11 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/v2"
|
|
"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 TestServerDeclinesSSLNegotiationInPlaintextMode(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()
|
|
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] != 'N' {
|
|
t.Fatalf("SSL response = %q, want N", response)
|
|
}
|
|
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 StartupMessage after SSL decline")
|
|
}
|
|
}
|
|
|
|
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}}
|
|
}
|