From 94db2b2de7a7c7ac53d56c56489b65143fb8aa5e Mon Sep 17 00:00:00 2001 From: Jing Chen Date: Mon, 4 Nov 2024 17:21:34 -0800 Subject: [PATCH] Implement a basic bridge FDB by leanrning. Bridge uses bridge FDB to decide which port(s) it forwards the received packets to. The bridge FDB, when the bridge receives a packet, records the packet's source MAC address and the bridge port where the packet is received. The received packet's source MAC address will be the lookup key to the bridge FDB. When deciding a forward port, the bridge looks for a FDB key which matches the packet's destination MAC address, the associated port from the FDB will forward the packet. Otherwise, the packet will flood to all available bridge ports. The FDB will never be garbage collected until the nic device is removed from the bridge. PiperOrigin-RevId: 693143851 --- pkg/tcpip/stack/bridge.go | 90 +++++++++++-- pkg/tcpip/stack/bridge_test.go | 225 ++++++++++++++++++++++++--------- 2 files changed, 244 insertions(+), 71 deletions(-) diff --git a/pkg/tcpip/stack/bridge.go b/pkg/tcpip/stack/bridge.go index 50c8f9640..5692d164c 100644 --- a/pkg/tcpip/stack/bridge.go +++ b/pkg/tcpip/stack/bridge.go @@ -28,6 +28,22 @@ type bridgePort struct { nic *nic } +// BridgeFDBKey is the MAC address of a device which a bridge port is associated with. +type BridgeFDBKey tcpip.LinkAddress + +// BridgeFDBEntry consists of all metadata for a FDB record. +type BridgeFDBEntry struct { + port *bridgePort +} + +// PortLinkAddress returns the mac address of the device that is bound to the bridge port. +func (e BridgeFDBEntry) PortLinkAddress() tcpip.LinkAddress { + if e.port == nil { + return "" + } + return e.port.nic.LinkAddress() +} + // ParseHeader implements stack.LinkEndpoint. func (p *bridgePort) ParseHeader(pkt *PacketBuffer) bool { _, ok := pkt.LinkHeader().Consume(header.EthernetMinimumSize) @@ -37,23 +53,49 @@ func (p *bridgePort) ParseHeader(pkt *PacketBuffer) bool { // DeliverNetworkPacket implements stack.NetworkDispatcher. func (p *bridgePort) DeliverNetworkPacket(protocol tcpip.NetworkProtocolNumber, pkt *PacketBuffer) { bridge := p.bridge + eth := header.Ethernet(pkt.LinkHeader().Slice()) + updateFDB := false bridge.mu.RLock() - - // Send the packet to all other ports. - for _, port := range bridge.ports { - if p == port { - continue + // Add an entry at the bridge FDB, it maps a MAC address + // to a bridge port where the traffic is received when + // the MAC address is not multicast. + // Network packets that are sent to the learned MAC address + // will be forwarded to the bridge port that is stored in + // the FDB table. + sourceAddress := eth.SourceAddress() + if _, hasSourceFDB := bridge.fdbTable[BridgeFDBKey(sourceAddress)]; !header.IsMulticastEthernetAddress(sourceAddress) && !hasSourceFDB { + updateFDB = true + } + if entry, exist := bridge.fdbTable[BridgeFDBKey(eth.DestinationAddress())]; !exist { + // When no FDB entry is found, send the packet to all ports. + for _, port := range bridge.ports { + if p == port { + continue + } + newPkt := NewPacketBuffer(PacketBufferOptions{ + ReserveHeaderBytes: int(port.nic.MaxHeaderLength()), + Payload: pkt.ToBuffer(), + }) + port.nic.writeRawPacket(newPkt) + newPkt.DecRef() } + } else if entry.port != p { + destPort := entry.port newPkt := NewPacketBuffer(PacketBufferOptions{ - ReserveHeaderBytes: int(port.nic.MaxHeaderLength()), + ReserveHeaderBytes: int(destPort.nic.MaxHeaderLength()), Payload: pkt.ToBuffer(), }) - port.nic.writeRawPacket(newPkt) + destPort.nic.writeRawPacket(newPkt) newPkt.DecRef() } d := bridge.dispatcher bridge.mu.RUnlock() + if updateFDB { + bridge.mu.Lock() + bridge.addFDBEntryLocked(eth.SourceAddress(), p, 0) + bridge.mu.Unlock() + } if d != nil { // The dispatcher may acquire Stack.mu in DeliverNetworkPacket(), which is // ordered above bridge.mu. So call DeliverNetworkPacket() without holding @@ -72,6 +114,7 @@ func NewBridgeEndpoint(mtu uint32) *BridgeEndpoint { addr: tcpip.GetRandMacAddr(), } b.ports = make(map[tcpip.NICID]*bridgePort) + b.fdbTable = make(map[BridgeFDBKey]BridgeFDBEntry) return b } @@ -89,7 +132,9 @@ type BridgeEndpoint struct { // +checklocks:mu attached bool // +checklocks:mu - mtu uint32 + mtu uint32 + // +checklocks:mu + fdbTable map[BridgeFDBKey]BridgeFDBEntry maxHeaderLength atomicbitops.Uint32 } @@ -143,6 +188,12 @@ func (b *BridgeEndpoint) DelNIC(nic *nic) tcpip.Error { b.mu.Lock() defer b.mu.Unlock() + port := b.ports[nic.id] + for k, e := range b.fdbTable { + if e.port == port { + delete(b.fdbTable, k) + } + } delete(b.ports, nic.id) nic.NetworkLinkEndpoint.Attach(nic) return nil @@ -198,6 +249,7 @@ func (b *BridgeEndpoint) Attach(dispatcher NetworkDispatcher) { } b.dispatcher = dispatcher b.ports = make(map[tcpip.NICID]*bridgePort) + b.fdbTable = make(map[BridgeFDBKey]BridgeFDBEntry) } // IsAttached implements stack.LinkEndpoint.IsAttached. @@ -230,3 +282,25 @@ func (b *BridgeEndpoint) Close() {} // SetOnCloseAction implements stack.LinkEndpoint.Close. func (b *BridgeEndpoint) SetOnCloseAction(func()) {} + +// Add a new FDBEntry by learning. The learning happens when a packaet +// is recevied by a bridge port, the bridge will use the port for the future +// deliveries to the NIC device. +// The addr is the key when it looks for the entry. +// +// +checklocks:b.mu +func (b *BridgeEndpoint) addFDBEntryLocked(addr tcpip.LinkAddress, source *bridgePort, flags uint64) bool { + // TODO(b/376924093): limit bridge FDB size. + b.fdbTable[BridgeFDBKey(addr)] = BridgeFDBEntry{ + port: source, + } + return true +} + +// FindFDBEntry find the FDB entry for the given address. If it doesn't exist, +// it will return an empty entry. +func (b *BridgeEndpoint) FindFDBEntry(addr tcpip.LinkAddress) BridgeFDBEntry { + b.mu.RLock() + defer b.mu.RUnlock() + return b.fdbTable[BridgeFDBKey(addr)] +} diff --git a/pkg/tcpip/stack/bridge_test.go b/pkg/tcpip/stack/bridge_test.go index a1efeb3c5..c066abd0b 100644 --- a/pkg/tcpip/stack/bridge_test.go +++ b/pkg/tcpip/stack/bridge_test.go @@ -30,48 +30,62 @@ import ( func TestWritePacketFromBridge(t *testing.T) { const ( - localLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x06") - remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x07") - bridgeLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x08") + channelLinkAddr1 = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x05") + channelLinkAddr2 = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x06") + remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x07") + bridgeLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x08") netProto = 55 - nicID = 5 - bridgeID = 6 + nicID1 = 5 + nicID2 = 6 + bridgeID = 7 ) - c := channel.New(1, header.EthernetMinimumSize, localLinkAddr) - - s := stack.New(stack.Options{}) - if err := s.CreateNIC(nicID, ethernet.New(c)); err != nil { - t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) - } + // Create two channel-based endpoints as receivers, they are both + // bound to the same bridge device and are both expected to receive the + // packets that are written by the bridge. + ch1 := channel.New(1, header.EthernetMinimumSize, channelLinkAddr1) + ch2 := channel.New(1, header.EthernetMinimumSize, channelLinkAddr2) bridgeEndpoint := stack.NewBridgeEndpoint(1500) bridgeEndpoint.SetLinkAddress(bridgeLinkAddr) + s := stack.New(stack.Options{}) + + if err := s.CreateNIC(nicID1, ethernet.New(ch1)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID1, err) + } + if err := s.CreateNIC(nicID2, ethernet.New(ch2)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID2, err) + } if err := s.CreateNIC(bridgeID, bridgeEndpoint); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", bridgeID, err) } - if err := s.SetNICCoordinator(nicID, bridgeID); err != nil { - t.Fatalf("s.SetNICCoordinator") + if err := s.SetNICCoordinator(nicID1, bridgeID); err != nil { + t.Fatalf("s.SetNICCoordinator(%d, %d)", nicID1, bridgeID) } - + if err := s.SetNICCoordinator(nicID2, bridgeID); err != nil { + t.Fatalf("s.SetNICCoordinator(%d, %d)", nicID2, bridgeID) + } + // When writing packets, the bridge will try all available bridge ports. if err := s.WritePacketToRemote(bridgeID, remoteLinkAddr, netProto, buffer.Buffer{}); err != nil { t.Fatalf("s.WritePacketToRemote(%d, %s, _): %s", bridgeID, remoteLinkAddr, err) } - pkt := c.Read() - if pkt == nil { - t.Fatal("expected to read a packet") - } + for _, c := range []*channel.Endpoint{ch1, ch2} { + pkt := c.Read() + if pkt == nil { + t.Fatal("expected to read a packet") + } - eth := header.Ethernet(pkt.LinkHeader().Slice()) - pkt.DecRef() - if got := eth.SourceAddress(); got != bridgeLinkAddr { - t.Errorf("got eth.SourceAddress() = %s, want = %s", got, bridgeLinkAddr) - } - if got := eth.DestinationAddress(); got != remoteLinkAddr { - t.Errorf("got eth.DestinationAddress() = %s, want = %s", got, remoteLinkAddr) - } - if got := eth.Type(); got != netProto { - t.Errorf("got eth.Type() = %d, want = %d", got, netProto) + eth := header.Ethernet(pkt.LinkHeader().Slice()) + pkt.DecRef() + if got := eth.SourceAddress(); got != bridgeLinkAddr { + t.Errorf("got eth.SourceAddress() = %s, want = %s", got, bridgeLinkAddr) + } + if got := eth.DestinationAddress(); got != remoteLinkAddr { + t.Errorf("got eth.DestinationAddress() = %s, want = %s", got, remoteLinkAddr) + } + if got := eth.Type(); got != netProto { + t.Errorf("got eth.Type() = %d, want = %d", got, netProto) + } } } @@ -83,70 +97,155 @@ func (n *testNotification) WriteNotify() { n.ch <- true } +// The test verifies that packates that are forwarded by +// a bridge will flooded to all bridge ports. func TestWritePacketBetweenDevices(t *testing.T) { const ( - localLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x06") - remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x07") - bridgeLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x08") + channelLinkAddr1 = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x04") + channelLinkAddr2 = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x05") + vethLinkAddr1 = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x06") + vethLinkAddr2 = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x07") + remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x08") + bridgeLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x09") netProto = 55 - nicID = 5 - vethID = 7 - bridgeID = 6 + nicID1 = 4 + nicID2 = 5 + vethID = 6 + bridgeID = 9 ) - + // Creates a pair of veth devices which will be attached to different + // network stacks. veth1, veth2 := veth.NewPair(1500) - secondStack := stack.New(stack.Options{}) - if err := secondStack.CreateNIC(vethID, ethernet.New(veth2)); err != nil { - t.Fatalf("s.CreateNIC(%d, _): %s", vethID, err) - } - veth2.SetLinkAddress(localLinkAddr) + veth1.SetLinkAddress(vethLinkAddr1) + veth2.SetLinkAddress(vethLinkAddr2) + ch1 := channel.New(1, header.EthernetMinimumSize, channelLinkAddr1) + ch2 := channel.New(1, header.EthernetMinimumSize, channelLinkAddr2) - s := stack.New(stack.Options{}) bridgeEndpoint := stack.NewBridgeEndpoint(1500) bridgeEndpoint.SetLinkAddress(bridgeLinkAddr) + s := stack.New(stack.Options{}) + secondStack := stack.New(stack.Options{}) if err := s.CreateNIC(bridgeID, bridgeEndpoint); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", bridgeID, err) } - - c := channel.New(1, header.EthernetMinimumSize, localLinkAddr) - c.SetLinkAddress(remoteLinkAddr) - if err := s.CreateNIC(nicID, ethernet.New(c)); err != nil { - t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) + if err := s.CreateNIC(nicID1, ethernet.New(ch1)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID1, err) } - if err := s.SetNICCoordinator(nicID, bridgeID); err != nil { - t.Fatalf("s.SetNICCoordinator") + if err := s.SetNICCoordinator(nicID1, bridgeID); err != nil { + t.Fatalf("s.SetNICCoordinator(%d, %d)", nicID1, bridgeID) + } + if err := s.CreateNIC(nicID2, ethernet.New(ch2)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID2, err) + } + if err := s.SetNICCoordinator(nicID2, bridgeID); err != nil { + t.Fatalf("s.SetNICCoordinator(%d, %d)", nicID2, bridgeID) } - if err := s.CreateNIC(vethID, ethernet.New(veth1)); err != nil { t.Fatalf("s.CreateNIC(%d, _): %s", vethID, err) } + // Attach one veth device to stack s, and attach the other + // veth to stack secondStack. if err := s.SetNICCoordinator(vethID, bridgeID); err != nil { t.Fatalf("s.SetNICCoordinator") } + if err := secondStack.CreateNIC(vethID, ethernet.New(veth2)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", vethID, err) + } - n := &testNotification{ch: make(chan bool, 1)} - c.AddNotify(n) + n1 := &testNotification{ch: make(chan bool, 1)} + n2 := &testNotification{ch: make(chan bool, 1)} + ch1.AddNotify(n1) + ch2.AddNotify(n2) + // Write a packte to the veth device at the stack secondStack, the packet + // will be available at the veth device at the stack s. if err := secondStack.WritePacketToRemote(vethID, remoteLinkAddr, netProto, buffer.Buffer{}); err != nil { t.Fatalf("s.WritePacketToRemote(%d, %s, _): %s", bridgeID, remoteLinkAddr, err) } + <-n1.ch + <-n2.ch + // No FDB entry, a package floods all bridge ports except the port that + // is attached to the veth device. + for _, c := range []*channel.Endpoint{ch1, ch2} { + pkt := c.Read() + if pkt == nil { + t.Fatal("expected to read a packet") + } + + pkt.LinkHeader().Consume(header.EthernetMinimumSize) + eth := header.Ethernet(pkt.LinkHeader().Slice()) + pkt.DecRef() + if got := eth.SourceAddress(); got != vethLinkAddr2 { + t.Errorf("got eth.SourceAddress() = %s, want = %s", got, vethLinkAddr2) + } + if got := eth.DestinationAddress(); got != remoteLinkAddr { + t.Errorf("got eth.DestinationAddress() = %s, want = %s", got, remoteLinkAddr) + } + if got := eth.Type(); got != netProto { + t.Errorf("got eth.Type() = %d, want = %d", got, netProto) + } + } +} + +func TestBridgeFDB(t *testing.T) { + const ( + channelLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x03") + remoteLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x08") + bridgeLinkAddr = tcpip.LinkAddress("\x02\x02\x03\x04\x05\x09") + + netProto = 55 + nicID = 4 + vethID1 = 7 + vethID2 = 8 + bridgeID = 9 + ) + veth1, veth2 := veth.NewPair(1500) + ch := channel.New(1, header.EthernetMinimumSize, channelLinkAddr) + bridgeEndpoint := stack.NewBridgeEndpoint(1500) + bridgeEndpoint.SetLinkAddress(bridgeLinkAddr) + s := stack.New(stack.Options{}) + secondStack := stack.New(stack.Options{}) + + if err := s.CreateNIC(bridgeID, bridgeEndpoint); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", bridgeID, err) + } + if err := s.CreateNIC(nicID, ethernet.New(ch)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", nicID, err) + } + if err := s.SetNICCoordinator(nicID, bridgeID); err != nil { + t.Fatalf("s.SetNICCoordinator(%d, %d)", nicID, bridgeID) + } + if err := s.CreateNIC(vethID1, ethernet.New(veth1)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", vethID1, err) + } + if err := s.SetNICCoordinator(vethID1, bridgeID); err != nil { + t.Fatalf("s.SetNICCoordinator(%d, %d)", vethID1, bridgeID) + } + if err := secondStack.CreateNIC(vethID2, ethernet.New(veth2)); err != nil { + t.Fatalf("s.CreateNIC(%d, _): %s", vethID2, err) + } + n := &testNotification{ch: make(chan bool, 1)} + ch.AddNotify(n) + // Write a packet to the veth device at secondStack, the + // packet will be available at stack s via veth. + if err := secondStack.WritePacketToRemote(vethID2, remoteLinkAddr, netProto, buffer.Buffer{}); err != nil { + t.Fatalf("s.WritePacketToRemote(%d, %s, _): %s", bridgeID, remoteLinkAddr, err) + } <-n.ch - pkt := c.Read() + pkt := ch.Read() + defer pkt.DecRef() if pkt == nil { t.Fatal("expected to read a packet") } - - pkt.LinkHeader().Consume(header.EthernetMinimumSize) - eth := header.Ethernet(pkt.LinkHeader().Slice()) - pkt.DecRef() - if got := eth.SourceAddress(); got != localLinkAddr { - t.Errorf("got eth.SourceAddress() = %s, want = %s", got, localLinkAddr) + // When forwarding the packet via the bridge device, the packet's + // source MAC address will be used as lookup key to the bridge + // FDB. + if e := bridgeEndpoint.FindFDBEntry(veth2.LinkAddress()); e.PortLinkAddress() != veth1.LinkAddress() { + t.Fatalf("bridgeEndpoint.FindFDBEntry(%s) = %s, want = %s", veth2.LinkAddress(), e.PortLinkAddress(), veth1.LinkAddress()) } - if got := eth.DestinationAddress(); got != remoteLinkAddr { - t.Errorf("got eth.DestinationAddress() = %s, want = %s", got, remoteLinkAddr) - } - if got := eth.Type(); got != netProto { - t.Errorf("got eth.Type() = %d, want = %d", got, netProto) + // No FDB entry is expected for devices other than veth1. + if e := bridgeEndpoint.FindFDBEntry(channelLinkAddr); len(e.PortLinkAddress()) != 0 { + t.Fatalf("bridgeEndpoint.FindFDBEntry(%s) = %s, want = \"\"", channelLinkAddr, e.PortLinkAddress()) } }