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:
199
internal/server/server.go
Normal file
199
internal/server/server.go
Normal file
@@ -0,0 +1,199 @@
|
||||
// Package server accepts PostgreSQL TLS sessions and routes them by SNI.
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/barkstack/pawsql/internal/pgwire"
|
||||
"github.com/barkstack/pawsql/internal/proxy"
|
||||
"github.com/barkstack/pawsql/internal/router"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultHandshakeTimeout = 15 * time.Second
|
||||
defaultResolveTimeout = 60 * time.Second
|
||||
defaultDialTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// Server terminates client TLS and transparently proxies PostgreSQL bytes.
|
||||
type Server struct {
|
||||
TLSConfig *tls.Config
|
||||
Resolver router.BackendResolver
|
||||
Logger *slog.Logger
|
||||
HandshakeTimeout time.Duration
|
||||
ResolveTimeout time.Duration
|
||||
DialTimeout time.Duration
|
||||
mu sync.Mutex
|
||||
listener net.Listener
|
||||
handlers sync.WaitGroup
|
||||
}
|
||||
|
||||
// New validates server dependencies and applies safe protocol defaults.
|
||||
func New(tlsConfig *tls.Config, resolver router.BackendResolver, logger *slog.Logger) (*Server, error) {
|
||||
if tlsConfig == nil {
|
||||
return nil, errors.New("TLS configuration is required")
|
||||
}
|
||||
if resolver == nil {
|
||||
return nil, errors.New("backend resolver is required")
|
||||
}
|
||||
if logger == nil {
|
||||
logger = slog.Default()
|
||||
}
|
||||
copy := tlsConfig.Clone()
|
||||
if copy.MinVersion == 0 {
|
||||
copy.MinVersion = tls.VersionTLS12
|
||||
}
|
||||
return &Server{TLSConfig: copy, Resolver: resolver, Logger: logger, HandshakeTimeout: defaultHandshakeTimeout, ResolveTimeout: defaultResolveTimeout, DialTimeout: defaultDialTimeout}, nil
|
||||
}
|
||||
|
||||
// Serve accepts connections until Shutdown closes its listener.
|
||||
func (s *Server) Serve(listener net.Listener) error {
|
||||
s.mu.Lock()
|
||||
if s.listener != nil {
|
||||
s.mu.Unlock()
|
||||
return errors.New("server is already serving")
|
||||
}
|
||||
s.listener = listener
|
||||
s.mu.Unlock()
|
||||
|
||||
for {
|
||||
connection, err := listener.Accept()
|
||||
if err != nil {
|
||||
if errors.Is(err, net.ErrClosed) {
|
||||
return nil
|
||||
}
|
||||
if temporary, ok := err.(interface{ Temporary() bool }); ok && temporary.Temporary() {
|
||||
s.Logger.Warn("temporary accept failure", "error", err)
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("accept connection: %w", err)
|
||||
}
|
||||
s.handlers.Add(1)
|
||||
go func() {
|
||||
defer s.handlers.Done()
|
||||
s.handle(connection)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// Shutdown stops accepting new connections. Existing sessions continue until they close.
|
||||
func (s *Server) Shutdown() error {
|
||||
s.mu.Lock()
|
||||
listener := s.listener
|
||||
s.mu.Unlock()
|
||||
if listener == nil {
|
||||
return nil
|
||||
}
|
||||
return listener.Close()
|
||||
}
|
||||
|
||||
// Wait waits for existing routed connections to finish.
|
||||
func (s *Server) Wait() { s.handlers.Wait() }
|
||||
|
||||
func (s *Server) handle(connection net.Conn) {
|
||||
started := time.Now()
|
||||
remote := connection.RemoteAddr().String()
|
||||
var sni string
|
||||
var backend router.Backend
|
||||
result := "rejected"
|
||||
defer func() {
|
||||
_ = connection.Close()
|
||||
s.Logger.Info("connection finished", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address, "duration", time.Since(started), "result", result)
|
||||
}()
|
||||
|
||||
handshakeTimeout := s.HandshakeTimeout
|
||||
if handshakeTimeout <= 0 {
|
||||
handshakeTimeout = defaultHandshakeTimeout
|
||||
}
|
||||
_ = connection.SetDeadline(time.Now().Add(handshakeTimeout))
|
||||
if err := pgwire.ReadSSLRequest(connection); err != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_ssl_request", "error", err)
|
||||
return
|
||||
}
|
||||
if _, err := connection.Write([]byte{'S'}); err != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "ssl_response_failed", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
tlsConnection := tls.Server(connection, s.TLSConfig)
|
||||
if err := tlsConnection.Handshake(); err != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "tls_handshake_failed", "error", err)
|
||||
return
|
||||
}
|
||||
_ = tlsConnection.SetDeadline(time.Time{})
|
||||
sni = normalizeHostname(tlsConnection.ConnectionState().ServerName)
|
||||
var (
|
||||
proxyClient net.Conn = tlsConnection
|
||||
err error
|
||||
)
|
||||
resolveTimeout := s.ResolveTimeout
|
||||
if resolveTimeout <= 0 {
|
||||
resolveTimeout = defaultResolveTimeout
|
||||
}
|
||||
resolveContext, cancelResolve := context.WithTimeout(context.Background(), resolveTimeout)
|
||||
defer cancelResolve()
|
||||
|
||||
if sni != "" {
|
||||
backend, err = s.Resolver.Resolve(resolveContext, sni)
|
||||
if err != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "sni", sni, "result", "unknown_hostname", "error", err)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
startup, startupErr := pgwire.ReadStartupMessage(tlsConnection)
|
||||
if startupErr != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_startup_message", "error", startupErr)
|
||||
return
|
||||
}
|
||||
backend, err = s.Resolver.ResolveDatabase(resolveContext, startup.Database)
|
||||
if err != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "unknown_database", "error", err)
|
||||
return
|
||||
}
|
||||
proxyClient = pgwire.Replay(tlsConnection, startup.Bytes())
|
||||
}
|
||||
|
||||
if leases, ok := s.Resolver.(router.ConnectionLeaseManager); ok {
|
||||
defer leases.ReleaseConnection(backend.DatabaseName)
|
||||
}
|
||||
|
||||
dialTimeout := s.DialTimeout
|
||||
if dialTimeout <= 0 {
|
||||
dialTimeout = defaultDialTimeout
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), dialTimeout)
|
||||
upstream, err := (&net.Dialer{}).DialContext(ctx, "tcp", backend.Address)
|
||||
cancel()
|
||||
if err != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address, "result", "upstream_connect_failed", "error", err)
|
||||
return
|
||||
}
|
||||
defer upstream.Close()
|
||||
s.Logger.Info("connection routed", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address)
|
||||
var proxyErr error
|
||||
if meter, ok := s.Resolver.(router.TrafficMeter); ok {
|
||||
proxyErr = proxy.BidirectionalWithTraffic(proxyClient, upstream, func(clientToBackend bool, bytes int64) {
|
||||
meter.RecordTraffic(backend.DatabaseName, clientToBackend, bytes)
|
||||
})
|
||||
} else {
|
||||
proxyErr = proxy.Bidirectional(proxyClient, upstream)
|
||||
}
|
||||
if proxyErr != nil {
|
||||
result = "proxy_error"
|
||||
s.Logger.Warn("connection proxy failed", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address, "error", proxyErr)
|
||||
return
|
||||
}
|
||||
result = "closed"
|
||||
}
|
||||
|
||||
func normalizeHostname(hostname string) string {
|
||||
return strings.ToLower(strings.TrimSuffix(hostname, "."))
|
||||
}
|
||||
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