Support unordered MGP record validation

This change updates the checker and network testutil packages to support
the validation of unordered MGP(v2) records which may be present across
one or more reports.

This allows us to share more code across the IP-generic and IP-specific
multicast group protocol tests while also preparing us for followup
changes which coalesces records in unpredictable order.

Updates #8346

PiperOrigin-RevId: 501910064
This commit is contained in:
Ghanan Gowripalan
2023-01-13 12:20:21 -08:00
committed by gVisor bot
parent 365a6f5490
commit 89062d5985
5 changed files with 401 additions and 657 deletions
+86 -51
View File
@@ -1223,27 +1223,47 @@ func MLDMaxRespDelay(want time.Duration) TransportChecker {
}
}
// MLDMulticastAddress creates a checker that checks the Multicast Address
// field of a MLD message.
// MLDMulticastAddressUnordered creates a checker that checks that the multicast
// address in the MLD message is expected to be seen.
//
// The seen address is removed from the expected groups map.
//
// The returned TransportChecker assumes that a valid ICMPv6 is passed to it
// containing a valid MLD message as far as the size is concerned.
func MLDMulticastAddress(want tcpip.Address) TransportChecker {
func MLDMulticastAddressUnordered(expectedGroups map[tcpip.Address]struct{}) TransportChecker {
return func(t *testing.T, h header.Transport) {
t.Helper()
icmp := h.(header.ICMPv6)
ns := header.MLD(icmp.MessageBody())
if got := ns.MulticastAddress(); got != want {
t.Errorf("got %T.MulticastAddress() = %s, want = %s", ns, got, want)
addr := ns.MulticastAddress()
if _, ok := expectedGroups[addr]; !ok {
t.Errorf("unexpected multicast group %s", addr)
} else {
delete(expectedGroups, addr)
}
}
}
// MLDMulticastAddress creates a checker that checks the Multicast Address
// field of a MLD message.
//
// The returned TransportChecker assumes that a valid ICMPv6 is passed to it
// containing a valid MLD message as far as the size is concerned.
func MLDMulticastAddress(want tcpip.Address) TransportChecker {
return MLDMulticastAddressUnordered(map[tcpip.Address]struct{}{
want: struct{}{},
})
}
// MLDv2Report creates a checker that checks that the packet contains a valid
// MLDv2 report with the specified records.
func MLDv2Report(expectedReport header.MLDv2ReportSerializer) NetworkChecker {
//
// Note that observed records are removed from expectedRecords. No error is
// logged if the report does not have all the records expected.
func MLDv2Report(expectedRecords map[tcpip.Address]header.MLDv2ReportRecordType) NetworkChecker {
return func(t *testing.T, h []header.Network) {
t.Helper()
@@ -1255,16 +1275,26 @@ func MLDv2Report(expectedReport header.MLDv2ReportSerializer) NetworkChecker {
last := h[len(h)-1]
icmp := header.ICMPv6(last.Payload())
report := header.MLDv2Report(icmp.MessageBody())
expectedRecords := expectedReport.Records
records := report.MulticastAddressRecords()
for len(expectedRecords) != 0 {
record, res := records.Next()
if res != header.MLDv2ReportMulticastAddressRecordIteratorNextOk {
t.Fatalf("got records.Next() = (%#v, %d), want = (_, %d)", record, res, header.MLDv2ReportMulticastAddressRecordIteratorNextOk)
switch res {
case header.MLDv2ReportMulticastAddressRecordIteratorNextOk:
case header.MLDv2ReportMulticastAddressRecordIteratorNextDone:
return
default:
t.Fatalf("unhandled res = %d", res)
}
if got, want := record.RecordType(), expectedRecords[0].RecordType; got != want {
addr := record.MulticastAddress()
expectedRecordType, ok := expectedRecords[addr]
if !ok {
t.Errorf("unexpected record for address %s", addr)
continue
}
if got, want := record.RecordType(), expectedRecordType; got != want {
t.Errorf("got record.RecordType() = %d, want = %d", got, want)
}
@@ -1272,30 +1302,17 @@ func MLDv2Report(expectedReport header.MLDv2ReportSerializer) NetworkChecker {
t.Errorf("got record.AuxDataLen() = %d, want = 0", got)
}
if got, want := record.MulticastAddress(), expectedRecords[0].MulticastAddress; got != want {
t.Errorf("got record.MulticastAddress() = %s, want = %s", got, want)
}
sources, ok := record.Sources()
if !ok {
t.Error("got record.Sources() = (_, false), want = (_, true)")
continue
}
expectedSources := expectedRecords[0].Sources
for len(expectedSources) != 0 {
source, ok := sources.Next()
if !ok {
t.Fatal("got sources.Next() = (_, false), want = (_, true)")
}
if source != expectedSources[0] {
t.Errorf("got sources.Next() = %s, want = %s", source, expectedSources[0])
}
expectedSources = expectedSources[1:]
if source, ok := sources.Next(); ok {
t.Fatalf("got sources.Next() = (%s, true), want = (_, false)", source)
}
expectedRecords = expectedRecords[1:]
delete(expectedRecords, addr)
}
if record, res := records.Next(); res != header.MLDv2ReportMulticastAddressRecordIteratorNextDone {
@@ -1580,8 +1597,14 @@ func IGMPMaxRespTime(want time.Duration) TransportChecker {
}
}
// IGMPGroupAddress creates a checker that checks the IGMP Group Address field.
func IGMPGroupAddress(want tcpip.Address) TransportChecker {
// IGMPGroupAddressUnordered creates a checker that checks that the group
// address in the IGMP message is expected to be seen.
//
// The seen address is removed from the expected groups map.
//
// The returned TransportChecker assumes that a valid IGMP is passed to it
// containing a valid IGMP message as far as the size is concerned.
func IGMPGroupAddressUnordered(expectedGroups map[tcpip.Address]struct{}) TransportChecker {
return func(t *testing.T, h header.Transport) {
t.Helper()
@@ -1589,15 +1612,30 @@ func IGMPGroupAddress(want tcpip.Address) TransportChecker {
if !ok {
t.Fatalf("got transport header = %T, want = header.IGMP", h)
}
if got := igmp.GroupAddress(); got != want {
t.Errorf("got igmp.GroupAddress() = %s, want = %s", got, want)
addr := igmp.GroupAddress()
if _, ok := expectedGroups[addr]; !ok {
t.Errorf("unexpected multicast group %s", addr)
} else {
delete(expectedGroups, addr)
}
}
}
// IGMPGroupAddress creates a checker that checks the IGMP Group Address field.
func IGMPGroupAddress(want tcpip.Address) TransportChecker {
return IGMPGroupAddressUnordered(map[tcpip.Address]struct{}{
want: struct{}{},
})
}
// IGMPv3Report creates a checker that checks that the packet contains a valid
// IGMPv3 report with the specified records.
func IGMPv3Report(expectedReport header.IGMPv3ReportSerializer) NetworkChecker {
//
// Note that observed records are removed from expectedRecords. No error is
// logged if the report does not have all the records expected.
func IGMPv3Report(expectedRecords map[tcpip.Address]header.IGMPv3ReportRecordType) NetworkChecker {
return func(t *testing.T, h []header.Network) {
t.Helper()
@@ -1612,16 +1650,26 @@ func IGMPv3Report(expectedReport header.IGMPv3ReportSerializer) NetworkChecker {
}
report := header.IGMPv3Report(igmp)
expectedRecords := expectedReport.Records
records := report.GroupAddressRecords()
for len(expectedRecords) != 0 {
record, res := records.Next()
if res != header.IGMPv3ReportGroupAddressRecordIteratorNextOk {
t.Fatalf("got records.Next() = (%#v, %d), want = (_, %d)", record, res, header.IGMPv3ReportGroupAddressRecordIteratorNextOk)
switch res {
case header.IGMPv3ReportGroupAddressRecordIteratorNextOk:
case header.IGMPv3ReportGroupAddressRecordIteratorNextDone:
return
default:
t.Fatalf("unhandled res = %d", res)
}
if got, want := record.RecordType(), expectedRecords[0].RecordType; got != want {
addr := record.GroupAddress()
expectedRecordType, ok := expectedRecords[addr]
if !ok {
t.Errorf("unexpected record for address %s", addr)
continue
}
if got, want := record.RecordType(), expectedRecordType; got != want {
t.Errorf("got record.RecordType() = %d, want = %d", got, want)
}
@@ -1629,30 +1677,17 @@ func IGMPv3Report(expectedReport header.IGMPv3ReportSerializer) NetworkChecker {
t.Errorf("got record.AuxDataLen() = %d, want = 0", got)
}
if got, want := record.GroupAddress(), expectedRecords[0].GroupAddress; got != want {
t.Errorf("got record.GroupAddress() = %s, want = %s", got, want)
}
sources, ok := record.Sources()
if !ok {
t.Error("got record.Sources() = (_, false), want = (_, true)")
continue
}
expectedSources := expectedRecords[0].Sources
for len(expectedSources) != 0 {
source, ok := sources.Next()
if !ok {
t.Fatal("got sources.Next() = (_, false), want = (_, true)")
}
if source != expectedSources[0] {
t.Errorf("got sources.Next() = %s, want = %s", source, expectedSources[0])
}
expectedSources = expectedSources[1:]
if source, ok := sources.Next(); ok {
t.Fatalf("got sources.Next() = (%s, true), want = (_, false)", source)
}
expectedRecords = expectedRecords[1:]
delete(expectedRecords, addr)
}
if record, res := records.Next(); res != header.IGMPv3ReportGroupAddressRecordIteratorNextDone {
@@ -19,6 +19,8 @@ go_library(
"//pkg/tcpip",
"//pkg/tcpip/checker",
"//pkg/tcpip/header",
"//pkg/tcpip/link/channel",
"//pkg/tcpip/stack",
"@com_github_google_go_cmp//cmp:go_default_library",
],
)
+176 -25
View File
@@ -21,10 +21,12 @@ import (
"math/rand"
"testing"
"github.com/google/go-cmp/cmp"
"gvisor.dev/gvisor/pkg/bufferv2"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/checker"
"gvisor.dev/gvisor/pkg/tcpip/header"
"gvisor.dev/gvisor/pkg/tcpip/link/channel"
"gvisor.dev/gvisor/pkg/tcpip/stack"
)
@@ -185,43 +187,111 @@ func CheckMLDv2Stats(t *testing.T, s *stack.Stack, reports, leaves, reportsV2 ui
checkMLDStats(t, s, 0 /* reports */, 0 /* leaves */, reports+leaves+reportsV2)
}
// ValidateIGMPv3Report validates an IGMPv3 report.
func ValidateIGMPv3Report(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address, addrs []tcpip.Address, recordType header.IGMPv3ReportRecordType) {
// ValidateIGMPv3ReportWithRecords validates an IGMPv3 report.
//
// Note that observed records are removed from expectedRecords. No error is
// logged if the report does not have all the records expected.
func ValidateIGMPv3ReportWithRecords(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address, expectedRecords map[tcpip.Address]header.IGMPv3ReportRecordType) {
t.Helper()
var records []header.IGMPv3ReportGroupAddressRecordSerializer
for _, addr := range addrs {
records = append(records, header.IGMPv3ReportGroupAddressRecordSerializer{
RecordType: recordType,
GroupAddress: addr,
Sources: nil,
})
}
checker.IPv4(t, v,
checker.SrcAddr(srcAddr),
checker.DstAddr(header.IGMPv3RoutersAddress),
checker.TTL(header.IGMPTTL),
checker.IPv4RouterAlert(),
checker.IGMPv3Report(header.IGMPv3ReportSerializer{
Records: records,
}),
checker.IGMPv3Report(expectedRecords),
)
}
// ValidateMLDv2Report validates an MLDv2 report.
func ValidateMLDv2Report(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address, addrs []tcpip.Address, recordType header.MLDv2ReportRecordType) {
// ValidateIGMPv3Report validates an IGMPv3 report.
func ValidateIGMPv3Report(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address, addrs []tcpip.Address, recordType header.IGMPv3ReportRecordType) {
t.Helper()
var records []header.MLDv2ReportMulticastAddressRecordSerializer
records := make(map[tcpip.Address]header.IGMPv3ReportRecordType)
for _, addr := range addrs {
records = append(records, header.MLDv2ReportMulticastAddressRecordSerializer{
RecordType: recordType,
MulticastAddress: addr,
Sources: nil,
})
records[addr] = recordType
}
ValidateIGMPv3ReportWithRecords(t, v, srcAddr, records)
if diff := cmp.Diff(map[tcpip.Address]header.IGMPv3ReportRecordType{}, records); diff != "" {
t.Errorf("post-validation records map mismatch (-want +got):\n%s", diff)
}
}
// ValidateIGMPv3RecordsAcrossReports validates IGMPv3 records across one or
// more reports.
func ValidateIGMPv3RecordsAcrossReports(t *testing.T, e *channel.Endpoint, srcAddr tcpip.Address, addrs []tcpip.Address, recordType header.IGMPv3ReportRecordType) {
t.Helper()
expectedRecords := make(map[tcpip.Address]header.IGMPv3ReportRecordType)
for _, addr := range addrs {
expectedRecords[addr] = recordType
}
for len(expectedRecords) != 0 {
p := e.Read()
if p.IsNil() {
t.Fatalf("expected IGMP message with expectedRecords = %#v", expectedRecords)
}
v := stack.PayloadSince(p.NetworkHeader())
ValidateIGMPv3ReportWithRecords(t, v, srcAddr, expectedRecords)
v.Release()
p.DecRef()
}
if diff := cmp.Diff(map[tcpip.Address]header.IGMPv3ReportRecordType{}, expectedRecords); diff != "" {
t.Errorf("post-validation records map mismatch (-want +got):\n%s", diff)
}
}
// ValidMultipleIGMPv2ReportLeaves validates the reception of multiple IGMPv2
// report/leave messages.
func ValidMultipleIGMPv2ReportLeaves(t *testing.T, e *channel.Endpoint, srcAddr tcpip.Address, addrs []tcpip.Address, leave bool) {
t.Helper()
expectedGroups := make(map[tcpip.Address]struct{})
for _, addr := range addrs {
expectedGroups[addr] = struct{}{}
}
igmpType := header.IGMPv2MembershipReport
if leave {
igmpType = header.IGMPLeaveGroup
}
for len(expectedGroups) != 0 {
p := e.Read()
if p.IsNil() {
t.Fatalf("expected IGMP message with expectedGroups = %#v", expectedGroups)
}
v := stack.PayloadSince(p.NetworkHeader())
checker.IPv4(t, v,
checker.SrcAddr(srcAddr),
checker.TTL(header.IGMPTTL),
checker.IPv4RouterAlert(),
checker.IGMP(
checker.IGMPType(igmpType),
checker.IGMPMaxRespTime(0),
checker.IGMPGroupAddressUnordered(expectedGroups),
),
)
v.Release()
p.DecRef()
}
if diff := cmp.Diff(map[tcpip.Address]struct{}{}, expectedGroups); diff != "" {
t.Errorf("post-validation groups map mismatch (-want +got):\n%s", diff)
}
}
// ValidateMLDv2ReportWithRecords validates an MLDv2 report.
//
// Note that observed records are removed from expectedRecords. No error is
// logged if the report does not have all the records expected.
func ValidateMLDv2ReportWithRecords(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address, expectedRecords map[tcpip.Address]header.MLDv2ReportRecordType) {
t.Helper()
checker.IPv6WithExtHdr(t, v,
checker.IPv6ExtHdr(
checker.IPv6HopByHopExtensionHeader(checker.IPv6RouterAlert(header.IPv6RouterAlertMLD)),
@@ -229,8 +299,89 @@ func ValidateMLDv2Report(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address,
checker.SrcAddr(srcAddr),
checker.DstAddr(header.MLDv2RoutersAddress),
checker.TTL(header.MLDHopLimit),
checker.MLDv2Report(header.MLDv2ReportSerializer{
Records: records,
}),
checker.MLDv2Report(expectedRecords),
)
}
// ValidateMLDv2Report validates an MLDv2 report.
func ValidateMLDv2Report(t *testing.T, v *bufferv2.View, srcAddr tcpip.Address, addrs []tcpip.Address, recordType header.MLDv2ReportRecordType) {
t.Helper()
records := make(map[tcpip.Address]header.MLDv2ReportRecordType)
for _, addr := range addrs {
records[addr] = recordType
}
ValidateMLDv2ReportWithRecords(t, v, srcAddr, records)
if diff := cmp.Diff(map[tcpip.Address]header.MLDv2ReportRecordType{}, records); diff != "" {
t.Errorf("post-validation records map mismatch (-want +got):\n%s", diff)
}
}
// ValidateMLDv2RecordsAcrossReports validates MLDv2 records across one or more
// reports.
func ValidateMLDv2RecordsAcrossReports(t *testing.T, e *channel.Endpoint, srcAddr tcpip.Address, addrs []tcpip.Address, recordType header.MLDv2ReportRecordType) {
t.Helper()
expectedRecords := make(map[tcpip.Address]header.MLDv2ReportRecordType)
for _, addr := range addrs {
expectedRecords[addr] = recordType
}
for len(expectedRecords) != 0 {
p := e.Read()
if p.IsNil() {
t.Fatalf("expected MLD Message with expectedRecords = %#v", expectedRecords)
}
v := stack.PayloadSince(p.NetworkHeader())
ValidateMLDv2ReportWithRecords(t, v, srcAddr, expectedRecords)
v.Release()
p.DecRef()
}
if diff := cmp.Diff(map[tcpip.Address]header.MLDv2ReportRecordType{}, expectedRecords); diff != "" {
t.Errorf("post-validation records map mismatch (-want +got):\n%s", diff)
}
}
// ValidMultipleMLDv1ReportLeaves validates the reception of multiple MLDv1
// report/leave messages.
func ValidMultipleMLDv1ReportLeaves(t *testing.T, e *channel.Endpoint, srcAddr tcpip.Address, addrs []tcpip.Address, leave bool) {
t.Helper()
expectedGroups := make(map[tcpip.Address]struct{})
for _, addr := range addrs {
expectedGroups[addr] = struct{}{}
}
mldType := header.ICMPv6MulticastListenerReport
if leave {
mldType = header.ICMPv6MulticastListenerDone
}
for len(expectedGroups) != 0 {
p := e.Read()
if p.IsNil() {
t.Fatalf("expected MLD Message with expectedGroups = %#v", expectedGroups)
}
v := stack.PayloadSince(p.NetworkHeader())
checker.IPv6WithExtHdr(t, v,
checker.IPv6ExtHdr(
checker.IPv6HopByHopExtensionHeader(checker.IPv6RouterAlert(header.IPv6RouterAlertMLD)),
),
checker.SrcAddr(srcAddr),
checker.TTL(header.MLDHopLimit),
checker.MLD(mldType, header.MLDMinimumSize,
checker.MLDMaxRespDelay(0),
checker.MLDMulticastAddressUnordered(expectedGroups),
),
)
v.Release()
p.DecRef()
}
if diff := cmp.Diff(map[tcpip.Address]struct{}{}, expectedGroups); diff != "" {
t.Errorf("post-validation groups map mismatch (-want +got):\n%s", diff)
}
}
+21 -173
View File
@@ -192,111 +192,6 @@ func TestSendQueuedMLDReports(t *testing.T) {
maxReports = 2
)
getAndCheckMLDv1MulticastAddress := func(t *testing.T, seen map[tcpip.Address]bool, p stack.PacketBufferPtr) tcpip.Address {
t.Helper()
payload := stack.PayloadSince(p.NetworkHeader())
defer payload.Release()
ipv6 := header.IPv6(payload.AsSlice())
ipv6HeaderIter := header.MakeIPv6PayloadIterator(
header.IPv6ExtensionHeaderIdentifier(ipv6.NextHeader()),
bufferv2.MakeWithData(ipv6.Payload()),
)
var transport header.IPv6RawPayloadHeader
for {
h, done, err := ipv6HeaderIter.Next()
if err != nil {
t.Fatalf("ipv6HeaderIter.Next(): %s", err)
}
if done {
t.Fatalf("ipv6HeaderIter.Next() = (%T, %t, _), want = (_, false, _)", h, done)
}
defer h.Release()
if t, ok := h.(header.IPv6RawPayloadHeader); ok {
transport = t
break
}
}
if got := tcpip.TransportProtocolNumber(transport.Identifier); got != header.ICMPv6ProtocolNumber {
t.Fatalf("got ipv6.NextHeader() = %d, want = %d", got, header.ICMPv6ProtocolNumber)
}
icmpv6 := header.ICMPv6(transport.Buf.Flatten())
if got := icmpv6.Type(); got != header.ICMPv6MulticastListenerReport && got != header.ICMPv6MulticastListenerDone {
t.Fatalf("got icmpv6.Type() = %d, want = %d or %d", got, header.ICMPv6MulticastListenerReport, header.ICMPv6MulticastListenerDone)
}
addr := header.MLD(icmpv6.MessageBody()).MulticastAddress()
s, ok := seen[addr]
if !ok {
t.Fatalf("unexpectedly got a packet for group %s", addr)
}
if s {
t.Fatalf("already saw packet for group %s", addr)
}
seen[addr] = true
return addr
}
getAndCheckMLDv2MulticastAddress := func(t *testing.T, seen map[tcpip.Address]bool, p stack.PacketBufferPtr) tcpip.Address {
t.Helper()
payload := stack.PayloadSince(p.NetworkHeader())
defer payload.Release()
ipv6 := header.IPv6(payload.AsSlice())
ipv6HeaderIter := header.MakeIPv6PayloadIterator(
header.IPv6ExtensionHeaderIdentifier(ipv6.NextHeader()),
bufferv2.MakeWithData(ipv6.Payload()),
)
var transport header.IPv6RawPayloadHeader
for {
h, done, err := ipv6HeaderIter.Next()
if err != nil {
t.Fatalf("ipv6HeaderIter.Next(): %s", err)
}
if done {
t.Fatalf("ipv6HeaderIter.Next() = (%T, %t, _), want = (_, false, _)", h, done)
}
defer h.Release()
if t, ok := h.(header.IPv6RawPayloadHeader); ok {
transport = t
break
}
}
if got := tcpip.TransportProtocolNumber(transport.Identifier); got != header.ICMPv6ProtocolNumber {
t.Fatalf("got ipv6.NextHeader() = %d, want = %d", got, header.ICMPv6ProtocolNumber)
}
icmpv6 := header.ICMPv6(transport.Buf.Flatten())
if got := icmpv6.Type(); got != header.ICMPv6MulticastListenerV2Report {
t.Fatalf("got icmpv6.Type() = %d, want = %d", got, header.ICMPv6MulticastListenerV2Report)
}
report := header.MLDv2Report(icmpv6.MessageBody())
records := report.MulticastAddressRecords()
record, res := records.Next()
if res != header.MLDv2ReportMulticastAddressRecordIteratorNextOk {
t.Fatalf("got records.Next() = %d, want = %d", res, header.MLDv2ReportMulticastAddressRecordIteratorNextOk)
}
addr := record.MulticastAddress()
s, ok := seen[addr]
if !ok {
t.Fatalf("unexpectedly got a packet for group %s", addr)
}
if s {
t.Fatalf("already saw packet for group %s", addr)
}
seen[addr] = true
if _, res := records.Next(); res != header.MLDv2ReportMulticastAddressRecordIteratorNextDone {
t.Errorf("got records.Next() = %d, want = %d", res, header.MLDv2ReportMulticastAddressRecordIteratorNextDone)
}
return addr
}
tests := []struct {
name string
dadTransmits uint8
@@ -315,34 +210,21 @@ func TestSendQueuedMLDReports(t *testing.T) {
}
subTests := []struct {
name string
v1Compatibility bool
validate func(t *testing.T, v *bufferv2.View, localAddress tcpip.Address, groupAddress tcpip.Address, leave bool)
checkStats func(*testing.T, *stack.Stack, uint64, uint64, uint64)
getAndCheckGroupAddress func(*testing.T, map[tcpip.Address]bool, stack.PacketBufferPtr) tcpip.Address
name string
v1Compatibility bool
validate func(t *testing.T, e *channel.Endpoint, localAddress tcpip.Address, groupAddresses []tcpip.Address, leave bool)
checkStats func(*testing.T, *stack.Stack, uint64, uint64, uint64)
}{
{
name: "V1 Compatibility",
v1Compatibility: true,
validate: func(t *testing.T, v *bufferv2.View, localAddress tcpip.Address, groupAddress tcpip.Address, leave bool) {
t.Helper()
remoteAddress := groupAddress
icmpType := header.ICMPv6MulticastListenerReport
if leave {
remoteAddress = header.IPv6AllRoutersLinkLocalMulticastAddress
icmpType = header.ICMPv6MulticastListenerDone
}
validateMLDPacket(t, v, localAddress, remoteAddress, icmpType, groupAddress)
},
checkStats: iptestutil.CheckMLDv1Stats,
getAndCheckGroupAddress: getAndCheckMLDv1MulticastAddress,
validate: iptestutil.ValidMultipleMLDv1ReportLeaves,
checkStats: iptestutil.CheckMLDv1Stats,
},
{
name: "V2",
v1Compatibility: false,
validate: func(t *testing.T, v *bufferv2.View, localAddress tcpip.Address, groupAddress tcpip.Address, leave bool) {
validate: func(t *testing.T, e *channel.Endpoint, localAddress tcpip.Address, groupAddresses []tcpip.Address, leave bool) {
t.Helper()
recordType := header.MLDv2ReportRecordChangeToExcludeMode
@@ -350,10 +232,9 @@ func TestSendQueuedMLDReports(t *testing.T) {
recordType = header.MLDv2ReportRecordChangeToIncludeMode
}
validateMLDv2ReportPacket(t, v, localAddress, groupAddress, recordType)
iptestutil.ValidateMLDv2RecordsAcrossReports(t, e, localAddress, groupAddresses, recordType)
},
checkStats: iptestutil.CheckMLDv2Stats,
getAndCheckGroupAddress: getAndCheckMLDv2MulticastAddress,
checkStats: iptestutil.CheckMLDv2Stats,
},
}
@@ -452,12 +333,7 @@ func TestSendQueuedMLDReports(t *testing.T) {
}
reportCounter++
subTest.checkStats(t, s, reportCounter, doneCounter, reportV2Counter)
if p := e.Read(); p.IsNil() {
t.Errorf("expected MLD report for %s", globalMulticastAddr)
} else {
subTest.validate(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, globalMulticastAddr, false /* leave */)
p.DecRef()
}
subTest.validate(t, e, header.IPv6Any, []tcpip.Address{globalMulticastAddr}, false /* leave */)
clock.Advance(time.Hour)
checkVersion()
if p := e.Read(); !p.IsNil() {
@@ -485,12 +361,7 @@ func TestSendQueuedMLDReports(t *testing.T) {
}
reportCounter++
subTest.checkStats(t, s, reportCounter, doneCounter, reportV2Counter)
if p := e.Read(); p.IsNil() {
t.Errorf("expected MLD report for %s", globalAddrSNMC)
} else {
subTest.validate(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, globalAddrSNMC, false /* leave */)
p.DecRef()
}
subTest.validate(t, e, header.IPv6Any, []tcpip.Address{globalAddrSNMC}, false /* leave */)
if dadResolutionTime != 0 {
// Reports should not be sent when the address resolves.
resolveDAD(globalAddr, globalAddrSNMC)
@@ -504,12 +375,7 @@ func TestSendQueuedMLDReports(t *testing.T) {
if !subTest.v1Compatibility {
doneCounter++
subTest.checkStats(t, s, reportCounter, doneCounter, reportV2Counter)
if p := e.Read(); p.IsNil() {
t.Errorf("expected MLD report for %s", globalAddrSNMC)
} else {
subTest.validate(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, globalAddrSNMC, true /* leave */)
p.DecRef()
}
subTest.validate(t, e, header.IPv6Any, []tcpip.Address{globalAddrSNMC}, true /* leave */)
}
subTest.checkStats(t, s, reportCounter, doneCounter, reportV2Counter)
if p := e.Read(); !p.IsNil() {
@@ -531,12 +397,7 @@ func TestSendQueuedMLDReports(t *testing.T) {
}
if dadResolutionTime != 0 {
reportCounter++
if p := e.Read(); p.IsNil() {
t.Errorf("expected MLD report for %s", linkLocalAddrSNMC)
} else {
subTest.validate(t, stack.PayloadSince(p.NetworkHeader()), header.IPv6Any, linkLocalAddrSNMC, false /* leave */)
p.DecRef()
}
subTest.validate(t, e, header.IPv6Any, []tcpip.Address{linkLocalAddrSNMC}, false /* leave */)
resolveDAD(linkLocalAddr, linkLocalAddrSNMC)
}
@@ -549,28 +410,15 @@ func TestSendQueuedMLDReports(t *testing.T) {
reportCounter += maxReports
subTest.checkStats(t, s, reportCounter, doneCounter, reportV2Counter)
addrs := map[tcpip.Address]bool{
globalMulticastAddr: false,
linkLocalAddrSNMC: false,
}
for range addrs {
p := e.Read()
if p.IsNil() {
t.Fatalf("expected MLD report for %s and %s; addrs = %#v", globalMulticastAddr, linkLocalAddrSNMC, addrs)
}
subTest.validate(
t,
e,
linkLocalAddr,
[]tcpip.Address{globalMulticastAddr, linkLocalAddrSNMC},
false, /* leave */
)
subTest.validate(
t,
stack.PayloadSince(p.NetworkHeader()),
linkLocalAddr,
subTest.getAndCheckGroupAddress(t, addrs, p),
false, /* leave */
)
p.DecRef()
clock.Advance(ipv6.UnsolicitedReportIntervalMax)
}
clock.Advance(ipv6.UnsolicitedReportIntervalMax)
}
// Should not send any more reports.
File diff suppressed because it is too large Load Diff