212 lines
6.7 KiB
Go
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, "."))
|
|
}
|