feat: add PawSQL docs examples and image CI
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 2m28s
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 2m28s
This commit is contained in:
261
internal/server/server_test.go
Normal file
261
internal/server/server_test.go
Normal file
@@ -0,0 +1,261 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/binary"
|
||||
"encoding/pem"
|
||||
"io"
|
||||
"math/big"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/barkstack/pawsql/internal/config"
|
||||
"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 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}}
|
||||
}
|
||||
Reference in New Issue
Block a user