From 6f5e475674d6c200e4f09f9f0edd1ef15d3fc0e0 Mon Sep 17 00:00:00 2001 From: Ghanan Gowripalan Date: Thu, 26 May 2022 14:04:28 -0700 Subject: [PATCH] Loop multicast packets from raw sockets Bug: https://fxbug.dev/101226 PiperOrigin-RevId: 451239264 --- pkg/tcpip/transport/BUILD | 1 + pkg/tcpip/transport/datagram_test.go | 146 +++++++++++++++++++++++++++ pkg/tcpip/transport/raw/endpoint.go | 1 + 3 files changed, 148 insertions(+) diff --git a/pkg/tcpip/transport/BUILD b/pkg/tcpip/transport/BUILD index b53118163..624381abd 100644 --- a/pkg/tcpip/transport/BUILD +++ b/pkg/tcpip/transport/BUILD @@ -30,5 +30,6 @@ go_test( "//pkg/tcpip/transport/raw", "//pkg/tcpip/transport/udp", "//pkg/waiter", + "@com_github_google_go_cmp//cmp:go_default_library", ], ) diff --git a/pkg/tcpip/transport/datagram_test.go b/pkg/tcpip/transport/datagram_test.go index 2905feeeb..a1e5c1851 100644 --- a/pkg/tcpip/transport/datagram_test.go +++ b/pkg/tcpip/transport/datagram_test.go @@ -21,6 +21,7 @@ import ( "math" "testing" + "github.com/google/go-cmp/cmp" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/buffer" "gvisor.dev/gvisor/pkg/tcpip/header" @@ -392,6 +393,7 @@ func TestDeviceReturnErrNoBufferSpace(t *testing.T) { name string createEndpoint func(*stack.Stack) (tcpip.Endpoint, error) }{ + // TODO(https://gvisor.dev/issues/7656): Also test ping sockets. { name: "UDP", createEndpoint: func(s *stack.Stack) (tcpip.Endpoint, error) { @@ -519,3 +521,147 @@ func TestDeviceReturnErrNoBufferSpace(t *testing.T) { }) } } + +func TestMulticastLoop(t *testing.T) { + const ( + nicID = 1 + port = 12345 + ) + + for _, netProto := range []struct { + name string + num tcpip.NetworkProtocolNumber + localAddr tcpip.AddressWithPrefix + destAddr tcpip.Address + rawSocketHdrLen int + }{ + { + name: "IPv4", + num: header.IPv4ProtocolNumber, + localAddr: testutil.MustParse4("1.2.3.4").WithPrefix(), + destAddr: header.IPv4AllSystems, + rawSocketHdrLen: header.IPv4MinimumSize, + }, + { + name: "IPv6", + num: header.IPv6ProtocolNumber, + localAddr: testutil.MustParse6("a::1").WithPrefix(), + destAddr: header.IPv6AllNodesMulticastAddress, + rawSocketHdrLen: 0, + }, + } { + t.Run(netProto.name, func(t *testing.T) { + for _, test := range []struct { + name string + createEndpoint func(*stack.Stack, *waiter.Queue) (tcpip.Endpoint, error) + includedHdrBytes int + }{ + { + name: "UDP", + createEndpoint: func(s *stack.Stack, wq *waiter.Queue) (tcpip.Endpoint, error) { + ep, err := s.NewEndpoint(udp.ProtocolNumber, netProto.num, wq) + if err != nil { + return nil, fmt.Errorf("s.NewEndpoint(%d, %d, _) failed: %s", udp.ProtocolNumber, netProto.num, err) + } + return ep, nil + }, + includedHdrBytes: 0, + }, + { + name: "RAW", + createEndpoint: func(s *stack.Stack, wq *waiter.Queue) (tcpip.Endpoint, error) { + ep, err := s.NewRawEndpoint(udp.ProtocolNumber, netProto.num, wq, true /* associated */) + if err != nil { + return nil, fmt.Errorf("s.NewRawEndpoint(%d, %d, _, true) failed: %s", udp.ProtocolNumber, netProto.num, err) + } + return ep, nil + }, + includedHdrBytes: netProto.rawSocketHdrLen, + }, + } { + t.Run(test.name, func(t *testing.T) { + s := stack.New(stack.Options{ + NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, + TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol}, + RawFactory: &raw.EndpointFactory{}, + }) + var e mockEndpoint + defer e.releasePackets() + if err := s.CreateNIC(nicID, &e); err != nil { + t.Fatalf("s.CreateNIC(%d, _) failed: %s", nicID, err) + } + var wq waiter.Queue + ep, err := test.createEndpoint(s, &wq) + if err != nil { + t.Fatalf("test.createEndpoint(_) failed: %s", err) + } + defer ep.Close() + + addr := tcpip.ProtocolAddress{ + Protocol: netProto.num, + AddressWithPrefix: netProto.localAddr, + } + if err := s.AddProtocolAddress(nicID, addr, stack.AddressProperties{}); err != nil { + t.Fatalf("AddProtocolAddress(%d, %#v, {}): %s", nicID, addr, err) + } + s.SetRouteTable([]tcpip.Route{ + { + Destination: header.IPv4EmptySubnet, + NIC: nicID, + }, + { + Destination: header.IPv6EmptySubnet, + NIC: nicID, + }, + }) + + bind := tcpip.FullAddress{Port: port} + if err := ep.Bind(bind); err != nil { + t.Fatalf("ep.Bind(%#v): %s", bind, err) + } + + to := tcpip.FullAddress{NIC: nicID, Addr: netProto.destAddr, Port: port} + checkWrite := func(buf []byte, withRead bool) { + t.Helper() + + { + var r bytes.Reader + r.Reset(buf[:]) + if n, err := ep.Write(&r, tcpip.WriteOptions{To: &to}); err != nil { + t.Fatalf("Write(...): %s", err) + } else if want := int64(len(buf)); n != want { + t.Fatalf("got Write(...) = %d, want = %d", n, want) + } + } + + var wantErr tcpip.Error + if !withRead { + wantErr = &tcpip.ErrWouldBlock{} + } + + var r bytes.Buffer + if _, err := ep.Read(&r, tcpip.ReadOptions{}); err != wantErr { + t.Fatalf("got Read(...) = %s, want = %s", err, wantErr) + } + if wantErr != nil { + return + } + + if diff := cmp.Diff(buf, r.Bytes()[test.includedHdrBytes:]); diff != "" { + t.Errorf("read data bytes mismatch (-want +got):\n%s", diff) + } + } + + checkWrite([]byte{1, 2, 3, 4}, true /* withRead */) + + ops := ep.SocketOptions() + ops.SetMulticastLoop(false) + checkWrite([]byte{5, 6, 7, 8}, false /* withRead */) + + ops.SetMulticastLoop(true) + checkWrite([]byte{9, 10, 11, 12}, true /* withRead */) + }) + } + }) + } +} diff --git a/pkg/tcpip/transport/raw/endpoint.go b/pkg/tcpip/transport/raw/endpoint.go index ba767abf5..496e972c7 100644 --- a/pkg/tcpip/transport/raw/endpoint.go +++ b/pkg/tcpip/transport/raw/endpoint.go @@ -131,6 +131,7 @@ func newEndpoint(s *stack.Stack, netProto tcpip.NetworkProtocolNumber, transProt ipv6ChecksumOffset: ipv6ChecksumOffset, } e.ops.InitHandler(e, e.stack, tcpip.GetStackSendBufferLimits, tcpip.GetStackReceiveBufferLimits) + e.ops.SetMulticastLoop(true) e.ops.SetHeaderIncluded(!associated) e.ops.SetSendBufferSize(32*1024, false /* notify */) e.ops.SetReceiveBufferSize(32*1024, false /* notify */)