diff --git a/pkg/refs/refs_map.go b/pkg/refs/refs_map.go index 6cdf0baf3..f94fea87c 100644 --- a/pkg/refs/refs_map.go +++ b/pkg/refs/refs_map.go @@ -145,7 +145,11 @@ type leakCheckDisabled interface { LeakCheckDisabled() bool } +// CleanupSync is used to wait for async cleanup actions. +var CleanupSync sync.WaitGroup + func doLeakCheck() { + CleanupSync.Wait() liveObjectsMu.Lock() defer liveObjectsMu.Unlock() leaked := len(liveObjects) diff --git a/pkg/sentry/socket/netstack/BUILD b/pkg/sentry/socket/netstack/BUILD index f772ebdc7..b5a036354 100644 --- a/pkg/sentry/socket/netstack/BUILD +++ b/pkg/sentry/socket/netstack/BUILD @@ -26,6 +26,7 @@ go_library( "//pkg/marshal", "//pkg/marshal/primitive", "//pkg/metric", + "//pkg/refs", "//pkg/sentry/arch", "//pkg/sentry/device", "//pkg/sentry/fsimpl/sockfs", diff --git a/pkg/sentry/socket/netstack/stack.go b/pkg/sentry/socket/netstack/stack.go index 958a96f81..c92a50ad7 100644 --- a/pkg/sentry/socket/netstack/stack.go +++ b/pkg/sentry/socket/netstack/stack.go @@ -21,6 +21,7 @@ import ( "gvisor.dev/gvisor/pkg/abi/linux" "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/log" + "gvisor.dev/gvisor/pkg/refs" "gvisor.dev/gvisor/pkg/sentry/inet" "gvisor.dev/gvisor/pkg/syserr" "gvisor.dev/gvisor/pkg/tcpip" @@ -41,6 +42,11 @@ type Stack struct { // Destroy implements inet.Stack.Destroy. func (s *Stack) Destroy() { s.Stack.Close() + refs.CleanupSync.Add(1) + go func() { + s.Stack.Wait() + refs.CleanupSync.Done() + }() } // SupportsIPv6 implements Stack.SupportsIPv6. diff --git a/pkg/tcpip/link/sharedmem/sharedmem.go b/pkg/tcpip/link/sharedmem/sharedmem.go index 5c0bfe9bd..1810bea5d 100644 --- a/pkg/tcpip/link/sharedmem/sharedmem.go +++ b/pkg/tcpip/link/sharedmem/sharedmem.go @@ -259,6 +259,10 @@ func (e *endpoint) Wait() { // Attach implements stack.LinkEndpoint.Attach. It launches the goroutine that // reads packets from the rx queue. func (e *endpoint) Attach(dispatcher stack.NetworkDispatcher) { + if dispatcher == nil { + e.Close() + return + } e.mu.Lock() if !e.workerStarted && e.stopRequested.Load() == 0 { e.workerStarted = true diff --git a/pkg/tcpip/network/ip_test.go b/pkg/tcpip/network/ip_test.go index d2d9a3059..dead13c60 100644 --- a/pkg/tcpip/network/ip_test.go +++ b/pkg/tcpip/network/ip_test.go @@ -637,6 +637,7 @@ func TestIPv4Send(t *testing.T) { if err != nil { t.Fatalf("could not find route: %v", err) } + defer r.Release() if err := ep.WritePacket(r, stack.NetworkHeaderParams{ Protocol: 123, TTL: 123, @@ -1091,6 +1092,7 @@ func TestIPv6Send(t *testing.T) { if err != nil { t.Fatalf("could not find route: %v", err) } + defer r.Release() if err := ep.WritePacket(r, stack.NetworkHeaderParams{ Protocol: 123, TTL: 123, diff --git a/pkg/tcpip/network/ipv4/ipv4_test.go b/pkg/tcpip/network/ipv4/ipv4_test.go index 8d3bfa10d..3b6c9ca8a 100644 --- a/pkg/tcpip/network/ipv4/ipv4_test.go +++ b/pkg/tcpip/network/ipv4/ipv4_test.go @@ -2117,6 +2117,7 @@ func TestFragmentationWritePacket(t *testing.T) { ep := iptestutil.NewMockLinkEndpoint(ft.mtu, nil, math.MaxInt32) defer ep.Close() r := buildRoute(t, ctx, ep) + defer r.Release() pkt := iptestutil.MakeRandPkt(ft.transportHeaderLength, extraHeaderReserve+header.IPv4MinimumSize, []int{ft.payloadSize}, header.IPv4ProtocolNumber) defer pkt.DecRef() source := pkt.Clone() @@ -2220,6 +2221,7 @@ func TestFragmentationErrors(t *testing.T) { ep := iptestutil.NewMockLinkEndpoint(ft.mtu, ft.mockError, ft.allowPackets) defer ep.Close() r := buildRoute(t, ctx, ep) + defer r.Release() pkt := iptestutil.MakeRandPkt(ft.transportHeaderLength, extraHeaderReserve+header.IPv4MinimumSize, []int{ft.payloadSize}, header.IPv4ProtocolNumber) defer pkt.DecRef() err := r.WritePacket(stack.NetworkHeaderParams{ @@ -3429,6 +3431,7 @@ func TestWriteStats(t *testing.T) { ep := iptestutil.NewMockLinkEndpoint(header.IPv4MinimumMTU, &tcpip.ErrInvalidEndpointState{}, test.allowPackets) defer ep.Close() rt := buildRoute(t, ctx, ep) + defer rt.Release() test.setup(t, rt.Stack()) nWritten := 0 diff --git a/pkg/tcpip/network/ipv6/ipv6_test.go b/pkg/tcpip/network/ipv6/ipv6_test.go index 6b4f315e5..a42f481d4 100644 --- a/pkg/tcpip/network/ipv6/ipv6_test.go +++ b/pkg/tcpip/network/ipv6/ipv6_test.go @@ -2586,6 +2586,7 @@ func TestWriteStats(t *testing.T) { defer ep.Close() rt := buildRoute(t, c, ep) + defer rt.Release() test.setup(t, rt.Stack()) nWritten := 0 @@ -2794,6 +2795,7 @@ func TestFragmentationWritePacket(t *testing.T) { defer ep.Close() r := buildRoute(t, c, ep) + defer r.Release() err := r.WritePacket(stack.NetworkHeaderParams{ Protocol: tcp.ProtocolNumber, TTL: ttl, @@ -2896,6 +2898,7 @@ func TestFragmentationErrors(t *testing.T) { defer ep.Close() r := buildRoute(t, c, ep) + defer r.Release() err := r.WritePacket(stack.NetworkHeaderParams{ Protocol: tcp.ProtocolNumber, TTL: ttl, diff --git a/pkg/tcpip/stack/addressable_endpoint_state.go b/pkg/tcpip/stack/addressable_endpoint_state.go index 5d8ea1c8f..4aa734377 100644 --- a/pkg/tcpip/stack/addressable_endpoint_state.go +++ b/pkg/tcpip/stack/addressable_endpoint_state.go @@ -684,12 +684,6 @@ func (a *AddressableEndpointState) Cleanup() { } } -// LeakCheckDisabled suppress reference leak warnings. -// FIXME(b/261201456): Re-enable after fixing the bug. -func (obj *addressStateRefs) LeakCheckDisabled() bool { - return true -} - var _ AddressEndpoint = (*addressState)(nil) // addressState holds state for an address. diff --git a/pkg/tcpip/stack/stack.go b/pkg/tcpip/stack/stack.go index 8e212b4ab..762dd1e7b 100644 --- a/pkg/tcpip/stack/stack.go +++ b/pkg/tcpip/stack/stack.go @@ -1764,6 +1764,12 @@ func (s *Stack) Wait() { } } +// Destroy destroys the stack with all endpoints. +func (s *Stack) Destroy() { + s.Close() + s.Wait() +} + // Pause pauses any protocol level background workers. func (s *Stack) Pause() { for _, p := range s.transportProtocols { diff --git a/pkg/tcpip/tests/integration/forward_test.go b/pkg/tcpip/tests/integration/forward_test.go index 7f5280ba7..82702bf4a 100644 --- a/pkg/tcpip/tests/integration/forward_test.go +++ b/pkg/tcpip/tests/integration/forward_test.go @@ -235,8 +235,11 @@ func TestForwarding(t *testing.T) { } host1Stack := stack.New(stackOpts) + defer host1Stack.Destroy() routerStack := stack.New(stackOpts) + defer routerStack.Destroy() host2Stack := stack.New(stackOpts) + defer host2Stack.Destroy() utils.SetupRoutedStacks(t, host1Stack, routerStack, host2Stack) epsAndAddrs := test.epAndAddrs(t, host1Stack, routerStack, host2Stack, subTest.proto) diff --git a/pkg/tcpip/tests/integration/iptables_test.go b/pkg/tcpip/tests/integration/iptables_test.go index 078a246df..26a858257 100644 --- a/pkg/tcpip/tests/integration/iptables_test.go +++ b/pkg/tcpip/tests/integration/iptables_test.go @@ -331,6 +331,7 @@ func TestIPTablesStatsForInput(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { s, e := test.setupStack(t) + defer s.Destroy() test.setupFilter(t, s) e.InjectInbound(test.proto, test.genPacket()) @@ -588,6 +589,7 @@ func TestIPTableWritePackets(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) + defer s.Destroy() e := channelEndpoint{ Endpoint: channel.New(4, header.IPv6MinimumMTU, linkAddr), t: t, @@ -843,6 +845,7 @@ func TestForwardingHook(t *testing.T) { s := stack.New(stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, }) + defer s.Destroy() subTest.setupFilter(t, s, test.netProto) @@ -1062,6 +1065,7 @@ func TestFilteringEchoPacketsWithLocalForwarding(t *testing.T) { s := stack.New(stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, }) + defer s.Destroy() subTest.setupFilter(t, s, test.netProto) @@ -1523,6 +1527,7 @@ func TestNATEcho(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{icmp.NewProtocol4, icmp.NewProtocol6}, }) + defer s.Destroy() ep1 := channel.New(1, header.IPv6MinimumMTU, "") ep2 := channel.New(1, header.IPv6MinimumMTU, "") @@ -1904,8 +1909,11 @@ func TestNAT(t *testing.T) { } host1Stack := stack.New(stackOpts) + defer host1Stack.Destroy() routerStack := stack.New(stackOpts) + defer routerStack.Destroy() host2Stack := stack.New(stackOpts) + defer host2Stack.Destroy() utils.SetupRoutedStacks(t, host1Stack, routerStack, host2Stack) epsAndAddrs := test.epAndAddrs(t, host1Stack, routerStack, host2Stack, subTest.proto) @@ -2455,6 +2463,7 @@ func TestNATICMPError(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol, tcp.NewProtocol}, }) + defer s.Destroy() ep1 := channel.New(1, header.IPv6MinimumMTU, "") ep2 := channel.New(1, header.IPv6MinimumMTU, "") @@ -2833,6 +2842,7 @@ func TestSNATHandlePortOrIdentConflicts(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol, tcp.NewProtocol}, }) + defer s.Destroy() ep1 := channel.New(1, header.IPv6MinimumMTU, "") ep2 := channel.New(1, header.IPv6MinimumMTU, "") @@ -2941,6 +2951,7 @@ func TestLocallyRoutedPackets(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) + defer s.Destroy() if err := s.CreateNIC(nicID, loopback.New()); err != nil { t.Fatalf("CreateNIC(%d, _) = %s", nicID, err) @@ -3274,6 +3285,7 @@ func TestRejectWith(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol, tcp.NewProtocol}, }) + defer s.Destroy() ep1 := channel.New(1, header.IPv6MinimumMTU, "") ep2 := channel.New(1, header.IPv6MinimumMTU, "") diff --git a/pkg/tcpip/tests/integration/istio_test.go b/pkg/tcpip/tests/integration/istio_test.go index cd74bd6d1..378965678 100644 --- a/pkg/tcpip/tests/integration/istio_test.go +++ b/pkg/tcpip/tests/integration/istio_test.go @@ -93,9 +93,9 @@ type testContext struct { func (ctx *testContext) cleanup() { ctx.localServerListener.Close() - ctx.localStack.Close() + ctx.localStack.Destroy() ctx.remoteServerListener.Close() - ctx.remoteStack.Close() + ctx.remoteStack.Destroy() ctx.wg.Wait() } @@ -120,7 +120,6 @@ func newTestContext(t *testing.T) *testContext { TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, HandleLocal: true, }) - remoteStack := stack.New(stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, diff --git a/pkg/tcpip/tests/integration/link_resolution_test.go b/pkg/tcpip/tests/integration/link_resolution_test.go index a0c6ae5e5..887227f7c 100644 --- a/pkg/tcpip/tests/integration/link_resolution_test.go +++ b/pkg/tcpip/tests/integration/link_resolution_test.go @@ -160,7 +160,9 @@ func TestPing(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{icmp.NewProtocol4, icmp.NewProtocol6}, } - host1Stack, _ := setupStack(t, stackOpts, host1NICID, host2NICID) + host1Stack, host2Stack := setupStack(t, stackOpts, host1NICID, host2NICID) + defer host1Stack.Destroy() + defer host2Stack.Destroy() var wq waiter.Queue we, waiterCH := waiter.NewChannelEntry(waiter.ReadableEvents) @@ -304,6 +306,8 @@ func TestTCPLinkResolutionFailure(t *testing.T) { } host1Stack, host2Stack := setupStack(t, stackOpts, host1NICID, host2NICID) + defer host1Stack.Destroy() + defer host2Stack.Destroy() var listenerWQ waiter.Queue listenerEP, err := host2Stack.NewEndpoint(tcp.ProtocolNumber, test.netProto, &listenerWQ) @@ -571,6 +575,7 @@ func TestForwardingWithLinkResolutionFailure(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{test.transportProtocol}, Clock: clock, }) + defer s.Destroy() // Set up endpoint through which we will receive packets. incomingEndpoint := channel.New(1, test.mtu, "") @@ -724,7 +729,9 @@ func TestGetLinkAddress(t *testing.T) { Clock: clock, } - host1Stack, _ := setupStack(t, stackOpts, host1NICID, host2NICID) + host1Stack, host2Stack := setupStack(t, stackOpts, host1NICID, host2NICID) + defer host1Stack.Destroy() + defer host2Stack.Destroy() ch := make(chan stack.LinkResolutionResult, 1) err := host1Stack.GetLinkAddress(host1NICID, test.remoteAddr, test.localAddr, test.netProto, func(r stack.LinkResolutionResult) { @@ -833,7 +840,9 @@ func TestRouteResolvedFields(t *testing.T) { Clock: clock, } - host1Stack, _ := setupStack(t, stackOpts, host1NICID, host2NICID) + host1Stack, host2Stack := setupStack(t, stackOpts, host1NICID, host2NICID) + defer host1Stack.Destroy() + defer host2Stack.Destroy() r, err := host1Stack.FindRoute(host1NICID, test.localAddr, test.remoteAddr, test.netProto, false /* multicastLoop */) if err != nil { t.Fatalf("host1Stack.FindRoute(%d, %s, %s, %d, false): %s", host1NICID, test.localAddr, test.remoteAddr, test.netProto, err) @@ -935,6 +944,8 @@ func TestWritePacketsLinkResolution(t *testing.T) { } host1Stack, host2Stack := setupStack(t, stackOpts, host1NICID, host2NICID) + defer host1Stack.Destroy() + defer host2Stack.Destroy() var serverWQ waiter.Queue serverWE, serverCH := waiter.NewChannelEntry(waiter.ReadableEvents) @@ -1333,8 +1344,11 @@ func TestTCPConfirmNeighborReachability(t *testing.T) { host1StackOpts.NUDDisp = &nudDisp host1Stack := stack.New(host1StackOpts) + defer host1Stack.Destroy() routerStack := stack.New(stackOpts) + defer routerStack.Destroy() host2Stack := stack.New(stackOpts) + defer host2Stack.Destroy() utils.SetupRoutedStacks(t, host1Stack, routerStack, host2Stack) // Add a reachable dynamic entry to our neighbor table for the remote. @@ -1592,7 +1606,9 @@ func TestDAD(t *testing.T) { }, } - host1Stack, _ := setupStack(t, stackOpts, utils.Host1NICID, utils.Host2NICID) + host1Stack, host2Stack := setupStack(t, stackOpts, utils.Host1NICID, utils.Host2NICID) + defer host1Stack.Destroy() + defer host2Stack.Destroy() // DAD should be disabled by default. if res, err := host1Stack.CheckDuplicateAddress(utils.Host1NICID, test.netProto, test.remoteAddr, func(r stack.DADResult) { @@ -1753,6 +1769,8 @@ func TestUpdateCachedNeighborEntry(t *testing.T) { host1Stack := stack.New(stackOpts) host2Stack := stack.New(stackOpts) + defer host1Stack.Destroy() + defer host2Stack.Destroy() host1Pipe, host2Pipe := pipe.New(utils.LinkAddr1, utils.LinkAddr2, maxFrameSize) diff --git a/pkg/tcpip/tests/integration/loopback_test.go b/pkg/tcpip/tests/integration/loopback_test.go index b3e7052a5..f09809e1d 100644 --- a/pkg/tcpip/tests/integration/loopback_test.go +++ b/pkg/tcpip/tests/integration/loopback_test.go @@ -85,6 +85,7 @@ func TestInitialLoopbackAddresses(t *testing.T) { }, })}, }) + defer s.Destroy() if err := s.CreateNIC(nicID, loopback.New()); err != nil { t.Fatalf("CreateNIC(%d, _): %s", nicID, err) @@ -193,6 +194,7 @@ func TestLoopbackAcceptAllInSubnetUDP(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) + defer s.Destroy() if err := s.CreateNIC(nicID, loopback.New()); err != nil { t.Fatalf("CreateNIC(%d, _): %s", nicID, err) } @@ -288,6 +290,7 @@ func TestLoopbackSubnetLifetimeBoundToAddr(t *testing.T) { s := stack.New(stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol}, }) + defer s.Destroy() if err := s.CreateNIC(nicID, loopback.New()); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) } @@ -429,6 +432,7 @@ func TestLoopbackAcceptAllInSubnetTCP(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, }) + defer s.Destroy() if err := s.CreateNIC(nicID, loopback.New()); err != nil { t.Fatalf("CreateNIC(%d, _): %s", nicID, err) } @@ -690,6 +694,7 @@ func TestExternalLoopbackTraffic(t *testing.T) { }, TransportProtocols: []stack.TransportProtocolFactory{icmp.NewProtocol4, icmp.NewProtocol6}, }) + defer s.Destroy() e := channel.New(1, header.IPv6MinimumMTU, "") if err := s.CreateNIC(nicID1, e); err != nil { t.Fatalf("CreateNIC(%d, _): %s", nicID1, err) diff --git a/pkg/tcpip/tests/integration/multicast_forward_test.go b/pkg/tcpip/tests/integration/multicast_forward_test.go index 7df740446..2179e7993 100644 --- a/pkg/tcpip/tests/integration/multicast_forward_test.go +++ b/pkg/tcpip/tests/integration/multicast_forward_test.go @@ -420,7 +420,7 @@ func TestAddMulticastRoute(t *testing.T) { s := stack.New(stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, }) - defer s.Close() + defer s.Destroy() endpoints := make(map[tcpip.NICID]*channel.Endpoint) for nicID, addrType := range endpointConfigs { @@ -550,7 +550,7 @@ func TestEnableMulticastForwardingE(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) - defer s.Close() + defer s.Destroy() for _, wantResult := range test.wantResult { alreadyEnabled, err := s.EnableMulticastForwardingForProtocol(protocol, test.eventDispatcher) @@ -641,7 +641,7 @@ func TestMulticastRouteLastUsedTime(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, Clock: clock, }) - defer s.Close() + defer s.Destroy() if _, err := s.EnableMulticastForwardingForProtocol(protocol, &fakeMulticastEventDispatcher{}); err != nil { t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", protocol, err) @@ -797,7 +797,7 @@ func TestRemoveMulticastRoute(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) - defer s.Close() + defer s.Destroy() if _, err := s.EnableMulticastForwardingForProtocol(protocol, &fakeMulticastEventDispatcher{}); err != nil { t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", protocol, err) @@ -1039,7 +1039,7 @@ func TestMulticastForwarding(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, }) - defer s.Close() + defer s.Destroy() eventDispatcher, ok := eventDispatchers[protocol] if !ok { diff --git a/pkg/tcpip/tests/integration/route_test.go b/pkg/tcpip/tests/integration/route_test.go index db5f447df..e968f91e7 100644 --- a/pkg/tcpip/tests/integration/route_test.go +++ b/pkg/tcpip/tests/integration/route_test.go @@ -179,6 +179,7 @@ func TestLocalPing(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{icmp.NewProtocol4, icmp.NewProtocol6}, HandleLocal: true, }) + defer s.Destroy() e := test.linkEndpoint() if err := s.CreateNIC(nicID, e); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) @@ -301,6 +302,7 @@ func TestLocalUDP(t *testing.T) { } s := stack.New(stackOpts) + defer s.Destroy() ep := channel.New(1, header.IPv6MinimumMTU, "") if err := s.CreateNIC(nicID, ep); err != nil { diff --git a/pkg/tcpip/transport/icmp/icmp_test.go b/pkg/tcpip/transport/icmp/icmp_test.go index 51626d448..7930d91a2 100644 --- a/pkg/tcpip/transport/icmp/icmp_test.go +++ b/pkg/tcpip/transport/icmp/icmp_test.go @@ -97,6 +97,7 @@ func TestWriteUnboundWithBindToDevice(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{icmp.NewProtocol4}, HandleLocal: true, }) + defer s.Destroy() // Add two NICs, both with default routes on the same subnet. The first NIC // added will be the default NIC for that subnet. diff --git a/pkg/tcpip/transport/internal/network/endpoint.go b/pkg/tcpip/transport/internal/network/endpoint.go index 11c76e920..fa062ea9d 100644 --- a/pkg/tcpip/transport/internal/network/endpoint.go +++ b/pkg/tcpip/transport/internal/network/endpoint.go @@ -648,6 +648,7 @@ func (e *Endpoint) ConnectAndThen(addr tcpip.FullAddress, f func(netProto tcpip. } if err := f(r.NetProto(), info.ID, id); err != nil { + r.Release() return err } diff --git a/pkg/tcpip/transport/internal/network/endpoint_test.go b/pkg/tcpip/transport/internal/network/endpoint_test.go index b5ea1882b..aeb1f507b 100644 --- a/pkg/tcpip/transport/internal/network/endpoint_test.go +++ b/pkg/tcpip/transport/internal/network/endpoint_test.go @@ -122,6 +122,7 @@ func TestEndpointStateTransitions(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, Clock: &faketime.NullClock{}, }) + defer s.Destroy() e := channel.New(1, header.IPv6MinimumMTU, "") if err := s.CreateNIC(nicID, e); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) @@ -271,6 +272,7 @@ func TestBindNICID(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, Clock: &faketime.NullClock{}, }) + defer s.Destroy() if err := s.CreateNIC(nicID, loopback.New()); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) } diff --git a/pkg/tcpip/transport/raw/endpoint.go b/pkg/tcpip/transport/raw/endpoint.go index 64b588403..dabe6587a 100644 --- a/pkg/tcpip/transport/raw/endpoint.go +++ b/pkg/tcpip/transport/raw/endpoint.go @@ -345,6 +345,7 @@ func (e *endpoint) write(p tcpip.Payloader, opts tcpip.WriteOptions) (int64, tcp if err != nil { return 0, err } + defer ctx.Release() if p.Len() > int(ctx.MTU()) { return 0, &tcpip.ErrMessageTooLong{} diff --git a/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go b/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go index 6e7339d76..073cda125 100644 --- a/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go +++ b/pkg/tcpip/transport/tcp/test/e2e/tcp_test.go @@ -285,6 +285,7 @@ func TestCloseWithoutConnect(t *testing.T) { } c.EP.Close() + c.EP = nil if got := c.Stack().Stats().TCP.CurrentConnected.Value(); got != 0 { t.Errorf("got stats.TCP.CurrentConnected.Value() = %d, want = 0", got) @@ -1576,6 +1577,7 @@ func TestListenerReadinessOnEvent(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol}, }) + defer s.Destroy() { ep := loopback.New() if testing.Verbose() { @@ -2428,6 +2430,7 @@ func TestSmallReceiveBufferReadiness(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, }) + defer s.Destroy() ep := loopback.New() if testing.Verbose() { @@ -5322,6 +5325,7 @@ func TestDefaultBufferSizes(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}, }) + defer s.Destroy() // Check the default values. ep, err := s.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &waiter.Queue{}) @@ -5385,6 +5389,7 @@ func TestBindToDeviceOption(t *testing.T) { NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol}}) + defer s.Destroy() ep, err := s.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &waiter.Queue{}) if err != nil { t.Fatalf("NewEndpoint failed; %s", err) @@ -5485,6 +5490,7 @@ func TestSelfConnect(t *testing.T) { if err != nil { t.Fatal(err) } + defer s.Destroy() var wq waiter.Queue ep, err := s.NewEndpoint(tcp.ProtocolNumber, ipv4.ProtocolNumber, &wq) @@ -5589,6 +5595,7 @@ func TestConnectAvoidsBoundPorts(t *testing.T) { if err != nil { t.Fatal(err) } + defer s.Destroy() var wq waiter.Queue var eps []tcpip.Endpoint @@ -8862,6 +8869,7 @@ func TestHandshakeRTT(t *testing.T) { func TestSetRTO(t *testing.T) { c := context.New(t, e2e.DefaultMTU) minRTO, maxRTO := tcpRTOMinMax(t, c) + c.Cleanup() for _, tt := range []struct { name string RTO time.Duration @@ -8890,6 +8898,7 @@ func TestSetRTO(t *testing.T) { } { t.Run(tt.name, func(t *testing.T) { c := context.New(t, e2e.DefaultMTU) + defer c.Cleanup() var opt tcpip.SettableTransportProtocolOption if tt.minRTO > 0 { min := tcpip.TCPMinRTOOption(tt.minRTO) diff --git a/pkg/tcpip/transport/testing/context/context.go b/pkg/tcpip/transport/testing/context/context.go index bfbd9da1b..6e45aabb0 100644 --- a/pkg/tcpip/transport/testing/context/context.go +++ b/pkg/tcpip/transport/testing/context/context.go @@ -156,6 +156,8 @@ func (c *Context) Cleanup() { if c.EP != nil { c.EP.Close() } + c.Stack.Destroy() + c.Stack = nil refs.DoRepeatedLeakCheck() } diff --git a/pkg/tcpip/transport/udp/endpoint.go b/pkg/tcpip/transport/udp/endpoint.go index 28f6feaf7..b3d448b2d 100644 --- a/pkg/tcpip/transport/udp/endpoint.go +++ b/pkg/tcpip/transport/udp/endpoint.go @@ -353,7 +353,10 @@ var _ tcpip.EndpointWithPreflight = (*endpoint)(nil) // is specified, binds the endpoint to that address. func (e *endpoint) Preflight(opts tcpip.WriteOptions) tcpip.Error { var r bytes.Reader - _, err := e.prepareForWrite(&r, opts) + udpInfo, err := e.prepareForWrite(&r, opts) + if err == nil { + udpInfo.ctx.Release() + } return err } diff --git a/pkg/tcpip/transport/udp/udp_test.go b/pkg/tcpip/transport/udp/udp_test.go index b76f2d68e..8ba4ceda1 100644 --- a/pkg/tcpip/transport/udp/udp_test.go +++ b/pkg/tcpip/transport/udp/udp_test.go @@ -85,6 +85,7 @@ func TestBindToDeviceOption(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, Clock: &faketime.NullClock{}, }) + defer s.Destroy() ep, err := s.NewEndpoint(udp.ProtocolNumber, ipv4.ProtocolNumber, &waiter.Queue{}) if err != nil { @@ -676,7 +677,6 @@ func TestDualWriteConnectedToV6(t *testing.T) { {writeOpSequence: []writeOperation{preflight, write}, expectedNoRouteErrCount: 1}, } { c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) @@ -693,116 +693,133 @@ func TestDualWriteConnectedToV6(t *testing.T) { if got := c.EP.Stats().(*tcpip.TransportEndpointStats).SendErrors.NoRoute.Value(); got != testCase.expectedNoRouteErrCount { c.T.Fatalf("Endpoint stat not updated. got %d want %d", got, testCase.expectedNoRouteErrCount) } + c.Cleanup() } } -var writeOpSequences [][]writeOperation = [][]writeOperation{ - []writeOperation{write}, - []writeOperation{preflight}, - []writeOperation{preflight, write}, +var writeOpSequences = map[string]([]writeOperation){ + "write": []writeOperation{write}, + "preflight": []writeOperation{preflight}, + "write|preflight": []writeOperation{preflight, write}, } func TestDualWriteConnectedToV4Mapped(t *testing.T) { - for _, writeOpSequence := range writeOpSequences { - c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() + for name, writeOpSequence := range writeOpSequences { + t.Run(name, func(t *testing.T) { + c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) + defer c.Cleanup() - c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) + c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) - // Connect to v4 mapped address. - if err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV4MappedAddr, Port: context.TestPort}); err != nil { - c.T.Fatalf("Bind failed: %s", err) - } + // Connect to v4 mapped address. + if err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV4MappedAddr, Port: context.TestPort}); err != nil { + c.T.Fatalf("Bind failed: %s", err) + } - testWriteOpSequenceSucceeds(c, context.UnicastV4in6, writeOpSequence) + testWriteOpSequenceSucceeds(c, context.UnicastV4in6, writeOpSequence) - // Write to v6 address. - testWriteOpSequenceFails(c, context.UnicastV6, writeOpSequence, &tcpip.ErrInvalidEndpointState{}) + // Write to v6 address. + testWriteOpSequenceFails(c, context.UnicastV6, writeOpSequence, &tcpip.ErrInvalidEndpointState{}) + }) } } func TestPreflightBindsEndpoint(t *testing.T) { - for _, ipProtocolNumber := range []tcpip.NetworkProtocolNumber{ipv6.ProtocolNumber, ipv4.ProtocolNumber} { - c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol}) - defer c.Cleanup() + protocols := map[string]tcpip.NetworkProtocolNumber{ + "ipv4": ipv4.ProtocolNumber, + "ipv6": ipv6.ProtocolNumber, + } + for name, ipProtocolNumber := range protocols { + t.Run(name, func(t *testing.T) { + c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol}) + defer c.Cleanup() - c.CreateEndpoint(ipProtocolNumber, udp.ProtocolNumber) + c.CreateEndpoint(ipProtocolNumber, udp.ProtocolNumber) - flow := context.UnicastV6 - h := flow.MakeHeader4Tuple(context.Outgoing) - writeDstAddr := flow.MapAddrIfApplicable(h.Dst.Addr) - writeOpts := tcpip.WriteOptions{ - To: &tcpip.FullAddress{Addr: writeDstAddr, Port: h.Dst.Port}, - } + flow := context.UnicastV6 + h := flow.MakeHeader4Tuple(context.Outgoing) + writeDstAddr := flow.MapAddrIfApplicable(h.Dst.Addr) + writeOpts := tcpip.WriteOptions{ + To: &tcpip.FullAddress{Addr: writeDstAddr, Port: h.Dst.Port}, + } - if err := getEndpointWithPreflight(c).Preflight(writeOpts); err != nil { - c.T.Fatalf("Preflight failed: %s", err) - } + if err := getEndpointWithPreflight(c).Preflight(writeOpts); err != nil { + c.T.Fatalf("Preflight failed: %s", err) + } - if c.EP.State() != uint32(transport.DatagramEndpointStateBound) { - c.T.Fatalf("Expect UDP endpoint in state %d, found %d", transport.DatagramEndpointStateBound, c.EP.State()) - } + if c.EP.State() != uint32(transport.DatagramEndpointStateBound) { + c.T.Fatalf("Expect UDP endpoint in state %d, found %d", transport.DatagramEndpointStateBound, c.EP.State()) + } + }) } } func TestV4WriteOnV6Only(t *testing.T) { - for _, writeOpSequence := range writeOpSequences { - c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() + for name, writeOpSequence := range writeOpSequences { + t.Run(name, func(t *testing.T) { + c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) + defer c.Cleanup() - c.CreateEndpointForFlow(context.UnicastV6Only, udp.ProtocolNumber) + c.CreateEndpointForFlow(context.UnicastV6Only, udp.ProtocolNumber) - // Write to V4 mapped address. - testWriteOpSequenceFails(c, context.UnicastV4in6, writeOpSequence, &tcpip.ErrHostUnreachable{}) + // Write to V4 mapped address. + testWriteOpSequenceFails(c, context.UnicastV4in6, writeOpSequence, &tcpip.ErrHostUnreachable{}) + }) } } func TestV6WriteOnBoundToV4Mapped(t *testing.T) { - for _, writeOpSequence := range writeOpSequences { - c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() + for name, writeOpSequence := range writeOpSequences { + t.Run(name, func(t *testing.T) { + c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) + defer c.Cleanup() - c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) + c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) - // Bind to v4 mapped address. - if err := c.EP.Bind(tcpip.FullAddress{Addr: context.StackV4MappedAddr, Port: context.StackPort}); err != nil { - c.T.Fatalf("Bind failed: %s", err) - } + // Bind to v4 mapped address. + if err := c.EP.Bind(tcpip.FullAddress{Addr: context.StackV4MappedAddr, Port: context.StackPort}); err != nil { + c.T.Fatalf("Bind failed: %s", err) + } - // Write to v6 address. - testWriteOpSequenceFails(c, context.UnicastV6, writeOpSequence, &tcpip.ErrInvalidEndpointState{}) + // Write to v6 address. + testWriteOpSequenceFails(c, context.UnicastV6, writeOpSequence, &tcpip.ErrInvalidEndpointState{}) + }) } } func TestV6WriteOnConnected(t *testing.T) { - for _, writeOpSequence := range writeOpSequences { - c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() + for name, writeOpSequence := range writeOpSequences { + t.Run(name, func(t *testing.T) { + c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) + defer c.Cleanup() - c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) + c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) - // Connect to v6 address. - if err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV6Addr, Port: context.TestPort}); err != nil { - c.T.Fatalf("Connect failed: %s", err) - } + // Connect to v6 address. + if err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV6Addr, Port: context.TestPort}); err != nil { + c.T.Fatalf("Connect failed: %s", err) + } - testWriteOpSequenceSucceedsNoDestination(c, context.UnicastV6, writeOpSequence) + testWriteOpSequenceSucceedsNoDestination(c, context.UnicastV6, writeOpSequence) + }) } } func TestV4WriteOnConnected(t *testing.T) { - for _, writeOpSequence := range writeOpSequences { - c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() + for name, writeOpSequence := range writeOpSequences { + t.Run(name, func(t *testing.T) { + c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) + defer c.Cleanup() - c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) + c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) - // Connect to v4 mapped address. - if err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV4MappedAddr, Port: context.TestPort}); err != nil { - c.T.Fatalf("Connect failed: %s", err) - } + // Connect to v4 mapped address. + if err := c.EP.Connect(tcpip.FullAddress{Addr: context.TestV4MappedAddr, Port: context.TestPort}); err != nil { + c.T.Fatalf("Connect failed: %s", err) + } - testWriteOpSequenceSucceedsNoDestination(c, context.UnicastV4, writeOpSequence) + testWriteOpSequenceSucceedsNoDestination(c, context.UnicastV4, writeOpSequence) + }) } } @@ -1928,7 +1945,6 @@ func TestShutdownRead(t *testing.T) { func TestShutdownWrite(t *testing.T) { for _, writeOpSequence := range writeOpSequences { c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) @@ -1941,6 +1957,7 @@ func TestShutdownWrite(t *testing.T) { } testWriteOpSequenceFails(c, context.UnicastV6, writeOpSequence, &tcpip.ErrClosedForSend{}) + c.Cleanup() } } @@ -2074,6 +2091,7 @@ func TestOutgoingSubnetBroadcast(t *testing.T) { TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, Clock: &faketime.NullClock{}, }) + defer s.Destroy() e := channel.New(0, context.DefaultMTU, "") defer e.Close() if err := s.CreateNIC(nicID1, e); err != nil { @@ -2243,7 +2261,6 @@ func TestChecksumWithZeroValueOnesComplementSum(t *testing.T) { func TestWritePayloadSizeTooBig(t *testing.T) { for _, writeOpSequence := range writeOpSequences { c := context.New(t, []stack.TransportProtocolFactory{udp.NewProtocol, icmp.NewProtocol6, icmp.NewProtocol4}) - defer c.Cleanup() c.CreateEndpoint(ipv6.ProtocolNumber, udp.ProtocolNumber) @@ -2261,6 +2278,7 @@ func TestWritePayloadSizeTooBig(t *testing.T) { testPreflightSucceeds(c, context.UnicastV6) } } + c.Cleanup() } }