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()) } }