diff --git a/pkg/tcpip/link/qdisc/fifo/fifo.go b/pkg/tcpip/link/qdisc/fifo/fifo.go index b90bb0a99..753562185 100644 --- a/pkg/tcpip/link/qdisc/fifo/fifo.go +++ b/pkg/tcpip/link/qdisc/fifo/fifo.go @@ -69,6 +69,13 @@ func New(lower stack.LinkWriter, n int, queueLen int) stack.QueueingDiscipline { go func() { defer d.wg.Done() qd.dispatchLoop() + qd.mu.Lock() + for qd.queue.Front() != nil { + p := qd.queue.Front() + qd.queue.Remove(p) + p.DecRef() + } + qd.mu.Unlock() }() } return d diff --git a/pkg/tcpip/link/sharedmem/BUILD b/pkg/tcpip/link/sharedmem/BUILD index 598719fe9..596c8c6e5 100644 --- a/pkg/tcpip/link/sharedmem/BUILD +++ b/pkg/tcpip/link/sharedmem/BUILD @@ -39,6 +39,8 @@ go_test( srcs = ["sharedmem_test.go"], library = ":sharedmem", deps = [ + "//pkg/refs", + "//pkg/refsvfs2", "//pkg/sync", "//pkg/tcpip", "//pkg/tcpip/buffer", @@ -57,6 +59,8 @@ go_test( deps = [ ":sharedmem", "//pkg/log", + "//pkg/refs", + "//pkg/refsvfs2", "//pkg/tcpip", "//pkg/tcpip/adapters/gonet", "//pkg/tcpip/header", diff --git a/pkg/tcpip/link/sharedmem/sharedmem_server_test.go b/pkg/tcpip/link/sharedmem/sharedmem_server_test.go index 8ebb9591e..886cc9257 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_server_test.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_server_test.go @@ -22,6 +22,7 @@ import ( "io" "net" "net/http" + "os" "strings" "syscall" "testing" @@ -29,6 +30,8 @@ import ( "golang.org/x/sync/errgroup" "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/log" + "gvisor.dev/gvisor/pkg/refs" + "gvisor.dev/gvisor/pkg/refsvfs2" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet" "gvisor.dev/gvisor/pkg/tcpip/header" @@ -287,3 +290,10 @@ func TestServerRoundTripStress(t *testing.T) { t.Fatalf("request failed: %s", err) } } + +func TestMain(m *testing.M) { + refs.SetLeakMode(refs.LeaksPanic) + code := m.Run() + refsvfs2.DoLeakCheck() + os.Exit(code) +} diff --git a/pkg/tcpip/link/sharedmem/sharedmem_test.go b/pkg/tcpip/link/sharedmem/sharedmem_test.go index cf58fe513..e12953df8 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem_test.go +++ b/pkg/tcpip/link/sharedmem/sharedmem_test.go @@ -20,11 +20,14 @@ package sharedmem import ( "bytes" "math/rand" + "os" "strings" "testing" "time" "golang.org/x/sys/unix" + "gvisor.dev/gvisor/pkg/refs" + "gvisor.dev/gvisor/pkg/refsvfs2" "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/buffer" @@ -235,6 +238,7 @@ func TestSimpleSend(t *testing.T) { pkt.NetworkProtocolNumber = proto var pkts stack.PacketBufferList pkts.PushBack(pkt) + defer pkts.DecRef() if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed: %s", err) } @@ -310,6 +314,7 @@ func TestPreserveSrcAddressInSend(t *testing.T) { pkt.NetworkProtocolNumber = proto var pkts stack.PacketBufferList + defer pkts.DecRef() pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed: %s", err) @@ -360,6 +365,8 @@ func TestFillTxQueue(t *testing.T) { // Each packet is uses no more than 40 bytes, so write that many packets // until the tx queue if full. + // Each packet uses no more than 40 bytes, so write that many packets + // until the tx queue if full. ids := make(map[uint64]struct{}) for i := queuePipeSize / 40; i > 0; i-- { pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{ @@ -372,8 +379,10 @@ func TestFillTxQueue(t *testing.T) { var pkts stack.PacketBufferList pkts.PushBack(pkt) if _, err := c.ep.WritePackets(pkts); err != nil { + pkts.DecRef() t.Fatalf("WritePackets failed unexpectedly: %s", err) } + pkts.DecRef() // Check that they have different IDs. desc := c.txq.tx.Pull() @@ -398,6 +407,7 @@ func TestFillTxQueue(t *testing.T) { if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } + pkts.DecRef() } // TestFillTxQueueAfterBadCompletion sends a bad completion, then sends packets @@ -431,6 +441,7 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } + pkts.DecRef() } // Complete the two writes twice. @@ -458,6 +469,7 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } + pkts.DecRef() // Check that they have different IDs. desc := c.txq.tx.Pull() @@ -481,6 +493,7 @@ func TestFillTxQueueAfterBadCompletion(t *testing.T) { if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } + pkts.DecRef() } // TestFillTxMemory sends packets until the we run out of shared memory. @@ -510,6 +523,7 @@ func TestFillTxMemory(t *testing.T) { if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } + pkts.DecRef() // Check that they have different IDs. desc := c.txq.tx.Pull() @@ -534,6 +548,7 @@ func TestFillTxMemory(t *testing.T) { if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } + pkts.DecRef() } // TestFillTxMemoryWithMultiBuffer sends packets until the we run out of @@ -564,6 +579,7 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } + pkts.DecRef() // Pull the posted buffer. c.txq.tx.Pull() @@ -584,6 +600,7 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { if _, ok := err.(*tcpip.ErrWouldBlock); !ok { t.Fatalf("got WritePackets(...) = %s, want %s", err, &tcpip.ErrWouldBlock{}) } + pkts.DecRef() } // Attempt to write the one-buffer packet again. It must succeed. @@ -599,6 +616,7 @@ func TestFillTxMemoryWithMultiBuffer(t *testing.T) { if _, err := c.ep.WritePackets(pkts); err != nil { t.Fatalf("WritePackets failed unexpectedly: %s", err) } + pkts.DecRef() } } @@ -815,3 +833,10 @@ func TestCloseWhileWaitingToPost(t *testing.T) { cleaned = true c.ep.Wait() } + +func TestMain(m *testing.M) { + refs.SetLeakMode(refs.LeaksPanic) + code := m.Run() + refsvfs2.DoLeakCheck() + os.Exit(code) +}