diff --git a/pkg/tcpip/network/ipv4/igmp.go b/pkg/tcpip/network/ipv4/igmp.go index 8ca69a790..356225a0f 100644 --- a/pkg/tcpip/network/ipv4/igmp.go +++ b/pkg/tcpip/network/ipv4/igmp.go @@ -19,7 +19,6 @@ import ( "math" "time" - "gvisor.dev/gvisor/pkg/atomicbitops" "gvisor.dev/gvisor/pkg/bufferv2" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/header" @@ -28,11 +27,6 @@ import ( ) const ( - // igmpV1PresentDefault is the initial state for igmpV1Present in the - // igmpState. As per RFC 2236 Page 9 says "No IGMPv1 Router Present ... is - // the initial state." - igmpV1PresentDefault = 0 - // v1RouterPresentTimeout from RFC 2236 Section 8.11, Page 18 // See note on igmpState.igmpV1Present for more detail. v1RouterPresentTimeout = 400 * time.Second @@ -51,6 +45,47 @@ const ( UnsolicitedReportIntervalMax = 10 * time.Second ) +type protocolMode int + +const ( + protocolModeV2OrV3 protocolMode = iota + protocolModeV1 + // protocolModeV1Compatibility is for maintaining compatibility with IGMPv1 + // Routers. + // + // Per RFC 2236 Section 4 Page 6: "The IGMPv1 router expects Version 1 + // Membership Reports in response to its Queries, and will not pay + // attention to Version 2 Membership Reports. Therefore, a state variable + // MUST be kept for each interface, describing whether the multicast + // Querier on that interface is running IGMPv1 or IGMPv2. This variable + // MUST be based upon whether or not an IGMPv1 query was heard in the last + // [Version 1 Router Present Timeout] seconds". + protocolModeV1Compatibility +) + +// IGMPVersion is the forced version of IGMP. +type IGMPVersion int + +const ( + _ IGMPVersion = iota + // IGMPVersion1 indicates IGMPv1. + IGMPVersion1 + // IGMPVersion2 indicates IGMPv2. Note that IGMP may still fallback to V1 + // compatibility mode as required by IGMPv2. + IGMPVersion2 + // IGMPVersion3 indicates IGMPv3. Note that IGMP may still fallback to V2 + // compatibility mode as required by IGMPv3. + IGMPVersion3 +) + +// IGMPEndpoint is a network endpoint that supports IGMP. +type IGMPEndpoint interface { + // Sets the IGMP version. + // + // Returns the previous IGMP version. + SetIGMPVersion(IGMPVersion) IGMPVersion +} + // IGMPOptions holds options for IGMP. type IGMPOptions struct { // Enabled indicates whether IGMP will be performed. @@ -75,17 +110,8 @@ type igmpState struct { genericMulticastProtocol ip.GenericMulticastProtocolState - // igmpV1Present is for maintaining compatibility with IGMPv1 Routers, from - // RFC 2236 Section 4 Page 6: "The IGMPv1 router expects Version 1 - // Membership Reports in response to its Queries, and will not pay - // attention to Version 2 Membership Reports. Therefore, a state variable - // MUST be kept for each interface, describing whether the multicast - // Querier on that interface is running IGMPv1 or IGMPv2. This variable - // MUST be based upon whether or not an IGMPv1 query was heard in the last - // [Version 1 Router Present Timeout] seconds". - // - // Holds a value of 1 when true, 0 when false. - igmpV1Present atomicbitops.Uint32 + // mode is used to configure the version of IGMP to perform. + mode protocolMode // igmpV1Job is scheduled when this interface receives an IGMPv1 style // message, upon expiration the igmpV1Present flag is cleared. @@ -105,8 +131,12 @@ func (igmp *igmpState) Enabled() bool { // +checklocksread:igmp.ep.mu func (igmp *igmpState) SendReport(groupAddress tcpip.Address) (bool, tcpip.Error) { igmpType := header.IGMPv2MembershipReport - if igmp.v1Present() { + switch igmp.mode { + case protocolModeV2OrV3: + case protocolModeV1, protocolModeV1Compatibility: igmpType = header.IGMPv1MembershipReport + default: + panic(fmt.Sprintf("unrecognized mode = %d", igmp.mode)) } return igmp.writePacket(groupAddress, groupAddress, igmpType) } @@ -119,11 +149,15 @@ func (igmp *igmpState) SendLeave(groupAddress tcpip.Address) tcpip.Error { // Querier is running IGMPv1, this action SHOULD be skipped. If the flag // saying we were the last host to report is cleared, this action MAY be // skipped." - if igmp.v1Present() { + switch igmp.mode { + case protocolModeV2OrV3: + _, err := igmp.writePacket(header.IPv4AllRoutersGroup, groupAddress, header.IGMPLeaveGroup) + return err + case protocolModeV1, protocolModeV1Compatibility: return nil + default: + panic(fmt.Sprintf("unrecognized mode = %d", igmp.mode)) } - _, err := igmp.writePacket(header.IPv4AllRoutersGroup, groupAddress, header.IGMPLeaveGroup) - return err } // ShouldPerformProtocol implements ip.MulticastGroupProtocol. @@ -251,9 +285,11 @@ func (igmp *igmpState) init(ep *endpoint) { Protocol: igmp, MaxUnsolicitedReportDelay: UnsolicitedReportIntervalMax, }) - igmp.igmpV1Present = atomicbitops.FromUint32(igmpV1PresentDefault) + // As per RFC 2236 Page 9 says "No IGMPv1 Router Present ... is + // the initial state. + igmp.mode = protocolModeV2OrV3 igmp.igmpV1Job = tcpip.NewJob(ep.protocol.stack.Clock(), &ep.mu, func() { - igmp.setV1Present(false) + igmp.mode = protocolModeV2OrV3 }) } @@ -381,21 +417,15 @@ func (igmp *igmpState) handleIGMP(pkt stack.PacketBufferPtr, hasRouterAlertOptio } } -func (igmp *igmpState) v1Present() bool { - return igmp.igmpV1Present.Load() == 1 -} - -func (igmp *igmpState) setV1Present(v bool) { - if v { - igmp.igmpV1Present.Store(1) - } else { - igmp.igmpV1Present.Store(0) - } -} - func (igmp *igmpState) resetV1Present() { igmp.igmpV1Job.Cancel() - igmp.setV1Present(false) + switch igmp.mode { + case protocolModeV2OrV3, protocolModeV1: + case protocolModeV1Compatibility: + igmp.mode = protocolModeV2OrV3 + default: + panic(fmt.Sprintf("unrecognized mode = %d", igmp.mode)) + } } // handleMembershipQuery handles a membership query. @@ -406,9 +436,16 @@ func (igmp *igmpState) handleMembershipQuery(groupAddress tcpip.Address, maxResp // then change the state to note that an IGMPv1 router is present and // schedule the query received Job. if maxRespTime == 0 && igmp.Enabled() { - igmp.igmpV1Job.Cancel() - igmp.igmpV1Job.Schedule(v1RouterPresentTimeout) - igmp.setV1Present(true) + switch igmp.mode { + case protocolModeV2OrV3, protocolModeV1Compatibility: + igmp.igmpV1Job.Cancel() + igmp.igmpV1Job.Schedule(v1RouterPresentTimeout) + igmp.mode = protocolModeV1Compatibility + case protocolModeV1: + default: + panic(fmt.Sprintf("unrecognized mode = %d", igmp.mode)) + } + maxRespTime = v1MaxRespTime } @@ -556,7 +593,44 @@ func (igmp *igmpState) initializeAll() { // sendQueuedReports attempts to send any reports that are queued for sending. // -// +checklocksread:igmp.ep.mu +// +checklocks:igmp.ep.mu func (igmp *igmpState) sendQueuedReports() { igmp.genericMulticastProtocol.SendQueuedReportsLocked() } + +// setVersion sets the IGMP version. +// +// +checklocks:igmp.ep.mu +func (igmp *igmpState) setVersion(v IGMPVersion) IGMPVersion { + prev := igmp.mode + igmp.igmpV1Job.Cancel() + + var prevGenericModeV1 bool + switch v { + case IGMPVersion3: + prevGenericModeV1 = igmp.genericMulticastProtocol.SetV1ModeLocked(false) + igmp.mode = protocolModeV2OrV3 + case IGMPVersion2: + // IGMPv1 and IGMPv2 map to V1 of the generic multicast protocol. + prevGenericModeV1 = igmp.genericMulticastProtocol.SetV1ModeLocked(true) + igmp.mode = protocolModeV2OrV3 + case IGMPVersion1: + // IGMPv1 and IGMPv2 map to V1 of the generic multicast protocol. + prevGenericModeV1 = igmp.genericMulticastProtocol.SetV1ModeLocked(true) + igmp.mode = protocolModeV1 + default: + panic(fmt.Sprintf("unrecognized version = %d", v)) + } + + switch prev { + case protocolModeV2OrV3, protocolModeV1Compatibility: + if prevGenericModeV1 { + return IGMPVersion2 + } + return IGMPVersion3 + case protocolModeV1: + return IGMPVersion1 + default: + panic(fmt.Sprintf("unrecognized mode = %d", igmp.mode)) + } +} diff --git a/pkg/tcpip/network/ipv4/igmp_test.go b/pkg/tcpip/network/ipv4/igmp_test.go index fb4e0ca00..c7e3f1e9c 100644 --- a/pkg/tcpip/network/ipv4/igmp_test.go +++ b/pkg/tcpip/network/ipv4/igmp_test.go @@ -43,7 +43,9 @@ var ( remoteAddr = testutil.MustParse4("10.0.0.2") multicastAddr1 = testutil.MustParse4("224.0.0.3") multicastAddr2 = testutil.MustParse4("224.0.0.4") - unusedMulticastAddr = testutil.MustParse4("224.0.0.5") + multicastAddr3 = testutil.MustParse4("224.0.0.5") + multicastAddr4 = testutil.MustParse4("224.0.0.6") + unusedMulticastAddr = testutil.MustParse4("224.0.0.7") ) // validateIgmpPacket checks that a passed packet is an IPv4 IGMP packet sent @@ -276,16 +278,17 @@ func TestSendQueuedIGMPReports(t *testing.T) { checkVersion := func() { if test.v2Compatibility { - createAndInjectIGMPPacket( - e, - header.IGMPMembershipQuery, - 1, /* maxRespTime */ - header.IGMPTTL, - remoteAddr, - header.IPv4AllSystems, - unusedMulticastAddr, - true, /* hasRouterAlertOption */ - ) + ep, err := s.GetNetworkEndpoint(nicID, header.IPv4ProtocolNumber) + if err != nil { + t.Fatalf("s.GetNetworkEndpoint(%d, %d): %s", nicID, header.IPv4ProtocolNumber, err) + } + + igmpEP, ok := ep.(ipv4.IGMPEndpoint) + if !ok { + t.Fatalf("got (%T).(%T) = (_, false), want = (_ true)", ep, igmpEP) + } + + igmpEP.SetIGMPVersion(ipv4.IGMPVersion2) } } protocolAddr := tcpip.ProtocolAddress{ @@ -505,3 +508,78 @@ func TestIGMPPacketValidation(t *testing.T) { }) } } + +func TestSetIGMPVersion(t *testing.T) { + const nicID = 1 + + c := newIGMPTestContext(t, true /* igmpEnabled */) + defer c.cleanup() + s := c.s + e := c.ep + + ep, err := s.GetNetworkEndpoint(nicID, header.IPv4ProtocolNumber) + if err != nil { + t.Fatalf("s.GetNetworkEndpoint(%d, %d): %s", nicID, header.IPv4ProtocolNumber, err) + } + igmpEP, ok := ep.(ipv4.IGMPEndpoint) + if !ok { + t.Fatalf("got (%T).(%T) = (_, false), want = (_ true)", ep, igmpEP) + } + + protocolAddr := tcpip.ProtocolAddress{ + Protocol: ipv4.ProtocolNumber, + AddressWithPrefix: tcpip.AddressWithPrefix{Address: stackAddr, PrefixLen: defaultPrefixLength}, + } + if err := s.AddProtocolAddress(nicID, protocolAddr, stack.AddressProperties{}); err != nil { + t.Fatalf("AddProtocolAddress(%d, %+v, {}): %s", nicID, protocolAddr, err) + } + + if err := s.JoinGroup(ipv4.ProtocolNumber, nicID, multicastAddr1); err != nil { + t.Fatalf("JoinGroup(ipv4, nic, %s) = %s", multicastAddr1, err) + } + if p := e.Read(); p.IsNil() { + t.Fatal("expected a report message to be sent") + } else { + validateIgmpv3ReportPacket(t, p, stackAddr, multicastAddr1) + p.DecRef() + } + + if got := igmpEP.SetIGMPVersion(ipv4.IGMPVersion2); got != ipv4.IGMPVersion3 { + t.Errorf("got igmpEP.SetIGMPVersion(%d) = %d, want = %d", ipv4.IGMPVersion2, got, ipv4.IGMPVersion3) + } + if err := s.JoinGroup(ipv4.ProtocolNumber, nicID, multicastAddr2); err != nil { + t.Fatalf("JoinGroup(ipv4, nic, %s) = %s", multicastAddr2, err) + } + if p := e.Read(); p.IsNil() { + t.Fatal("expected a report message to be sent") + } else { + validateIgmpPacket(t, p, header.IGMPv2MembershipReport, 0, stackAddr, multicastAddr2, multicastAddr2) + p.DecRef() + } + + if got := igmpEP.SetIGMPVersion(ipv4.IGMPVersion1); got != ipv4.IGMPVersion2 { + t.Errorf("got igmpEP.SetIGMPVersion(%d) = %d, want = %d", ipv4.IGMPVersion1, got, ipv4.IGMPVersion2) + } + if err := s.JoinGroup(ipv4.ProtocolNumber, nicID, multicastAddr3); err != nil { + t.Fatalf("JoinGroup(ipv4, nic, %s) = %s", multicastAddr3, err) + } + if p := e.Read(); p.IsNil() { + t.Fatal("expected a report message to be sent") + } else { + validateIgmpPacket(t, p, header.IGMPv1MembershipReport, 0, stackAddr, multicastAddr3, multicastAddr3) + p.DecRef() + } + + if got := igmpEP.SetIGMPVersion(ipv4.IGMPVersion3); got != ipv4.IGMPVersion1 { + t.Errorf("got igmpEP.SetIGMPVersion(%d) = %d, want = %d", ipv4.IGMPVersion3, got, ipv4.IGMPVersion1) + } + if err := s.JoinGroup(ipv4.ProtocolNumber, nicID, multicastAddr4); err != nil { + t.Fatalf("JoinGroup(ipv4, nic, %s) = %s", multicastAddr4, err) + } + if p := e.Read(); p.IsNil() { + t.Fatal("expected a report message to be sent") + } else { + validateIgmpv3ReportPacket(t, p, stackAddr, multicastAddr4) + p.DecRef() + } +} diff --git a/pkg/tcpip/network/ipv4/ipv4.go b/pkg/tcpip/network/ipv4/ipv4.go index b70e3bc9d..53930ff82 100644 --- a/pkg/tcpip/network/ipv4/ipv4.go +++ b/pkg/tcpip/network/ipv4/ipv4.go @@ -77,6 +77,7 @@ var _ stack.MulticastForwardingNetworkEndpoint = (*endpoint)(nil) var _ stack.GroupAddressableEndpoint = (*endpoint)(nil) var _ stack.AddressableEndpoint = (*endpoint)(nil) var _ stack.NetworkEndpoint = (*endpoint)(nil) +var _ IGMPEndpoint = (*endpoint)(nil) type endpoint struct { nic stack.NetworkInterface @@ -109,6 +110,19 @@ type endpoint struct { igmp igmpState } +// SetIGMPVersion implements IGMPEndpoint. +func (e *endpoint) SetIGMPVersion(v IGMPVersion) IGMPVersion { + e.mu.Lock() + defer e.mu.Unlock() + return e.setIGMPVersionLocked(v) +} + +// +checklocks:e.mu +// +checklocksalias:e.igmp.ep.mu=e.mu +func (e *endpoint) setIGMPVersionLocked(v IGMPVersion) IGMPVersion { + return e.igmp.setVersion(v) +} + // HandleLinkResolutionFailure implements stack.LinkResolvableNetworkEndpoint. func (e *endpoint) HandleLinkResolutionFailure(pkt stack.PacketBufferPtr) { // If we are operating as a router, return an ICMP error to the original