From be5314c5a6aa678ef615b22f106feda58bea0dfb Mon Sep 17 00:00:00 2001 From: Ghanan Gowripalan Date: Tue, 7 Feb 2023 11:06:13 -0800 Subject: [PATCH] Allow forcing multicast group protocol mode Updates #8346 PiperOrigin-RevId: 507832493 --- .../internal/ip/generic_multicast_protocol.go | 119 +++++++++++------- .../ip/generic_multicast_protocol_test.go | 55 ++++++-- 2 files changed, 121 insertions(+), 53 deletions(-) diff --git a/pkg/tcpip/network/internal/ip/generic_multicast_protocol.go b/pkg/tcpip/network/internal/ip/generic_multicast_protocol.go index 6addafbdb..798a75f48 100644 --- a/pkg/tcpip/network/internal/ip/generic_multicast_protocol.go +++ b/pkg/tcpip/network/internal/ip/generic_multicast_protocol.go @@ -244,6 +244,7 @@ type protocolMode int const ( protocolModeV2 protocolMode = iota + protocolModeV1Forced protocolModeV1Compatibility ) @@ -291,6 +292,35 @@ type GenericMulticastProtocolState struct { stateChangedReportV2TimerSet bool } +// SetForcedV1ModeLocked sets the V1 forced configuration. +// +// Precondition: g.protocolMU must be locked. +func (g *GenericMulticastProtocolState) SetForcedV1ModeLocked(v bool) { + if v { + switch g.mode { + case protocolModeV2: + g.cancelV2ReportTimers() + case protocolModeV1Compatibility: + g.modeTimer.Stop() + case protocolModeV1Forced: + // Already in V1 forced mode; nothing to do. + default: + panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) + } + g.mode = protocolModeV1Forced + return + } + + switch g.mode { + case protocolModeV2, protocolModeV1Compatibility: + // Not in V1 forced mode; nothing to do. + case protocolModeV1Forced: + g.mode = protocolModeV2 + default: + panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) + } +} + func (g *GenericMulticastProtocolState) cancelV2ReportTimers() { if g.generalQueryV2Timer != nil { g.generalQueryV2Timer.Stop() @@ -357,7 +387,7 @@ func (g *GenericMulticastProtocolState) MakeAllNonMemberLocked() { groupAddress, ) } - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: handler = g.transitionToNonMemberLocked default: panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) @@ -404,7 +434,7 @@ func (g *GenericMulticastProtocolState) InitializeGroupsLocked() { switch g.mode { case protocolModeV2: v2ReportBuilder = g.opts.Protocol.NewReportV2Builder() - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: default: panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) } @@ -451,7 +481,7 @@ func (g *GenericMulticastProtocolState) SendQueuedReportsLocked() { switch g.mode { case protocolModeV2: g.sendV2ReportAndMaybeScheduleChangedTimer(groupAddress, &info, MulticastGroupProtocolV2ReportRecordChangeToExcludeMode) - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: g.maybeSendReportLocked(groupAddress, &info) default: panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) @@ -498,7 +528,7 @@ func (g *GenericMulticastProtocolState) JoinGroupLocked(groupAddress tcpip.Addre // Nothing meaningful we can do with the error here - we only try to // send a delayed report once. _, _ = reportBuilder.Send() - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: g.maybeSendReportLocked(groupAddress, &info) default: panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) @@ -644,7 +674,7 @@ func (g *GenericMulticastProtocolState) LeaveGroupLocked(groupAddress tcpip.Addr } else { delete(g.memberships, groupAddress) } - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: g.transitionToNonMemberLocked(groupAddress, &info) delete(g.memberships, groupAddress) default: @@ -663,7 +693,7 @@ func (g *GenericMulticastProtocolState) HandleQueryV2Locked(groupAddress tcpip.A } switch g.mode { - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: g.handleQueryInnerLocked(groupAddress, g.opts.Protocol.V2QueryMaxRespCodeToV1Delay(maxResponseCode)) return case protocolModeV2: @@ -847,42 +877,47 @@ func (g *GenericMulticastProtocolState) HandleQueryLocked(groupAddress tcpip.Add return } - // As per 3376 section 8.12 (for IGMPv3), - // - // The Older Version Querier Interval is the time-out for transitioning - // a host back to IGMPv3 mode once an older version query is heard. - // When an older version query is received, hosts set their Older - // Version Querier Present Timer to Older Version Querier Interval. - // - // This value MUST be ((the Robustness Variable) times (the Query - // Interval in the last Query received)) plus (one Query Response - // Interval). - // - // As per RFC 3810 section 9.12 (for MLDv2), - // - // The Older Version Querier Present Timeout is the time-out for - // transitioning a host back to MLDv2 Host Compatibility Mode. When an - // MLDv1 query is received, MLDv2 hosts set their Older Version Querier - // Present Timer to [Older Version Querier Present Timeout]. - // - // This value MUST be ([Robustness Variable] times (the [Query Interval] - // in the last Query received)) plus ([Query Response Interval]). - modeRevertDelay := time.Duration(g.robustnessVariable) * g.queryInterval - if g.modeTimer == nil { - // TODO(https://issuetracker.google.com/264799098): Create timer on - // initialization instead of lazily creating the timer since the timer - // does not change after being created. - g.modeTimer = g.opts.Clock.AfterFunc(modeRevertDelay, func() { - g.protocolMU.Lock() - defer g.protocolMU.Unlock() - g.mode = protocolModeV2 - }) - } else { - g.modeTimer.Reset(modeRevertDelay) + switch g.mode { + case protocolModeV2, protocolModeV1Compatibility: + // As per 3376 section 8.12 (for IGMPv3), + // + // The Older Version Querier Interval is the time-out for transitioning + // a host back to IGMPv3 mode once an older version query is heard. + // When an older version query is received, hosts set their Older + // Version Querier Present Timer to Older Version Querier Interval. + // + // This value MUST be ((the Robustness Variable) times (the Query + // Interval in the last Query received)) plus (one Query Response + // Interval). + // + // As per RFC 3810 section 9.12 (for MLDv2), + // + // The Older Version Querier Present Timeout is the time-out for + // transitioning a host back to MLDv2 Host Compatibility Mode. When an + // MLDv1 query is received, MLDv2 hosts set their Older Version Querier + // Present Timer to [Older Version Querier Present Timeout]. + // + // This value MUST be ([Robustness Variable] times (the [Query Interval] + // in the last Query received)) plus ([Query Response Interval]). + modeRevertDelay := time.Duration(g.robustnessVariable) * g.queryInterval + if g.modeTimer == nil { + // TODO(https://issuetracker.google.com/264799098): Create timer on + // initialization instead of lazily creating the timer since the timer + // does not change after being created. + g.modeTimer = g.opts.Clock.AfterFunc(modeRevertDelay, func() { + g.protocolMU.Lock() + defer g.protocolMU.Unlock() + g.mode = protocolModeV2 + }) + } else { + g.modeTimer.Reset(modeRevertDelay) + } + g.mode = protocolModeV1Compatibility + g.cancelV2ReportTimers() + case protocolModeV1Forced: + default: + panic(fmt.Sprintf("unrecognized mode = %d", g.mode)) } - g.mode = protocolModeV1Compatibility - g.cancelV2ReportTimers() - g.handleQueryInnerLocked(groupAddress, maxResponseTime) } @@ -961,7 +996,7 @@ func (g *GenericMulticastProtocolState) initializeNewMemberLocked(groupAddress t callersV2ReportBuilder.AddRecord(MulticastGroupProtocolV2ReportRecordChangeToExcludeMode, groupAddress) info.transmissionLeft-- } - case protocolModeV1Compatibility: + case protocolModeV1Compatibility, protocolModeV1Forced: info.transmissionLeft = unsolicitedTransmissionCount g.maybeSendReportLocked(groupAddress, info) default: diff --git a/pkg/tcpip/network/internal/ip/generic_multicast_protocol_test.go b/pkg/tcpip/network/internal/ip/generic_multicast_protocol_test.go index ad0337752..ca45db31f 100644 --- a/pkg/tcpip/network/internal/ip/generic_multicast_protocol_test.go +++ b/pkg/tcpip/network/internal/ip/generic_multicast_protocol_test.go @@ -17,7 +17,6 @@ package ip_test import ( "bytes" "fmt" - "math" "math/rand" "testing" "time" @@ -61,11 +60,7 @@ func (m *mockMulticastGroupProtocol) init(opts ip.GenericMulticastProtocolOption m.mu.genericMulticastGroup.Init(&m.mu.RWMutex, opts) if v1Compatibility { - // A General V1 query should make us drop into V1 compatibility mode. - // - // Since we just init-ed, we know we don't have any groups to send - // reports for so this won't break any tests looking at packets. - m.mu.genericMulticastGroup.HandleQueryLocked("", math.MaxInt64) + m.mu.genericMulticastGroup.SetForcedV1ModeLocked(true) } } @@ -87,6 +82,12 @@ func (m *mockMulticastGroupProtocol) setQueuePackets(v bool) { m.mu.makeQueuePackets = v } +func (m *mockMulticastGroupProtocol) setForcedV1Mode(v bool) { + m.mu.Lock() + defer m.mu.Unlock() + m.mu.genericMulticastGroup.SetForcedV1ModeLocked(v) +} + func (m *mockMulticastGroupProtocol) joinGroup(addr tcpip.Address) { m.mu.Lock() defer m.mu.Unlock() @@ -1416,11 +1417,6 @@ func TestQueuedPackets(t *testing.T) { // The delayed report timer should have been cancelled since we did not send // the initial report earlier. clock.Advance(time.Hour) - if test.v1Compatibility { - // V1 query targetting an unjoined group should drop us into V1 - // compatibility mode without sending any packets, affecting tests. - mgp.handleQuery(addr3, 0) - } if diff := mgp.check(checkFields{}); diff != "" { t.Fatalf("mockMulticastGroupProtocol mismatch (-want +got):\n%s", diff) } @@ -1509,3 +1505,40 @@ func TestQueuedPackets(t *testing.T) { }) } } + +func TestV1Compatibility(t *testing.T) { + clock := faketime.NewManualClock() + mgp := mockMulticastGroupProtocol{t: t} + mgp.init(ip.GenericMulticastProtocolOptions{ + Rand: rand.New(rand.NewSource(4)), + Clock: clock, + MaxUnsolicitedReportDelay: maxUnsolicitedReportDelay, + }, false /* v1Compatibility */) + + mgp.joinGroup(addr1) + if diff := mgp.check(checkFields{sentV2Reports: []mockReportV2{{records: []mockReportV2Record{ + { + recordType: ip.MulticastGroupProtocolV2ReportRecordChangeToExcludeMode, + groupAddress: addr1, + }, + }}}}); diff != "" { + t.Fatalf("mockMulticastGroupProtocol mismatch (-want +got):\n%s", diff) + } + + mgp.setForcedV1Mode(true) + mgp.joinGroup(addr2) + if diff := mgp.check(checkFields{sendReportGroupAddresses: []tcpip.Address{addr2}}); diff != "" { + t.Fatalf("mockMulticastGroupProtocol mismatch (-want +got):\n%s", diff) + } + + mgp.setForcedV1Mode(false) + mgp.joinGroup(addr3) + if diff := mgp.check(checkFields{sentV2Reports: []mockReportV2{{records: []mockReportV2Record{ + { + recordType: ip.MulticastGroupProtocolV2ReportRecordChangeToExcludeMode, + groupAddress: addr3, + }, + }}}}); diff != "" { + t.Fatalf("mockMulticastGroupProtocol mismatch (-want +got):\n%s", diff) + } +}