Loop multicast packets from raw sockets

Bug: https://fxbug.dev/101226
PiperOrigin-RevId: 451239264
This commit is contained in:
Ghanan Gowripalan
2022-05-26 14:07:18 -07:00
committed by gVisor bot
parent 3dfcffdb15
commit 6f5e475674
3 changed files with 148 additions and 0 deletions
+1
View File
@@ -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",
],
)
+146
View File
@@ -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 */)
})
}
})
}
}
+1
View File
@@ -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 */)