mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Enable multicast forwarding per protocol.
After this change, multicast packets will only be forwarded if multicast forwarding is enabled on the relevant NIC AND it is enabled for the relevant protocol (IPv4 or IPv6). Multicast forwarding at the protocol level will be tied to the lifetime of the IPv4/IPv6RoutingTableControllers. In particular, this change enables us to satisfy the following requirements: https://cs.opensource.google/fuchsia/fuchsia/+/main:sdk/fidl/fuchsia.net.multicast.admin/ipv4.fidl;l=20-21;drc=acd6519f50c92e30f18cc1816bd1c400759b533c Updates #7338. PiperOrigin-RevId: 455155793
This commit is contained in:
@@ -413,6 +413,16 @@ func (r *RouteTable) RemoveInstalledRoute(key stack.UnicastSourceAndMulticastDes
|
||||
return false
|
||||
}
|
||||
|
||||
// RemoveAllInstalledRoutes removes all installed routes from the table.
|
||||
func (r *RouteTable) RemoveAllInstalledRoutes() {
|
||||
r.installedMu.Lock()
|
||||
defer r.installedMu.Unlock()
|
||||
|
||||
for key := range r.installedRoutes {
|
||||
delete(r.installedRoutes, key)
|
||||
}
|
||||
}
|
||||
|
||||
// GetLastUsedTimestamp returns a monotonic timestamp that represents the last
|
||||
// time the route that matches the provided key was used or updated.
|
||||
//
|
||||
|
||||
@@ -407,6 +407,44 @@ func TestRemoveInstalledRouteWithNoMatchingRoute(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveAllInstalledRoutes(t *testing.T) {
|
||||
otherAddress := testutil.MustParse4("192.168.2.1")
|
||||
|
||||
table := RouteTable{}
|
||||
defer table.Close()
|
||||
config := defaultConfig()
|
||||
if err := table.Init(config); err != nil {
|
||||
t.Fatalf("table.Init(%#v): %s", config, err)
|
||||
}
|
||||
|
||||
routes := map[stack.UnicastSourceAndMulticastDestination]stack.MulticastRoute{
|
||||
defaultRouteKey: defaultRoute,
|
||||
stack.UnicastSourceAndMulticastDestination{otherAddress, otherAddress}: defaultRoute,
|
||||
}
|
||||
|
||||
for key, route := range routes {
|
||||
installedRoute := table.NewInstalledRoute(route)
|
||||
table.AddInstalledRoute(key, installedRoute)
|
||||
}
|
||||
|
||||
table.RemoveAllInstalledRoutes()
|
||||
|
||||
for key := range routes {
|
||||
pkt := newPacketBuffer("hello")
|
||||
defer pkt.DecRef()
|
||||
|
||||
result, hasBufferSpace := table.GetRouteOrInsertPending(key, pkt)
|
||||
|
||||
if !hasBufferSpace {
|
||||
t.Fatalf("table.GetRouteOrInsertPending(%#v, %#v): false", key, pkt)
|
||||
}
|
||||
|
||||
if result.InstalledRoute != nil {
|
||||
t.Errorf("result.InstalledRoute = %v, want = nil", result.InstalledRoute)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLastUsedTimestampWithNoMatchingRoute(t *testing.T) {
|
||||
table := RouteTable{}
|
||||
defer table.Close()
|
||||
|
||||
@@ -366,16 +366,15 @@ func (e *endpoint) disableLocked() {
|
||||
}
|
||||
}
|
||||
|
||||
// multicastEventDispatcher returns the multicast forwarding event dispatcher.
|
||||
//
|
||||
// Panics if a multicast forwarding event dispatcher does not exist. This
|
||||
// indicates that multicast forwarding is enabled, but no dispatcher was
|
||||
// provided.
|
||||
func (e *endpoint) multicastEventDispatcher() stack.MulticastForwardingEventDispatcher {
|
||||
if mcastDisp := e.protocol.options.MulticastForwardingDisp; mcastDisp != nil {
|
||||
return mcastDisp
|
||||
// emitMulticastEvent emits a multicast forwarding event using the provided
|
||||
// generator if a valid event dispatcher exists.
|
||||
func (e *endpoint) emitMulticastEvent(eventGenerator func(stack.MulticastForwardingEventDispatcher)) {
|
||||
e.protocol.mu.RLock()
|
||||
defer e.protocol.mu.RUnlock()
|
||||
|
||||
if mcastDisp := e.protocol.multicastForwardingDisp; mcastDisp != nil {
|
||||
eventGenerator(mcastDisp)
|
||||
}
|
||||
panic("e.procotol.options.MulticastForwardingDisp unexpectedly nil")
|
||||
}
|
||||
|
||||
// DefaultTTL is the default time-to-live value for this endpoint.
|
||||
@@ -902,9 +901,11 @@ func (e *endpoint) forwardMulticastPacket(h header.IPv4, pkt *stack.PacketBuffer
|
||||
// Attempt to forward the pkt using an existing route.
|
||||
return e.forwardValidatedMulticastPacket(pkt, result.InstalledRoute)
|
||||
case multicast.NoRouteFoundAndPendingInserted:
|
||||
e.multicastEventDispatcher().OnMissingRoute(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
e.emitMulticastEvent(func(disp stack.MulticastForwardingEventDispatcher) {
|
||||
disp.OnMissingRoute(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
})
|
||||
})
|
||||
case multicast.PacketQueuedInPendingRoute:
|
||||
default:
|
||||
@@ -957,10 +958,12 @@ func (e *endpoint) forwardValidatedMulticastPacket(pkt *stack.PacketBuffer, inst
|
||||
// dropped silently.
|
||||
if e.nic.ID() != installedRoute.ExpectedInputInterface {
|
||||
h := header.IPv4(pkt.NetworkHeader().View())
|
||||
e.multicastEventDispatcher().OnUnexpectedInputInterface(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
}, installedRoute.ExpectedInputInterface)
|
||||
e.emitMulticastEvent(func(disp stack.MulticastForwardingEventDispatcher) {
|
||||
disp.OnUnexpectedInputInterface(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
}, installedRoute.ExpectedInputInterface)
|
||||
})
|
||||
return &ip.ErrUnexpectedMulticastInputInterface{}
|
||||
}
|
||||
|
||||
@@ -1046,7 +1049,7 @@ func (e *endpoint) handleValidatedPacket(h header.IPv4, pkt *stack.PacketBuffer,
|
||||
// RFC 1812 section 5.2.3 for details regarding the forwarding/local
|
||||
// delivery decision.
|
||||
|
||||
multicastForwarding := e.MulticastForwarding()
|
||||
multicastForwarding := e.MulticastForwarding() && e.protocol.multicastForwarding()
|
||||
|
||||
if multicastForwarding {
|
||||
e.handleForwardingError(e.forwardMulticastPacket(h, pkt))
|
||||
@@ -1432,6 +1435,11 @@ type protocol struct {
|
||||
options Options
|
||||
|
||||
multicastRouteTable multicast.RouteTable
|
||||
// multicastForwardingDisp is the multicast forwarding event dispatcher that
|
||||
// an integrator can provide to receive multicast forwarding events. Note
|
||||
// that multicast packets will only be forwarded if this is non-nil.
|
||||
// +checklocks:mu
|
||||
multicastForwardingDisp stack.MulticastForwardingEventDispatcher
|
||||
}
|
||||
|
||||
// Number returns the ipv4 protocol number.
|
||||
@@ -1505,6 +1513,12 @@ func (p *protocol) validateUnicastSourceAndMulticastDestination(addresses stack.
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *protocol) multicastForwarding() bool {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
return p.multicastForwardingDisp != nil
|
||||
}
|
||||
|
||||
func (p *protocol) newInstalledRoute(route stack.MulticastRoute) (*multicast.InstalledRoute, tcpip.Error) {
|
||||
if len(route.OutgoingInterfaces) == 0 {
|
||||
return nil, &tcpip.ErrMissingRequiredFields{}
|
||||
@@ -1528,6 +1542,10 @@ func (p *protocol) newInstalledRoute(route stack.MulticastRoute) (*multicast.Ins
|
||||
|
||||
// AddMulticastRoute implements stack.MulticastForwardingNetworkProtocol.
|
||||
func (p *protocol) AddMulticastRoute(addresses stack.UnicastSourceAndMulticastDestination, route stack.MulticastRoute) tcpip.Error {
|
||||
if !p.multicastForwarding() {
|
||||
return &tcpip.ErrNotPermitted{}
|
||||
}
|
||||
|
||||
if err := p.validateUnicastSourceAndMulticastDestination(addresses); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -1559,6 +1577,34 @@ func (p *protocol) RemoveMulticastRoute(addresses stack.UnicastSourceAndMulticas
|
||||
return nil
|
||||
}
|
||||
|
||||
// EnableMulticastForwarding implements
|
||||
// stack.MulticastForwardingNetworkProtocol.EnableMulticastForwarding.
|
||||
func (p *protocol) EnableMulticastForwarding(disp stack.MulticastForwardingEventDispatcher) (bool, tcpip.Error) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
if p.multicastForwardingDisp != nil {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
if disp == nil {
|
||||
return false, &tcpip.ErrInvalidOptionValue{}
|
||||
}
|
||||
|
||||
p.multicastForwardingDisp = disp
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// DisableMulticastForwarding implements
|
||||
// stack.MulticastForwardingNetworkProtocol.DisableMulticastForwarding.
|
||||
func (p *protocol) DisableMulticastForwarding() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
p.multicastForwardingDisp = nil
|
||||
p.multicastRouteTable.RemoveAllInstalledRoutes()
|
||||
}
|
||||
|
||||
// MulticastRouteLastUsedTime implements
|
||||
// stack.MulticastForwardingNetworkProtocol.
|
||||
func (p *protocol) MulticastRouteLastUsedTime(addresses stack.UnicastSourceAndMulticastDestination) (tcpip.MonotonicTime, tcpip.Error) {
|
||||
@@ -1589,6 +1635,11 @@ func (p *protocol) forwardPendingMulticastPacket(pkt *stack.PacketBuffer, instal
|
||||
// drop the pkt.
|
||||
return
|
||||
}
|
||||
|
||||
if !ep.MulticastForwarding() {
|
||||
return
|
||||
}
|
||||
|
||||
ep.handleForwardingError(ep.forwardValidatedMulticastPacket(pkt, installedRoute))
|
||||
}
|
||||
|
||||
@@ -1771,10 +1822,6 @@ type Options struct {
|
||||
// AllowExternalLoopbackTraffic indicates that inbound loopback packets (i.e.
|
||||
// martian loopback packets) should be accepted.
|
||||
AllowExternalLoopbackTraffic bool
|
||||
|
||||
// MulticastForwardingDisp is the multicast forwarding event dispatcher that
|
||||
// an integrator can provide to receive multicast forwarding events.
|
||||
MulticastForwardingDisp stack.MulticastForwardingEventDispatcher
|
||||
}
|
||||
|
||||
// NewProtocolWithOptions returns an IPv4 network protocol.
|
||||
|
||||
@@ -58,6 +58,15 @@ type testContext struct {
|
||||
clock *faketime.ManualClock
|
||||
}
|
||||
|
||||
var _ stack.MulticastForwardingEventDispatcher = (*fakeMulticastEventDispatcher)(nil)
|
||||
|
||||
type fakeMulticastEventDispatcher struct{}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnMissingRoute(context stack.MulticastPacketContext) {}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnUnexpectedInputInterface(context stack.MulticastPacketContext, expectedInputInterface tcpip.NICID) {
|
||||
}
|
||||
|
||||
func newTestContext() testContext {
|
||||
clock := faketime.NewManualClock()
|
||||
s := stack.New(stack.Options{
|
||||
@@ -203,6 +212,10 @@ func TestAddMulticastRouteIPv4Errors(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(ipv4.ProtocolNumber, &fakeMulticastEventDispatcher{}); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", ipv4.ProtocolNumber, err)
|
||||
}
|
||||
|
||||
outgoingInterfaces := []stack.MulticastRouteOutgoingInterface{{ID: outgoingNICID, MinTTL: 1}}
|
||||
|
||||
addresses := stack.UnicastSourceAndMulticastDestination{
|
||||
@@ -814,6 +827,10 @@ func TestMulticastFragmentForwarding(t *testing.T) {
|
||||
defer ctx.cleanup()
|
||||
s := ctx.s
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(ipv4.ProtocolNumber, &fakeMulticastEventDispatcher{}); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", ipv4.ProtocolNumber, err)
|
||||
}
|
||||
|
||||
endpoints := make(map[tcpip.NICID]*channel.Endpoint)
|
||||
for nicID, addr := range defaultEndpointConfigs {
|
||||
// For the input interface, we expect at most a single packet in
|
||||
@@ -983,6 +1000,10 @@ func TestMulticastForwardingOptions(t *testing.T) {
|
||||
// it give a more recognisable signature than 00,00,00,00.
|
||||
clock.Advance(time.Millisecond * randomTimeOffset)
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(ipv4.ProtocolNumber, &fakeMulticastEventDispatcher{}); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", ipv4.ProtocolNumber, err)
|
||||
}
|
||||
|
||||
endpoints := make(map[tcpip.NICID]*channel.Endpoint)
|
||||
for nicID, addr := range defaultEndpointConfigs {
|
||||
ep := channel.New(1, ipv4.MaxTotalSize, "")
|
||||
|
||||
@@ -199,15 +199,6 @@ func handleICMPInIPv6(ep stack.NetworkEndpoint, src, dst tcpip.Address, icmp hea
|
||||
pkt.DecRef()
|
||||
}
|
||||
|
||||
var _ stack.MulticastForwardingEventDispatcher = (*fakeMulticastEventDispatcher)(nil)
|
||||
|
||||
type fakeMulticastEventDispatcher struct{}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnMissingRoute(context stack.MulticastPacketContext) {}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnUnexpectedInputInterface(context stack.MulticastPacketContext, expectedInputInterface tcpip.NICID) {
|
||||
}
|
||||
|
||||
type testContext struct {
|
||||
s *stack.Stack
|
||||
clock *faketime.ManualClock
|
||||
@@ -216,7 +207,7 @@ type testContext struct {
|
||||
func newTestContext() testContext {
|
||||
clock := faketime.NewManualClock()
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{NewProtocolWithOptions(Options{MulticastForwardingDisp: &fakeMulticastEventDispatcher{}})},
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{NewProtocol},
|
||||
TransportProtocols: []stack.TransportProtocolFactory{icmp.NewProtocol6, udp.NewProtocol},
|
||||
Clock: clock,
|
||||
})
|
||||
|
||||
@@ -1128,9 +1128,11 @@ func (e *endpoint) forwardMulticastPacket(h header.IPv6, pkt *stack.PacketBuffer
|
||||
// Attempt to forward the pkt using an existing route.
|
||||
return e.forwardValidatedMulticastPacket(pkt, result.InstalledRoute)
|
||||
case multicast.NoRouteFoundAndPendingInserted:
|
||||
e.multicastEventDispatcher().OnMissingRoute(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
e.emitMulticastEvent(func(disp stack.MulticastForwardingEventDispatcher) {
|
||||
disp.OnMissingRoute(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
})
|
||||
})
|
||||
case multicast.PacketQueuedInPendingRoute:
|
||||
default:
|
||||
@@ -1152,10 +1154,12 @@ func (e *endpoint) forwardValidatedMulticastPacket(pkt *stack.PacketBuffer, inst
|
||||
// dropped silently.
|
||||
if e.nic.ID() != installedRoute.ExpectedInputInterface {
|
||||
h := header.IPv6(pkt.NetworkHeader().View())
|
||||
e.multicastEventDispatcher().OnUnexpectedInputInterface(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
}, installedRoute.ExpectedInputInterface)
|
||||
e.emitMulticastEvent(func(disp stack.MulticastForwardingEventDispatcher) {
|
||||
disp.OnUnexpectedInputInterface(stack.MulticastPacketContext{
|
||||
stack.UnicastSourceAndMulticastDestination{h.SourceAddress(), h.DestinationAddress()},
|
||||
e.nic.ID(),
|
||||
}, installedRoute.ExpectedInputInterface)
|
||||
})
|
||||
return &ip.ErrUnexpectedMulticastInputInterface{}
|
||||
}
|
||||
|
||||
@@ -1258,7 +1262,7 @@ func (e *endpoint) handleValidatedPacket(h header.IPv6, pkt *stack.PacketBuffer,
|
||||
// RFC 1812 section 5.2.3 for details regarding the forwarding/local
|
||||
// delivery decision.
|
||||
|
||||
multicastForwading := e.MulticastForwarding()
|
||||
multicastForwading := e.MulticastForwarding() && e.protocol.multicastForwarding()
|
||||
|
||||
if multicastForwading {
|
||||
e.handleForwardingError(e.forwardMulticastPacket(h, pkt))
|
||||
@@ -2126,6 +2130,11 @@ type protocol struct {
|
||||
|
||||
// ICMP types for which the stack's global rate limiting must apply.
|
||||
icmpRateLimitedTypes map[header.ICMPv6Type]struct{}
|
||||
|
||||
// multicastForwardingDisp is the multicast forwarding event dispatcher that
|
||||
// an integrator can provide to receive multicast forwarding events. Note
|
||||
// that multicast packets will only be forwarded if this is non-nil.
|
||||
multicastForwardingDisp stack.MulticastForwardingEventDispatcher
|
||||
}
|
||||
|
||||
ids []atomicbitops.Uint32
|
||||
@@ -2267,16 +2276,14 @@ func (p *protocol) DefaultTTL() uint8 {
|
||||
return uint8(p.defaultTTL.Load())
|
||||
}
|
||||
|
||||
// multicastEventDispatcher returns the multicast forwarding event dispatcher.
|
||||
//
|
||||
// Panics if a multicast forwarding event dispatcher does not exist. This
|
||||
// indicates that multicast forwarding is enabled, but no dispatcher was
|
||||
// provided.
|
||||
func (e *endpoint) multicastEventDispatcher() stack.MulticastForwardingEventDispatcher {
|
||||
if mcastDisp := e.protocol.options.MulticastForwardingDisp; mcastDisp != nil {
|
||||
return mcastDisp
|
||||
// emitMulticastEvent emits a multicast forwarding event using the provided
|
||||
// generator if a valid event dispatcher exists.
|
||||
func (e *endpoint) emitMulticastEvent(eventGenerator func(stack.MulticastForwardingEventDispatcher)) {
|
||||
e.protocol.mu.RLock()
|
||||
defer e.protocol.mu.RUnlock()
|
||||
if mcastDisp := e.protocol.mu.multicastForwardingDisp; mcastDisp != nil {
|
||||
eventGenerator(mcastDisp)
|
||||
}
|
||||
panic("e.procotol.options.MulticastForwardingDisp unexpectedly nil")
|
||||
}
|
||||
|
||||
// Close implements stack.TransportProtocol.
|
||||
@@ -2297,6 +2304,12 @@ func validateUnicastSourceAndMulticastDestination(addresses stack.UnicastSourceA
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *protocol) multicastForwarding() bool {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
return p.mu.multicastForwardingDisp != nil
|
||||
}
|
||||
|
||||
func (p *protocol) newInstalledRoute(route stack.MulticastRoute) (*multicast.InstalledRoute, tcpip.Error) {
|
||||
if len(route.OutgoingInterfaces) == 0 {
|
||||
return nil, &tcpip.ErrMissingRequiredFields{}
|
||||
@@ -2320,6 +2333,10 @@ func (p *protocol) newInstalledRoute(route stack.MulticastRoute) (*multicast.Ins
|
||||
|
||||
// AddMulticastRoute implements stack.MulticastForwardingNetworkProtocol.
|
||||
func (p *protocol) AddMulticastRoute(addresses stack.UnicastSourceAndMulticastDestination, route stack.MulticastRoute) tcpip.Error {
|
||||
if !p.multicastForwarding() {
|
||||
return &tcpip.ErrNotPermitted{}
|
||||
}
|
||||
|
||||
if err := validateUnicastSourceAndMulticastDestination(addresses); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -2367,6 +2384,33 @@ func (p *protocol) MulticastRouteLastUsedTime(addresses stack.UnicastSourceAndMu
|
||||
return timestamp, nil
|
||||
}
|
||||
|
||||
// EnableMulticastForwarding implements
|
||||
// stack.MulticastForwardingNetworkProtocol.EnableMulticastForwarding.
|
||||
func (p *protocol) EnableMulticastForwarding(disp stack.MulticastForwardingEventDispatcher) (bool, tcpip.Error) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
if p.mu.multicastForwardingDisp != nil {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
if disp == nil {
|
||||
return false, &tcpip.ErrInvalidOptionValue{}
|
||||
}
|
||||
|
||||
p.mu.multicastForwardingDisp = disp
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// DisableMulticastForwarding implements
|
||||
// stack.MulticastForwardingNetworkProtocol.DisableMulticastForwarding.
|
||||
func (p *protocol) DisableMulticastForwarding() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.mu.multicastForwardingDisp = nil
|
||||
p.multicastRouteTable.RemoveAllInstalledRoutes()
|
||||
}
|
||||
|
||||
func (p *protocol) forwardPendingMulticastPacket(pkt *stack.PacketBuffer, installedRoute *multicast.InstalledRoute) {
|
||||
defer pkt.DecRef()
|
||||
|
||||
@@ -2382,6 +2426,10 @@ func (p *protocol) forwardPendingMulticastPacket(pkt *stack.PacketBuffer, instal
|
||||
return
|
||||
}
|
||||
|
||||
if !ep.MulticastForwarding() {
|
||||
return
|
||||
}
|
||||
|
||||
ep.handleForwardingError(ep.forwardValidatedMulticastPacket(pkt, installedRoute))
|
||||
}
|
||||
|
||||
@@ -2543,10 +2591,6 @@ type Options struct {
|
||||
// AllowExternalLoopbackTraffic indicates that inbound loopback packets (i.e.
|
||||
// martian loopback packets) should be accepted.
|
||||
AllowExternalLoopbackTraffic bool
|
||||
|
||||
// MulticastForwardingDisp is the multicast forwarding event dispatcher that
|
||||
// an integrator can provide to receive multicast forwarding events.
|
||||
MulticastForwardingDisp stack.MulticastForwardingEventDispatcher
|
||||
}
|
||||
|
||||
// NewProtocolWithOptions returns an IPv6 network protocol.
|
||||
|
||||
@@ -59,6 +59,15 @@ const (
|
||||
extraHeaderReserve = 50
|
||||
)
|
||||
|
||||
var _ stack.MulticastForwardingEventDispatcher = (*fakeMulticastEventDispatcher)(nil)
|
||||
|
||||
type fakeMulticastEventDispatcher struct{}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnMissingRoute(context stack.MulticastPacketContext) {}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnUnexpectedInputInterface(context stack.MulticastPacketContext, expectedInputInterface tcpip.NICID) {
|
||||
}
|
||||
|
||||
// testReceiveICMP tests receiving an ICMP packet from src to dst. want is the
|
||||
// expected Neighbor Advertisement received count after receiving the packet.
|
||||
func testReceiveICMP(t *testing.T, s *stack.Stack, e *channel.Endpoint, src, dst tcpip.Address, want uint64) {
|
||||
@@ -3440,6 +3449,10 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
defer c.cleanup()
|
||||
s := c.s
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(ProtocolNumber, &fakeMulticastEventDispatcher{}); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", ProtocolNumber, err)
|
||||
}
|
||||
|
||||
endpoints := make(map[tcpip.NICID]*channel.Endpoint)
|
||||
for nicID, addr := range defaultEndpointConfigs {
|
||||
ep := channel.New(1, header.IPv6MinimumMTU, "")
|
||||
|
||||
@@ -816,6 +816,16 @@ type MulticastForwardingNetworkProtocol interface {
|
||||
// Returns an error if the addresses are invalid or a matching route was not
|
||||
// found.
|
||||
MulticastRouteLastUsedTime(UnicastSourceAndMulticastDestination) (tcpip.MonotonicTime, tcpip.Error)
|
||||
|
||||
// EnableMulticastForwarding enables multicast forwarding for the protocol.
|
||||
//
|
||||
// Returns an error if the provided multicast forwarding event dispatcher is
|
||||
// nil. Otherwise, returns true if the multicast forwarding was already
|
||||
// enabled.
|
||||
EnableMulticastForwarding(MulticastForwardingEventDispatcher) (bool, tcpip.Error)
|
||||
|
||||
// DisableMulticastForwarding disables multicast forwarding for the protocol.
|
||||
DisableMulticastForwarding()
|
||||
}
|
||||
|
||||
// MulticastPacketContext is the context in which a multicast packet triggered
|
||||
|
||||
@@ -616,6 +616,52 @@ func (s *Stack) MulticastRouteLastUsedTime(protocol tcpip.NetworkProtocolNumber,
|
||||
return forwardingNetProto.MulticastRouteLastUsedTime(addresses)
|
||||
}
|
||||
|
||||
// EnableMulticastForwardingForProtocol enables multicast forwarding for the
|
||||
// provided protocol.
|
||||
//
|
||||
// Returns true if forwarding was already enabled on the protocol.
|
||||
// Additionally, returns an error if:
|
||||
//
|
||||
// - The protocol is not found.
|
||||
// - The protocol doesn't support multicast forwarding.
|
||||
// - The multicast forwarding event dispatcher is nil.
|
||||
//
|
||||
// If successful, future multicast forwarding events will be sent to the
|
||||
// provided event dispatcher.
|
||||
func (s *Stack) EnableMulticastForwardingForProtocol(protocol tcpip.NetworkProtocolNumber, disp MulticastForwardingEventDispatcher) (bool, tcpip.Error) {
|
||||
netProto, ok := s.networkProtocols[protocol]
|
||||
if !ok {
|
||||
return false, &tcpip.ErrUnknownProtocol{}
|
||||
}
|
||||
|
||||
forwardingNetProto, ok := netProto.(MulticastForwardingNetworkProtocol)
|
||||
if !ok {
|
||||
return false, &tcpip.ErrNotSupported{}
|
||||
}
|
||||
|
||||
return forwardingNetProto.EnableMulticastForwarding(disp)
|
||||
}
|
||||
|
||||
// DisableMulticastForwardingForProtocol disables multicast forwarding for the
|
||||
// provided protocol.
|
||||
//
|
||||
// Returns an error if the provided protocol is not found or if it does not
|
||||
// support multicast forwarding.
|
||||
func (s *Stack) DisableMulticastForwardingForProtocol(protocol tcpip.NetworkProtocolNumber) tcpip.Error {
|
||||
netProto, ok := s.networkProtocols[protocol]
|
||||
if !ok {
|
||||
return &tcpip.ErrUnknownProtocol{}
|
||||
}
|
||||
|
||||
forwardingNetProto, ok := netProto.(MulticastForwardingNetworkProtocol)
|
||||
if !ok {
|
||||
return &tcpip.ErrNotSupported{}
|
||||
}
|
||||
|
||||
forwardingNetProto.DisableMulticastForwarding()
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetNICMulticastForwarding enables or disables multicast packet forwarding on
|
||||
// the specified NIC for the passed protocol.
|
||||
//
|
||||
|
||||
@@ -231,6 +231,11 @@ type addMulticastRouteData struct {
|
||||
route stack.MulticastRoute
|
||||
}
|
||||
|
||||
type enableMulticastForwardingForProtocolResult struct {
|
||||
AlreadyEnabled bool
|
||||
Err tcpip.Error
|
||||
}
|
||||
|
||||
// fakeNetworkProtocol is a network-layer protocol descriptor. It aggregates the
|
||||
// number of packets sent and received via endpoints of this protocol. The index
|
||||
// where packets are added is given by the packet's destination address MOD 10.
|
||||
@@ -244,6 +249,9 @@ type fakeNetworkProtocol struct {
|
||||
addMulticastRouteData addMulticastRouteData
|
||||
multicastRouteLastUsedTimeData stack.UnicastSourceAndMulticastDestination
|
||||
removeMulticastRouteData stack.UnicastSourceAndMulticastDestination
|
||||
|
||||
enableMulticastForwardingForProtocolResult enableMulticastForwardingForProtocolResult
|
||||
disableMulticastForwardingForProtocolCalled bool
|
||||
}
|
||||
|
||||
func (*fakeNetworkProtocol) Number() tcpip.NetworkProtocolNumber {
|
||||
@@ -329,6 +337,18 @@ func (f *fakeNetworkProtocol) MulticastRouteLastUsedTime(addresses stack.Unicast
|
||||
return tcpip.MonotonicTime{}, nil
|
||||
}
|
||||
|
||||
// EnableMulticastForwarding implements
|
||||
// MulticastForwardingNetworkProtocol.EnableMulticastForwarding.
|
||||
func (f *fakeNetworkProtocol) EnableMulticastForwarding(stack.MulticastForwardingEventDispatcher) (bool, tcpip.Error) {
|
||||
return f.enableMulticastForwardingForProtocolResult.AlreadyEnabled, f.enableMulticastForwardingForProtocolResult.Err
|
||||
}
|
||||
|
||||
// DisableMulticastForwarding implements
|
||||
// MulticastForwardingNetworkProtocol.DisableMulticastForwarding.
|
||||
func (f *fakeNetworkProtocol) DisableMulticastForwarding() {
|
||||
f.disableMulticastForwardingForProtocolCalled = true
|
||||
}
|
||||
|
||||
// Forwarding implements stack.ForwardingNetworkEndpoint.
|
||||
func (f *fakeNetworkEndpoint) Forwarding() bool {
|
||||
f.mu.RLock()
|
||||
@@ -382,6 +402,17 @@ func (l *linkEPWithMockedAttach) isAttached() bool {
|
||||
return l.attached
|
||||
}
|
||||
|
||||
var _ stack.MulticastForwardingEventDispatcher = (*fakeMulticastEventDispatcher)(nil)
|
||||
|
||||
type fakeMulticastEventDispatcher struct {
|
||||
}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnMissingRoute(context stack.MulticastPacketContext) {
|
||||
}
|
||||
|
||||
func (m *fakeMulticastEventDispatcher) OnUnexpectedInputInterface(context stack.MulticastPacketContext, expectedInputInterface tcpip.NICID) {
|
||||
}
|
||||
|
||||
// Checks to see if list contains an address.
|
||||
func containsAddr(list []tcpip.ProtocolAddress, item tcpip.ProtocolAddress) bool {
|
||||
for _, i := range list {
|
||||
@@ -4860,6 +4891,116 @@ func TestMulticastRouteLastUsedTime(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnableMulticastForwardingForProtocol(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
netProto tcpip.NetworkProtocolNumber
|
||||
factory stack.NetworkProtocolFactory
|
||||
delegateOutput enableMulticastForwardingForProtocolResult
|
||||
wantResult enableMulticastForwardingForProtocolResult
|
||||
}{
|
||||
{
|
||||
name: "impl returns previously enabled",
|
||||
netProto: fakeNetNumber,
|
||||
factory: fakeNetFactory,
|
||||
delegateOutput: enableMulticastForwardingForProtocolResult{true, nil},
|
||||
wantResult: enableMulticastForwardingForProtocolResult{true, nil},
|
||||
},
|
||||
{
|
||||
name: "impl returns previously disabled",
|
||||
netProto: fakeNetNumber,
|
||||
factory: fakeNetFactory,
|
||||
delegateOutput: enableMulticastForwardingForProtocolResult{false, nil},
|
||||
wantResult: enableMulticastForwardingForProtocolResult{false, nil},
|
||||
},
|
||||
{
|
||||
name: "impl returns error",
|
||||
netProto: fakeNetNumber,
|
||||
factory: fakeNetFactory,
|
||||
delegateOutput: enableMulticastForwardingForProtocolResult{false, &tcpip.ErrUnknownDevice{}},
|
||||
wantResult: enableMulticastForwardingForProtocolResult{false, &tcpip.ErrUnknownDevice{}},
|
||||
},
|
||||
{
|
||||
name: "unknown protocol",
|
||||
factory: fakeNetFactory,
|
||||
netProto: arp.ProtocolNumber,
|
||||
wantResult: enableMulticastForwardingForProtocolResult{false, &tcpip.ErrUnknownProtocol{}},
|
||||
},
|
||||
{
|
||||
name: "not supported",
|
||||
factory: arp.NewProtocol,
|
||||
netProto: arp.ProtocolNumber,
|
||||
wantResult: enableMulticastForwardingForProtocolResult{false, &tcpip.ErrNotSupported{}},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{test.factory},
|
||||
})
|
||||
|
||||
if test.netProto == fakeNetNumber {
|
||||
fakeNet := s.NetworkProtocolInstance(fakeNetNumber).(*fakeNetworkProtocol)
|
||||
fakeNet.enableMulticastForwardingForProtocolResult = test.delegateOutput
|
||||
}
|
||||
|
||||
alreadyEnabled, err := s.EnableMulticastForwardingForProtocol(test.netProto, &fakeMulticastEventDispatcher{})
|
||||
|
||||
if !cmp.Equal(enableMulticastForwardingForProtocolResult{alreadyEnabled, err}, test.wantResult, cmpopts.EquateErrors()) {
|
||||
t.Errorf("s.EnableMulticastForwardingForProtocol(%d, _) = (%t, %s), want = (%t, %s)", test.netProto, alreadyEnabled, err, test.wantResult.AlreadyEnabled, test.wantResult.Err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisableMulticastForwardingForProtocol(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
netProto tcpip.NetworkProtocolNumber
|
||||
factory stack.NetworkProtocolFactory
|
||||
wantErr tcpip.Error
|
||||
}{
|
||||
{
|
||||
name: "valid",
|
||||
netProto: fakeNetNumber,
|
||||
factory: fakeNetFactory,
|
||||
wantErr: nil,
|
||||
},
|
||||
{
|
||||
name: "unknown protocol",
|
||||
factory: fakeNetFactory,
|
||||
netProto: arp.ProtocolNumber,
|
||||
wantErr: &tcpip.ErrUnknownProtocol{},
|
||||
},
|
||||
{
|
||||
name: "not supported",
|
||||
factory: arp.NewProtocol,
|
||||
netProto: arp.ProtocolNumber,
|
||||
wantErr: &tcpip.ErrNotSupported{},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{test.factory},
|
||||
})
|
||||
|
||||
err := s.DisableMulticastForwardingForProtocol(test.netProto)
|
||||
|
||||
if !cmp.Equal(err, test.wantErr, cmpopts.EquateErrors()) {
|
||||
t.Errorf("s.DisableMulticastForwardingForProtocol(%d) = %s, want = %s", test.netProto, err, test.wantErr)
|
||||
}
|
||||
|
||||
if err == nil {
|
||||
fakeNet := s.NetworkProtocolInstance(fakeNetNumber).(*fakeNetworkProtocol)
|
||||
if !fakeNet.disableMulticastForwardingForProtocolCalled {
|
||||
t.Errorf("fakeNet.disableMulticastForwardingForProtocolCalled = false, want = true")
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNICForwarding(t *testing.T) {
|
||||
const nicID = 1
|
||||
|
||||
|
||||
@@ -231,15 +231,26 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
otherNICID: otherEndpointAddr,
|
||||
}
|
||||
|
||||
type multicastForwardingEvent int
|
||||
const (
|
||||
enabledForProtocol multicastForwardingEvent = iota
|
||||
enabledForNIC
|
||||
injectPendingPacket
|
||||
)
|
||||
|
||||
type multicastForwardingStateBeforeAddRouteCalled struct {
|
||||
multicastForwardingEvents []multicastForwardingEvent
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
srcAddr, dstAddr addrType
|
||||
routeIncomingNICID tcpip.NICID
|
||||
routeOutgoingNICID tcpip.NICID
|
||||
omitOutgoingInterfaces bool
|
||||
injectPendingPacket bool
|
||||
expectForward bool
|
||||
wantErr tcpip.Error
|
||||
name string
|
||||
srcAddr, dstAddr addrType
|
||||
routeIncomingNICID tcpip.NICID
|
||||
routeOutgoingNICID tcpip.NICID
|
||||
omitOutgoingInterfaces bool
|
||||
multicastForwardingEventsBeforeAddRouteCalled []multicastForwardingEvent
|
||||
expectForward bool
|
||||
wantErr tcpip.Error
|
||||
}{
|
||||
{
|
||||
name: "no pending packets",
|
||||
@@ -247,16 +258,26 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: nil,
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: nil,
|
||||
},
|
||||
{
|
||||
name: "pending packet forwarded",
|
||||
srcAddr: remoteUnicastAddr,
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
injectPendingPacket: true,
|
||||
expectForward: true,
|
||||
name: "packet arrived after forwarding enabled but before add route called",
|
||||
srcAddr: remoteUnicastAddr,
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol, injectPendingPacket},
|
||||
expectForward: true,
|
||||
},
|
||||
{
|
||||
name: "packet arrived before multicast forwarding enabled",
|
||||
srcAddr: remoteUnicastAddr,
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, injectPendingPacket, enabledForProtocol},
|
||||
expectForward: false,
|
||||
},
|
||||
{
|
||||
name: "unexpected input interface",
|
||||
@@ -264,9 +285,28 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
// The added route's incoming NICID does not match the pending packet's
|
||||
// incoming NICID. As a result, the packet should not be forwarded.
|
||||
routeIncomingNICID: otherNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
injectPendingPacket: true,
|
||||
routeIncomingNICID: otherNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
},
|
||||
{
|
||||
name: "multicast forwarding disabled for NIC",
|
||||
srcAddr: remoteUnicastAddr,
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForProtocol},
|
||||
expectForward: false,
|
||||
wantErr: nil,
|
||||
},
|
||||
{
|
||||
name: "multicast forwarding disabled for protocol",
|
||||
srcAddr: remoteUnicastAddr,
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC},
|
||||
wantErr: &tcpip.ErrNotPermitted{},
|
||||
},
|
||||
{
|
||||
name: "multicast source",
|
||||
@@ -274,7 +314,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "any source",
|
||||
@@ -282,7 +323,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "link-local unicast source",
|
||||
@@ -290,7 +332,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "empty source",
|
||||
@@ -298,7 +341,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "unicast destination",
|
||||
@@ -306,7 +350,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: remoteUnicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "empty destination",
|
||||
@@ -314,7 +359,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: emptyAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "link-local multicast destination",
|
||||
@@ -322,7 +368,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: linkLocalMulticastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrBadAddress{},
|
||||
},
|
||||
{
|
||||
name: "unknown input NICID",
|
||||
@@ -330,7 +377,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: unknownNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
wantErr: &tcpip.ErrUnknownNICID{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrUnknownNICID{},
|
||||
},
|
||||
{
|
||||
name: "unknown output NICID",
|
||||
@@ -338,7 +386,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: unknownNICID,
|
||||
wantErr: &tcpip.ErrUnknownNICID{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrUnknownNICID{},
|
||||
},
|
||||
{
|
||||
name: "input NIC matches output NIC",
|
||||
@@ -346,7 +395,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
dstAddr: multicastAddr,
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: incomingNICID,
|
||||
wantErr: &tcpip.ErrMulticastInputCannotBeOutput{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrMulticastInputCannotBeOutput{},
|
||||
},
|
||||
{
|
||||
name: "empty outgoing interfaces",
|
||||
@@ -355,7 +405,8 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
routeIncomingNICID: incomingNICID,
|
||||
routeOutgoingNICID: outgoingNICID,
|
||||
omitOutgoingInterfaces: true,
|
||||
wantErr: &tcpip.ErrMissingRequiredFields{},
|
||||
multicastForwardingEventsBeforeAddRouteCalled: []multicastForwardingEvent{enabledForNIC, enabledForProtocol},
|
||||
wantErr: &tcpip.ErrMissingRequiredFields{},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -364,10 +415,7 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
t.Run(fmt.Sprintf("%s %d", test.name, protocol), func(t *testing.T) {
|
||||
eventDispatcher := &fakeMulticastEventDispatcher{}
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{
|
||||
ipv4.NewProtocolWithOptions(ipv4.Options{MulticastForwardingDisp: eventDispatcher}),
|
||||
ipv6.NewProtocolWithOptions(ipv6.Options{MulticastForwardingDisp: eventDispatcher}),
|
||||
},
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
||||
})
|
||||
defer s.Close()
|
||||
|
||||
@@ -386,25 +434,37 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
if err := s.AddProtocolAddress(nicID, addr, stack.AddressProperties{}); err != nil {
|
||||
t.Fatalf("s.AddProtocolAddress(%d, %#v, {}): %s", nicID, addr, err)
|
||||
}
|
||||
s.SetNICMulticastForwarding(nicID, protocol, true /* enabled */)
|
||||
endpoints[nicID] = ep
|
||||
}
|
||||
|
||||
srcAddr := getAddr(protocol, test.srcAddr)
|
||||
dstAddr := getAddr(protocol, test.dstAddr)
|
||||
|
||||
if test.injectPendingPacket {
|
||||
incomingEp, ok := endpoints[incomingNICID]
|
||||
if !ok {
|
||||
t.Fatalf("got endpoints[%d] = (_, false), want (_, true)", incomingNICID)
|
||||
}
|
||||
for _, event := range test.multicastForwardingEventsBeforeAddRouteCalled {
|
||||
switch event {
|
||||
case enabledForNIC:
|
||||
for nicID := range endpoints {
|
||||
s.SetNICMulticastForwarding(nicID, protocol, true /* enable */)
|
||||
}
|
||||
case enabledForProtocol:
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(protocol, eventDispatcher); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", protocol, err)
|
||||
}
|
||||
case injectPendingPacket:
|
||||
incomingEp, ok := endpoints[incomingNICID]
|
||||
if !ok {
|
||||
t.Fatalf("got endpoints[%d] = (_, false), want (_, true)", incomingNICID)
|
||||
}
|
||||
|
||||
injectPacket(incomingEp, protocol, srcAddr, dstAddr, packetTTL)
|
||||
p := incomingEp.Read()
|
||||
injectPacket(incomingEp, protocol, srcAddr, dstAddr, packetTTL)
|
||||
p := incomingEp.Read()
|
||||
|
||||
if p != nil {
|
||||
// An ICMP error should never be sent in response to a multicast packet.
|
||||
t.Fatalf("got incomingEp.Read() = %#v, want = nil", p)
|
||||
if p != nil {
|
||||
// An ICMP error should never be sent in response to a multicast packet.
|
||||
t.Fatalf("got incomingEp.Read() = %#v, want = nil", p)
|
||||
}
|
||||
default:
|
||||
panic(fmt.Sprintf("unsupported multicastForwardingEvent: %d", event))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -451,6 +511,56 @@ func TestAddMulticastRoute(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnableMulticastForwardingE(t *testing.T) {
|
||||
eventDispatcher := &fakeMulticastEventDispatcher{}
|
||||
|
||||
type enableMulticastForwardingResult struct {
|
||||
AlreadyEnabled bool
|
||||
Err tcpip.Error
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
eventDispatcher stack.MulticastForwardingEventDispatcher
|
||||
wantResult []enableMulticastForwardingResult
|
||||
}{
|
||||
{
|
||||
name: "success",
|
||||
eventDispatcher: eventDispatcher,
|
||||
wantResult: []enableMulticastForwardingResult{{false, nil}},
|
||||
},
|
||||
{
|
||||
name: "already enabled",
|
||||
eventDispatcher: eventDispatcher,
|
||||
wantResult: []enableMulticastForwardingResult{{false, nil}, {true, nil}},
|
||||
},
|
||||
{
|
||||
name: "invalid event dispatcher",
|
||||
eventDispatcher: nil,
|
||||
wantResult: []enableMulticastForwardingResult{{false, &tcpip.ErrInvalidOptionValue{}}},
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
for _, protocol := range []tcpip.NetworkProtocolNumber{ipv4.ProtocolNumber, ipv6.ProtocolNumber} {
|
||||
t.Run(fmt.Sprintf("%s %d", test.name, protocol), func(t *testing.T) {
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
||||
TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol},
|
||||
})
|
||||
defer s.Close()
|
||||
|
||||
for _, wantResult := range test.wantResult {
|
||||
alreadyEnabled, err := s.EnableMulticastForwardingForProtocol(protocol, test.eventDispatcher)
|
||||
result := enableMulticastForwardingResult{alreadyEnabled, err}
|
||||
if !cmp.Equal(result, wantResult, cmpopts.EquateErrors()) {
|
||||
t.Errorf("s.EnableMulticastForwardingForProtocol(%d, %#v) = (%t, %s), want = (%t, %s)", protocol, test.eventDispatcher, alreadyEnabled, err, wantResult.AlreadyEnabled, wantResult.Err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMulticastRouteLastUsedTime(t *testing.T) {
|
||||
endpointConfigs := map[tcpip.NICID]endpointAddrType{
|
||||
incomingNICID: incomingEndpointAddr,
|
||||
@@ -530,6 +640,10 @@ func TestMulticastRouteLastUsedTime(t *testing.T) {
|
||||
})
|
||||
defer s.Close()
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(protocol, &fakeMulticastEventDispatcher{}); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", protocol, err)
|
||||
}
|
||||
|
||||
endpoints := make(map[tcpip.NICID]*channel.Endpoint)
|
||||
for nicID, addrType := range endpointConfigs {
|
||||
ep := channel.New(1, ipv4.MaxTotalSize, "")
|
||||
@@ -676,16 +790,16 @@ func TestRemoveMulticastRoute(t *testing.T) {
|
||||
for _, test := range tests {
|
||||
for _, protocol := range []tcpip.NetworkProtocolNumber{ipv4.ProtocolNumber, ipv6.ProtocolNumber} {
|
||||
t.Run(fmt.Sprintf("%s %d", test.name, protocol), func(t *testing.T) {
|
||||
eventDispatcher := &fakeMulticastEventDispatcher{}
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{
|
||||
ipv4.NewProtocolWithOptions(ipv4.Options{MulticastForwardingDisp: eventDispatcher}),
|
||||
ipv6.NewProtocolWithOptions(ipv6.Options{MulticastForwardingDisp: eventDispatcher}),
|
||||
},
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
||||
TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol},
|
||||
})
|
||||
defer s.Close()
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(protocol, &fakeMulticastEventDispatcher{}); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", protocol, err)
|
||||
}
|
||||
|
||||
endpoints := make(map[tcpip.NICID]*channel.Endpoint)
|
||||
for nicID, addrType := range endpointConfigs {
|
||||
ep := channel.New(1, ipv4.MaxTotalSize, "")
|
||||
@@ -791,16 +905,17 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
dstAddr addrType
|
||||
ttl uint8
|
||||
routeInputInterface tcpip.NICID
|
||||
disableMulticastForwarding bool
|
||||
removeOutputInterface tcpip.NICID
|
||||
expectMissingRouteEvent bool
|
||||
expectUnexpectedInputInterfaceEvent bool
|
||||
joinMulticastGroup bool
|
||||
expectedForwardingInterfaces []tcpip.NICID
|
||||
name string
|
||||
dstAddr addrType
|
||||
ttl uint8
|
||||
routeInputInterface tcpip.NICID
|
||||
disableMulticastForwardingForNIC bool
|
||||
updateMulticastForwardingForProtocol func(*testing.T, *stack.Stack, tcpip.NetworkProtocolNumber, stack.MulticastForwardingEventDispatcher)
|
||||
removeOutputInterface tcpip.NICID
|
||||
expectMissingRouteEvent bool
|
||||
expectUnexpectedInputInterfaceEvent bool
|
||||
joinMulticastGroup bool
|
||||
expectedForwardingInterfaces []tcpip.NICID
|
||||
}{
|
||||
{
|
||||
name: "forward only",
|
||||
@@ -826,13 +941,39 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
expectedForwardingInterfaces: []tcpip.NICID{},
|
||||
},
|
||||
{
|
||||
name: "multicast forwarding disabled",
|
||||
disableMulticastForwarding: true,
|
||||
dstAddr: multicastAddr,
|
||||
name: "multicast forwarding disabled for NIC",
|
||||
disableMulticastForwardingForNIC: true,
|
||||
dstAddr: multicastAddr,
|
||||
ttl: packetTTL,
|
||||
routeInputInterface: incomingNICID,
|
||||
expectedForwardingInterfaces: []tcpip.NICID{},
|
||||
},
|
||||
{
|
||||
name: "multicast forwarding disabled for protocol",
|
||||
dstAddr: multicastAddr,
|
||||
updateMulticastForwardingForProtocol: func(t *testing.T, s *stack.Stack, protocol tcpip.NetworkProtocolNumber, disp stack.MulticastForwardingEventDispatcher) {
|
||||
s.DisableMulticastForwardingForProtocol(protocol)
|
||||
},
|
||||
ttl: packetTTL,
|
||||
routeInputInterface: incomingNICID,
|
||||
expectedForwardingInterfaces: []tcpip.NICID{},
|
||||
},
|
||||
{
|
||||
name: "route table cleared after multicast forwarding disabled for protocol",
|
||||
dstAddr: multicastAddr,
|
||||
updateMulticastForwardingForProtocol: func(t *testing.T, s *stack.Stack, protocol tcpip.NetworkProtocolNumber, disp stack.MulticastForwardingEventDispatcher) {
|
||||
t.Helper()
|
||||
|
||||
s.DisableMulticastForwardingForProtocol(protocol)
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(protocol, disp); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, _): (_, %s)", protocol, err)
|
||||
}
|
||||
},
|
||||
ttl: packetTTL,
|
||||
routeInputInterface: incomingNICID,
|
||||
expectMissingRouteEvent: true,
|
||||
expectedForwardingInterfaces: []tcpip.NICID{},
|
||||
},
|
||||
{
|
||||
name: "unexpected input interface",
|
||||
dstAddr: multicastAddr,
|
||||
@@ -892,14 +1033,20 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
|
||||
t.Run(fmt.Sprintf("%s %d", test.name, protocol), func(t *testing.T) {
|
||||
s := stack.New(stack.Options{
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{
|
||||
ipv4.NewProtocolWithOptions(ipv4.Options{MulticastForwardingDisp: ipv4EventDispatcher}),
|
||||
ipv6.NewProtocolWithOptions(ipv6.Options{MulticastForwardingDisp: ipv6EventDispatcher}),
|
||||
},
|
||||
NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol},
|
||||
TransportProtocols: []stack.TransportProtocolFactory{udp.NewProtocol},
|
||||
})
|
||||
defer s.Close()
|
||||
|
||||
eventDispatcher, ok := eventDispatchers[protocol]
|
||||
if !ok {
|
||||
t.Fatalf("eventDispatchers[%d] = (_, false), want (_, true)", protocol)
|
||||
}
|
||||
|
||||
if _, err := s.EnableMulticastForwardingForProtocol(protocol, eventDispatcher); err != nil {
|
||||
t.Fatalf("s.EnableMulticastForwardingForProtocol(%d, %#v): (_, %s)", protocol, eventDispatcher, err)
|
||||
}
|
||||
|
||||
endpoints := make(map[tcpip.NICID]*channel.Endpoint)
|
||||
for nicID, addrType := range endpointConfigs {
|
||||
ep := channel.New(1, ipv4.MaxTotalSize, "")
|
||||
@@ -916,7 +1063,7 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
t.Fatalf("s.AddProtocolAddress(%d, %+v, {}): %s", nicID, addr, err)
|
||||
}
|
||||
|
||||
s.SetNICMulticastForwarding(nicID, protocol, !test.disableMulticastForwarding)
|
||||
s.SetNICMulticastForwarding(nicID, protocol, true /* enable */)
|
||||
endpoints[nicID] = ep
|
||||
}
|
||||
|
||||
@@ -945,6 +1092,16 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
t.Fatalf("AddMulticastRoute(%d, %#v, %#v): %s", protocol, addresses, route, err)
|
||||
}
|
||||
|
||||
if test.disableMulticastForwardingForNIC {
|
||||
for nicID := range endpoints {
|
||||
s.SetNICMulticastForwarding(nicID, protocol, false /* enable */)
|
||||
}
|
||||
}
|
||||
|
||||
if test.updateMulticastForwardingForProtocol != nil {
|
||||
test.updateMulticastForwardingForProtocol(t, s, protocol, eventDispatcher)
|
||||
}
|
||||
|
||||
if test.removeOutputInterface != 0 {
|
||||
if err := s.RemoveNIC(test.removeOutputInterface); err != nil {
|
||||
t.Fatalf("RemoveNIC(%d): %s", test.removeOutputInterface, err)
|
||||
@@ -1024,11 +1181,6 @@ func TestMulticastForwarding(t *testing.T) {
|
||||
p.DecRef()
|
||||
}
|
||||
|
||||
eventDispatcher, ok := eventDispatchers[protocol]
|
||||
if !ok {
|
||||
t.Fatalf("eventDispatchers[%d] = (_, false), want (_, true)", protocol)
|
||||
}
|
||||
|
||||
wantUnexpectedInputInterfaceEvent := func() *onUnexpectedInputInterfaceData {
|
||||
if test.expectUnexpectedInputInterfaceEvent {
|
||||
return &onUnexpectedInputInterfaceData{stack.MulticastPacketContext{stack.UnicastSourceAndMulticastDestination{srcAddr, dstAddr}, incomingNICID}, test.routeInputInterface}
|
||||
|
||||
Reference in New Issue
Block a user