diff --git a/internal/smdclient/SMDclient.go b/internal/smdclient/SMDclient.go index 6a3cb000..0e126c30 100644 --- a/internal/smdclient/SMDclient.go +++ b/internal/smdclient/SMDclient.go @@ -230,62 +230,62 @@ func (s *SMDClient) getSMD(ep string, smd interface{}) error { // PopulateNodes fetches the Ethernet interface data from the SMD server and populates the nodes map // with the corresponding node information, including MAC addresses, IP addresses, descriptions, and group membership. func (s *SMDClient) PopulateNodes() { - s.nodesMutex.Lock() - defer s.nodesMutex.Unlock() var ethIfaceArray []sm.CompEthInterfaceV2 ep := "/hsm/v2/Inventory/EthernetInterfaces/" if err := s.getSMD(ep, ðIfaceArray); err != nil { log.Error().Err(err).Msg("Failed to get SMD data") return } + + nextNodes := make(map[string]NodeMapping) log.Debug().Msgf("Populating nodes with %d Ethernet interfaces", len(ethIfaceArray)) - for _, ep := range ethIfaceArray { - if existingNode, exists := s.nodes[ep.CompID]; exists { + for _, ethIface := range ethIfaceArray { + if existingNode, exists := nextNodes[ethIface.CompID]; exists { found := false for index, existingInterface := range existingNode.Interfaces { - if strings.EqualFold(existingInterface.MAC, ep.MACAddr) { + if strings.EqualFold(existingInterface.MAC, ethIface.MACAddr) { // found the interface. Update the IP and Description found = true // Update the IP and Description - if len(ep.IPAddrs) > 0 { - existingInterface.IP = ep.IPAddrs[0].IPAddr + if len(ethIface.IPAddrs) > 0 { + existingInterface.IP = ethIface.IPAddrs[0].IPAddr } - existingInterface.Desc = ep.Desc + existingInterface.Desc = ethIface.Desc existingNode.Interfaces[index] = existingInterface } } if !found { // This is a new interface. Add it to the map newInterface := NodeInterface{ - MAC: ep.MACAddr, - Desc: ep.Desc, + MAC: ethIface.MACAddr, + Desc: ethIface.Desc, } - if len(ep.IPAddrs) > 0 { - newInterface.IP = ep.IPAddrs[0].IPAddr + if len(ethIface.IPAddrs) > 0 { + newInterface.IP = ethIface.IPAddrs[0].IPAddr } existingNode.Interfaces = append(existingNode.Interfaces, newInterface) - s.nodes[ep.CompID] = existingNode } + nextNodes[ethIface.CompID] = existingNode } else { // This is a new node newNode := NodeMapping{ - Xname: ep.CompID, + Xname: ethIface.CompID, } newInterface := NodeInterface{ - MAC: ep.MACAddr, - Desc: ep.Desc, + MAC: ethIface.MACAddr, + Desc: ethIface.Desc, } - log.Debug().Msgf("Adding new node %s with MAC %s and IPs: %v", ep.CompID, ep.MACAddr, ep.IPAddrs) - if len(ep.IPAddrs) > 0 { - newInterface.IP = ep.IPAddrs[0].IPAddr + log.Debug().Msgf("Adding new node %s with MAC %s and IPs: %v", ethIface.CompID, ethIface.MACAddr, ethIface.IPAddrs) + if len(ethIface.IPAddrs) > 0 { + newInterface.IP = ethIface.IPAddrs[0].IPAddr } newNode.Interfaces = append(newNode.Interfaces, newInterface) - s.nodes[ep.CompID] = newNode + nextNodes[ethIface.CompID] = newNode } } // Populate group membership for all nodes log.Debug().Msg("Fetching group membership for all nodes") - for xname, node := range s.nodes { + for xname, node := range nextNodes { ml := new(sm.Membership) membershipEp := "/hsm/v2/memberships/" + xname if err := s.getSMD(membershipEp, ml); err != nil { @@ -294,32 +294,57 @@ func (s *SMDClient) PopulateNodes() { } else { node.Groups = ml.GroupLabels } - s.nodes[xname] = node + nextNodes[xname] = node } // Build reverse indexes for O(1) lookups log.Debug().Msg("Building reverse indexes") - s.ipToXname = make(map[string]string) - s.macToXname = make(map[string]string) - s.wgipToXname = make(map[string]string) + nextIPToXname := make(map[string]string) + nextMACToXname := make(map[string]string) + nextWGIPToXname := make(map[string]string) - for xname, node := range s.nodes { + for xname, node := range nextNodes { for _, iface := range node.Interfaces { if iface.IP != "" { - s.ipToXname[strings.ToLower(iface.IP)] = xname + nextIPToXname[strings.ToLower(iface.IP)] = xname } if iface.MAC != "" { - s.macToXname[strings.ToLower(iface.MAC)] = xname + nextMACToXname[strings.ToLower(iface.MAC)] = xname } if iface.WGIP != "" { - s.wgipToXname[strings.ToLower(iface.WGIP)] = xname + nextWGIPToXname[strings.ToLower(iface.WGIP)] = xname + } + } + } + + s.nodesMutex.Lock() + defer s.nodesMutex.Unlock() + + for xname, nextNode := range nextNodes { + currentNode, found := s.nodes[xname] + if !found { + continue + } + for nextIndex, nextInterface := range nextNode.Interfaces { + for _, currentInterface := range currentNode.Interfaces { + if currentInterface.WGIP == "" || !strings.EqualFold(currentInterface.MAC, nextInterface.MAC) { + continue + } + nextNode.Interfaces[nextIndex].WGIP = currentInterface.WGIP + nextWGIPToXname[strings.ToLower(currentInterface.WGIP)] = xname + break } } + nextNodes[xname] = nextNode } + s.nodes = nextNodes + s.ipToXname = nextIPToXname + s.macToXname = nextMACToXname + s.wgipToXname = nextWGIPToXname s.nodes_last_update = time.Now() log.Debug().Msgf("Nodes map populated with %d nodes, %d IP mappings, %d MAC mappings", - len(s.nodes), len(s.ipToXname), len(s.macToXname)) + len(nextNodes), len(nextIPToXname), len(nextMACToXname)) } // IDfromMAC returns the ID of the xname that has the MAC address diff --git a/internal/smdclient/SMDclient_performance_test.go b/internal/smdclient/SMDclient_performance_test.go index c0a03372..4c147fed 100644 --- a/internal/smdclient/SMDclient_performance_test.go +++ b/internal/smdclient/SMDclient_performance_test.go @@ -5,6 +5,7 @@ import ( "net/http" "net/http/httptest" "sync" + "sync/atomic" "testing" "time" @@ -12,6 +13,114 @@ import ( "github.com/stretchr/testify/require" ) +func TestPopulateNodesBlockedRefreshDoesNotBlockCachedOperations(t *testing.T) { + var blockRefresh atomic.Bool + refreshStarted := make(chan struct{}) + releaseRefresh := make(chan struct{}) + var signalRefreshStarted sync.Once + var releaseRefreshOnce sync.Once + + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/hsm/v2/Inventory/EthernetInterfaces/" && blockRefresh.Load() { + signalRefreshStarted.Do(func() { close(refreshStarted) }) + <-releaseRefresh + } + + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + switch r.URL.Path { + case "/hsm/v2/Inventory/EthernetInterfaces/": + _, _ = w.Write([]byte(`[{ + "ComponentID": "x1000", + "MACAddress": "00:11:22:33:44:55", + "IPAddresses": [{"IPAddress": "192.168.1.1"}], + "Description": "Test Node" + }]`)) + case "/hsm/v2/memberships/x1000": + _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + } + }) + server := httptest.NewServer(handler) + defer server.Close() + release := func() { + releaseRefreshOnce.Do(func() { close(releaseRefresh) }) + } + defer release() + + client := &SMDClient{ + smdClient: server.Client(), + smdBaseURL: server.URL, + nodesMutex: &sync.RWMutex{}, + nodes: make(map[string]NodeMapping), + ipToXname: make(map[string]string), + macToXname: make(map[string]string), + wgipToXname: make(map[string]string), + } + client.PopulateNodes() + + blockRefresh.Store(true) + refreshDone := make(chan struct{}) + go func() { + client.PopulateNodes() + close(refreshDone) + }() + + select { + case <-refreshStarted: + case <-time.After(time.Second): + t.Fatal("refresh did not reach blocked SMD handler") + } + + lookupResult := make(chan struct { + xname string + err error + }, 1) + go func() { + xname, err := client.IDfromIP("192.168.1.1") + lookupResult <- struct { + xname string + err error + }{xname: xname, err: err} + }() + + select { + case result := <-lookupResult: + require.NoError(t, result.err) + assert.Equal(t, "x1000", result.xname) + case <-time.After(time.Second): + t.Fatal("IDfromIP blocked on slow PopulateNodes network I/O") + } + + addWGIPResult := make(chan error, 1) + go func() { + addWGIPResult <- client.AddWGIP("x1000", "10.99.0.2") + }() + + select { + case err := <-addWGIPResult: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("AddWGIP blocked on slow PopulateNodes network I/O") + } + + release() + select { + case <-refreshDone: + case <-time.After(time.Second): + t.Fatal("PopulateNodes did not finish after SMD response was released") + } + + xname, err := client.IDfromIP("10.99.0.2") + require.NoError(t, err) + assert.Equal(t, "x1000", xname) + wgip, err := client.WGIPfromID("x1000") + require.NoError(t, err) + assert.Equal(t, "10.99.0.2", wgip) + groups, err := client.GroupMembership("x1000") + require.NoError(t, err) + assert.Equal(t, []string{"compute"}, groups) +} + // TestGroupMembershipCached verifies that Bug #1 is fixed: // GroupMembership should use the cache instead of making HTTP requests func TestGroupMembershipCached(t *testing.T) {