From 45e8e4162096a302c85604de71a38c791f0c69b4 Mon Sep 17 00:00:00 2001 From: Amit Katyal Date: Sat, 1 Apr 2023 04:13:01 +0000 Subject: [PATCH] On net stack close, wake up all the waiters --- pkg/tcpip/adapters/gonet/gonet_test.go | 59 ++++++++++++++++++++++++++ pkg/tcpip/transport/tcp/endpoint.go | 1 + 2 files changed, 60 insertions(+) diff --git a/pkg/tcpip/adapters/gonet/gonet_test.go b/pkg/tcpip/adapters/gonet/gonet_test.go index 3370b529d..38c718927 100644 --- a/pkg/tcpip/adapters/gonet/gonet_test.go +++ b/pkg/tcpip/adapters/gonet/gonet_test.go @@ -370,6 +370,65 @@ func TestCloseWrite(t *testing.T) { } } + +// TestCloseStack tests that stack.Close wakes TCPConn.Read when +// using tcp.Forwarder. +func TestCloseStack(t *testing.T) { + s, err := newLoopbackStack() + if err != nil { + t.Fatalf("newLoopbackStack() = %v", err) + } + + 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) + } + + done := make(chan struct{}) + + fwd := tcp.NewForwarder(s, 30000, 10, func(r *tcp.ForwarderRequest) { + defer close(done) + + var wq waiter.Queue + ep, err := r.CreateEndpoint(&wq) + if err != nil { + t.Fatalf("r.CreateEndpoint() = %v", err) + } + r.Complete(false) + + c := NewTCPConn(&wq, ep) + + // Give c.Read() a chance to block before closing the stack. + time.AfterFunc(time.Second*1, func() { + s.Close() + s.Wait() + }) + + buf := make([]byte, 256) + n, e := c.Read(buf) + if n != 0 || !strings.Contains(e.Error(), "operation aborted") { + t.Errorf("c.Read() = (%d, %v), want (0, operation aborted)", n, e) + } + }) + s.SetTransportProtocolHandler(tcp.ProtocolNumber, fwd.HandlePacket) + + sender, err := connect(s, addr) + if err != nil { + t.Fatalf("connect() = %v", err) + } + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Errorf("c.Read() didn't unblock") + } + sender.close() +} + func TestUDPForwarder(t *testing.T) { s, terr := newLoopbackStack() if terr != nil { diff --git a/pkg/tcpip/transport/tcp/endpoint.go b/pkg/tcpip/transport/tcp/endpoint.go index 20419ee0f..6a8ad1b64 100644 --- a/pkg/tcpip/transport/tcp/endpoint.go +++ b/pkg/tcpip/transport/tcp/endpoint.go @@ -1012,6 +1012,7 @@ func (e *endpoint) Abort() { switch state := e.EndpointState(); { case state.connected(): e.resetConnectionLocked(&tcpip.ErrAborted{}) + e.waiterQueue.Notify(waiter.EventHUp | waiter.EventErr | waiter.ReadableEvents | waiter.WritableEvents) return } e.closeLocked()