Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 29 additions & 19 deletions p2p/transport/tcpreuse/listener.go
Original file line number Diff line number Diff line change
Expand Up @@ -109,19 +109,19 @@ func (t *ConnMgr) DemultiplexedListen(laddr ma.Multiaddr, connType Demultiplexed
}

ctx, cancel := context.WithCancel(context.Background())
cancelFunc := func() error {
cancel()
removeFunc := func() error {
t.mx.Lock()
defer t.mx.Unlock()
delete(t.listeners, laddr.String())
delete(t.listeners, gmal.Multiaddr().String())
return gmal.Close()
return nil
}
ml = &multiplexedListener{
GatedMaListener: gmal,
listeners: make(map[DemultiplexedConnType]*demultiplexedListener),
ctx: ctx,
closeFn: cancelFunc,
cancel: cancel,
closeFn: removeFunc,
}
t.listeners[laddr.String()] = ml
t.listeners[gmal.Multiaddr().String()] = ml
Expand All @@ -146,8 +146,12 @@ type multiplexedListener struct {
mx sync.RWMutex

ctx context.Context
cancel context.CancelFunc
closeFn func() error
wg sync.WaitGroup

closeOnce sync.Once
closeErr error
}

var ErrListenerExists = errors.New("listener already exists for this conn type on this address")
Expand All @@ -159,6 +163,9 @@ func (m *multiplexedListener) DemultiplexedListen(connType DemultiplexedConnType

m.mx.Lock()
defer m.mx.Unlock()
if m.ctx.Err() != nil {
return nil, transport.ErrListenerClosed
}
if _, ok := m.listeners[connType]; ok {
return nil, ErrListenerExists
}
Expand Down Expand Up @@ -249,14 +256,20 @@ func (m *multiplexedListener) run() error {
}

func (m *multiplexedListener) Close() error {
m.mx.Lock()
for _, l := range m.listeners {
l.cancelFunc()
}
err := m.closeListener()
m.mx.Unlock()
m.wg.Wait()
return err
m.closeOnce.Do(func() {
m.mx.Lock()
m.cancel()
for _, l := range m.listeners {
l.cancelFunc()
}
m.mx.Unlock()

// closeFn acquires the ConnMgr lock. Call it without holding m.mx,
// since DemultiplexedListen acquires those locks in the opposite order.
m.closeErr = m.closeListener()
m.wg.Wait()
})
return m.closeErr
}

func (m *multiplexedListener) closeListener() error {
Expand All @@ -267,14 +280,11 @@ func (m *multiplexedListener) closeListener() error {

func (m *multiplexedListener) removeDemultiplexedListener(c DemultiplexedConnType) {
m.mx.Lock()
defer m.mx.Unlock()

delete(m.listeners, c)
if len(m.listeners) == 0 {
m.closeListener()
m.mx.Unlock()
m.wg.Wait()
m.mx.Lock()
empty := len(m.listeners) == 0
m.mx.Unlock()
if empty {
_ = m.Close()
}
}

Expand Down
89 changes: 89 additions & 0 deletions p2p/transport/tcpreuse/listener_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"errors"
"fmt"
"math/big"
"net"
Expand Down Expand Up @@ -67,6 +68,94 @@ func (ml *maListener) Accept() (manet.Conn, error) {
return c, err
}

type blockingCloseGatedMaListener struct {
transport.GatedMaListener
closeStarted chan struct{}
continueClose chan struct{}
}

func (l *blockingCloseGatedMaListener) Close() error {
close(l.closeStarted)
<-l.continueClose
return l.GatedMaListener.Close()
}

func TestConcurrentListenAndCloseDoesNotDeadlock(t *testing.T) {
cm := NewConnMgr(false, upgrader(t))
listenAddr := ma.StringCast("/ip4/127.0.0.1/tcp/0")
gmal, err := cm.gatedMaListen(listenAddr)
require.NoError(t, err)

closeStarted := make(chan struct{})
continueClose := make(chan struct{})
ctx, cancel := context.WithCancel(context.Background())
ml := &multiplexedListener{
GatedMaListener: &blockingCloseGatedMaListener{
GatedMaListener: gmal,
closeStarted: closeStarted,
continueClose: continueClose,
},
listeners: make(map[DemultiplexedConnType]*demultiplexedListener),
ctx: ctx,
cancel: cancel,
}
ml.closeFn = func() error {
cm.mx.Lock()
defer cm.mx.Unlock()
delete(cm.listeners, listenAddr.String())
delete(cm.listeners, gmal.Multiaddr().String())
return nil
}
cm.mx.Lock()
cm.listeners[listenAddr.String()] = ml
cm.listeners[gmal.Multiaddr().String()] = ml
cm.mx.Unlock()

closeDone := make(chan error, 1)
go func() { closeDone <- ml.Close() }()
<-closeStarted // Pause closure while starting a concurrent listener registration.

listenDone := make(chan error, 1)
go func() {
_, err := cm.DemultiplexedListen(gmal.Multiaddr(), DemultiplexedConnType_HTTP)
listenDone <- err
}()

// On the buggy path, DemultiplexedListen holds cm.mx and blocks on ml.mx.
// With the fix, it may instead observe the canceled listener and return.
var listenErr error
listenReturned := false
require.Eventually(t, func() bool {
select {
case listenErr = <-listenDone:
listenReturned = true
return true
default:
}
if cm.mx.TryLock() {
cm.mx.Unlock()
return false
}
return true
}, time.Second, time.Millisecond)
close(continueClose)

select {
case err := <-closeDone:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("multiplexed listener Close deadlocked with DemultiplexedListen")
}
if !listenReturned {
select {
case listenErr = <-listenDone:
case <-time.After(time.Second):
t.Fatal("DemultiplexedListen deadlocked with multiplexed listener Close")
}
}
require.True(t, listenErr == nil || errors.Is(listenErr, transport.ErrListenerClosed), listenErr)
}

type wsHandler struct{ conns chan *websocket.Conn }

func (wh wsHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
Expand Down