// 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, ".")) }