Files
pawsql/internal/server/server.go
Shaun Campbell cae9419cd5
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 5m7s
Test and Release PawSQL / test (push) Successful in 44s
Test and Release PawSQL / release (push) Successful in 10s
fix: decline SSL negotiation in plaintext mode
2026-09-15 22:28:16 -04:00

212 lines
6.7 KiB
Go

// 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. A nil
// TLS configuration selects plaintext mode: PostgreSQL SSL negotiation is not
// offered and routes are selected by database name only.
func New(tlsConfig *tls.Config, resolver router.BackendResolver, logger *slog.Logger) (*Server, error) {
if resolver == nil {
return nil, errors.New("backend resolver is required")
}
if logger == nil {
logger = slog.Default()
}
if tlsConfig != nil {
copy := tlsConfig.Clone()
if copy.MinVersion == 0 {
copy.MinVersion = tls.VersionTLS12
}
tlsConfig = copy
}
return &Server{TLSConfig: tlsConfig, 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))
var (
proxyClient net.Conn = connection
err error
)
if s.TLSConfig != nil {
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)
proxyClient = tlsConnection
}
resolveTimeout := s.ResolveTimeout
if resolveTimeout <= 0 {
resolveTimeout = defaultResolveTimeout
}
resolveContext, cancelResolve := context.WithTimeout(context.Background(), resolveTimeout)
defer cancelResolve()
if s.TLSConfig == nil {
negotiated, negotiateErr := pgwire.NegotiatePlainPostgreSQL(connection)
if negotiateErr != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_negotiation", "error", negotiateErr)
return
}
proxyClient = negotiated
}
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(proxyClient)
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(proxyClient, 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, "."))
}