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 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}} }