diff --git a/lib/multiplexer/multiplexer.go b/lib/multiplexer/multiplexer.go index 1ff9c585b5b..67a6468a26d 100644 --- a/lib/multiplexer/multiplexer.go +++ b/lib/multiplexer/multiplexer.go @@ -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) { diff --git a/lib/multiplexer/tls.go b/lib/multiplexer/tls.go index db93ef2be9e..bfa13587c34 100644 --- a/lib/multiplexer/tls.go +++ b/lib/multiplexer/tls.go @@ -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 } } diff --git a/lib/multiplexer/web.go b/lib/multiplexer/web.go index b8d200e3a2e..b3de7d3c168 100644 --- a/lib/multiplexer/web.go +++ b/lib/multiplexer/web.go @@ -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. diff --git a/lib/multiplexer/wrappers.go b/lib/multiplexer/wrappers.go index 0ca9a0ddc77..6bb1805f8b3 100644 --- a/lib/multiplexer/wrappers.go +++ b/lib/multiplexer/wrappers.go @@ -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 diff --git a/lib/service/service.go b/lib/service/service.go index dff09093260..3b95e521cae 100644 --- a/lib/service/service.go +++ b/lib/service/service.go @@ -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() { diff --git a/lib/service/signals.go b/lib/service/signals.go index 0193fed16ab..09a753b4f97 100644 --- a/lib/service/signals.go +++ b/lib/service/signals.go @@ -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{}, }) diff --git a/lib/srv/alpnproxy/listener.go b/lib/srv/alpnproxy/listener.go index 06a19748ec3..dba33dcc29d 100644 --- a/lib/srv/alpnproxy/listener.go +++ b/lib/srv/alpnproxy/listener.go @@ -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 {