diff --git a/pkg/tcpip/adapters/gonet/gonet.go b/pkg/tcpip/adapters/gonet/gonet.go index a93f51b8f..b5426811b 100644 --- a/pkg/tcpip/adapters/gonet/gonet.go +++ b/pkg/tcpip/adapters/gonet/gonet.go @@ -47,10 +47,11 @@ func (e *timeoutError) Temporary() bool { return true } // A TCPListener is a wrapper around a TCP tcpip.Endpoint that implements // net.Listener. type TCPListener struct { - stack *stack.Stack - ep tcpip.Endpoint - wq *waiter.Queue - cancel chan struct{} + stack *stack.Stack + ep tcpip.Endpoint + wq *waiter.Queue + cancelOnce sync.Once + cancel chan struct{} } // NewTCPListener creates a new TCPListener from a listening tcpip.Endpoint. @@ -112,7 +113,9 @@ func (l *TCPListener) Close() error { // Shutdown stops the HTTP server. func (l *TCPListener) Shutdown() { l.ep.Shutdown(tcpip.ShutdownWrite | tcpip.ShutdownRead) - close(l.cancel) // broadcast cancellation + l.cancelOnce.Do(func() { + close(l.cancel) // broadcast cancellation + }) } // Addr implements net.Listener.Addr. diff --git a/pkg/tcpip/adapters/gonet/gonet_test.go b/pkg/tcpip/adapters/gonet/gonet_test.go index 9f4e593c5..e48bea4f9 100644 --- a/pkg/tcpip/adapters/gonet/gonet_test.go +++ b/pkg/tcpip/adapters/gonet/gonet_test.go @@ -807,6 +807,105 @@ func TestDialContextTCPTimeout(t *testing.T) { } } +// TestInterruptListender tests that (*TCPListener).Accept can be interrupted. +func TestInterruptListender(t *testing.T) { + for _, test := range []struct { + name string + stop func(l *TCPListener) error + }{ + { + "Close", + (*TCPListener).Close, + }, + { + "Shutdown", + func(l *TCPListener) error { + l.Shutdown() + return nil + }, + }, + { + "Double Shutdown", + func(l *TCPListener) error { + l.Shutdown() + l.Shutdown() + return nil + }, + }, + } { + t.Run(test.name, func(t *testing.T) { + s, err := newLoopbackStack() + if err != nil { + t.Fatalf("newLoopbackStack() = %v", err) + } + defer func() { + s.Close() + s.Wait() + }() + + addr := tcpip.FullAddress{NICID, tcpip.Address(net.IPv4(169, 254, 10, 1).To4()), 11211} + + protocolAddr := tcpip.ProtocolAddress{ + Protocol: ipv4.ProtocolNumber, + AddressWithPrefix: addr.Addr.WithPrefix(), + } + if err := s.AddProtocolAddress(NICID, protocolAddr, stack.AddressProperties{}); err != nil { + t.Fatalf("AddProtocolAddress(%d, %+v, {}): %s", NICID, protocolAddr, err) + } + + l, e := ListenTCP(s, addr, ipv4.ProtocolNumber) + if e != nil { + t.Fatalf("NewListener() = %v", e) + } + defer l.Close() + done := make(chan struct{}) + go func() { + defer close(done) + c, err := l.Accept() + if err != nil { + // Accept is expected to return an error. + t.Log("Accept #1:", err) + return + } + t.Errorf("Accept #1 returned a connection: %v -> %v", c.LocalAddr(), c.RemoteAddr()) + c.Close() + }() + + // Give l.Accept a chance to block before stopping it. + time.Sleep(time.Millisecond * 50) + + if err := test.stop(l); err != nil { + t.Error("stop:", err) + } + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Errorf("c.Accept didn't unblock") + } + + done = make(chan struct{}) + go func() { + defer close(done) + c, err := l.Accept() + if err != nil { + // Accept is expected to return an error. + t.Log("Accept #2:", err) + return + } + t.Errorf("Accept #2 returned a connection: %v -> %v", c.LocalAddr(), c.RemoteAddr()) + c.Close() + }() + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Errorf("c.Accept didn't unblock a second time") + } + }) + } +} + func TestNetTest(t *testing.T) { nettest.TestConn(t, makePipe) }