mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Allow forcing multicast group protocol mode
Updates #8346 PiperOrigin-RevId: 507832493
This commit is contained in:
committed by
gVisor bot
parent
2b16029543
commit
be5314c5a6
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user