feat: serve plaintext PostgreSQL when TLS is not configured
This commit is contained in:
@@ -36,22 +36,24 @@ type Server struct {
|
||||
handlers sync.WaitGroup
|
||||
}
|
||||
|
||||
// New validates server dependencies and applies safe protocol defaults.
|
||||
// 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 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
|
||||
if tlsConfig != nil {
|
||||
copy := tlsConfig.Clone()
|
||||
if copy.MinVersion == 0 {
|
||||
copy.MinVersion = tls.VersionTLS12
|
||||
}
|
||||
tlsConfig = copy
|
||||
}
|
||||
return &Server{TLSConfig: copy, Resolver: resolver, Logger: logger, HandshakeTimeout: defaultHandshakeTimeout, ResolveTimeout: defaultResolveTimeout, DialTimeout: defaultDialTimeout}, nil
|
||||
return &Server{TLSConfig: tlsConfig, Resolver: resolver, Logger: logger, HandshakeTimeout: defaultHandshakeTimeout, ResolveTimeout: defaultResolveTimeout, DialTimeout: defaultDialTimeout}, nil
|
||||
}
|
||||
|
||||
// Serve accepts connections until Shutdown closes its listener.
|
||||
@@ -114,26 +116,28 @@ func (s *Server) handle(connection net.Conn) {
|
||||
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
|
||||
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
|
||||
@@ -148,7 +152,7 @@ func (s *Server) handle(connection net.Conn) {
|
||||
return
|
||||
}
|
||||
} else {
|
||||
startup, startupErr := pgwire.ReadStartupMessage(tlsConnection)
|
||||
startup, startupErr := pgwire.ReadStartupMessage(proxyClient)
|
||||
if startupErr != nil {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_startup_message", "error", startupErr)
|
||||
return
|
||||
@@ -158,7 +162,7 @@ func (s *Server) handle(connection net.Conn) {
|
||||
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "unknown_database", "error", err)
|
||||
return
|
||||
}
|
||||
proxyClient = pgwire.Replay(tlsConnection, startup.Bytes())
|
||||
proxyClient = pgwire.Replay(proxyClient, startup.Bytes())
|
||||
}
|
||||
|
||||
if leases, ok := s.Resolver.(router.ConnectionLeaseManager); ok {
|
||||
|
||||
Reference in New Issue
Block a user