Listener hygiene (#12540)

This commit is contained in:
Edoardo Spadolini
2022-05-17 08:16:31 +00:00
committed by GitHub
parent 740e626943
commit 875dcf7ebc
7 changed files with 86 additions and 47 deletions
+1 -5
View File
@@ -243,11 +243,7 @@ func (m *Mux) detectAndForward(conn net.Conn) {
return
}
select {
case listener.connC <- connWrapper:
case <-m.context.Done():
connWrapper.Close()
}
listener.HandleConnection(m.context, connWrapper)
}
func detect(conn net.Conn, enableProxyProtocol bool) (*Conn, error) {
+2 -13
View File
@@ -162,23 +162,12 @@ func (l *TLSListener) detectAndForward(conn *tls.Conn) {
switch conn.ConnectionState().NegotiatedProtocol {
case http2.NextProtoTLS:
select {
case l.http2Listener.connC <- conn:
case <-l.context.Done():
conn.Close()
return
}
l.http2Listener.HandleConnection(l.context, conn)
case teleport.HTTPNextProtoTLS, "":
select {
case l.httpListener.connC <- conn:
case <-l.context.Done():
conn.Close()
return
}
l.httpListener.HandleConnection(l.context, conn)
default:
conn.Close()
l.log.WithError(err).Errorf("unsupported protocol: %v", conn.ConnectionState().NegotiatedProtocol)
return
}
}
+2 -11
View File
@@ -154,20 +154,11 @@ func (l *WebListener) detectAndForward(conn *tls.Conn) {
l.log.WithError(err).Debug("Failed to check if connection is database connection.")
}
if isDatabaseConnection {
select {
case l.dbListener.connC <- conn:
case <-l.context.Done():
conn.Close()
}
l.dbListener.HandleConnection(l.context, conn)
return
}
select {
case l.webListener.connC <- conn:
case <-l.context.Done():
conn.Close()
return
}
l.webListener.HandleConnection(l.context, conn)
}
// Close closes the listener.
+15 -4
View File
@@ -92,6 +92,8 @@ func (c *Conn) ReadProxyLine() (*ProxyLine, error) {
return proxyLine, nil
}
// returns a Listener that pretends to be listening on addr, closed whenever the
// parent context is done.
func newListener(parent context.Context, addr net.Addr) *Listener {
context, cancel := context.WithCancel(parent)
return &Listener{
@@ -122,14 +124,23 @@ func (l *Listener) Accept() (net.Conn, error) {
case <-l.context.Done():
return nil, trace.ConnectionProblem(net.ErrClosed, "listener is closed")
case conn := <-l.connC:
if conn == nil {
return nil, trace.ConnectionProblem(net.ErrClosed, "listener is closed")
}
return conn, nil
}
}
// Close closes the listener, connections to multiplexer will hang
// HandleConnection injects the connection into the Listener, blocking until the
// context expires, the connection is accepted or the Listener is closed.
func (l *Listener) HandleConnection(ctx context.Context, conn net.Conn) {
select {
case <-ctx.Done():
conn.Close()
case <-l.context.Done():
conn.Close()
case l.connC <- conn:
}
}
// Close closes the listener.
func (l *Listener) Close() error {
l.cancel()
return nil
+11
View File
@@ -294,6 +294,10 @@ type TeleportProcess struct {
// importedDescriptors is a list of imported file descriptors
// passed by the parent process
importedDescriptors []FileDescriptor
// listenersClosed is a flag that indicates that the process should not open
// new listeners (for instance, because we're shutting down and we've already
// closed all the listeners)
listenersClosed bool
// forkedPIDs is a collection of a teleport processes forked
// during restart used to collect their status in case if the
@@ -3837,6 +3841,13 @@ func (process *TeleportProcess) WaitWithContext(ctx context.Context) {
// StartShutdown launches non-blocking graceful shutdown process that signals
// completion, returns context that will be closed once the shutdown is done
func (process *TeleportProcess) StartShutdown(ctx context.Context) context.Context {
// by the time we get here we've already extracted the parent pipe, which is
// the only potential imported file descriptor that's not a listening
// socket, so closing every imported FD with a prefix of "" will close all
// imported listeners that haven't been used so far
warnOnErr(process.closeImportedDescriptors(""), process.log)
warnOnErr(process.stopListeners(), process.log)
process.BroadcastEvent(Event{Name: TeleportExitEvent, Payload: ctx})
localCtx, cancel := context.WithCancel(ctx)
go func() {
+50 -13
View File
@@ -19,6 +19,7 @@ package service
import (
"context"
"encoding/json"
"fmt"
"net"
"os"
"os/exec"
@@ -184,13 +185,16 @@ func (process *TeleportProcess) closeImportedDescriptors(prefix string) error {
defer process.Unlock()
var errors []error
for i := range process.importedDescriptors {
d := process.importedDescriptors[i]
openDescriptors := make([]FileDescriptor, 0, len(process.importedDescriptors))
for _, d := range process.importedDescriptors {
if strings.HasPrefix(d.Type, prefix) {
process.log.Infof("Closing imported but unused descriptor %v %v.", d.Type, d.Address)
errors = append(errors, d.File.Close())
} else {
openDescriptors = append(openDescriptors, d)
}
}
process.importedDescriptors = openDescriptors
return trace.NewAggregate(errors...)
}
@@ -213,10 +217,10 @@ func (process *TeleportProcess) importSignalPipe() (*os.File, error) {
process.Lock()
defer process.Unlock()
for i := range process.importedDescriptors {
d := process.importedDescriptors[i]
for i, d := range process.importedDescriptors {
if d.Type == signalPipeName {
process.importedDescriptors = append(process.importedDescriptors[:i], process.importedDescriptors[i+1:]...)
process.importedDescriptors[i] = process.importedDescriptors[len(process.importedDescriptors)-1]
process.importedDescriptors = process.importedDescriptors[:len(process.importedDescriptors)-1]
return d.File, nil
}
}
@@ -230,16 +234,17 @@ func (process *TeleportProcess) importListener(typ listenerType, address string)
process.Lock()
defer process.Unlock()
for i := range process.importedDescriptors {
d := process.importedDescriptors[i]
for i, d := range process.importedDescriptors {
if d.Type == string(typ) && d.Address == address {
l, err := d.ToListener()
listener, err := d.ToListener()
if err != nil {
return nil, trace.Wrap(err)
}
process.importedDescriptors = append(process.importedDescriptors[:i], process.importedDescriptors[i+1:]...)
process.registeredListeners = append(process.registeredListeners, registeredListener{typ: typ, address: address, listener: l})
return l, nil
process.importedDescriptors[i] = process.importedDescriptors[len(process.importedDescriptors)-1]
process.importedDescriptors = process.importedDescriptors[:len(process.importedDescriptors)-1]
r := registeredListener{typ: typ, address: address, listener: listener}
process.registeredListeners = append(process.registeredListeners, r)
return listener, nil
}
}
@@ -248,17 +253,48 @@ func (process *TeleportProcess) importListener(typ listenerType, address string)
// createListener creates listener and adds to a list of tracked listeners
func (process *TeleportProcess) createListener(typ listenerType, address string) (net.Listener, error) {
listenersClosed := func() bool {
process.Lock()
defer process.Unlock()
return process.listenersClosed
}
if listenersClosed() {
process.log.Debug("Listening is blocked, not opening listener for type %v and address %v.", typ, address)
return nil, trace.BadParameter("listening is blocked")
}
listener, err := net.Listen("tcp", address)
if err != nil {
return nil, trace.Wrap(err)
}
process.Lock()
defer process.Unlock()
// check this again in case we stopped allowing new listeners halfway
// through the net.Listen (which can block, if the address is a hostname and
// needs a dns lookup, so we can't do it while holding the lock)
if process.listenersClosed {
listener.Close()
process.log.Debug("Listening is blocked, closing newly-created listener for type %v and address %v.", typ, address)
return nil, trace.BadParameter("listening is blocked")
}
r := registeredListener{typ: typ, address: address, listener: listener}
process.registeredListeners = append(process.registeredListeners, r)
return listener, nil
}
func (process *TeleportProcess) stopListeners() error {
process.Lock()
defer process.Unlock()
process.listenersClosed = true
errors := make([]error, 0, len(process.registeredListeners))
for _, r := range process.registeredListeners {
errors = append(errors, r.listener.Close())
}
process.registeredListeners = nil
return trace.NewAggregate(errors...)
}
// ExportFileDescriptors exports file descriptors to be passed to child process
func (process *TeleportProcess) ExportFileDescriptors() ([]FileDescriptor, error) {
var out []FileDescriptor
@@ -278,6 +314,7 @@ func (process *TeleportProcess) ExportFileDescriptors() ([]FileDescriptor, error
func importFileDescriptors(log logrus.FieldLogger) ([]FileDescriptor, error) {
// These files may be passed in by the parent process
filesString := os.Getenv(teleportFilesEnvVar)
os.Unsetenv(teleportFilesEnvVar)
if filesString == "" {
return nil, nil
}
@@ -437,11 +474,11 @@ func (process *TeleportProcess) forkChild() error {
}
log.Infof("Passing %s to child", vals)
os.Setenv(teleportFilesEnvVar, vals)
env := append(os.Environ(), fmt.Sprintf("%s=%s", teleportFilesEnvVar, vals))
p, err := os.StartProcess(path, os.Args, &os.ProcAttr{
Dir: workingDir,
Env: os.Environ(),
Env: env,
Files: files,
Sys: &syscall.SysProcAttr{},
})
+5 -1
View File
@@ -21,6 +21,8 @@ import (
"net"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/utils"
)
// ListenerMuxWrapper wraps the net.Listener and multiplex incoming connection from serviceListener and connection
@@ -88,7 +90,9 @@ func (l *ListenerMuxWrapper) startAcceptingConnectionServiceListener() {
for {
conn, err := l.Listener.Accept()
if err != nil {
l.errC <- err
if !utils.IsUseOfClosedNetworkError(err) {
l.errC <- err
}
return
}
select {