feat: add PawSQL docs examples and image CI
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 2m28s

This commit is contained in:
2026-09-15 18:49:26 -04:00
commit 865f7c26c9
25 changed files with 2811 additions and 0 deletions

199
internal/server/server.go Normal file
View 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, "."))
}

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