diff --git a/pkg/tcpip/checker/checker.go b/pkg/tcpip/checker/checker.go index 49ba2d916..0e338d59c 100644 --- a/pkg/tcpip/checker/checker.go +++ b/pkg/tcpip/checker/checker.go @@ -1650,6 +1650,9 @@ func IGMPv3Report(expectedRecords map[tcpip.Address]header.IGMPv3ReportRecordTyp } report := header.IGMPv3Report(igmp) + if got, want := report.Checksum(), header.IGMPCalculateChecksum(igmp); got != want { + t.Errorf("got report.Checksum() = %d, want = %d", got, want) + } records := report.GroupAddressRecords() for len(expectedRecords) != 0 { diff --git a/pkg/tcpip/header/igmp_test.go b/pkg/tcpip/header/igmp_test.go index 8d6b08cdb..e0dc62bde 100644 --- a/pkg/tcpip/header/igmp_test.go +++ b/pkg/tcpip/header/igmp_test.go @@ -336,6 +336,11 @@ func TestIGMPv3Report(t *testing.T) { test.serializer.SerializeInto(b) report := header.IGMPv3Report(b) + + if got, want := report.Checksum(), header.IGMPCalculateChecksum(header.IGMP(report)); got != want { + t.Errorf("got report.Checksum() = %d, want = %d", got, want) + } + expectedRecords := test.serializer.Records records := report.GroupAddressRecords() diff --git a/pkg/tcpip/header/igmpv3.go b/pkg/tcpip/header/igmpv3.go index 0ffb4cd43..3e00b6669 100644 --- a/pkg/tcpip/header/igmpv3.go +++ b/pkg/tcpip/header/igmpv3.go @@ -309,12 +309,13 @@ func (s *IGMPv3ReportSerializer) SerializeInto(b []byte) { b[igmpv3ReportReserved1Offset] = 0 binary.BigEndian.PutUint16(b[igmpv3ReportReserved2Offset:], 0) binary.BigEndian.PutUint16(b[igmpv3ReportNumberOfGroupAddressRecordsOffset:], uint16(len(s.Records))) - b = b[igmpv3ReportGroupAddressRecordsOffset:] + recordsBytes := b[igmpv3ReportGroupAddressRecordsOffset:] for _, record := range s.Records { len := record.Length() - record.SerializeInto(b[:len]) - b = b[len:] + record.SerializeInto(recordsBytes[:len]) + recordsBytes = recordsBytes[len:] } + binary.BigEndian.PutUint16(b[igmpChecksumOffset:], IGMPCalculateChecksum(b)) } // IGMPv3ReportGroupAddressRecord is an IGMPv3 record. @@ -436,6 +437,11 @@ func (r IGMPv3ReportGroupAddressRecord) Sources() (AddressIterator, bool) { // +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ type IGMPv3Report []byte +// Checksum returns the checksum. +func (i IGMPv3Report) Checksum() uint16 { + return binary.BigEndian.Uint16(i[igmpChecksumOffset:]) +} + // IGMPv3ReportGroupAddressRecordIterator is an iterator over IGMPv3 Multicast // Address Records. type IGMPv3ReportGroupAddressRecordIterator struct {