mirror of
https://github.com/soypat/lneto.git
synced 2026-09-08 15:59:10 +00:00
xnet.TCPListener guard racy reading of closed bool (#200)
This commit is contained in:
+13
-29
@@ -290,9 +290,9 @@ func (u *udppktconn) Write(b []byte) (int, error) {
|
||||
}
|
||||
func (u *udppktconn) RemoteAddr() net.Addr { return nil }
|
||||
|
||||
// tcplistener adapts [tcp.Listener] to [net.Listener].
|
||||
type tcplistener struct {
|
||||
l tcp.Listener
|
||||
closed bool
|
||||
sleep lneto.BackoffStrategy
|
||||
localAddr net.Addr
|
||||
}
|
||||
@@ -308,42 +308,26 @@ func (l *tcplistener) Addr() net.Addr {
|
||||
|
||||
func (l *tcplistener) Shutdown() { l.Close() }
|
||||
|
||||
// Accept blocks until the next connection is accepted and returned or until it is closed.
|
||||
// It ignores [lneto.ErrExhausted] which cause dropped connections.
|
||||
func (l *tcplistener) Accept() (net.Conn, error) {
|
||||
if l.closed {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
var backoffs uint
|
||||
for {
|
||||
if l.closed {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
n := l.l.NumberOfReadyToAccept()
|
||||
if n == 0 {
|
||||
backoff(l.sleep, backoffs)
|
||||
backoffs++
|
||||
continue
|
||||
}
|
||||
backoffs = 0
|
||||
c, _, err := l.l.TryAccept()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
if err == nil {
|
||||
return tcpconn{
|
||||
Conn: c,
|
||||
localAddr: l.localAddr,
|
||||
}, nil
|
||||
} else if err != lneto.ErrExhausted {
|
||||
return nil, err // net.ErrClosed or failure.
|
||||
}
|
||||
cc := tcpconn{
|
||||
Conn: c,
|
||||
localAddr: l.localAddr,
|
||||
}
|
||||
return cc, nil
|
||||
backoff(l.sleep, backoffs)
|
||||
backoffs++
|
||||
}
|
||||
}
|
||||
|
||||
func (l *tcplistener) Close() error {
|
||||
if l.closed {
|
||||
return net.ErrClosed
|
||||
}
|
||||
err := l.l.Close()
|
||||
l.closed = true
|
||||
return err
|
||||
}
|
||||
func (l *tcplistener) Close() error { return l.l.Close() }
|
||||
|
||||
type tcpconn struct {
|
||||
*tcp.Conn
|
||||
|
||||
@@ -277,6 +277,67 @@ func TestListener_Close(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestTCPListener_CloseUnblocksAccept covers the net.Listener wrapper's Accept
|
||||
// poll loop being ended by a Close from another goroutine.
|
||||
func TestTCPListener_CloseUnblocksAccept(t *testing.T) {
|
||||
const svPort uint16 = 80
|
||||
|
||||
pool, err := NewTCPPool(TCPPoolConfig{
|
||||
PoolSize: 1,
|
||||
QueueSize: 4,
|
||||
TxBufSize: 512,
|
||||
RxBufSize: 512,
|
||||
EstablishedTimeout: 10e9,
|
||||
ClosingTimeout: 10e9,
|
||||
NewBackoff: func() lneto.BackoffStrategy { return backoffYield },
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var l tcplistener
|
||||
// Sleep rather than yield between polls so Accept is genuinely parked in the
|
||||
// loop when Close lands, instead of spinning a core for the whole test.
|
||||
l.sleep = func(consecutiveBackoffs uint) time.Duration { return time.Millisecond }
|
||||
l.localAddr = net.TCPAddrFromAddrPort(netip.AddrPortFrom(netip.AddrFrom4([4]byte{10, 0, 0, 1}), svPort))
|
||||
err = l.l.Reset(svPort, pool)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// No stack is driving this listener, so Accept can only ever block: nothing
|
||||
// will become ready and the sole way out is the Close below.
|
||||
accepted := make(chan error, 1)
|
||||
go func() {
|
||||
c, err := l.Accept()
|
||||
if c != nil {
|
||||
c.Close()
|
||||
}
|
||||
accepted <- err
|
||||
}()
|
||||
time.Sleep(20 * time.Millisecond) // Let Accept reach its poll loop.
|
||||
|
||||
if err := l.Close(); err != nil {
|
||||
t.Fatal("Close while Accept is blocked:", err)
|
||||
}
|
||||
select {
|
||||
case err := <-accepted:
|
||||
if err != net.ErrClosed {
|
||||
t.Fatalf("blocked Accept: want net.ErrClosed, got %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Accept did not return after Close")
|
||||
}
|
||||
|
||||
// A later Accept reports the close rather than blocking again.
|
||||
if _, err := l.Accept(); err != net.ErrClosed {
|
||||
t.Fatalf("Accept after Close: want net.ErrClosed, got %v", err)
|
||||
}
|
||||
if err := l.Close(); err != net.ErrClosed {
|
||||
t.Fatalf("double Close: want net.ErrClosed, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListener_ResetAfterClose(t *testing.T) {
|
||||
const svPort uint16 = 80
|
||||
|
||||
|
||||
Reference in New Issue
Block a user