mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Listener hygiene (#12540)
This commit is contained in:
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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{},
|
||||
})
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user