From 2c22d460766d32b8587863710b49b68d9fae533d Mon Sep 17 00:00:00 2001 From: Alex Lovell-Troy Date: Mon, 3 Aug 2026 16:33:10 -0400 Subject: [PATCH 1/2] feat: enhance SMD client with membership listing and improve tests for node population Signed-off-by: Alex Lovell-Troy --- .gitignore | 3 + internal/smdclient/FakeSMDClient.go | 23 ++ internal/smdclient/FakeSMDClient_test.go | 51 ++++ internal/smdclient/SMDclient.go | 21 +- .../smdclient/SMDclient_performance_test.go | 253 ++++++++++-------- internal/smdclient/SMDclient_test.go | 211 ++++++++++++--- 6 files changed, 409 insertions(+), 153 deletions(-) diff --git a/.gitignore b/.gitignore index 194267600..bfd0c7f58 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,6 @@ dist # Our compiled binary /cloud-init-server + +.omo/ +graphify-out/cache/ \ No newline at end of file diff --git a/internal/smdclient/FakeSMDClient.go b/internal/smdclient/FakeSMDClient.go index 8dadaacd8..492b4dc6f 100644 --- a/internal/smdclient/FakeSMDClient.go +++ b/internal/smdclient/FakeSMDClient.go @@ -11,6 +11,7 @@ import ( "github.com/rs/zerolog/log" base "github.com/Cray-HPE/hms-base" + "github.com/OpenCHAMI/smd/v2/pkg/sm" "github.com/OpenCHAMI/cloud-init/pkg/cistore" ) @@ -288,6 +289,28 @@ func (f *FakeSMDClient) PopulateNodes() { // ***** Simulated SMD Client functions. Not part of the SMDClientInterface ***** +// ListMemberships returns memberships derived from the simulator's group state. +func (f *FakeSMDClient) ListMemberships() []sm.Membership { + memberships := make([]sm.Membership, 0, len(f.rosetta_mapping)) + for _, component := range f.rosetta_mapping { + groupLabels := make([]string, 0) + for group, componentIDs := range f.groups { + for _, componentID := range componentIDs { + if componentID == component.ComponentID { + groupLabels = append(groupLabels, group) + break + } + } + } + memberships = append(memberships, sm.Membership{ + ID: component.ComponentID, + GroupLabels: groupLabels, + PartitionName: "", + }) + } + return memberships +} + // AddNodeToInventory adds a node to the inventory. This is not part of the SMDClient Interface and only useful as part of the simulator func (f *FakeSMDClient) AddNodeToInventory(node cistore.OpenCHAMIComponent) error { log.Debug().Msgf("FakeSMDClient: AddNodeToInventory(%s)", node.ID) diff --git a/internal/smdclient/FakeSMDClient_test.go b/internal/smdclient/FakeSMDClient_test.go index 33ee69ebd..65e16a626 100644 --- a/internal/smdclient/FakeSMDClient_test.go +++ b/internal/smdclient/FakeSMDClient_test.go @@ -4,6 +4,10 @@ import ( "net" "strings" "testing" + + "github.com/OpenCHAMI/smd/v2/pkg/sm" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestIncrementXname(t *testing.T) { @@ -130,3 +134,50 @@ func TestNewFakeSMDClient(t *testing.T) { t.Errorf("expected 9 components in io group, got %d", len(client.groups["io"])) } } + +func TestFakeSMDClientListMemberships(t *testing.T) { + client := NewFakeSMDClient("fake", 12) + representativeIDs := []string{ + client.rosetta_mapping[0].ComponentID, + client.rosetta_mapping[len(client.rosetta_mapping)-1].ComponentID, + } + + memberships := client.ListMemberships() + require.Len(t, memberships, len(client.components)) + for _, id := range representativeIDs { + expected, err := client.GroupMembership(id) + require.NoError(t, err) + actual := membershipByID(t, memberships, id) + assert.ElementsMatch(t, expected, actual.GroupLabels) + assert.Empty(t, actual.PartitionName) + } + + ungroupedID := representativeIDs[0] + for group, componentIDs := range client.groups { + filtered := make([]string, 0, len(componentIDs)) + for _, id := range componentIDs { + if id != ungroupedID { + filtered = append(filtered, id) + } + } + client.groups[group] = filtered + } + ungrouped := membershipByID(t, client.ListMemberships(), ungroupedID) + assert.Empty(t, ungrouped.GroupLabels) + assert.NotNil(t, ungrouped.GroupLabels) + + require.NoError(t, client.AddNodeToGroups(ungroupedID, []string{"new-group"})) + updated := membershipByID(t, client.ListMemberships(), ungroupedID) + assert.Equal(t, []string{"new-group"}, updated.GroupLabels) +} + +func membershipByID(t *testing.T, memberships []sm.Membership, id string) sm.Membership { + t.Helper() + for _, membership := range memberships { + if membership.ID == id { + return membership + } + } + t.Fatalf("membership for %s not found", id) + return sm.Membership{} +} diff --git a/internal/smdclient/SMDclient.go b/internal/smdclient/SMDclient.go index 0e126c306..0c4084367 100644 --- a/internal/smdclient/SMDclient.go +++ b/internal/smdclient/SMDclient.go @@ -283,16 +283,21 @@ func (s *SMDClient) PopulateNodes() { } } - // Populate group membership for all nodes log.Debug().Msg("Fetching group membership for all nodes") + memberships := make([]sm.Membership, 0) + if err := s.getSMD("/hsm/v2/memberships?type=node", &memberships); err != nil { + log.Error().Err(err).Msg("Failed to get SMD node memberships") + return + } + + groupsByXname := make(map[string][]string, len(memberships)) + for _, membership := range memberships { + groupsByXname[membership.ID] = membership.GroupLabels + } for xname, node := range nextNodes { - ml := new(sm.Membership) - membershipEp := "/hsm/v2/memberships/" + xname - if err := s.getSMD(membershipEp, ml); err != nil { - log.Debug().Err(err).Msgf("Failed to get group membership for %s", xname) - node.Groups = []string{} // Empty groups if fetch fails - } else { - node.Groups = ml.GroupLabels + node.Groups = []string{} + if groups, found := groupsByXname[xname]; found && groups != nil { + node.Groups = groups } nextNodes[xname] = node } diff --git a/internal/smdclient/SMDclient_performance_test.go b/internal/smdclient/SMDclient_performance_test.go index 4c147fed5..45d4f3df6 100644 --- a/internal/smdclient/SMDclient_performance_test.go +++ b/internal/smdclient/SMDclient_performance_test.go @@ -14,111 +14,121 @@ import ( ) 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) }) + tests := []struct { + name string + blockedPath string + }{ + {name: "blocked inventory request", blockedPath: "/hsm/v2/Inventory/EthernetInterfaces/"}, + {name: "blocked bulk membership request", blockedPath: "/hsm/v2/memberships"}, } - 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() + for _, tt := range tests { + t.Run(tt.name, func(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 == tt.blockedPath && blockRefresh.Load() { + signalRefreshStarted.Do(func() { close(refreshStarted) }) + <-releaseRefresh + } - 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") - } + 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": + if got := r.URL.Query().Get("type"); got != "node" { + t.Errorf("membership type query = %q, want node", got) + } + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + } + }) + 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") - } + 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") - }() + 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") - } + 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") - } + 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) + 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: @@ -140,8 +150,8 @@ func TestGroupMembershipCached(t *testing.T) { "Description": "Test Node 1" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute", "cabinet1"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute","cabinet1"],"partitionName":""}]`)) } }) server := httptest.NewServer(handler) @@ -161,6 +171,7 @@ func TestGroupMembershipCached(t *testing.T) { initialRequests := requestCount client.PopulateNodes() populateRequests := requestCount - initialRequests + assert.Equal(t, 2, populateRequests) // Verify group membership was cached groups, err := client.GroupMembership("x1000") @@ -207,10 +218,11 @@ func TestConcurrentReads(t *testing.T) { "Description": "Test Node 2" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1001": - _, _ = w.Write([]byte(`{"GroupLabels": ["io"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[ + {"id":"x1000","groupLabels":["compute"],"partitionName":""}, + {"id":"x1001","groupLabels":["io"],"partitionName":""} + ]`)) } }) server := httptest.NewServer(handler) @@ -289,8 +301,10 @@ func TestReverseIndexPerformance(t *testing.T) { if r.URL.Path == "/hsm/v2/Inventory/EthernetInterfaces/" { _, _ = w.Write([]byte(ethInterfaces)) - } else if len(r.URL.Path) >= len("/hsm/v2/memberships/") && r.URL.Path[:len("/hsm/v2/memberships/")] == "/hsm/v2/memberships/" { - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + } else if r.URL.Path == "/hsm/v2/memberships" { + _, _ = w.Write([]byte(bulkMembershipsJSON(nodeCount))) + } else { + http.NotFound(w, r) } }) server := httptest.NewServer(handler) @@ -369,8 +383,8 @@ func TestCaseInsensitiveLookup(t *testing.T) { "Description": "Test Node" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) } }) server := httptest.NewServer(handler) @@ -430,8 +444,8 @@ func TestAddWGIPUpdatesReverseIndex(t *testing.T) { "Description": "Test Node" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) } }) server := httptest.NewServer(handler) @@ -490,8 +504,10 @@ func BenchmarkIDfromIP(b *testing.B) { } ethInterfaces += "]" _, _ = w.Write([]byte(ethInterfaces)) - } else if len(r.URL.Path) >= len("/hsm/v2/memberships/") && r.URL.Path[:len("/hsm/v2/memberships/")] == "/hsm/v2/memberships/" { - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + } else if r.URL.Path == "/hsm/v2/memberships" { + _, _ = w.Write([]byte(bulkMembershipsJSON(1000))) + } else { + http.NotFound(w, r) } }) server := httptest.NewServer(handler) @@ -532,8 +548,8 @@ func BenchmarkGroupMembership(b *testing.B) { "Description": "Test Node" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute", "cabinet1", "rack1"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute","cabinet1","rack1"],"partitionName":""}]`)) } }) server := httptest.NewServer(handler) @@ -556,3 +572,14 @@ func BenchmarkGroupMembership(b *testing.B) { _, _ = client.GroupMembership("x1000") } } + +func bulkMembershipsJSON(nodeCount int) string { + memberships := "[" + for i := 0; i < nodeCount; i++ { + if i > 0 { + memberships += "," + } + memberships += fmt.Sprintf(`{"id":"x%d","groupLabels":["compute"],"partitionName":""}`, i) + } + return memberships + "]" +} diff --git a/internal/smdclient/SMDclient_test.go b/internal/smdclient/SMDclient_test.go index ddfe4f806..fa48ce9c8 100644 --- a/internal/smdclient/SMDclient_test.go +++ b/internal/smdclient/SMDclient_test.go @@ -4,16 +4,26 @@ import ( "errors" "net/http" "net/http/httptest" + "strings" "sync" + "sync/atomic" "testing" "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestPopulateNodes(t *testing.T) { + var requestsMutex sync.Mutex + requests := make([]string, 0, 2) + // Mock SMD server handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestsMutex.Lock() + requests = append(requests, r.URL.RequestURI()) + requestsMutex.Unlock() + w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) @@ -51,14 +61,14 @@ func TestPopulateNodes(t *testing.T) { "Description": "Test Node 4 Interface 2" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1001": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute", "io"]}`)) - case "/hsm/v2/memberships/x1002": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1003": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute", "cabinet1"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[ + {"id":"x1000","groupLabels":["compute"],"partitionName":"ignored"}, + {"id":"x1001","groupLabels":["compute","io"],"partitionName":""}, + {"id":"x1002","groupLabels":["compute"],"partitionName":""}, + {"id":"x1003","groupLabels":["compute","cabinet1"],"partitionName":""}, + {"id":"x9999","groupLabels":["unrelated"],"partitionName":""} + ]`)) } }) server := httptest.NewServer(handler) @@ -111,6 +121,17 @@ func TestPopulateNodes(t *testing.T) { assert.True(t, exists) assert.Equal(t, "x1003", node4.Xname) assert.Equal(t, 2, len(node4.Interfaces)) + assert.Equal(t, []string{"compute", "cabinet1"}, node4.Groups) + + requestsMutex.Lock() + defer requestsMutex.Unlock() + require.Equal(t, []string{ + "/hsm/v2/Inventory/EthernetInterfaces/", + "/hsm/v2/memberships?type=node", + }, requests) + for _, request := range requests { + assert.False(t, strings.HasPrefix(request, "/hsm/v2/memberships/"), "unexpected per-node membership request: %s", request) + } } func TestIPfromID(t *testing.T) { // Mock SMD server @@ -152,14 +173,13 @@ func TestIPfromID(t *testing.T) { "Description": "Test Node 4 Interface 2" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1001": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1002": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1003": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[ + {"id":"x1000","groupLabels":["compute"],"partitionName":""}, + {"id":"x1001","groupLabels":["compute"],"partitionName":""}, + {"id":"x1002","groupLabels":["compute"],"partitionName":""}, + {"id":"x1003","groupLabels":["compute"],"partitionName":""} + ]`)) } }) server := httptest.NewServer(handler) @@ -245,14 +265,13 @@ func TestIDfromIP(t *testing.T) { "Description": "Test Node 4 Interface 2" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1001": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1002": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1003": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[ + {"id":"x1000","groupLabels":["compute"],"partitionName":""}, + {"id":"x1001","groupLabels":["compute"],"partitionName":""}, + {"id":"x1002","groupLabels":["compute"],"partitionName":""}, + {"id":"x1003","groupLabels":["compute"],"partitionName":""} + ]`)) } }) server := httptest.NewServer(handler) @@ -338,14 +357,13 @@ func TestIDfromMAC(t *testing.T) { "Description": "Test Node 4 Interface 2" } ]`)) - case "/hsm/v2/memberships/x1000": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1001": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1002": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) - case "/hsm/v2/memberships/x1003": - _, _ = w.Write([]byte(`{"GroupLabels": ["compute"]}`)) + case "/hsm/v2/memberships": + _, _ = w.Write([]byte(`[ + {"id":"x1000","groupLabels":["compute"],"partitionName":""}, + {"id":"x1001","groupLabels":["compute"],"partitionName":""}, + {"id":"x1002","groupLabels":["compute"],"partitionName":""}, + {"id":"x1003","groupLabels":["compute"],"partitionName":""} + ]`)) } }) server := httptest.NewServer(handler) @@ -391,3 +409,132 @@ func TestIDfromMAC(t *testing.T) { }) } } + +func TestPopulateNodesMissingMembershipUsesEmptyGroups(t *testing.T) { + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + 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"}]}, + {"ComponentID":"x1001","MACAddress":"00:11:22:33:44:66","IPAddresses":[{"IPAddress":"192.168.1.2"}]} + ]`)) + case "/hsm/v2/memberships": + if got := r.URL.Query().Get("type"); got != "node" { + t.Errorf("membership type query = %q, want node", got) + } + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":"ignored"}]`)) + default: + t.Errorf("unexpected SMD request: %s", r.URL.RequestURI()) + w.WriteHeader(http.StatusNotFound) + } + }) + server := httptest.NewServer(handler) + defer server.Close() + + client := newTestSMDClient(server) + client.PopulateNodes() + + groups, err := client.GroupMembership("x1001") + require.NoError(t, err) + assert.Empty(t, groups) + assert.NotNil(t, groups) +} + +func TestPopulateNodesBulkMembershipFailurePreservesCache(t *testing.T) { + tests := []struct { + name string + writeFailure func(http.ResponseWriter) + }{ + { + name: "HTTP failure", + writeFailure: func(w http.ResponseWriter) { + w.WriteHeader(http.StatusServiceUnavailable) + _, _ = w.Write([]byte(`{"error":"unavailable"}`)) + }, + }, + { + name: "malformed JSON", + writeFailure: func(w http.ResponseWriter) { + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":`)) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var failMemberships atomic.Bool + var perNodeRequests atomic.Int32 + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + 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"}]}]`)) + case "/hsm/v2/memberships": + if failMemberships.Load() { + tt.writeFailure(w) + return + } + _, _ = w.Write([]byte(`[{"id":"x1000","groupLabels":["compute"],"partitionName":""}]`)) + default: + if strings.HasPrefix(r.URL.Path, "/hsm/v2/memberships/") { + perNodeRequests.Add(1) + } + w.WriteHeader(http.StatusNotFound) + } + }) + server := httptest.NewServer(handler) + defer server.Close() + + client := newTestSMDClient(server) + client.PopulateNodes() + require.NoError(t, client.AddWGIP("x1000", "10.99.0.1")) + + client.nodesMutex.RLock() + oldTimestamp := client.nodes_last_update + client.nodesMutex.RUnlock() + failMemberships.Store(true) + client.PopulateNodes() + + groups, err := client.GroupMembership("x1000") + require.NoError(t, err) + assert.Equal(t, []string{"compute"}, groups) + assert.Equal(t, "x1000", mustIDfromIP(t, client, "192.168.1.1")) + assert.Equal(t, "x1000", mustIDfromIP(t, client, "10.99.0.1")) + assert.Equal(t, "x1000", mustIDfromMAC(t, client, "00:11:22:33:44:55")) + wgip, err := client.WGIPfromID("x1000") + require.NoError(t, err) + assert.Equal(t, "10.99.0.1", wgip) + client.nodesMutex.RLock() + assert.Equal(t, oldTimestamp, client.nodes_last_update) + client.nodesMutex.RUnlock() + assert.Zero(t, perNodeRequests.Load()) + }) + } +} + +func newTestSMDClient(server *httptest.Server) *SMDClient { + return &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), + } +} + +func mustIDfromIP(t *testing.T, client *SMDClient, ip string) string { + t.Helper() + id, err := client.IDfromIP(ip) + require.NoError(t, err) + return id +} + +func mustIDfromMAC(t *testing.T, client *SMDClient, mac string) string { + t.Helper() + id, err := client.IDfromMAC(mac) + require.NoError(t, err) + return id +} From 5d9d01420285432f6083ddb1670ecbb51b2a844b Mon Sep 17 00:00:00 2001 From: Alex Lovell-Troy Date: Mon, 3 Aug 2026 16:42:46 -0400 Subject: [PATCH 2/2] refactor: replace if-else with switch for URL path handling in performance tests Signed-off-by: Alex Lovell-Troy --- internal/smdclient/SMDclient_performance_test.go | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/internal/smdclient/SMDclient_performance_test.go b/internal/smdclient/SMDclient_performance_test.go index 45d4f3df6..108d9081b 100644 --- a/internal/smdclient/SMDclient_performance_test.go +++ b/internal/smdclient/SMDclient_performance_test.go @@ -299,11 +299,12 @@ func TestReverseIndexPerformance(t *testing.T) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) - if r.URL.Path == "/hsm/v2/Inventory/EthernetInterfaces/" { + switch r.URL.Path { + case "/hsm/v2/Inventory/EthernetInterfaces/": _, _ = w.Write([]byte(ethInterfaces)) - } else if r.URL.Path == "/hsm/v2/memberships" { + case "/hsm/v2/memberships": _, _ = w.Write([]byte(bulkMembershipsJSON(nodeCount))) - } else { + default: http.NotFound(w, r) } }) @@ -488,7 +489,8 @@ func BenchmarkIDfromIP(b *testing.B) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusOK) - if r.URL.Path == "/hsm/v2/Inventory/EthernetInterfaces/" { + switch r.URL.Path { + case "/hsm/v2/Inventory/EthernetInterfaces/": // Generate 1000 nodes ethInterfaces := "[" for i := 0; i < 1000; i++ { @@ -504,9 +506,9 @@ func BenchmarkIDfromIP(b *testing.B) { } ethInterfaces += "]" _, _ = w.Write([]byte(ethInterfaces)) - } else if r.URL.Path == "/hsm/v2/memberships" { + case "/hsm/v2/memberships": _, _ = w.Write([]byte(bulkMembershipsJSON(1000))) - } else { + default: http.NotFound(w, r) } })